Skip to content

Commit 98fc036

Browse files
GWealecopybara-github
authored andcommitted
fix: grow eligible Gemini cache prefixes
When a cache expired, renewal recreated the original prefix, so completed turns could stay outside every refreshed cache. It also assumed a 4,096-token floor for every model, which skipped eligible Gemini 2.5 prefixes between 2,048 and 4,095 tokens. This grows the cache to the latest validated prefix and applies the documented 2,048-token floor for Gemini 2.5 while keeping 4,096 for Gemini 3. Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 949034094
1 parent 221bad9 commit 98fc036

3 files changed

Lines changed: 151 additions & 19 deletions

File tree

src/google/adk/agents/context_cache_config.py

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -36,8 +36,9 @@ class ContextCacheConfig(BaseModel):
3636
by reusing previously processed context across multiple requests.
3737
3838
Caching begins on the second turn of a session at the earliest and requires
39-
the prior request to reach Gemini's hard 4096-token minimum, so short or
40-
single-turn sessions are never cached.
39+
the cacheable prefix to reach the model-specific minimum: 2048 tokens for
40+
Gemini 2.5 or 4096 tokens for Gemini 3. Short or single-turn sessions are
41+
therefore never cached.
4142
4243
Attributes:
4344
cache_intervals: Maximum number of invocations to reuse the same cache before refreshing it
@@ -71,11 +72,11 @@ class ContextCacheConfig(BaseModel):
7172
description=(
7273
"Minimum prior-request tokens required to enable caching. This gates"
7374
" on the previous request's actual prompt token count, not an"
74-
" estimate of the current request. Gemini enforces a hard 4096-token"
75-
" minimum that always applies, so values below 4096 have no"
76-
" additional effect. No cache is created on the first request of a"
77-
" session; caching begins on the second turn once a previous token"
78-
" count is known. Set higher to avoid caching small requests where"
75+
" estimate of the current request. Gemini's model-specific minimum"
76+
" always applies: 2048 tokens for Gemini 2.5 and 4096 tokens for"
77+
" Gemini 3. No cache is created on the first request of a session;"
78+
" caching begins on the second turn once a previous token count is"
79+
" known. Set this higher to avoid caching small requests where"
7980
" storage overhead may exceed benefits."
8081
),
8182
)

src/google/adk/models/gemini_context_cache_manager.py

Lines changed: 35 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -34,13 +34,25 @@
3434

3535
logger = logging.getLogger("google_adk." + __name__)
3636

37-
# Gemini API requires a minimum of 4096 tokens for cached content.
38-
_GEMINI_MIN_CACHE_TOKENS = 4096
37+
# Named Gemini model families have documented explicit-cache floors. For
38+
# opaque tuned-model and endpoint IDs, the server remains authoritative.
39+
_GEMINI_2_5_MIN_CACHE_TOKENS = 2048
40+
_GEMINI_3_MIN_CACHE_TOKENS = 4096
3941

4042
if TYPE_CHECKING:
4143
from google.genai import Client
4244

4345

46+
def _minimum_cache_tokens(model: Optional[str]) -> Optional[int]:
47+
"""Return the explicit-cache token floor for a named Gemini model."""
48+
model_name = (model or "").rsplit("/", maxsplit=1)[-1]
49+
if model_name.startswith("gemini-2.5-"):
50+
return _GEMINI_2_5_MIN_CACHE_TOKENS
51+
if model_name.startswith("gemini-3"):
52+
return _GEMINI_3_MIN_CACHE_TOKENS
53+
return None
54+
55+
4456
@experimental
4557
class GeminiContextCacheManager:
4658
"""Manages context cache lifecycle for Gemini models.
@@ -104,17 +116,27 @@ async def handle_context_caching(
104116
)
105117
await self.cleanup_cache(old_cache_metadata.cache_name)
106118

107-
# Calculate current fingerprint using contents count from old metadata
108-
cache_contents_count = old_cache_metadata.contents_count
119+
# Validate the previously fingerprinted prefix before growing it.
120+
previous_cache_contents_count = old_cache_metadata.contents_count
109121
current_fingerprint = self._generate_cache_fingerprint(
110-
llm_request, cache_contents_count
122+
llm_request, previous_cache_contents_count
111123
)
112124

113125
# If fingerprints match, create new cache (expired but same content)
114126
if current_fingerprint == old_cache_metadata.fingerprint:
115127
logger.debug(
116128
"Fingerprints match after invalidation, creating new cache"
117129
)
130+
current_cacheable_contents_count = (
131+
self._find_count_of_contents_to_cache(llm_request.contents)
132+
)
133+
cache_contents_count = max(
134+
previous_cache_contents_count,
135+
current_cacheable_contents_count,
136+
)
137+
current_fingerprint = self._generate_cache_fingerprint(
138+
llm_request, cache_contents_count
139+
)
118140
cache_metadata = await self._create_new_cache_with_contents(
119141
llm_request, cache_contents_count
120142
)
@@ -124,9 +146,8 @@ async def handle_context_caching(
124146
)
125147
return cache_metadata
126148

127-
# Cache creation failed (e.g., below Gemini's 4096 token minimum).
128-
# Preserve the original contents_count so the fingerprint stays
129-
# stable for subsequent calls instead of resetting to total.
149+
# Cache creation failed (for example, below the model's minimum).
150+
# Preserve the largest stable prefix for the next attempt.
130151
logger.debug(
131152
"Cache creation failed, preserving prefix fingerprint "
132153
"(contents_count=%d)",
@@ -358,11 +379,15 @@ async def _create_new_cache_with_contents(
358379
cacheable_prefix_tokens = self._estimate_cacheable_prefix_tokens(
359380
llm_request, cache_contents_count
360381
)
361-
if cacheable_prefix_tokens < _GEMINI_MIN_CACHE_TOKENS:
382+
minimum_cache_tokens = _minimum_cache_tokens(llm_request.model)
383+
if (
384+
minimum_cache_tokens is not None
385+
and cacheable_prefix_tokens < minimum_cache_tokens
386+
):
362387
logger.info(
363388
"Cacheable prefix below Gemini minimum cache size (%d < %d tokens)",
364389
cacheable_prefix_tokens,
365-
_GEMINI_MIN_CACHE_TOKENS,
390+
minimum_cache_tokens,
366391
)
367392
return None
368393

tests/unittests/agents/test_gemini_context_cache_manager.py

Lines changed: 108 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -304,6 +304,109 @@ async def test_create_cache_gates_on_prefix_not_full_prompt(self):
304304
assert result is None
305305
self.manager.genai_client.aio.caches.create.assert_not_called()
306306

307+
async def test_completed_turn_grows_cacheable_prefix(self):
308+
"""A completed turn becomes part of the next explicit cache."""
309+
first_user = types.Content(
310+
role="user", parts=[types.Part(text="First question")]
311+
)
312+
first_model = types.Content(
313+
role="model", parts=[types.Part(text="First answer")]
314+
)
315+
next_user = types.Content(
316+
role="user", parts=[types.Part(text="Next question")]
317+
)
318+
first_request = self.create_llm_request(contents_count=0)
319+
first_request.contents = [first_user]
320+
321+
first_metadata = await self.manager.handle_context_caching(first_request)
322+
323+
assert first_metadata is not None
324+
assert first_metadata.contents_count == 0
325+
326+
next_request = self.create_llm_request(
327+
cache_metadata=first_metadata, contents_count=0
328+
)
329+
next_request.contents = [first_user, first_model, next_user]
330+
next_request.cacheable_contents_token_count = 30_000
331+
cached_content = AsyncMock()
332+
cached_content.name = "cachedContents/grown-prefix"
333+
self.manager.genai_client.aio.caches.create = AsyncMock(
334+
return_value=cached_content
335+
)
336+
337+
next_metadata = await self.manager.handle_context_caching(next_request)
338+
339+
assert next_metadata is not None
340+
assert next_metadata.cache_name == "cachedContents/grown-prefix"
341+
assert next_metadata.contents_count == 2
342+
create_config = (
343+
self.manager.genai_client.aio.caches.create.call_args.kwargs["config"]
344+
)
345+
assert create_config.contents == [first_user, first_model]
346+
assert next_request.contents == [next_user]
347+
348+
async def test_gemini_25_creates_cache_above_2048_token_minimum(self):
349+
"""Gemini 2.5 creates an explicit cache above its 2,048-token floor."""
350+
llm_request = self.create_llm_request(contents_count=0)
351+
llm_request.config.system_instruction = "x" * 12_000
352+
llm_request.cacheable_contents_token_count = 3_000
353+
llm_request.cache_metadata = CacheMetadata(
354+
fingerprint=self.manager._generate_cache_fingerprint(llm_request, 0),
355+
contents_count=0,
356+
)
357+
cached_content = AsyncMock()
358+
cached_content.name = "cachedContents/gemini-25"
359+
self.manager.genai_client.aio.caches.create = AsyncMock(
360+
return_value=cached_content
361+
)
362+
363+
result = await self.manager.handle_context_caching(llm_request)
364+
365+
assert result is not None
366+
assert result.cache_name == "cachedContents/gemini-25"
367+
self.manager.genai_client.aio.caches.create.assert_awaited_once()
368+
369+
async def test_gemini_3_skips_cache_below_4096_token_minimum(self):
370+
"""Gemini 3 skips an explicit cache below its 4,096-token floor."""
371+
llm_request = self.create_llm_request(contents_count=0)
372+
llm_request.model = "gemini-3.1-pro-preview"
373+
llm_request.config.system_instruction = "x" * 12_000
374+
llm_request.cacheable_contents_token_count = 3_000
375+
llm_request.cache_metadata = CacheMetadata(
376+
fingerprint=self.manager._generate_cache_fingerprint(llm_request, 0),
377+
contents_count=0,
378+
)
379+
380+
result = await self.manager.handle_context_caching(llm_request)
381+
382+
assert result is not None
383+
assert result.cache_name is None
384+
self.manager.genai_client.aio.caches.create.assert_not_called()
385+
386+
async def test_opaque_model_does_not_apply_guessed_token_minimum(self):
387+
"""Opaque tuned-model IDs let the server enforce the cache floor."""
388+
llm_request = self.create_llm_request(contents_count=0)
389+
llm_request.model = (
390+
"projects/test/locations/us-central1/endpoints/tuned-model"
391+
)
392+
llm_request.config.system_instruction = "x" * 12_000
393+
llm_request.cacheable_contents_token_count = 3_000
394+
llm_request.cache_metadata = CacheMetadata(
395+
fingerprint=self.manager._generate_cache_fingerprint(llm_request, 0),
396+
contents_count=0,
397+
)
398+
cached_content = AsyncMock()
399+
cached_content.name = "cachedContents/tuned-model"
400+
self.manager.genai_client.aio.caches.create = AsyncMock(
401+
return_value=cached_content
402+
)
403+
404+
result = await self.manager.handle_context_caching(llm_request)
405+
406+
assert result is not None
407+
assert result.cache_name == "cachedContents/tuned-model"
408+
self.manager.genai_client.aio.caches.create.assert_awaited_once()
409+
307410
async def test_handle_context_caching_invalid_cache_fingerprint_mismatch(
308411
self,
309412
):
@@ -1155,9 +1258,12 @@ async def test_dynamic_instruction_does_not_break_initial_cache_fingerprint(
11551258
assert result_2.cache_name == (
11561259
"projects/test/locations/us-central1/cachedContents/new789"
11571260
)
1158-
assert result_2.contents_count == 0
1261+
assert result_2.contents_count == 2
11591262
assert result_2.invocations_used == 1
1160-
self.manager.genai_client.aio.caches.create.assert_called_once()
1263+
create_config = (
1264+
self.manager.genai_client.aio.caches.create.call_args.kwargs["config"]
1265+
)
1266+
assert create_config.contents == [user_msg, model_tool_call]
11611267

11621268
async def test_create_cache_uses_server_expire_time(self):
11631269
"""The server-reported expiry is authoritative when it is available."""

0 commit comments

Comments
 (0)