chore: Create the context cache based on the token count of previous request

before this change, we estimate the token count of the contents to cache and use it to compare with the threshold user set. but that's not precise , so we use the actual prompt token count of previous llm request.

We won't create cache for the very initial request

PiperOrigin-RevId: 814484840
This commit is contained in:
Xiang (Sean) Zhou
2025-10-02 19:22:00 -07:00
committed by Copybara-Service
parent 420df25f58
commit c5b976b306
5 changed files with 323 additions and 22 deletions
@@ -62,9 +62,11 @@ class ContextCacheRequestProcessor(BaseLlmRequestProcessor):
# 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
# Find latest cache metadata and previous token count from session events
latest_cache_metadata, previous_token_count = (
self._find_cache_info_from_events(
invocation_context, agent.name, invocation_context.invocation_id
)
)
if latest_cache_metadata:
@@ -77,51 +79,78 @@ class ContextCacheRequestProcessor(BaseLlmRequestProcessor):
latest_cache_metadata.cached_contents_count,
)
if previous_token_count is not None:
llm_request.cacheable_contents_token_count = previous_token_count
logger.debug(
'Found previous prompt token count for agent %s: %d',
agent.name,
previous_token_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(
def _find_cache_info_from_events(
self,
invocation_context: 'InvocationContext',
agent_name: str,
current_invocation_id: str,
) -> Optional[CacheMetadata]:
"""Find the latest cache metadata from session events.
) -> tuple[Optional[CacheMetadata], Optional[int]]:
"""Find cache metadata and previous token count from session events.
Args:
invocation_context: Context containing session with events
agent_name: Name of agent to find cache metadata for
agent_name: Name of agent to find cache info 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
Tuple of (cache_metadata, previous_prompt_token_count)
cache_metadata: Latest cache metadata with updated invocations_used if needed
previous_prompt_token_count: Most recent prompt token count from LLM response
"""
if not invocation_context.session or not invocation_context.session.events:
return None
return None, None
cache_metadata = None
previous_token_count = 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
if event.author != agent_name:
continue
# Look for cache metadata (only in actual LLM response events)
if cache_metadata is None and event.cache_metadata is not None:
# 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}
cache_metadata = event.cache_metadata.model_copy(
update={
'invocations_used': event.cache_metadata.invocations_used + 1
}
)
else:
# Same invocation or no invocation_id - return as-is
return cache_metadata
cache_metadata = event.cache_metadata
return None
# Look for previous prompt token count (from actual LLM response events)
if (
previous_token_count is None
and event.usage_metadata
and event.usage_metadata.prompt_token_count is not None
):
previous_token_count = event.usage_metadata.prompt_token_count
# Stop early if we found both pieces of information
if cache_metadata is not None and previous_token_count is not None:
break
return cache_metadata, previous_token_count
# Create processor instance for use in flows
@@ -257,12 +257,21 @@ class GeminiContextCacheManager:
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:
# Check if we have token count from previous response for cache size validation
if llm_request.cacheable_contents_token_count is None:
logger.info(
"Request too small for caching (%d < %d tokens)",
estimated_tokens,
"No previous token count available, skipping cache creation for"
" initial request"
)
return None
if (
llm_request.cacheable_contents_token_count
< llm_request.cache_config.min_tokens
):
logger.info(
"Previous request too small for caching (%d < %d tokens)",
llm_request.cacheable_contents_token_count,
llm_request.cache_config.min_tokens,
)
return None
+3
View File
@@ -88,6 +88,9 @@ class LlmRequest(BaseModel):
cache_metadata: Optional[CacheMetadata] = None
"""Cache metadata from previous requests, used for cache management."""
cacheable_contents_token_count: Optional[int] = None
"""Token count from previous request's prompt, used for cache size validation."""
def append_instructions(
self, instructions: Union[list[str], types.Content]
) -> list[types.Content]: