mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: support context caching
1. add a context cache config in app level which will apply to all agents in the app 2. pass on cache config through invocation context to llm_reqeust 3. store cache metadata in llm_response 4. lookup old cache metadata from latest event for reusing old cache 5. create new cache if old cache cannot be reused PiperOrigin-RevId: 809158578
This commit is contained in:
committed by
Copybara-Service
parent
13a95c463d
commit
c66245a3b8
@@ -0,0 +1,84 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from pydantic import ConfigDict
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
|
from ..utils.feature_decorator import experimental
|
||||||
|
|
||||||
|
|
||||||
|
@experimental
|
||||||
|
class ContextCacheConfig(BaseModel):
|
||||||
|
"""Configuration for context caching across all agents in an app.
|
||||||
|
|
||||||
|
This configuration enables and controls context caching behavior for
|
||||||
|
all LLM agents in an app. When this config is present on an app, context
|
||||||
|
caching is enabled for all agents. When absent (None), context caching
|
||||||
|
is disabled.
|
||||||
|
|
||||||
|
Context caching can significantly reduce costs and improve response times
|
||||||
|
by reusing previously processed context across multiple requests.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
cache_intervals: Maximum number of invocations to reuse the same cache before refreshing it
|
||||||
|
ttl_seconds: Time-to-live for cache in seconds
|
||||||
|
min_tokens: Minimum tokens required to enable caching
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_config = ConfigDict(
|
||||||
|
extra="forbid",
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_intervals: int = Field(
|
||||||
|
default=10,
|
||||||
|
ge=1,
|
||||||
|
le=100,
|
||||||
|
description=(
|
||||||
|
"Maximum number of invocations to reuse the same cache before"
|
||||||
|
" refreshing it"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
ttl_seconds: int = Field(
|
||||||
|
default=1800, # 30 minutes
|
||||||
|
gt=0,
|
||||||
|
description="Time-to-live for cache in seconds",
|
||||||
|
)
|
||||||
|
|
||||||
|
min_tokens: int = Field(
|
||||||
|
default=0,
|
||||||
|
ge=0,
|
||||||
|
description=(
|
||||||
|
"Minimum estimated request tokens required to enable caching. This"
|
||||||
|
" compares against the estimated total tokens of the request (system"
|
||||||
|
" instruction + tools + contents). Context cache storage may have"
|
||||||
|
" cost. Set higher to avoid caching small requests where overhead may"
|
||||||
|
" exceed benefits."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def ttl_string(self) -> str:
|
||||||
|
"""Get TTL as string format for cache creation."""
|
||||||
|
return f"{self.ttl_seconds}s"
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
"""String representation for logging."""
|
||||||
|
return (
|
||||||
|
f"ContextCacheConfig(cache_intervals={self.cache_intervals}, "
|
||||||
|
f"ttl={self.ttl_seconds}s, min_tokens={self.min_tokens})"
|
||||||
|
)
|
||||||
@@ -15,6 +15,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
from google.genai import types
|
from google.genai import types
|
||||||
@@ -33,6 +34,7 @@ from ..sessions.session import Session
|
|||||||
from ..utils.feature_decorator import working_in_progress
|
from ..utils.feature_decorator import working_in_progress
|
||||||
from .active_streaming_tool import ActiveStreamingTool
|
from .active_streaming_tool import ActiveStreamingTool
|
||||||
from .base_agent import BaseAgent
|
from .base_agent import BaseAgent
|
||||||
|
from .context_cache_config import ContextCacheConfig
|
||||||
from .live_request_queue import LiveRequestQueue
|
from .live_request_queue import LiveRequestQueue
|
||||||
from .run_config import RunConfig
|
from .run_config import RunConfig
|
||||||
from .transcription_entry import TranscriptionEntry
|
from .transcription_entry import TranscriptionEntry
|
||||||
@@ -141,6 +143,7 @@ class InvocationContext(BaseModel):
|
|||||||
session_service: BaseSessionService
|
session_service: BaseSessionService
|
||||||
memory_service: Optional[BaseMemoryService] = None
|
memory_service: Optional[BaseMemoryService] = None
|
||||||
credential_service: Optional[BaseCredentialService] = None
|
credential_service: Optional[BaseCredentialService] = None
|
||||||
|
context_cache_config: Optional[ContextCacheConfig] = None
|
||||||
|
|
||||||
invocation_id: str
|
invocation_id: str
|
||||||
"""The id of this invocation context. Readonly."""
|
"""The id of this invocation context. Readonly."""
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from pydantic import ConfigDict
|
|||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
from ..agents.base_agent import BaseAgent
|
from ..agents.base_agent import BaseAgent
|
||||||
|
from ..agents.context_cache_config import ContextCacheConfig
|
||||||
from ..apps.base_events_compactor import BaseEventsCompactor
|
from ..apps.base_events_compactor import BaseEventsCompactor
|
||||||
from ..plugins.base_plugin import BasePlugin
|
from ..plugins.base_plugin import BasePlugin
|
||||||
from ..utils.feature_decorator import experimental
|
from ..utils.feature_decorator import experimental
|
||||||
@@ -53,3 +54,6 @@ class App(BaseModel):
|
|||||||
|
|
||||||
event_compactor: Optional[BaseEventsCompactor] = None
|
event_compactor: Optional[BaseEventsCompactor] = None
|
||||||
"""The event compactor strategy for the application."""
|
"""The event compactor strategy for the application."""
|
||||||
|
|
||||||
|
context_cache_config: Optional[ContextCacheConfig] = None
|
||||||
|
"""Context cache configuration that applies to all LLM agents in the app."""
|
||||||
|
|||||||
@@ -0,0 +1,128 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Context cache processor for LLM requests."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import AsyncGenerator
|
||||||
|
from typing import Optional
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from ...events.event import Event
|
||||||
|
from ...models.cache_metadata import CacheMetadata
|
||||||
|
from ._base_llm_processor import BaseLlmRequestProcessor
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ...agents.invocation_context import InvocationContext
|
||||||
|
from ...models.llm_request import LlmRequest
|
||||||
|
|
||||||
|
logger = logging.getLogger('google_adk.' + __name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ContextCacheRequestProcessor(BaseLlmRequestProcessor):
|
||||||
|
"""Request processor that enables context caching for LLM requests.
|
||||||
|
|
||||||
|
This processor sets up context caching configuration for agents that have
|
||||||
|
context caching enabled and finds the latest cache metadata from session
|
||||||
|
events. The actual cache management is handled by the model-specific cache
|
||||||
|
managers (e.g., GeminiContextCacheManager).
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def run_async(
|
||||||
|
self, invocation_context: 'InvocationContext', llm_request: 'LlmRequest'
|
||||||
|
) -> AsyncGenerator[Event, None]:
|
||||||
|
"""Process LLM request to enable context caching.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
invocation_context: Invocation context containing agent and session info
|
||||||
|
llm_request: Request to process for caching
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
Event: No events are yielded by this processor
|
||||||
|
"""
|
||||||
|
agent = invocation_context.agent
|
||||||
|
|
||||||
|
# Return early if no cache config
|
||||||
|
if not invocation_context.context_cache_config:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Set cache config to request
|
||||||
|
llm_request.cache_config = invocation_context.context_cache_config
|
||||||
|
|
||||||
|
# Find latest cache metadata from session events
|
||||||
|
latest_cache_metadata = self._find_latest_cache_metadata(
|
||||||
|
invocation_context, agent.name, invocation_context.invocation_id
|
||||||
|
)
|
||||||
|
|
||||||
|
if latest_cache_metadata:
|
||||||
|
llm_request.cache_metadata = latest_cache_metadata
|
||||||
|
logger.debug(
|
||||||
|
'Found cache metadata for agent %s: invocations_used=%d, '
|
||||||
|
'cached_contents=%d',
|
||||||
|
agent.name,
|
||||||
|
latest_cache_metadata.invocations_used,
|
||||||
|
latest_cache_metadata.cached_contents_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.debug('Context caching enabled for agent %s', agent.name)
|
||||||
|
|
||||||
|
# This processor yields no events
|
||||||
|
return
|
||||||
|
yield # AsyncGenerator requires a yield in function body
|
||||||
|
|
||||||
|
def _find_latest_cache_metadata(
|
||||||
|
self,
|
||||||
|
invocation_context: 'InvocationContext',
|
||||||
|
agent_name: str,
|
||||||
|
current_invocation_id: str,
|
||||||
|
) -> Optional[CacheMetadata]:
|
||||||
|
"""Find the latest cache metadata from session events.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
invocation_context: Context containing session with events
|
||||||
|
agent_name: Name of agent to find cache metadata for
|
||||||
|
current_invocation_id: Current invocation ID to compare for increment
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Latest cache metadata for the agent (with updated invocations_used
|
||||||
|
if needed), or None if not found
|
||||||
|
"""
|
||||||
|
if not invocation_context.session or not invocation_context.session.events:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Search events from most recent to oldest using index traversal
|
||||||
|
events = invocation_context.session.events
|
||||||
|
for i in range(len(events) - 1, -1, -1):
|
||||||
|
event = events[i]
|
||||||
|
if event.cache_metadata is not None and event.author == agent_name:
|
||||||
|
|
||||||
|
cache_metadata = event.cache_metadata
|
||||||
|
|
||||||
|
# Check if this is a different invocation - increment invocations_used
|
||||||
|
if event.invocation_id and event.invocation_id != current_invocation_id:
|
||||||
|
# Different invocation - increment invocations_used
|
||||||
|
return cache_metadata.model_copy(
|
||||||
|
update={'invocations_used': cache_metadata.invocations_used + 1}
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Same invocation or no invocation_id - return as-is
|
||||||
|
return cache_metadata
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# Create processor instance for use in flows
|
||||||
|
request_processor = ContextCacheRequestProcessor()
|
||||||
@@ -23,6 +23,7 @@ from . import _nl_planning
|
|||||||
from . import _output_schema_processor
|
from . import _output_schema_processor
|
||||||
from . import basic
|
from . import basic
|
||||||
from . import contents
|
from . import contents
|
||||||
|
from . import context_cache_processor
|
||||||
from . import identity
|
from . import identity
|
||||||
from . import instructions
|
from . import instructions
|
||||||
from . import request_confirmation
|
from . import request_confirmation
|
||||||
@@ -48,6 +49,8 @@ class SingleFlow(BaseLlmFlow):
|
|||||||
instructions.request_processor,
|
instructions.request_processor,
|
||||||
identity.request_processor,
|
identity.request_processor,
|
||||||
contents.request_processor,
|
contents.request_processor,
|
||||||
|
# Context cache processor sets up cache config and finds existing cache metadata
|
||||||
|
context_cache_processor.request_processor,
|
||||||
# Some implementations of NL Planning mark planning contents as thoughts
|
# Some implementations of NL Planning mark planning contents as thoughts
|
||||||
# in the post processor. Since these need to be unmarked, NL Planning
|
# in the post processor. Since these need to be unmarked, NL Planning
|
||||||
# should be after contents.
|
# should be after contents.
|
||||||
|
|||||||
@@ -0,0 +1,90 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from pydantic import ConfigDict
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
|
|
||||||
|
class CacheMetadata(BaseModel):
|
||||||
|
"""Metadata for context cache associated with LLM responses.
|
||||||
|
|
||||||
|
This class stores cache identification, usage tracking, and lifecycle
|
||||||
|
information for a particular cache instance.
|
||||||
|
|
||||||
|
Token counts (cached and total) are available in the LlmResponse.usage_metadata
|
||||||
|
and should be accessed from there to avoid duplication.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
cache_name: The full resource name of the cached content (e.g.,
|
||||||
|
'projects/123/locations/us-central1/cachedContents/456')
|
||||||
|
expire_time: Unix timestamp when the cache expires
|
||||||
|
fingerprint: Hash of agent configuration (instruction + tools + model)
|
||||||
|
invocations_used: Number of invocations this cache has been used for
|
||||||
|
cached_contents_count: Number of contents stored in this cache
|
||||||
|
created_at: Unix timestamp when the cache was created
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_config = ConfigDict(
|
||||||
|
extra="forbid",
|
||||||
|
frozen=True, # Cache metadata should be immutable
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_name: str = Field(
|
||||||
|
description="Full resource name of the cached content"
|
||||||
|
)
|
||||||
|
|
||||||
|
expire_time: float = Field(description="Unix timestamp when cache expires")
|
||||||
|
|
||||||
|
fingerprint: str = Field(
|
||||||
|
description="Hash of agent configuration used to detect changes"
|
||||||
|
)
|
||||||
|
|
||||||
|
invocations_used: int = Field(
|
||||||
|
ge=0,
|
||||||
|
description="Number of invocations this cache has been used for",
|
||||||
|
)
|
||||||
|
|
||||||
|
cached_contents_count: int = Field(
|
||||||
|
ge=0,
|
||||||
|
description="Number of contents stored in this cache",
|
||||||
|
)
|
||||||
|
|
||||||
|
created_at: Optional[float] = Field(
|
||||||
|
default=None,
|
||||||
|
description=(
|
||||||
|
"Unix timestamp when cache was created (None if reused existing)"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def expire_soon(self) -> bool:
|
||||||
|
"""Check if the cache will expire soon (with 2-minute buffer)."""
|
||||||
|
buffer_seconds = 120 # 2 minutes buffer for processing time
|
||||||
|
return time.time() > (self.expire_time - buffer_seconds)
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
"""String representation for logging and debugging."""
|
||||||
|
cache_id = self.cache_name.split("/")[-1]
|
||||||
|
time_until_expiry_minutes = (self.expire_time - time.time()) / 60
|
||||||
|
return (
|
||||||
|
f"Cache {cache_id}: used {self.invocations_used} invocations, "
|
||||||
|
f"cached {self.cached_contents_count} contents, "
|
||||||
|
f"expires in {time_until_expiry_minutes:.1f}min"
|
||||||
|
)
|
||||||
@@ -0,0 +1,391 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Manages context cache lifecycle for Gemini models."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from google.genai import Client
|
||||||
|
from google.genai import types
|
||||||
|
|
||||||
|
from .cache_metadata import CacheMetadata
|
||||||
|
from .llm_request import LlmRequest
|
||||||
|
from .llm_response import LlmResponse
|
||||||
|
|
||||||
|
logger = logging.getLogger("google_adk." + __name__)
|
||||||
|
|
||||||
|
|
||||||
|
class GeminiContextCacheManager:
|
||||||
|
"""Manages context cache lifecycle for Gemini models.
|
||||||
|
|
||||||
|
This manager handles cache creation, validation, cleanup, and metadata
|
||||||
|
population for Gemini context caching. It uses content hashing to determine
|
||||||
|
cache compatibility and implements efficient caching strategies.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, genai_client: Client):
|
||||||
|
"""Initialize cache manager with shared client.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
genai_client: The GenAI client to use for cache operations.
|
||||||
|
"""
|
||||||
|
self.genai_client = genai_client
|
||||||
|
|
||||||
|
async def handle_context_caching(
|
||||||
|
self, llm_request: LlmRequest
|
||||||
|
) -> Optional[CacheMetadata]:
|
||||||
|
"""Handle context caching for Gemini models.
|
||||||
|
|
||||||
|
Validates existing cache or creates a new one if needed. Applies
|
||||||
|
the cache to the request by setting cached_content and removing cached
|
||||||
|
contents from the request.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request that may contain cache config and metadata.
|
||||||
|
Modified in-place to use the cache.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Cache metadata to be included in response, or None if caching failed
|
||||||
|
"""
|
||||||
|
# Check if we have existing cache metadata and if it's valid
|
||||||
|
if llm_request.cache_metadata:
|
||||||
|
if await self._is_cache_valid(llm_request):
|
||||||
|
# Valid cache found - use it
|
||||||
|
cache_name = llm_request.cache_metadata.cache_name
|
||||||
|
cache_contents_count = llm_request.cache_metadata.cached_contents_count
|
||||||
|
self._apply_cache_to_request(
|
||||||
|
llm_request, cache_name, cache_contents_count
|
||||||
|
)
|
||||||
|
return llm_request.cache_metadata.model_copy()
|
||||||
|
else:
|
||||||
|
# Invalid cache - clean it up
|
||||||
|
await self.cleanup_cache(llm_request.cache_metadata.cache_name)
|
||||||
|
llm_request.cache_metadata = None
|
||||||
|
|
||||||
|
# Find contents to cache for new cache creation
|
||||||
|
cache_contents_count = self._find_count_of_contents_to_cache(
|
||||||
|
llm_request.contents
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create new cache with the determined contents
|
||||||
|
cache_metadata = await self._create_new_cache_with_contents(
|
||||||
|
llm_request, cache_contents_count
|
||||||
|
)
|
||||||
|
if not cache_metadata:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Set up request to use the new cache
|
||||||
|
self._apply_cache_to_request(
|
||||||
|
llm_request, cache_metadata.cache_name, cache_contents_count
|
||||||
|
)
|
||||||
|
|
||||||
|
return cache_metadata
|
||||||
|
|
||||||
|
def _find_count_of_contents_to_cache(
|
||||||
|
self, contents: list[types.Content]
|
||||||
|
) -> int:
|
||||||
|
"""Find the number of contents to cache based on user content strategy.
|
||||||
|
|
||||||
|
Strategy: Find the last continuous batch of user contents and cache
|
||||||
|
all contents before them.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
contents: List of contents from the LLM request
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Number of contents to cache (can be 0 if all contents are user contents)
|
||||||
|
"""
|
||||||
|
if not contents:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
# Find the last continuous batch of user contents
|
||||||
|
last_user_batch_start = len(contents)
|
||||||
|
|
||||||
|
# Scan backwards to find the start of the last user content batch
|
||||||
|
for i in range(len(contents) - 1, -1, -1):
|
||||||
|
if contents[i].role == "user":
|
||||||
|
last_user_batch_start = i
|
||||||
|
else:
|
||||||
|
# Found non-user content, stop the batch
|
||||||
|
break
|
||||||
|
|
||||||
|
# Cache all contents before the last user batch
|
||||||
|
# This ensures we always have some user content to send to the API
|
||||||
|
return last_user_batch_start
|
||||||
|
|
||||||
|
async def _is_cache_valid(self, llm_request: LlmRequest) -> bool:
|
||||||
|
"""Check if the cache from request metadata is still valid.
|
||||||
|
|
||||||
|
Validates expiry, cache intervals, and fingerprint compatibility.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request containing cache metadata to validate
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if cache is valid, False otherwise
|
||||||
|
"""
|
||||||
|
cache_metadata = llm_request.cache_metadata
|
||||||
|
if not cache_metadata:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Check if cache has expired
|
||||||
|
if time.time() >= cache_metadata.expire_time:
|
||||||
|
logger.info("Cache expired: %s", cache_metadata.cache_name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Check if cache has been used for too many invocations
|
||||||
|
if (
|
||||||
|
cache_metadata.invocations_used
|
||||||
|
> llm_request.cache_config.cache_intervals
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"Cache exceeded cache intervals: %s (%d > %d intervals)",
|
||||||
|
cache_metadata.cache_name,
|
||||||
|
cache_metadata.invocations_used,
|
||||||
|
llm_request.cache_config.cache_intervals,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Check if fingerprint matches using cached contents count
|
||||||
|
current_fingerprint = self._generate_cache_fingerprint(
|
||||||
|
llm_request, cache_metadata.cached_contents_count
|
||||||
|
)
|
||||||
|
if current_fingerprint != cache_metadata.fingerprint:
|
||||||
|
logger.debug("Cache content fingerprint mismatch")
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _generate_cache_fingerprint(
|
||||||
|
self, llm_request: LlmRequest, cache_contents_count: int
|
||||||
|
) -> str:
|
||||||
|
"""Generate a fingerprint for cache validation.
|
||||||
|
|
||||||
|
Includes system instruction, tools, tool_config, and first N contents.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request to generate fingerprint for
|
||||||
|
cache_contents_count: Number of contents to include in fingerprint
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
16-character hexadecimal fingerprint representing the cached state
|
||||||
|
"""
|
||||||
|
# Create fingerprint from system instruction, tools, tool_config, and first N contents
|
||||||
|
fingerprint_data = {}
|
||||||
|
|
||||||
|
if llm_request.config and llm_request.config.system_instruction:
|
||||||
|
fingerprint_data["system_instruction"] = (
|
||||||
|
llm_request.config.system_instruction
|
||||||
|
)
|
||||||
|
|
||||||
|
if llm_request.config and llm_request.config.tools:
|
||||||
|
# Simplified: just dump types.Tool instances to JSON
|
||||||
|
tools_data = []
|
||||||
|
for tool in llm_request.config.tools:
|
||||||
|
if isinstance(tool, types.Tool):
|
||||||
|
tools_data.append(tool.model_dump())
|
||||||
|
fingerprint_data["tools"] = tools_data
|
||||||
|
|
||||||
|
if llm_request.config and llm_request.config.tool_config:
|
||||||
|
fingerprint_data["tool_config"] = (
|
||||||
|
llm_request.config.tool_config.model_dump()
|
||||||
|
)
|
||||||
|
|
||||||
|
# Include first N contents in fingerprint
|
||||||
|
if cache_contents_count > 0 and llm_request.contents:
|
||||||
|
contents_data = []
|
||||||
|
for i in range(min(cache_contents_count, len(llm_request.contents))):
|
||||||
|
content = llm_request.contents[i]
|
||||||
|
contents_data.append(content.model_dump())
|
||||||
|
fingerprint_data["cached_contents"] = contents_data
|
||||||
|
|
||||||
|
# Generate hash
|
||||||
|
fingerprint_str = json.dumps(fingerprint_data, sort_keys=True)
|
||||||
|
return hashlib.sha256(fingerprint_str.encode()).hexdigest()[:16]
|
||||||
|
|
||||||
|
async def _create_new_cache_with_contents(
|
||||||
|
self, llm_request: LlmRequest, cache_contents_count: int
|
||||||
|
) -> Optional[CacheMetadata]:
|
||||||
|
"""Create a new cache with specified number of contents.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request to create cache for
|
||||||
|
cache_contents_count: Number of contents to include in cache
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Cache metadata if successful, None otherwise
|
||||||
|
"""
|
||||||
|
# Estimate token count for minimum cache size check
|
||||||
|
estimated_tokens = self._estimate_request_tokens(llm_request)
|
||||||
|
if estimated_tokens < llm_request.cache_config.min_tokens:
|
||||||
|
logger.info(
|
||||||
|
"Request too small for caching (%d < %d tokens)",
|
||||||
|
estimated_tokens,
|
||||||
|
llm_request.cache_config.min_tokens,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Create cache using Gemini API directly
|
||||||
|
return await self._create_gemini_cache(llm_request, cache_contents_count)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to create cache: %s", e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _estimate_request_tokens(self, llm_request: LlmRequest) -> int:
|
||||||
|
"""Estimate token count for the request.
|
||||||
|
|
||||||
|
This is a rough estimation based on content text length.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request to estimate tokens for
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Estimated token count
|
||||||
|
"""
|
||||||
|
total_chars = 0
|
||||||
|
|
||||||
|
# System instruction
|
||||||
|
if llm_request.config and llm_request.config.system_instruction:
|
||||||
|
total_chars += len(llm_request.config.system_instruction)
|
||||||
|
|
||||||
|
# Tools
|
||||||
|
if llm_request.config and llm_request.config.tools:
|
||||||
|
for tool in llm_request.config.tools:
|
||||||
|
if isinstance(tool, types.Tool):
|
||||||
|
tool_str = json.dumps(tool.model_dump())
|
||||||
|
total_chars += len(tool_str)
|
||||||
|
|
||||||
|
# Contents
|
||||||
|
for content in llm_request.contents:
|
||||||
|
for part in content.parts:
|
||||||
|
if part.text:
|
||||||
|
total_chars += len(part.text)
|
||||||
|
|
||||||
|
# Rough estimate: 4 characters per token
|
||||||
|
return total_chars // 4
|
||||||
|
|
||||||
|
async def _create_gemini_cache(
|
||||||
|
self, llm_request: LlmRequest, cache_contents_count: int
|
||||||
|
) -> CacheMetadata:
|
||||||
|
"""Create cache using Gemini API.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request to create cache for
|
||||||
|
cache_contents_count: Number of contents to cache
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Cache metadata with precise creation timestamp
|
||||||
|
"""
|
||||||
|
# Prepare cache contents (first N contents + system instruction + tools)
|
||||||
|
cache_contents = llm_request.contents[:cache_contents_count]
|
||||||
|
|
||||||
|
cache_config = types.CreateCachedContentConfig(
|
||||||
|
contents=cache_contents,
|
||||||
|
ttl=llm_request.cache_config.ttl_string,
|
||||||
|
display_name=(
|
||||||
|
f"adk-cache-{int(time.time())}-{cache_contents_count}contents"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Add system instruction if present
|
||||||
|
if llm_request.config and llm_request.config.system_instruction:
|
||||||
|
cache_config.system_instruction = llm_request.config.system_instruction
|
||||||
|
|
||||||
|
# Add tools if present
|
||||||
|
if llm_request.config and llm_request.config.tools:
|
||||||
|
cache_config.tools = llm_request.config.tools
|
||||||
|
|
||||||
|
# Add tool config if present
|
||||||
|
if llm_request.config and llm_request.config.tool_config:
|
||||||
|
cache_config.tool_config = llm_request.config.tool_config
|
||||||
|
|
||||||
|
try:
|
||||||
|
cached_content = await self.genai_client.aio.caches.create(
|
||||||
|
model=llm_request.model,
|
||||||
|
config=cache_config,
|
||||||
|
)
|
||||||
|
# Set precise creation timestamp right after cache creation
|
||||||
|
created_at = time.time()
|
||||||
|
logger.info("Cache created successfully: %s", cached_content.name)
|
||||||
|
|
||||||
|
# Return complete cache metadata with precise timing
|
||||||
|
return CacheMetadata(
|
||||||
|
cache_name=cached_content.name,
|
||||||
|
expire_time=created_at + llm_request.cache_config.ttl_seconds,
|
||||||
|
fingerprint=self._generate_cache_fingerprint(
|
||||||
|
llm_request, cache_contents_count
|
||||||
|
),
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=cache_contents_count,
|
||||||
|
created_at=created_at,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to create Gemini cache: %s", e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def cleanup_cache(self, cache_name: str) -> None:
|
||||||
|
"""Clean up cache by deleting it.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
cache_name: Name of cache to delete
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
await self.genai_client.aio.caches.delete(name=cache_name)
|
||||||
|
logger.info("Cache cleaned up: %s", cache_name)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to cleanup cache %s: %s", cache_name, e)
|
||||||
|
|
||||||
|
def _apply_cache_to_request(
|
||||||
|
self,
|
||||||
|
llm_request: LlmRequest,
|
||||||
|
cache_name: str,
|
||||||
|
cache_contents_count: int,
|
||||||
|
) -> None:
|
||||||
|
"""Apply cache to the request by modifying it to use cached content.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_request: Request to modify
|
||||||
|
cache_name: Name of cache to use
|
||||||
|
cache_contents_count: Number of contents that are cached
|
||||||
|
"""
|
||||||
|
# Remove system instruction, tools, and tool config from request config since they're in cache
|
||||||
|
if llm_request.config:
|
||||||
|
llm_request.config.system_instruction = None
|
||||||
|
llm_request.config.tools = None
|
||||||
|
llm_request.config.tool_config = None
|
||||||
|
|
||||||
|
# Set cached content reference
|
||||||
|
llm_request.config.cached_content = cache_name
|
||||||
|
|
||||||
|
# Remove cached contents from the request (keep only uncached contents)
|
||||||
|
llm_request.contents = llm_request.contents[cache_contents_count:]
|
||||||
|
|
||||||
|
def populate_cache_metadata_in_response(
|
||||||
|
self, llm_response: LlmResponse, cache_metadata: CacheMetadata
|
||||||
|
) -> None:
|
||||||
|
"""Populate cache metadata in LLM response.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
llm_response: Response to populate metadata in
|
||||||
|
cache_metadata: Cache metadata to copy into response
|
||||||
|
"""
|
||||||
|
# Create a copy of cache metadata for the response
|
||||||
|
llm_response.cache_metadata = cache_metadata.model_copy()
|
||||||
@@ -28,7 +28,6 @@ from typing import Union
|
|||||||
|
|
||||||
from google.genai import Client
|
from google.genai import Client
|
||||||
from google.genai import types
|
from google.genai import types
|
||||||
from google.genai.types import FinishReason
|
|
||||||
from typing_extensions import override
|
from typing_extensions import override
|
||||||
|
|
||||||
from .. import version
|
from .. import version
|
||||||
@@ -110,6 +109,16 @@ class Gemini(BaseLlm):
|
|||||||
"""
|
"""
|
||||||
await self._preprocess_request(llm_request)
|
await self._preprocess_request(llm_request)
|
||||||
self._maybe_append_user_content(llm_request)
|
self._maybe_append_user_content(llm_request)
|
||||||
|
|
||||||
|
# Handle context caching if configured
|
||||||
|
cache_metadata = None
|
||||||
|
cache_manager = None
|
||||||
|
if llm_request.cache_config:
|
||||||
|
from .gemini_context_cache_manager import GeminiContextCacheManager
|
||||||
|
|
||||||
|
cache_manager = GeminiContextCacheManager(self.api_client)
|
||||||
|
cache_metadata = await cache_manager.handle_context_caching(llm_request)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
'Sending out request, model: %s, backend: %s, stream: %s',
|
'Sending out request, model: %s, backend: %s, stream: %s',
|
||||||
llm_request.model,
|
llm_request.model,
|
||||||
@@ -150,6 +159,11 @@ class Gemini(BaseLlm):
|
|||||||
async for llm_response in aggregator_gen:
|
async for llm_response in aggregator_gen:
|
||||||
yield llm_response
|
yield llm_response
|
||||||
if (close_result := aggregator.close()) is not None:
|
if (close_result := aggregator.close()) is not None:
|
||||||
|
# Populate cache metadata in the final aggregated response for streaming
|
||||||
|
if cache_metadata:
|
||||||
|
cache_manager.populate_cache_metadata_in_response(
|
||||||
|
close_result, cache_metadata
|
||||||
|
)
|
||||||
yield close_result
|
yield close_result
|
||||||
|
|
||||||
else:
|
else:
|
||||||
@@ -160,7 +174,13 @@ class Gemini(BaseLlm):
|
|||||||
)
|
)
|
||||||
logger.info('Response received from the model.')
|
logger.info('Response received from the model.')
|
||||||
logger.debug(_build_response_log(response))
|
logger.debug(_build_response_log(response))
|
||||||
yield LlmResponse.create(response)
|
|
||||||
|
llm_response = LlmResponse.create(response)
|
||||||
|
if cache_metadata:
|
||||||
|
cache_manager.populate_cache_metadata_in_response(
|
||||||
|
llm_response, cache_metadata
|
||||||
|
)
|
||||||
|
yield llm_response
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def api_client(self) -> Client:
|
def api_client(self) -> Client:
|
||||||
|
|||||||
@@ -21,7 +21,9 @@ from pydantic import BaseModel
|
|||||||
from pydantic import ConfigDict
|
from pydantic import ConfigDict
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
|
from ..agents.context_cache_config import ContextCacheConfig
|
||||||
from ..tools.base_tool import BaseTool
|
from ..tools.base_tool import BaseTool
|
||||||
|
from .cache_metadata import CacheMetadata
|
||||||
|
|
||||||
|
|
||||||
def _find_tool_with_function_declarations(
|
def _find_tool_with_function_declarations(
|
||||||
@@ -52,6 +54,8 @@ class LlmRequest(BaseModel):
|
|||||||
contents: The contents to send to the model.
|
contents: The contents to send to the model.
|
||||||
config: Additional config for the generate content request.
|
config: Additional config for the generate content request.
|
||||||
tools_dict: The tools dictionary.
|
tools_dict: The tools dictionary.
|
||||||
|
cache_config: Context cache configuration for this request.
|
||||||
|
cache_metadata: Cache metadata from previous requests, used for cache management.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||||
@@ -76,6 +80,12 @@ class LlmRequest(BaseModel):
|
|||||||
tools_dict: dict[str, BaseTool] = Field(default_factory=dict, exclude=True)
|
tools_dict: dict[str, BaseTool] = Field(default_factory=dict, exclude=True)
|
||||||
"""The tools dictionary."""
|
"""The tools dictionary."""
|
||||||
|
|
||||||
|
cache_config: Optional[ContextCacheConfig] = None
|
||||||
|
"""Context cache configuration for this request."""
|
||||||
|
|
||||||
|
cache_metadata: Optional[CacheMetadata] = None
|
||||||
|
"""Cache metadata from previous requests, used for cache management."""
|
||||||
|
|
||||||
def append_instructions(self, instructions: list[str]) -> None:
|
def append_instructions(self, instructions: list[str]) -> None:
|
||||||
"""Appends instructions to the system instruction.
|
"""Appends instructions to the system instruction.
|
||||||
|
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ from pydantic import alias_generators
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from pydantic import ConfigDict
|
from pydantic import ConfigDict
|
||||||
|
|
||||||
|
from .cache_metadata import CacheMetadata
|
||||||
|
|
||||||
|
|
||||||
class LlmResponse(BaseModel):
|
class LlmResponse(BaseModel):
|
||||||
"""LLM response class that provides the first candidate response from the
|
"""LLM response class that provides the first candidate response from the
|
||||||
@@ -117,6 +119,13 @@ class LlmResponse(BaseModel):
|
|||||||
logprobs_result: Optional[types.LogprobsResult] = None
|
logprobs_result: Optional[types.LogprobsResult] = None
|
||||||
"""Detailed log probabilities for chosen and top candidate tokens."""
|
"""Detailed log probabilities for chosen and top candidate tokens."""
|
||||||
|
|
||||||
|
cache_metadata: Optional[CacheMetadata] = None
|
||||||
|
"""Context cache metadata if caching was used for this response.
|
||||||
|
|
||||||
|
Contains cache identification, usage tracking, and lifecycle information.
|
||||||
|
This field is automatically populated when context caching is enabled.
|
||||||
|
"""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create(
|
def create(
|
||||||
generate_content_response: types.GenerateContentResponse,
|
generate_content_response: types.GenerateContentResponse,
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ from google.genai import types
|
|||||||
|
|
||||||
from .agents.active_streaming_tool import ActiveStreamingTool
|
from .agents.active_streaming_tool import ActiveStreamingTool
|
||||||
from .agents.base_agent import BaseAgent
|
from .agents.base_agent import BaseAgent
|
||||||
|
from .agents.context_cache_config import ContextCacheConfig
|
||||||
from .agents.invocation_context import InvocationContext
|
from .agents.invocation_context import InvocationContext
|
||||||
from .agents.invocation_context import new_invocation_context_id
|
from .agents.invocation_context import new_invocation_context_id
|
||||||
from .agents.live_request_queue import LiveRequestQueue
|
from .agents.live_request_queue import LiveRequestQueue
|
||||||
@@ -125,8 +126,8 @@ class Runner:
|
|||||||
ValueError: If `app` is provided along with `app_name` or `plugins`, or
|
ValueError: If `app` is provided along with `app_name` or `plugins`, or
|
||||||
if `app` is not provided but either `app_name` or `agent` is missing.
|
if `app` is not provided but either `app_name` or `agent` is missing.
|
||||||
"""
|
"""
|
||||||
self.app_name, self.agent, plugins = self._validate_runner_params(
|
self.app_name, self.agent, self.context_cache_config, plugins = (
|
||||||
app, app_name, agent, plugins
|
self._validate_runner_params(app, app_name, agent, plugins)
|
||||||
)
|
)
|
||||||
self.artifact_service = artifact_service
|
self.artifact_service = artifact_service
|
||||||
self.session_service = session_service
|
self.session_service = session_service
|
||||||
@@ -140,7 +141,9 @@ class Runner:
|
|||||||
app_name: Optional[str],
|
app_name: Optional[str],
|
||||||
agent: Optional[BaseAgent],
|
agent: Optional[BaseAgent],
|
||||||
plugins: Optional[List[BasePlugin]],
|
plugins: Optional[List[BasePlugin]],
|
||||||
) -> tuple[str, BaseAgent, Optional[List[BasePlugin]]]:
|
) -> tuple[
|
||||||
|
str, BaseAgent, Optional[ContextCacheConfig], Optional[List[BasePlugin]]
|
||||||
|
]:
|
||||||
"""Validates and extracts runner parameters.
|
"""Validates and extracts runner parameters.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -150,7 +153,7 @@ class Runner:
|
|||||||
plugins: A list of plugins for the runner.
|
plugins: A list of plugins for the runner.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing (app_name, agent, plugins).
|
A tuple containing (app_name, agent, context_cache_config, plugins).
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ValueError: If parameters are invalid.
|
ValueError: If parameters are invalid.
|
||||||
@@ -170,10 +173,13 @@ class Runner:
|
|||||||
app_name = app.name
|
app_name = app.name
|
||||||
agent = app.root_agent
|
agent = app.root_agent
|
||||||
plugins = app.plugins
|
plugins = app.plugins
|
||||||
|
context_cache_config = app.context_cache_config
|
||||||
elif not app_name or not agent:
|
elif not app_name or not agent:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
'Either app or both app_name and agent must be provided.'
|
'Either app or both app_name and agent must be provided.'
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
context_cache_config = None
|
||||||
|
|
||||||
if plugins:
|
if plugins:
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
@@ -181,7 +187,7 @@ class Runner:
|
|||||||
' to provide plugins instead.',
|
' to provide plugins instead.',
|
||||||
DeprecationWarning,
|
DeprecationWarning,
|
||||||
)
|
)
|
||||||
return app_name, agent, plugins
|
return app_name, agent, context_cache_config, plugins
|
||||||
|
|
||||||
def run(
|
def run(
|
||||||
self,
|
self,
|
||||||
@@ -659,6 +665,7 @@ class Runner:
|
|||||||
memory_service=self.memory_service,
|
memory_service=self.memory_service,
|
||||||
credential_service=self.credential_service,
|
credential_service=self.credential_service,
|
||||||
plugin_manager=self.plugin_manager,
|
plugin_manager=self.plugin_manager,
|
||||||
|
context_cache_config=self.context_cache_config,
|
||||||
invocation_id=invocation_id,
|
invocation_id=invocation_id,
|
||||||
agent=self.agent,
|
agent=self.agent,
|
||||||
session=session,
|
session=session,
|
||||||
|
|||||||
@@ -0,0 +1,166 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Cache performance analysis utilities for ADK context caching system.
|
||||||
|
|
||||||
|
This module provides tools to analyze cache performance metrics from event
|
||||||
|
history, including hit ratios, cost savings, and cache refresh patterns.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from typing import Dict
|
||||||
|
from typing import List
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from google.adk.models.cache_metadata import CacheMetadata
|
||||||
|
from google.adk.sessions.base_session_service import BaseSessionService
|
||||||
|
|
||||||
|
|
||||||
|
class CachePerformanceAnalyzer:
|
||||||
|
"""Analyzes cache performance through event history."""
|
||||||
|
|
||||||
|
def __init__(self, session_service: BaseSessionService):
|
||||||
|
self.session_service = session_service
|
||||||
|
|
||||||
|
async def _get_agent_cache_history(
|
||||||
|
self,
|
||||||
|
session_id: str,
|
||||||
|
user_id: str,
|
||||||
|
app_name: str,
|
||||||
|
agent_name: Optional[str] = None,
|
||||||
|
) -> List[CacheMetadata]:
|
||||||
|
"""Get cache usage history for agent from events.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session_id: Session to analyze
|
||||||
|
user_id: User ID for session lookup
|
||||||
|
app_name: App name for session lookup
|
||||||
|
agent_name: Agent to get history for. If None, gets all cache events.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of cache metadata in chronological order
|
||||||
|
"""
|
||||||
|
session = await self.session_service.get_session(
|
||||||
|
session_id=session_id,
|
||||||
|
app_name=app_name,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
cache_history = []
|
||||||
|
|
||||||
|
for event in session.events:
|
||||||
|
# Check if event has cache metadata and optionally filter by agent
|
||||||
|
if event.cache_metadata is not None and (
|
||||||
|
agent_name is None or event.author == agent_name
|
||||||
|
):
|
||||||
|
cache_history.append(event.cache_metadata)
|
||||||
|
|
||||||
|
return cache_history
|
||||||
|
|
||||||
|
async def analyze_agent_cache_performance(
|
||||||
|
self, session_id: str, user_id: str, app_name: str, agent_name: str
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Analyze cache performance for agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session_id: Session to analyze
|
||||||
|
user_id: User ID for session lookup
|
||||||
|
app_name: App name for session lookup
|
||||||
|
agent_name: Agent to analyze
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Performance analysis dictionary containing:
|
||||||
|
- status: "active" if cache data found, "no_cache_data" if none
|
||||||
|
- requests_with_cache: Number of requests that used caching
|
||||||
|
- avg_invocations_used: Average number of invocations each cache was used
|
||||||
|
- latest_cache: Resource name of most recent cache used
|
||||||
|
- cache_refreshes: Number of unique cache instances created
|
||||||
|
- total_invocations: Total number of invocations across all caches
|
||||||
|
- total_prompt_tokens: Total prompt tokens across all requests
|
||||||
|
- total_cached_tokens: Total cached content tokens across all requests
|
||||||
|
- cache_hit_ratio_percent: Percentage of tokens served from cache
|
||||||
|
- cache_utilization_ratio_percent: Percentage of requests with cache hits
|
||||||
|
- avg_cached_tokens_per_request: Average cached tokens per request
|
||||||
|
- total_requests: Total number of requests processed
|
||||||
|
- requests_with_cache_hits: Number of requests that had cache hits
|
||||||
|
"""
|
||||||
|
cache_history = await self._get_agent_cache_history(
|
||||||
|
session_id, user_id, app_name, agent_name
|
||||||
|
)
|
||||||
|
|
||||||
|
if not cache_history:
|
||||||
|
return {"status": "no_cache_data"}
|
||||||
|
|
||||||
|
# Get all events for token analysis
|
||||||
|
session = await self.session_service.get_session(
|
||||||
|
session_id=session_id,
|
||||||
|
app_name=app_name,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Collect token metrics from events
|
||||||
|
total_prompt_tokens = 0
|
||||||
|
total_cached_tokens = 0
|
||||||
|
requests_with_cache_hits = 0
|
||||||
|
total_requests = 0
|
||||||
|
|
||||||
|
for event in session.events:
|
||||||
|
if event.author == agent_name and event.usage_metadata:
|
||||||
|
total_requests += 1
|
||||||
|
if event.usage_metadata.prompt_token_count:
|
||||||
|
total_prompt_tokens += event.usage_metadata.prompt_token_count
|
||||||
|
if event.usage_metadata.cached_content_token_count:
|
||||||
|
total_cached_tokens += event.usage_metadata.cached_content_token_count
|
||||||
|
requests_with_cache_hits += 1
|
||||||
|
|
||||||
|
# Calculate cache metrics
|
||||||
|
cache_hit_ratio_percent = (
|
||||||
|
(total_cached_tokens / total_prompt_tokens) * 100
|
||||||
|
if total_prompt_tokens > 0
|
||||||
|
else 0.0
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_utilization_ratio_percent = (
|
||||||
|
(requests_with_cache_hits / total_requests) * 100
|
||||||
|
if total_requests > 0
|
||||||
|
else 0.0
|
||||||
|
)
|
||||||
|
|
||||||
|
avg_cached_tokens_per_request = (
|
||||||
|
total_cached_tokens / total_requests if total_requests > 0 else 0.0
|
||||||
|
)
|
||||||
|
|
||||||
|
invocations_used = [c.invocations_used for c in cache_history]
|
||||||
|
total_invocations = sum(invocations_used)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"status": "active",
|
||||||
|
"requests_with_cache": len(cache_history),
|
||||||
|
"avg_invocations_used": (
|
||||||
|
sum(invocations_used) / len(invocations_used)
|
||||||
|
if invocations_used
|
||||||
|
else 0
|
||||||
|
),
|
||||||
|
"latest_cache": cache_history[-1].cache_name,
|
||||||
|
"cache_refreshes": len(set(c.cache_name for c in cache_history)),
|
||||||
|
"total_invocations": total_invocations,
|
||||||
|
"total_prompt_tokens": total_prompt_tokens,
|
||||||
|
"total_cached_tokens": total_cached_tokens,
|
||||||
|
"cache_hit_ratio_percent": cache_hit_ratio_percent,
|
||||||
|
"cache_utilization_ratio_percent": cache_utilization_ratio_percent,
|
||||||
|
"avg_cached_tokens_per_request": avg_cached_tokens_per_request,
|
||||||
|
"total_requests": total_requests,
|
||||||
|
"requests_with_cache_hits": requests_with_cache_hits,
|
||||||
|
}
|
||||||
@@ -0,0 +1,178 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Tests for ContextCacheConfig."""
|
||||||
|
|
||||||
|
from google.adk.agents.context_cache_config import ContextCacheConfig
|
||||||
|
from pydantic import ValidationError
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
class TestContextCacheConfig:
|
||||||
|
"""Test suite for ContextCacheConfig."""
|
||||||
|
|
||||||
|
def test_default_values(self):
|
||||||
|
"""Test that default values are set correctly."""
|
||||||
|
config = ContextCacheConfig()
|
||||||
|
|
||||||
|
assert config.cache_intervals == 10
|
||||||
|
assert config.ttl_seconds == 1800 # 30 minutes
|
||||||
|
assert config.min_tokens == 0
|
||||||
|
|
||||||
|
def test_custom_values(self):
|
||||||
|
"""Test creating config with custom values."""
|
||||||
|
config = ContextCacheConfig(
|
||||||
|
cache_intervals=15, ttl_seconds=3600, min_tokens=1024
|
||||||
|
)
|
||||||
|
|
||||||
|
assert config.cache_intervals == 15
|
||||||
|
assert config.ttl_seconds == 3600
|
||||||
|
assert config.min_tokens == 1024
|
||||||
|
|
||||||
|
def test_cache_intervals_validation(self):
|
||||||
|
"""Test cache_intervals validation constraints."""
|
||||||
|
# Valid range
|
||||||
|
config = ContextCacheConfig(cache_intervals=1)
|
||||||
|
assert config.cache_intervals == 1
|
||||||
|
|
||||||
|
config = ContextCacheConfig(cache_intervals=100)
|
||||||
|
assert config.cache_intervals == 100
|
||||||
|
|
||||||
|
# Invalid: too low
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(cache_intervals=0)
|
||||||
|
assert "greater than or equal to 1" in str(exc_info.value)
|
||||||
|
|
||||||
|
# Invalid: too high
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(cache_intervals=101)
|
||||||
|
assert "less than or equal to 100" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_ttl_seconds_validation(self):
|
||||||
|
"""Test ttl_seconds validation constraints."""
|
||||||
|
# Valid range
|
||||||
|
config = ContextCacheConfig(ttl_seconds=1)
|
||||||
|
assert config.ttl_seconds == 1
|
||||||
|
|
||||||
|
config = ContextCacheConfig(ttl_seconds=86400) # 24 hours
|
||||||
|
assert config.ttl_seconds == 86400
|
||||||
|
|
||||||
|
# Invalid: zero or negative
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(ttl_seconds=0)
|
||||||
|
assert "greater than 0" in str(exc_info.value)
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(ttl_seconds=-1)
|
||||||
|
assert "greater than 0" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_min_tokens_validation(self):
|
||||||
|
"""Test min_tokens validation constraints."""
|
||||||
|
# Valid values
|
||||||
|
config = ContextCacheConfig(min_tokens=0)
|
||||||
|
assert config.min_tokens == 0
|
||||||
|
|
||||||
|
config = ContextCacheConfig(min_tokens=1024)
|
||||||
|
assert config.min_tokens == 1024
|
||||||
|
|
||||||
|
# Invalid: negative
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(min_tokens=-1)
|
||||||
|
assert "greater than or equal to 0" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_ttl_string_property(self):
|
||||||
|
"""Test ttl_string property returns correct format."""
|
||||||
|
config = ContextCacheConfig(ttl_seconds=1800)
|
||||||
|
assert config.ttl_string == "1800s"
|
||||||
|
|
||||||
|
config = ContextCacheConfig(ttl_seconds=3600)
|
||||||
|
assert config.ttl_string == "3600s"
|
||||||
|
|
||||||
|
def test_str_representation(self):
|
||||||
|
"""Test string representation for logging."""
|
||||||
|
config = ContextCacheConfig(
|
||||||
|
cache_intervals=15, ttl_seconds=3600, min_tokens=1024
|
||||||
|
)
|
||||||
|
|
||||||
|
expected = (
|
||||||
|
"ContextCacheConfig(cache_intervals=15, ttl=3600s, min_tokens=1024)"
|
||||||
|
)
|
||||||
|
assert str(config) == expected
|
||||||
|
|
||||||
|
def test_str_representation_defaults(self):
|
||||||
|
"""Test string representation with default values."""
|
||||||
|
config = ContextCacheConfig()
|
||||||
|
|
||||||
|
expected = "ContextCacheConfig(cache_intervals=10, ttl=1800s, min_tokens=0)"
|
||||||
|
assert str(config) == expected
|
||||||
|
|
||||||
|
def test_pydantic_model_validation(self):
|
||||||
|
"""Test that Pydantic model validation works correctly."""
|
||||||
|
# Test extra fields are forbidden
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
ContextCacheConfig(cache_intervals=10, extra_field="not_allowed")
|
||||||
|
assert "extra" in str(exc_info.value).lower()
|
||||||
|
|
||||||
|
def test_field_descriptions(self):
|
||||||
|
"""Test that fields have proper descriptions."""
|
||||||
|
config = ContextCacheConfig()
|
||||||
|
schema = config.model_json_schema()
|
||||||
|
|
||||||
|
assert "cache_intervals" in schema["properties"]
|
||||||
|
assert (
|
||||||
|
"Maximum number of invocations"
|
||||||
|
in schema["properties"]["cache_intervals"]["description"]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "ttl_seconds" in schema["properties"]
|
||||||
|
assert (
|
||||||
|
"Time-to-live for cache"
|
||||||
|
in schema["properties"]["ttl_seconds"]["description"]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "min_tokens" in schema["properties"]
|
||||||
|
assert (
|
||||||
|
"Minimum estimated request tokens"
|
||||||
|
in schema["properties"]["min_tokens"]["description"]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_immutability_config(self):
|
||||||
|
"""Test that the model config is set correctly."""
|
||||||
|
config = ContextCacheConfig()
|
||||||
|
assert config.model_config["extra"] == "forbid"
|
||||||
|
|
||||||
|
def test_realistic_scenarios(self):
|
||||||
|
"""Test realistic configuration scenarios."""
|
||||||
|
# Quick caching for development
|
||||||
|
dev_config = ContextCacheConfig(
|
||||||
|
cache_intervals=5, ttl_seconds=600, min_tokens=0 # 10 minutes
|
||||||
|
)
|
||||||
|
assert dev_config.cache_intervals == 5
|
||||||
|
assert dev_config.ttl_seconds == 600
|
||||||
|
|
||||||
|
# Production caching
|
||||||
|
prod_config = ContextCacheConfig(
|
||||||
|
cache_intervals=20, ttl_seconds=7200, min_tokens=2048 # 2 hours
|
||||||
|
)
|
||||||
|
assert prod_config.cache_intervals == 20
|
||||||
|
assert prod_config.ttl_seconds == 7200
|
||||||
|
assert prod_config.min_tokens == 2048
|
||||||
|
|
||||||
|
# Conservative caching
|
||||||
|
conservative_config = ContextCacheConfig(
|
||||||
|
cache_intervals=3, ttl_seconds=300, min_tokens=4096 # 5 minutes
|
||||||
|
)
|
||||||
|
assert conservative_config.cache_intervals == 3
|
||||||
|
assert conservative_config.ttl_seconds == 300
|
||||||
|
assert conservative_config.min_tokens == 4096
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -15,6 +15,7 @@
|
|||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
from google.adk.agents.base_agent import BaseAgent
|
from google.adk.agents.base_agent import BaseAgent
|
||||||
|
from google.adk.agents.context_cache_config import ContextCacheConfig
|
||||||
from google.adk.apps.app import App
|
from google.adk.apps.app import App
|
||||||
from google.adk.plugins.base_plugin import BasePlugin
|
from google.adk.plugins.base_plugin import BasePlugin
|
||||||
|
|
||||||
@@ -38,3 +39,82 @@ class TestApp:
|
|||||||
assert app.name == "test_app"
|
assert app.name == "test_app"
|
||||||
assert app.root_agent == mock_agent
|
assert app.root_agent == mock_agent
|
||||||
assert app.plugins == [mock_plugin]
|
assert app.plugins == [mock_plugin]
|
||||||
|
|
||||||
|
def test_app_initialization_without_cache_config(self):
|
||||||
|
"""Test that the app is initialized correctly without context cache config."""
|
||||||
|
mock_agent = Mock(spec=BaseAgent)
|
||||||
|
app = App(name="test_app", root_agent=mock_agent)
|
||||||
|
assert app.name == "test_app"
|
||||||
|
assert app.root_agent == mock_agent
|
||||||
|
assert app.context_cache_config is None
|
||||||
|
|
||||||
|
def test_app_initialization_with_cache_config(self):
|
||||||
|
"""Test that the app is initialized correctly with context cache config."""
|
||||||
|
mock_agent = Mock(spec=BaseAgent)
|
||||||
|
cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=15, ttl_seconds=3600, min_tokens=1024
|
||||||
|
)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="test_app",
|
||||||
|
root_agent=mock_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert app.name == "test_app"
|
||||||
|
assert app.root_agent == mock_agent
|
||||||
|
assert app.context_cache_config == cache_config
|
||||||
|
assert app.context_cache_config.cache_intervals == 15
|
||||||
|
assert app.context_cache_config.ttl_seconds == 3600
|
||||||
|
assert app.context_cache_config.min_tokens == 1024
|
||||||
|
|
||||||
|
def test_app_with_all_components(self):
|
||||||
|
"""Test app with all components: agent, plugins, and cache config."""
|
||||||
|
mock_agent = Mock(spec=BaseAgent)
|
||||||
|
mock_plugin = Mock(spec=BasePlugin)
|
||||||
|
cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=20, ttl_seconds=7200, min_tokens=2048
|
||||||
|
)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="full_test_app",
|
||||||
|
root_agent=mock_agent,
|
||||||
|
plugins=[mock_plugin],
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert app.name == "full_test_app"
|
||||||
|
assert app.root_agent == mock_agent
|
||||||
|
assert app.plugins == [mock_plugin]
|
||||||
|
assert app.context_cache_config == cache_config
|
||||||
|
|
||||||
|
def test_app_cache_config_defaults(self):
|
||||||
|
"""Test that cache config has proper defaults when created."""
|
||||||
|
mock_agent = Mock(spec=BaseAgent)
|
||||||
|
cache_config = ContextCacheConfig() # Use defaults
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="default_cache_app",
|
||||||
|
root_agent=mock_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert app.context_cache_config.cache_intervals == 10 # Default
|
||||||
|
assert app.context_cache_config.ttl_seconds == 1800 # Default 30 minutes
|
||||||
|
assert app.context_cache_config.min_tokens == 0 # Default
|
||||||
|
|
||||||
|
def test_app_context_cache_config_is_optional(self):
|
||||||
|
"""Test that context_cache_config is truly optional."""
|
||||||
|
mock_agent = Mock(spec=BaseAgent)
|
||||||
|
|
||||||
|
# Should work without context_cache_config
|
||||||
|
app = App(name="no_cache_app", root_agent=mock_agent)
|
||||||
|
assert app.context_cache_config is None
|
||||||
|
|
||||||
|
# Should work with explicit None
|
||||||
|
app = App(
|
||||||
|
name="explicit_none_app",
|
||||||
|
root_agent=mock_agent,
|
||||||
|
context_cache_config=None,
|
||||||
|
)
|
||||||
|
assert app.context_cache_config is None
|
||||||
|
|||||||
@@ -0,0 +1,454 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Tests for ContextCacheRequestProcessor."""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from google.adk.agents.context_cache_config import ContextCacheConfig
|
||||||
|
from google.adk.agents.invocation_context import InvocationContext
|
||||||
|
from google.adk.agents.llm_agent import LlmAgent
|
||||||
|
from google.adk.events.event import Event
|
||||||
|
from google.adk.flows.llm_flows.context_cache_processor import ContextCacheRequestProcessor
|
||||||
|
from google.adk.models.cache_metadata import CacheMetadata
|
||||||
|
from google.adk.models.llm_request import LlmRequest
|
||||||
|
from google.adk.sessions.base_session_service import BaseSessionService
|
||||||
|
from google.adk.sessions.session import Session
|
||||||
|
from google.genai import types
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
class TestContextCacheRequestProcessor:
|
||||||
|
"""Test suite for ContextCacheRequestProcessor."""
|
||||||
|
|
||||||
|
def setup_method(self):
|
||||||
|
"""Set up test fixtures."""
|
||||||
|
self.processor = ContextCacheRequestProcessor()
|
||||||
|
self.cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=10, ttl_seconds=1800, min_tokens=1024
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_invocation_context(
|
||||||
|
self,
|
||||||
|
agent,
|
||||||
|
context_cache_config=None,
|
||||||
|
session_events=None,
|
||||||
|
invocation_id="test_invocation",
|
||||||
|
):
|
||||||
|
"""Helper to create InvocationContext."""
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=session_events or [],
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_session_service = MagicMock(spec=BaseSessionService)
|
||||||
|
|
||||||
|
return InvocationContext(
|
||||||
|
agent=agent,
|
||||||
|
session=mock_session,
|
||||||
|
session_service=mock_session_service,
|
||||||
|
context_cache_config=context_cache_config,
|
||||||
|
invocation_id=invocation_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_cache_metadata(
|
||||||
|
self, invocations_used=1, cache_name="test-cache", cached_contents_count=3
|
||||||
|
):
|
||||||
|
"""Helper to create CacheMetadata."""
|
||||||
|
return CacheMetadata(
|
||||||
|
cache_name=(
|
||||||
|
f"projects/test/locations/us-central1/cachedContents/{cache_name}"
|
||||||
|
),
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="test_fingerprint",
|
||||||
|
invocations_used=invocations_used,
|
||||||
|
cached_contents_count=cached_contents_count,
|
||||||
|
created_at=time.time() - 600,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_no_cache_config(self):
|
||||||
|
"""Test processor with no cache config."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent, context_cache_config=None
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process should complete without adding cache config
|
||||||
|
events = []
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
events.append(event)
|
||||||
|
|
||||||
|
assert len(events) == 0 # No events yielded
|
||||||
|
assert llm_request.cache_config is None
|
||||||
|
|
||||||
|
async def test_with_cache_config_no_session_events(self):
|
||||||
|
"""Test processor with cache config but no session events."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent, context_cache_config=self.cache_config
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process should add cache config but no metadata
|
||||||
|
events = []
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
events.append(event)
|
||||||
|
|
||||||
|
assert len(events) == 0 # No events yielded
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata is None
|
||||||
|
|
||||||
|
async def test_with_cache_metadata_same_invocation(self):
|
||||||
|
"""Test processor finds cache metadata from same invocation."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
cache_metadata = self.create_cache_metadata(invocations_used=5)
|
||||||
|
|
||||||
|
# Event with same invocation ID
|
||||||
|
events = [
|
||||||
|
Event(
|
||||||
|
author="test_agent",
|
||||||
|
cache_metadata=cache_metadata,
|
||||||
|
invocation_id="test_invocation",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
invocation_id="test_invocation",
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process should add cache config and metadata (same invocation, no increment)
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata == cache_metadata
|
||||||
|
assert llm_request.cache_metadata.invocations_used == 5 # No increment
|
||||||
|
|
||||||
|
async def test_with_cache_metadata_different_invocation(self):
|
||||||
|
"""Test processor finds cache metadata from different invocation."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
cache_metadata = self.create_cache_metadata(invocations_used=5)
|
||||||
|
|
||||||
|
# Event with different invocation ID
|
||||||
|
events = [
|
||||||
|
Event(
|
||||||
|
author="test_agent",
|
||||||
|
cache_metadata=cache_metadata,
|
||||||
|
invocation_id="previous_invocation",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
invocation_id="current_invocation",
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process should add cache config and increment invocations_used
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata is not None
|
||||||
|
assert llm_request.cache_metadata.invocations_used == 6 # Incremented
|
||||||
|
|
||||||
|
async def test_cache_metadata_agent_filtering(self):
|
||||||
|
"""Test that cache metadata is filtered by agent name."""
|
||||||
|
agent = LlmAgent(name="target_agent")
|
||||||
|
target_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=3, cache_name="target"
|
||||||
|
)
|
||||||
|
other_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=7, cache_name="other"
|
||||||
|
)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
Event(
|
||||||
|
author="other_agent",
|
||||||
|
cache_metadata=other_cache,
|
||||||
|
invocation_id="other_invocation",
|
||||||
|
),
|
||||||
|
Event(
|
||||||
|
author="target_agent",
|
||||||
|
cache_metadata=target_cache,
|
||||||
|
invocation_id="target_invocation",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
invocation_id="current_invocation",
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should only use target_agent's cache metadata
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_metadata is not None
|
||||||
|
assert llm_request.cache_metadata.cache_name == target_cache.cache_name
|
||||||
|
assert llm_request.cache_metadata.invocations_used == 4 # target_cache + 1
|
||||||
|
|
||||||
|
async def test_latest_cache_metadata_selected(self):
|
||||||
|
"""Test that the latest cache metadata is selected."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
older_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=2, cache_name="older"
|
||||||
|
)
|
||||||
|
newer_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=5, cache_name="newer"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Events in chronological order (older first)
|
||||||
|
events = [
|
||||||
|
Event(
|
||||||
|
author="test_agent",
|
||||||
|
cache_metadata=older_cache,
|
||||||
|
invocation_id="older_invocation",
|
||||||
|
),
|
||||||
|
Event(
|
||||||
|
author="test_agent",
|
||||||
|
cache_metadata=newer_cache,
|
||||||
|
invocation_id="newer_invocation",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
invocation_id="current_invocation",
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should use the newer (latest) cache metadata
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_metadata is not None
|
||||||
|
assert llm_request.cache_metadata.cache_name == newer_cache.cache_name
|
||||||
|
assert llm_request.cache_metadata.invocations_used == 6 # newer_cache + 1
|
||||||
|
|
||||||
|
async def test_no_cache_metadata_events(self):
|
||||||
|
"""Test when session has events but no cache metadata."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
|
||||||
|
events = [
|
||||||
|
Event(author="test_agent", cache_metadata=None),
|
||||||
|
Event(author="other_agent", cache_metadata=None),
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should add cache config but no metadata
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata is None
|
||||||
|
|
||||||
|
async def test_empty_session(self):
|
||||||
|
"""Test with empty session."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should add cache config but no metadata
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata is None
|
||||||
|
|
||||||
|
async def test_processor_yields_no_events(self):
|
||||||
|
"""Test that processor yields no events."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent, context_cache_config=self.cache_config
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
events = []
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
events.append(event)
|
||||||
|
|
||||||
|
# Processor should never yield events
|
||||||
|
assert len(events) == 0
|
||||||
|
|
||||||
|
async def test_mixed_events_scenario(self):
|
||||||
|
"""Test complex scenario with mixed events."""
|
||||||
|
agent = LlmAgent(name="test_agent")
|
||||||
|
cache_metadata = self.create_cache_metadata(invocations_used=10)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
Event(author="other_agent", cache_metadata=None),
|
||||||
|
Event(author="test_agent", cache_metadata=None), # No cache metadata
|
||||||
|
Event(
|
||||||
|
author="different_agent", cache_metadata=cache_metadata
|
||||||
|
), # Wrong agent
|
||||||
|
Event(
|
||||||
|
author="test_agent",
|
||||||
|
cache_metadata=cache_metadata,
|
||||||
|
invocation_id="prev",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
invocation_context = self.create_invocation_context(
|
||||||
|
agent,
|
||||||
|
context_cache_config=self.cache_config,
|
||||||
|
session_events=events,
|
||||||
|
invocation_id="current",
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_request = LlmRequest(
|
||||||
|
model="gemini-2.0-flash",
|
||||||
|
contents=[
|
||||||
|
types.Content(
|
||||||
|
role="user",
|
||||||
|
parts=[types.Part(text="Hello")],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
async for event in self.processor.run_async(
|
||||||
|
invocation_context, llm_request
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Should find the test_agent's cache metadata and increment it
|
||||||
|
assert llm_request.cache_config == self.cache_config
|
||||||
|
assert llm_request.cache_metadata is not None
|
||||||
|
assert llm_request.cache_metadata.invocations_used == 11 # 10 + 1
|
||||||
@@ -0,0 +1,314 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Tests for CacheMetadata."""
|
||||||
|
|
||||||
|
import time
|
||||||
|
|
||||||
|
from google.adk.models.cache_metadata import CacheMetadata
|
||||||
|
from pydantic import ValidationError
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
class TestCacheMetadata:
|
||||||
|
"""Test suite for CacheMetadata."""
|
||||||
|
|
||||||
|
def test_required_fields(self):
|
||||||
|
"""Test that all required fields must be provided."""
|
||||||
|
# Valid creation with all required fields
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=5,
|
||||||
|
cached_contents_count=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
metadata.cache_name
|
||||||
|
== "projects/123/locations/us-central1/cachedContents/456"
|
||||||
|
)
|
||||||
|
assert metadata.expire_time > time.time()
|
||||||
|
assert metadata.fingerprint == "abc123"
|
||||||
|
assert metadata.invocations_used == 5
|
||||||
|
assert metadata.cached_contents_count == 3
|
||||||
|
assert metadata.created_at is None # Optional field
|
||||||
|
|
||||||
|
def test_optional_created_at(self):
|
||||||
|
"""Test that created_at is optional."""
|
||||||
|
current_time = time.time()
|
||||||
|
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=3,
|
||||||
|
cached_contents_count=2,
|
||||||
|
created_at=current_time,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert metadata.created_at == current_time
|
||||||
|
|
||||||
|
def test_invocations_used_validation(self):
|
||||||
|
"""Test invocations_used validation constraints."""
|
||||||
|
# Valid: zero or positive
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=0,
|
||||||
|
cached_contents_count=1,
|
||||||
|
)
|
||||||
|
assert metadata.invocations_used == 0
|
||||||
|
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=10,
|
||||||
|
cached_contents_count=1,
|
||||||
|
)
|
||||||
|
assert metadata.invocations_used == 10
|
||||||
|
|
||||||
|
# Invalid: negative
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=-1,
|
||||||
|
cached_contents_count=1,
|
||||||
|
)
|
||||||
|
assert "greater than or equal to 0" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_cached_contents_count_validation(self):
|
||||||
|
"""Test cached_contents_count validation constraints."""
|
||||||
|
# Valid: zero or positive
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=0,
|
||||||
|
)
|
||||||
|
assert metadata.cached_contents_count == 0
|
||||||
|
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=10,
|
||||||
|
)
|
||||||
|
assert metadata.cached_contents_count == 10
|
||||||
|
|
||||||
|
# Invalid: negative
|
||||||
|
with pytest.raises(ValidationError) as exc_info:
|
||||||
|
CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=-1,
|
||||||
|
)
|
||||||
|
assert "greater than or equal to 0" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_expire_soon_property(self):
|
||||||
|
"""Test expire_soon property."""
|
||||||
|
# Cache that expires in 10 minutes (should not expire soon)
|
||||||
|
future_time = time.time() + 600 # 10 minutes
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=future_time,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=1,
|
||||||
|
)
|
||||||
|
assert not metadata.expire_soon
|
||||||
|
|
||||||
|
# Cache that expires in 1 minute (should expire soon)
|
||||||
|
soon_time = time.time() + 60 # 1 minute
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=soon_time,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=1,
|
||||||
|
)
|
||||||
|
assert metadata.expire_soon
|
||||||
|
|
||||||
|
def test_str_representation(self):
|
||||||
|
"""Test string representation."""
|
||||||
|
current_time = time.time()
|
||||||
|
expire_time = current_time + 1800 # 30 minutes
|
||||||
|
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/test456",
|
||||||
|
expire_time=expire_time,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=7,
|
||||||
|
cached_contents_count=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
str_repr = str(metadata)
|
||||||
|
assert "test456" in str_repr # Cache ID
|
||||||
|
assert "used 7 invocations" in str_repr
|
||||||
|
assert "cached 4 contents" in str_repr
|
||||||
|
assert "expires in" in str_repr
|
||||||
|
|
||||||
|
def test_immutability(self):
|
||||||
|
"""Test that CacheMetadata is immutable (frozen)."""
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=5,
|
||||||
|
cached_contents_count=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should not be able to modify fields
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
metadata.invocations_used = 10
|
||||||
|
|
||||||
|
def test_model_config(self):
|
||||||
|
"""Test that model config is set correctly."""
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=5,
|
||||||
|
cached_contents_count=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert metadata.model_config["extra"] == "forbid"
|
||||||
|
assert metadata.model_config["frozen"] == True
|
||||||
|
|
||||||
|
def test_field_descriptions(self):
|
||||||
|
"""Test that fields have proper descriptions."""
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=5,
|
||||||
|
cached_contents_count=3,
|
||||||
|
)
|
||||||
|
schema = metadata.model_json_schema()
|
||||||
|
|
||||||
|
assert "invocations_used" in schema["properties"]
|
||||||
|
assert (
|
||||||
|
"Number of invocations"
|
||||||
|
in schema["properties"]["invocations_used"]["description"]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "cached_contents_count" in schema["properties"]
|
||||||
|
assert (
|
||||||
|
"Number of contents"
|
||||||
|
in schema["properties"]["cached_contents_count"]["description"]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_realistic_cache_scenarios(self):
|
||||||
|
"""Test realistic cache scenarios."""
|
||||||
|
current_time = time.time()
|
||||||
|
|
||||||
|
# Fresh cache
|
||||||
|
fresh_cache = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/fresh123",
|
||||||
|
expire_time=current_time + 1800,
|
||||||
|
fingerprint="fresh_fingerprint",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=5,
|
||||||
|
created_at=current_time,
|
||||||
|
)
|
||||||
|
assert fresh_cache.invocations_used == 1
|
||||||
|
assert not fresh_cache.expire_soon
|
||||||
|
|
||||||
|
# Well-used cache
|
||||||
|
used_cache = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/used456",
|
||||||
|
expire_time=current_time + 600,
|
||||||
|
fingerprint="used_fingerprint",
|
||||||
|
invocations_used=8,
|
||||||
|
cached_contents_count=3,
|
||||||
|
created_at=current_time - 1200,
|
||||||
|
)
|
||||||
|
assert used_cache.invocations_used == 8
|
||||||
|
|
||||||
|
# Expiring cache
|
||||||
|
expiring_cache = CacheMetadata(
|
||||||
|
cache_name=(
|
||||||
|
"projects/123/locations/us-central1/cachedContents/expiring789"
|
||||||
|
),
|
||||||
|
expire_time=current_time + 60, # 1 minute
|
||||||
|
fingerprint="expiring_fingerprint",
|
||||||
|
invocations_used=15,
|
||||||
|
cached_contents_count=10,
|
||||||
|
)
|
||||||
|
assert expiring_cache.expire_soon
|
||||||
|
|
||||||
|
def test_cache_name_extraction(self):
|
||||||
|
"""Test cache name ID extraction in string representation."""
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name=(
|
||||||
|
"projects/123/locations/us-central1/cachedContents/extracted_id"
|
||||||
|
),
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=1,
|
||||||
|
cached_contents_count=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
str_repr = str(metadata)
|
||||||
|
assert "extracted_id" in str_repr
|
||||||
|
|
||||||
|
def test_no_performance_metrics(self):
|
||||||
|
"""Test that performance metrics are not in CacheMetadata."""
|
||||||
|
metadata = CacheMetadata(
|
||||||
|
cache_name="projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="abc123",
|
||||||
|
invocations_used=5,
|
||||||
|
cached_contents_count=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify that token counts are NOT in CacheMetadata
|
||||||
|
# (they should be in LlmResponse.usage_metadata)
|
||||||
|
assert not hasattr(metadata, "cached_tokens")
|
||||||
|
assert not hasattr(metadata, "total_tokens")
|
||||||
|
assert not hasattr(metadata, "prompt_tokens")
|
||||||
|
|
||||||
|
def test_missing_required_fields(self):
|
||||||
|
"""Test validation when required fields are missing."""
|
||||||
|
# Test each required field
|
||||||
|
required_fields = [
|
||||||
|
"cache_name",
|
||||||
|
"expire_time",
|
||||||
|
"fingerprint",
|
||||||
|
"invocations_used",
|
||||||
|
"cached_contents_count",
|
||||||
|
]
|
||||||
|
|
||||||
|
base_args = {
|
||||||
|
"cache_name": "projects/123/locations/us-central1/cachedContents/456",
|
||||||
|
"expire_time": time.time() + 1800,
|
||||||
|
"fingerprint": "abc123",
|
||||||
|
"invocations_used": 1,
|
||||||
|
"cached_contents_count": 2,
|
||||||
|
}
|
||||||
|
|
||||||
|
for field in required_fields:
|
||||||
|
args = base_args.copy()
|
||||||
|
del args[field]
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
CacheMetadata(**args)
|
||||||
@@ -19,6 +19,8 @@ from unittest import mock
|
|||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
from google.adk import version as adk_version
|
from google.adk import version as adk_version
|
||||||
|
from google.adk.agents.context_cache_config import ContextCacheConfig
|
||||||
|
from google.adk.models.cache_metadata import CacheMetadata
|
||||||
from google.adk.models.gemini_llm_connection import GeminiLlmConnection
|
from google.adk.models.gemini_llm_connection import GeminiLlmConnection
|
||||||
from google.adk.models.google_llm import _AGENT_ENGINE_TELEMETRY_ENV_VARIABLE_NAME
|
from google.adk.models.google_llm import _AGENT_ENGINE_TELEMETRY_ENV_VARIABLE_NAME
|
||||||
from google.adk.models.google_llm import _AGENT_ENGINE_TELEMETRY_TAG
|
from google.adk.models.google_llm import _AGENT_ENGINE_TELEMETRY_TAG
|
||||||
@@ -84,6 +86,37 @@ def llm_request():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def cache_metadata():
|
||||||
|
import time
|
||||||
|
|
||||||
|
return CacheMetadata(
|
||||||
|
cache_name="projects/test/locations/us-central1/cachedContents/test123",
|
||||||
|
expire_time=time.time() + 3600,
|
||||||
|
fingerprint="test_fingerprint",
|
||||||
|
invocations_used=2,
|
||||||
|
cached_contents_count=3,
|
||||||
|
created_at=time.time() - 600,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def llm_request_with_cache(cache_metadata):
|
||||||
|
return LlmRequest(
|
||||||
|
model="gemini-1.5-flash",
|
||||||
|
contents=[Content(role="user", parts=[Part.from_text(text="Hello")])],
|
||||||
|
config=types.GenerateContentConfig(
|
||||||
|
temperature=0.1,
|
||||||
|
response_modalities=[types.Modality.TEXT],
|
||||||
|
system_instruction="You are a helpful assistant",
|
||||||
|
),
|
||||||
|
cache_config=ContextCacheConfig(
|
||||||
|
cache_intervals=10, ttl_seconds=3600, min_tokens=100
|
||||||
|
),
|
||||||
|
cache_metadata=cache_metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def llm_request_with_computer_use():
|
def llm_request_with_computer_use():
|
||||||
return LlmRequest(
|
return LlmRequest(
|
||||||
@@ -1600,3 +1633,96 @@ async def test_adapt_computer_use_tool_no_wait():
|
|||||||
# Verify tools_dict is unchanged
|
# Verify tools_dict is unchanged
|
||||||
assert llm_request.tools_dict == original_tools_dict
|
assert llm_request.tools_dict == original_tools_dict
|
||||||
assert "wait_5_seconds" not in llm_request.tools_dict
|
assert "wait_5_seconds" not in llm_request.tools_dict
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_generate_content_async_with_cache_metadata_integration(
|
||||||
|
gemini_llm, llm_request_with_cache, cache_metadata
|
||||||
|
):
|
||||||
|
"""Test integration between Google LLM and cache manager with proper parameter order.
|
||||||
|
|
||||||
|
This test specifically validates that the cache manager's populate_cache_metadata_in_response
|
||||||
|
method is called with the correct parameter order: (llm_response, cache_metadata).
|
||||||
|
|
||||||
|
This test would have caught the parameter order bug where cache_metadata and llm_response
|
||||||
|
were passed in the wrong order, causing 'CacheMetadata' object has no attribute 'usage_metadata' errors.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Create a mock response with usage metadata including cached tokens
|
||||||
|
generate_content_response = types.GenerateContentResponse(
|
||||||
|
candidates=[
|
||||||
|
types.Candidate(
|
||||||
|
content=Content(
|
||||||
|
role="model",
|
||||||
|
parts=[Part.from_text(text="Hello, how can I help you?")],
|
||||||
|
),
|
||||||
|
finish_reason=types.FinishReason.STOP,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
usage_metadata=types.GenerateContentResponseUsageMetadata(
|
||||||
|
prompt_token_count=1500,
|
||||||
|
candidates_token_count=150,
|
||||||
|
cached_content_token_count=800, # This is the key field that was always 0 due to the bug
|
||||||
|
total_token_count=1650,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
with mock.patch.object(gemini_llm, "api_client") as mock_client:
|
||||||
|
# Create a mock coroutine that returns the generate_content_response
|
||||||
|
async def mock_coro():
|
||||||
|
return generate_content_response
|
||||||
|
|
||||||
|
mock_client.aio.models.generate_content.return_value = mock_coro()
|
||||||
|
|
||||||
|
# Mock the cache manager module to verify correct method call
|
||||||
|
with mock.patch(
|
||||||
|
"google.adk.models.gemini_context_cache_manager.GeminiContextCacheManager"
|
||||||
|
) as MockCacheManagerClass:
|
||||||
|
mock_cache_manager = MockCacheManagerClass.return_value
|
||||||
|
# Configure cache manager to handle context caching
|
||||||
|
mock_cache_manager.handle_context_caching = AsyncMock(
|
||||||
|
return_value=cache_metadata
|
||||||
|
)
|
||||||
|
|
||||||
|
responses = [
|
||||||
|
resp
|
||||||
|
async for resp in gemini_llm.generate_content_async(
|
||||||
|
llm_request_with_cache, stream=False
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Verify the response was processed
|
||||||
|
assert len(responses) == 1
|
||||||
|
response = responses[0]
|
||||||
|
assert isinstance(response, LlmResponse)
|
||||||
|
assert response.content.parts[0].text == "Hello, how can I help you?"
|
||||||
|
|
||||||
|
# CRITICAL TEST: Verify populate_cache_metadata_in_response was called with correct parameter order
|
||||||
|
mock_cache_manager.populate_cache_metadata_in_response.assert_called_once()
|
||||||
|
call_args = (
|
||||||
|
mock_cache_manager.populate_cache_metadata_in_response.call_args
|
||||||
|
)
|
||||||
|
|
||||||
|
# The first argument should be the LlmResponse (not CacheMetadata)
|
||||||
|
first_arg = call_args[0][0] # First positional argument
|
||||||
|
second_arg = call_args[0][1] # Second positional argument
|
||||||
|
|
||||||
|
# Verify correct parameter order: (llm_response, cache_metadata)
|
||||||
|
assert isinstance(first_arg, LlmResponse), (
|
||||||
|
f"First parameter should be LlmResponse, got {type(first_arg)}. "
|
||||||
|
"This indicates parameters are in wrong order."
|
||||||
|
)
|
||||||
|
assert isinstance(second_arg, CacheMetadata), (
|
||||||
|
f"Second parameter should be CacheMetadata, got {type(second_arg)}. "
|
||||||
|
"This indicates parameters are in wrong order."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify the LlmResponse has the expected usage metadata
|
||||||
|
assert first_arg.usage_metadata is not None
|
||||||
|
assert first_arg.usage_metadata.cached_content_token_count == 800
|
||||||
|
assert first_arg.usage_metadata.prompt_token_count == 1500
|
||||||
|
assert first_arg.usage_metadata.candidates_token_count == 150
|
||||||
|
|
||||||
|
# Verify cache metadata is preserved
|
||||||
|
assert second_arg.cache_name == cache_metadata.cache_name
|
||||||
|
assert second_arg.invocations_used == cache_metadata.invocations_used
|
||||||
|
|||||||
@@ -15,6 +15,7 @@
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from google.adk.agents.base_agent import BaseAgent
|
from google.adk.agents.base_agent import BaseAgent
|
||||||
|
from google.adk.agents.context_cache_config import ContextCacheConfig
|
||||||
from google.adk.agents.invocation_context import InvocationContext
|
from google.adk.agents.invocation_context import InvocationContext
|
||||||
from google.adk.agents.llm_agent import LlmAgent
|
from google.adk.agents.llm_agent import LlmAgent
|
||||||
from google.adk.apps.app import App
|
from google.adk.apps.app import App
|
||||||
@@ -467,5 +468,191 @@ class TestRunnerWithPlugins:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRunnerCacheConfig:
|
||||||
|
"""Tests for Runner cache config extraction and handling."""
|
||||||
|
|
||||||
|
def setup_method(self):
|
||||||
|
"""Set up test fixtures."""
|
||||||
|
self.session_service = InMemorySessionService()
|
||||||
|
self.artifact_service = InMemoryArtifactService()
|
||||||
|
self.root_agent = MockLlmAgent("root_agent")
|
||||||
|
|
||||||
|
def test_runner_extracts_cache_config_from_app(self):
|
||||||
|
"""Test that Runner extracts cache config from App."""
|
||||||
|
cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=15, ttl_seconds=3600, min_tokens=1024
|
||||||
|
)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="test_app",
|
||||||
|
root_agent=self.root_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert runner.context_cache_config == cache_config
|
||||||
|
assert runner.context_cache_config.cache_intervals == 15
|
||||||
|
assert runner.context_cache_config.ttl_seconds == 3600
|
||||||
|
assert runner.context_cache_config.min_tokens == 1024
|
||||||
|
|
||||||
|
def test_runner_with_app_without_cache_config(self):
|
||||||
|
"""Test Runner with App that has no cache config."""
|
||||||
|
app = App(
|
||||||
|
name="test_app", root_agent=self.root_agent, context_cache_config=None
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert runner.context_cache_config is None
|
||||||
|
|
||||||
|
def test_runner_without_app_has_no_cache_config(self):
|
||||||
|
"""Test Runner created without App has no cache config."""
|
||||||
|
runner = Runner(
|
||||||
|
app_name="test_app",
|
||||||
|
agent=self.root_agent,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert runner.context_cache_config is None
|
||||||
|
|
||||||
|
def test_runner_cache_config_passed_to_invocation_context(self):
|
||||||
|
"""Test that cache config is passed to InvocationContext."""
|
||||||
|
cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=20, ttl_seconds=7200, min_tokens=2048
|
||||||
|
)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="test_app",
|
||||||
|
root_agent=self.root_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create a mock session
|
||||||
|
mock_session = Session(
|
||||||
|
id=TEST_SESSION_ID,
|
||||||
|
app_name=TEST_APP_ID,
|
||||||
|
user_id=TEST_USER_ID,
|
||||||
|
events=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create invocation context using runner's method
|
||||||
|
invocation_context = runner._new_invocation_context(mock_session)
|
||||||
|
|
||||||
|
assert invocation_context.context_cache_config == cache_config
|
||||||
|
assert invocation_context.context_cache_config.cache_intervals == 20
|
||||||
|
|
||||||
|
def test_runner_validate_params_return_order(self):
|
||||||
|
"""Test that _validate_runner_params returns values in correct order."""
|
||||||
|
cache_config = ContextCacheConfig(cache_intervals=25)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="order_test_app",
|
||||||
|
root_agent=self.root_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Test the validation method directly
|
||||||
|
app_name, agent, context_cache_config, plugins = (
|
||||||
|
runner._validate_runner_params(app, None, None, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert app_name == "order_test_app"
|
||||||
|
assert agent == self.root_agent
|
||||||
|
assert context_cache_config == cache_config
|
||||||
|
assert context_cache_config.cache_intervals == 25
|
||||||
|
assert plugins == []
|
||||||
|
|
||||||
|
def test_runner_validate_params_without_app(self):
|
||||||
|
"""Test _validate_runner_params without App returns None for cache config."""
|
||||||
|
runner = Runner(
|
||||||
|
app_name="test_app",
|
||||||
|
agent=self.root_agent,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
app_name, agent, context_cache_config, plugins = (
|
||||||
|
runner._validate_runner_params(None, "test_app", self.root_agent, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert app_name == "test_app"
|
||||||
|
assert agent == self.root_agent
|
||||||
|
assert context_cache_config is None
|
||||||
|
assert plugins is None
|
||||||
|
|
||||||
|
def test_runner_app_name_and_agent_extracted_correctly(self):
|
||||||
|
"""Test that app_name and agent are correctly extracted from App."""
|
||||||
|
cache_config = ContextCacheConfig()
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="extracted_app",
|
||||||
|
root_agent=self.root_agent,
|
||||||
|
context_cache_config=cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert runner.app_name == "extracted_app"
|
||||||
|
assert runner.agent == self.root_agent
|
||||||
|
assert runner.context_cache_config == cache_config
|
||||||
|
|
||||||
|
def test_runner_realistic_cache_config_scenario(self):
|
||||||
|
"""Test realistic scenario with production-like cache config."""
|
||||||
|
# Production cache config
|
||||||
|
production_cache_config = ContextCacheConfig(
|
||||||
|
cache_intervals=30, ttl_seconds=14400, min_tokens=4096 # 4 hours
|
||||||
|
)
|
||||||
|
|
||||||
|
app = App(
|
||||||
|
name="production_app",
|
||||||
|
root_agent=self.root_agent,
|
||||||
|
context_cache_config=production_cache_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = Runner(
|
||||||
|
app=app,
|
||||||
|
session_service=self.session_service,
|
||||||
|
artifact_service=self.artifact_service,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify all settings are preserved
|
||||||
|
assert runner.context_cache_config.cache_intervals == 30
|
||||||
|
assert runner.context_cache_config.ttl_seconds == 14400
|
||||||
|
assert runner.context_cache_config.ttl_string == "14400s"
|
||||||
|
assert runner.context_cache_config.min_tokens == 4096
|
||||||
|
|
||||||
|
# Verify string representation
|
||||||
|
expected_str = (
|
||||||
|
"ContextCacheConfig(cache_intervals=30, ttl=14400s, min_tokens=4096)"
|
||||||
|
)
|
||||||
|
assert str(runner.context_cache_config) == expected_str
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
|
|||||||
@@ -0,0 +1,450 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Tests for CachePerformanceAnalyzer."""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from unittest.mock import AsyncMock
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from google.adk.events.event import Event
|
||||||
|
from google.adk.models.cache_metadata import CacheMetadata
|
||||||
|
from google.adk.sessions.base_session_service import BaseSessionService
|
||||||
|
from google.adk.sessions.session import Session
|
||||||
|
from google.adk.utils.cache_performance_analyzer import CachePerformanceAnalyzer
|
||||||
|
from google.genai import types
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
class TestCachePerformanceAnalyzer:
|
||||||
|
"""Test suite for CachePerformanceAnalyzer."""
|
||||||
|
|
||||||
|
def setup_method(self):
|
||||||
|
"""Set up test fixtures."""
|
||||||
|
self.mock_session_service = MagicMock(spec=BaseSessionService)
|
||||||
|
self.analyzer = CachePerformanceAnalyzer(self.mock_session_service)
|
||||||
|
|
||||||
|
def create_cache_metadata(
|
||||||
|
self, invocations_used=1, cache_name="test-cache", cached_contents_count=5
|
||||||
|
):
|
||||||
|
"""Helper to create test CacheMetadata."""
|
||||||
|
return CacheMetadata(
|
||||||
|
cache_name=(
|
||||||
|
f"projects/test/locations/us-central1/cachedContents/{cache_name}"
|
||||||
|
),
|
||||||
|
expire_time=time.time() + 1800,
|
||||||
|
fingerprint="test_fingerprint",
|
||||||
|
invocations_used=invocations_used,
|
||||||
|
cached_contents_count=cached_contents_count,
|
||||||
|
created_at=time.time() - 600,
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_mock_usage_metadata(
|
||||||
|
self, prompt_tokens=1000, cached_tokens=500, candidates_tokens=100
|
||||||
|
):
|
||||||
|
"""Helper to create mock usage metadata."""
|
||||||
|
return types.GenerateContentResponseUsageMetadata(
|
||||||
|
prompt_token_count=prompt_tokens,
|
||||||
|
cached_content_token_count=cached_tokens,
|
||||||
|
candidates_token_count=candidates_tokens,
|
||||||
|
total_token_count=prompt_tokens + candidates_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_mock_event(
|
||||||
|
self, author="test_agent", cache_metadata=None, usage_metadata=None
|
||||||
|
):
|
||||||
|
"""Helper to create mock event."""
|
||||||
|
event = Event(author=author, cache_metadata=cache_metadata)
|
||||||
|
if usage_metadata:
|
||||||
|
event.usage_metadata = usage_metadata
|
||||||
|
return event
|
||||||
|
|
||||||
|
def test_init(self):
|
||||||
|
"""Test analyzer initialization."""
|
||||||
|
assert self.analyzer.session_service == self.mock_session_service
|
||||||
|
|
||||||
|
async def test_get_agent_cache_history_empty_session(self):
|
||||||
|
"""Test getting cache history from empty session."""
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=[],
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer._get_agent_cache_history(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
async def test_get_agent_cache_history_no_cache_events(self):
|
||||||
|
"""Test getting cache history when no events have cache metadata."""
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(author="test_agent"),
|
||||||
|
self.create_mock_event(author="other_agent"),
|
||||||
|
self.create_mock_event(author="test_agent"),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer._get_agent_cache_history(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
async def test_get_agent_cache_history_specific_agent(self):
|
||||||
|
"""Test getting cache history for specific agent."""
|
||||||
|
cache1 = self.create_cache_metadata(invocations_used=1, cache_name="cache1")
|
||||||
|
cache2 = self.create_cache_metadata(invocations_used=3, cache_name="cache2")
|
||||||
|
cache3 = self.create_cache_metadata(invocations_used=5, cache_name="cache3")
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(author="test_agent", cache_metadata=cache1),
|
||||||
|
self.create_mock_event(author="other_agent", cache_metadata=cache2),
|
||||||
|
self.create_mock_event(author="test_agent", cache_metadata=cache3),
|
||||||
|
self.create_mock_event(author="test_agent"), # No cache metadata
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer._get_agent_cache_history(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should only return cache metadata for test_agent
|
||||||
|
assert len(result) == 2
|
||||||
|
assert result[0] == cache1
|
||||||
|
assert result[1] == cache3
|
||||||
|
|
||||||
|
async def test_get_agent_cache_history_all_agents(self):
|
||||||
|
"""Test getting cache history for all agents."""
|
||||||
|
cache1 = self.create_cache_metadata(invocations_used=1, cache_name="cache1")
|
||||||
|
cache2 = self.create_cache_metadata(invocations_used=3, cache_name="cache2")
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(author="agent1", cache_metadata=cache1),
|
||||||
|
self.create_mock_event(author="agent2", cache_metadata=cache2),
|
||||||
|
self.create_mock_event(author="agent1"), # No cache metadata
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
# Pass None for agent_name to get all agents
|
||||||
|
result = await self.analyzer._get_agent_cache_history(
|
||||||
|
"test_session", "test_user", "test_app", None
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should return cache metadata for all agents
|
||||||
|
assert len(result) == 2
|
||||||
|
assert result[0] == cache1
|
||||||
|
assert result[1] == cache2
|
||||||
|
|
||||||
|
async def test_analyze_agent_cache_performance_no_cache_data(self):
|
||||||
|
"""Test analysis with no cache data."""
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=[],
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["status"] == "no_cache_data"
|
||||||
|
|
||||||
|
async def test_analyze_agent_cache_performance_with_cache_data(self):
|
||||||
|
"""Test comprehensive analysis with cache data and token metrics."""
|
||||||
|
cache1 = self.create_cache_metadata(invocations_used=2, cache_name="cache1")
|
||||||
|
cache2 = self.create_cache_metadata(invocations_used=5, cache_name="cache2")
|
||||||
|
cache3 = self.create_cache_metadata(invocations_used=8, cache_name="cache3")
|
||||||
|
|
||||||
|
usage1 = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=1000, cached_tokens=800
|
||||||
|
)
|
||||||
|
usage2 = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=1500, cached_tokens=1200
|
||||||
|
)
|
||||||
|
usage3 = self.create_mock_usage_metadata(prompt_tokens=800, cached_tokens=0)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(
|
||||||
|
author="test_agent", cache_metadata=cache1, usage_metadata=usage1
|
||||||
|
),
|
||||||
|
self.create_mock_event(author="other_agent", cache_metadata=cache2),
|
||||||
|
self.create_mock_event(
|
||||||
|
author="test_agent", cache_metadata=cache2, usage_metadata=usage2
|
||||||
|
),
|
||||||
|
self.create_mock_event(
|
||||||
|
author="test_agent", cache_metadata=cache3, usage_metadata=usage3
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Basic cache metrics
|
||||||
|
assert result["status"] == "active"
|
||||||
|
assert result["requests_with_cache"] == 3
|
||||||
|
assert result["cache_refreshes"] == 3 # 3 unique cache names
|
||||||
|
assert result["total_invocations"] == 15 # 2 + 5 + 8
|
||||||
|
|
||||||
|
expected_avg_invocations = (2 + 5 + 8) / 3 # 5.0
|
||||||
|
assert result["avg_invocations_used"] == expected_avg_invocations
|
||||||
|
|
||||||
|
# Token metrics
|
||||||
|
assert result["total_prompt_tokens"] == 3300 # 1000 + 1500 + 800
|
||||||
|
assert result["total_cached_tokens"] == 2000 # 800 + 1200 + 0
|
||||||
|
assert result["total_requests"] == 3
|
||||||
|
assert (
|
||||||
|
result["requests_with_cache_hits"] == 2
|
||||||
|
) # Only first two have cached tokens
|
||||||
|
|
||||||
|
# Calculated metrics
|
||||||
|
expected_hit_ratio = (2000 / 3300) * 100 # ~60.6%
|
||||||
|
expected_utilization = (2 / 3) * 100 # ~66.7%
|
||||||
|
expected_avg_cached = 2000 / 3 # ~666.7
|
||||||
|
|
||||||
|
assert abs(result["cache_hit_ratio_percent"] - expected_hit_ratio) < 0.01
|
||||||
|
assert (
|
||||||
|
abs(result["cache_utilization_ratio_percent"] - expected_utilization)
|
||||||
|
< 0.01
|
||||||
|
)
|
||||||
|
assert (
|
||||||
|
abs(result["avg_cached_tokens_per_request"] - expected_avg_cached)
|
||||||
|
< 0.01
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_analyze_agent_cache_performance_single_cache(self):
|
||||||
|
"""Test analysis with single cache instance."""
|
||||||
|
cache = self.create_cache_metadata(
|
||||||
|
invocations_used=10, cache_name="single_cache"
|
||||||
|
)
|
||||||
|
usage = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=2000, cached_tokens=1500
|
||||||
|
)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(
|
||||||
|
author="test_agent", cache_metadata=cache, usage_metadata=usage
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["status"] == "active"
|
||||||
|
assert result["requests_with_cache"] == 1
|
||||||
|
assert result["avg_invocations_used"] == 10.0
|
||||||
|
assert result["cache_refreshes"] == 1
|
||||||
|
assert result["total_invocations"] == 10
|
||||||
|
assert result["latest_cache"] == cache.cache_name
|
||||||
|
|
||||||
|
# Token metrics for single request
|
||||||
|
assert result["total_prompt_tokens"] == 2000
|
||||||
|
assert result["total_cached_tokens"] == 1500
|
||||||
|
assert result["cache_hit_ratio_percent"] == 75.0 # 1500/2000 * 100
|
||||||
|
assert result["cache_utilization_ratio_percent"] == 100.0 # 1/1 * 100
|
||||||
|
assert result["avg_cached_tokens_per_request"] == 1500.0
|
||||||
|
|
||||||
|
async def test_analyze_agent_cache_performance_no_token_data(self):
|
||||||
|
"""Test analysis when events have no usage_metadata."""
|
||||||
|
cache = self.create_cache_metadata(invocations_used=5)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(author="test_agent", cache_metadata=cache),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should still work but with zero token metrics
|
||||||
|
assert result["status"] == "active"
|
||||||
|
assert result["requests_with_cache"] == 1
|
||||||
|
assert result["total_prompt_tokens"] == 0
|
||||||
|
assert result["total_cached_tokens"] == 0
|
||||||
|
assert result["cache_hit_ratio_percent"] == 0.0
|
||||||
|
assert result["cache_utilization_ratio_percent"] == 0.0
|
||||||
|
assert result["avg_cached_tokens_per_request"] == 0.0
|
||||||
|
|
||||||
|
async def test_analyze_agent_cache_performance_zero_invocations(self):
|
||||||
|
"""Test analysis with zero invocations."""
|
||||||
|
cache = self.create_cache_metadata(
|
||||||
|
invocations_used=0, cache_name="zero_cache"
|
||||||
|
)
|
||||||
|
usage = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=1000, cached_tokens=500
|
||||||
|
)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(
|
||||||
|
author="test_agent", cache_metadata=cache, usage_metadata=usage
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "test_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["status"] == "active"
|
||||||
|
assert result["avg_invocations_used"] == 0.0
|
||||||
|
assert result["total_invocations"] == 0
|
||||||
|
|
||||||
|
# Token metrics should still work
|
||||||
|
assert result["total_prompt_tokens"] == 1000
|
||||||
|
assert result["total_cached_tokens"] == 500
|
||||||
|
|
||||||
|
async def test_session_service_integration(self):
|
||||||
|
"""Test integration with session service."""
|
||||||
|
cache_metadata = self.create_cache_metadata(invocations_used=7)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(
|
||||||
|
author="integration_agent", cache_metadata=cache_metadata
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="integration_session",
|
||||||
|
app_name="integration_app",
|
||||||
|
user_id="integration_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Configure the mock to return the session
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"integration_session",
|
||||||
|
"integration_user",
|
||||||
|
"integration_app",
|
||||||
|
"integration_agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify the session service was called with correct parameters (twice internally)
|
||||||
|
assert self.mock_session_service.get_session.call_count == 2
|
||||||
|
self.mock_session_service.get_session.assert_called_with(
|
||||||
|
session_id="integration_session",
|
||||||
|
app_name="integration_app",
|
||||||
|
user_id="integration_user",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["status"] == "active"
|
||||||
|
assert result["requests_with_cache"] == 1
|
||||||
|
|
||||||
|
async def test_mixed_agents_filtering(self):
|
||||||
|
"""Test that analysis correctly filters by agent name."""
|
||||||
|
target_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=3, cache_name="target"
|
||||||
|
)
|
||||||
|
other_cache = self.create_cache_metadata(
|
||||||
|
invocations_used=5, cache_name="other"
|
||||||
|
)
|
||||||
|
|
||||||
|
target_usage = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=1000, cached_tokens=800
|
||||||
|
)
|
||||||
|
other_usage = self.create_mock_usage_metadata(
|
||||||
|
prompt_tokens=2000, cached_tokens=1600
|
||||||
|
)
|
||||||
|
|
||||||
|
events = [
|
||||||
|
self.create_mock_event(
|
||||||
|
author="target_agent",
|
||||||
|
cache_metadata=target_cache,
|
||||||
|
usage_metadata=target_usage,
|
||||||
|
),
|
||||||
|
self.create_mock_event(
|
||||||
|
author="other_agent",
|
||||||
|
cache_metadata=other_cache,
|
||||||
|
usage_metadata=other_usage,
|
||||||
|
),
|
||||||
|
self.create_mock_event(author="target_agent"), # No cache data
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_session = Session(
|
||||||
|
id="test_session",
|
||||||
|
app_name="test_app",
|
||||||
|
user_id="test_user",
|
||||||
|
events=events,
|
||||||
|
)
|
||||||
|
self.mock_session_service.get_session = AsyncMock(return_value=mock_session)
|
||||||
|
|
||||||
|
result = await self.analyzer.analyze_agent_cache_performance(
|
||||||
|
"test_session", "test_user", "test_app", "target_agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should only include target_agent's data
|
||||||
|
assert result["requests_with_cache"] == 1
|
||||||
|
assert result["total_invocations"] == 3
|
||||||
|
assert result["total_prompt_tokens"] == 1000 # Only target_agent's tokens
|
||||||
|
assert result["total_cached_tokens"] == 800 # Only target_agent's tokens
|
||||||
Reference in New Issue
Block a user