From f22b0380963c78f16c4d339e7f62e20394656aa2 Mon Sep 17 00:00:00 2001 From: George Weale Date: Mon, 17 Aug 2026 23:03:52 +0000 Subject: [PATCH] fix(models): Port Gemini context cache identity scoping to v1 The explicit-cache fingerprint was computed from the system instruction, tools, tool config and content prefix alone. It did not include the model, and it did not include the backend the cache lives in. An identical prompt prefix under a different model, or under a different Vertex project or location, produced the same fingerprint, so a stored cache_name belonging to another context was treated as reusable and attached to a request that had no relationship to it. The fingerprint is now seeded with the model id and a cache scope. The scope records whether the client talks to Vertex or the Developer API, and for Vertex it records the project, the location and the base URL. Model dumps run in JSON mode and the hash input is canonical JSON with sorted keys, rather than the repr of a Python dict. A system instruction the SDK accepts but cannot serialize to JSON falls back to its string form. Behaviour change: every fingerprint value changes, so on upgrade a user with an active explicit context cache takes one cache miss and one cache re-creation. That cost is unavoidable while fixing the identity, because the old fingerprints did not identify what they were caching. In the other direction, argument mappings that differ only in key insertion order now hash the same, which removes cache misses that were happening for no reason. --- .../models/gemini_context_cache_manager.py | 57 +++++++++-- .../test_gemini_context_cache_manager.py | 95 +++++++++++++++++++ 2 files changed, 142 insertions(+), 10 deletions(-) diff --git a/src/google/adk/models/gemini_context_cache_manager.py b/src/google/adk/models/gemini_context_cache_manager.py index bc2cc8d8415..7a883c2cba4 100644 --- a/src/google/adk/models/gemini_context_cache_manager.py +++ b/src/google/adk/models/gemini_context_cache_manager.py @@ -20,6 +20,7 @@ import json import logging import time +from typing import Any from typing import Optional from typing import TYPE_CHECKING @@ -256,25 +257,37 @@ def _generate_cache_fingerprint( Returns: 16-character hexadecimal fingerprint representing the cached state """ - # Create fingerprint from system instruction, tools, tool_config, and first N contents - fingerprint_data = {} + # Explicit caches are model-specific, so the model is part of their + # compatibility boundary along with the cached request fields. + fingerprint_data: dict[str, Any] = { + "model": llm_request.model, + "cache_scope": self._cache_scope(), + } if llm_request.config and llm_request.config.system_instruction: - fingerprint_data["system_instruction"] = ( - llm_request.config.system_instruction - ) + try: + fingerprint_data["system_instruction"] = llm_request.config.model_dump( + mode="json", include={"system_instruction"} + )["system_instruction"] + except Exception: # pylint: disable=broad-except + # Preserve support for SDK-accepted objects without a JSON serializer + # (for example PIL images). Their string form is the best available + # compatibility boundary. + fingerprint_data["system_instruction"] = str( + 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()) + tools_data.append(tool.model_dump(mode="json")) 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() + llm_request.config.tool_config.model_dump(mode="json") ) # Include first N contents in fingerprint @@ -282,11 +295,18 @@ def _generate_cache_fingerprint( 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()) + contents_data.append(content.model_dump(mode="json")) fingerprint_data["cached_contents"] = contents_data - # Generate hash using str() instead of json.dumps() to handle bytes - fingerprint_str = str(fingerprint_data) + # Canonical JSON makes semantically identical mappings produce the same + # cache identity regardless of their insertion order. SDK model dumps in + # JSON mode also encode binary parts deterministically. + fingerprint_str = json.dumps( + fingerprint_data, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ) return hashlib.sha256(fingerprint_str.encode()).hexdigest()[:16] async def _create_new_cache_with_contents( @@ -336,6 +356,23 @@ async def _create_new_cache_with_contents( logger.warning("Failed to create cache: %s", e) return None + def _cache_scope(self) -> dict[str, Any]: + """Return the backend namespace that owns explicit cache resources.""" + is_vertex = bool(self.genai_client.vertexai) + scope: dict[str, Any] = { + "backend": "vertex" if is_vertex else "gemini", + } + api_client = getattr(self.genai_client, "_api_client", None) + if is_vertex and api_client is not None: + scope["project"] = getattr(api_client, "project", None) + scope["location"] = getattr(api_client, "location", None) + + http_options = getattr(api_client, "_http_options", None) + base_url = getattr(http_options, "base_url", None) + if base_url: + scope["base_url"] = base_url + return scope + def _estimate_request_tokens(self, llm_request: LlmRequest) -> int: """Estimate token count for the request. diff --git a/tests/unittests/agents/test_gemini_context_cache_manager.py b/tests/unittests/agents/test_gemini_context_cache_manager.py index 337495b1243..d612a57d9e0 100644 --- a/tests/unittests/agents/test_gemini_context_cache_manager.py +++ b/tests/unittests/agents/test_gemini_context_cache_manager.py @@ -34,6 +34,7 @@ class TestGeminiContextCacheManager: def setup_method(self): """Set up test fixtures.""" mock_client = AsyncMock(spec=Client) + mock_client.vertexai = False self.manager = GeminiContextCacheManager(mock_client) self.cache_config = ContextCacheConfig( cache_intervals=10, @@ -237,6 +238,68 @@ async def test_handle_context_caching_invalid_cache_fingerprint_mismatch( mock_cleanup.assert_called_once_with(existing_cache.cache_name) self.manager.genai_client.aio.caches.create.assert_not_called() + async def test_model_change_invalidates_active_cache(self): + """A cache created for one model is not reused by another model.""" + flash_request = self.create_llm_request(contents_count=0) + flash_metadata = await self.manager.handle_context_caching(flash_request) + assert flash_metadata is not None + active_metadata = CacheMetadata( + cache_name="cachedContents/flash-cache", + expire_time=time.time() + 1_800, + fingerprint=flash_metadata.fingerprint, + invocations_used=1, + contents_count=flash_metadata.contents_count, + created_at=time.time(), + ) + pro_request = self.create_llm_request( + cache_metadata=active_metadata, contents_count=0 + ) + pro_request.model = "gemini-2.5-pro" + self.manager.genai_client.aio.caches.delete = AsyncMock() + + pro_metadata = await self.manager.handle_context_caching(pro_request) + + assert pro_metadata is not None + assert pro_metadata.cache_name is None + assert pro_metadata.fingerprint != active_metadata.fingerprint + self.manager.genai_client.aio.caches.delete.assert_awaited_once_with( + name="cachedContents/flash-cache" + ) + + async def test_backend_change_invalidates_active_cache(self): + """A Developer API cache is not reused by a Vertex client.""" + developer_request = self.create_llm_request(contents_count=0) + developer_metadata = await self.manager.handle_context_caching( + developer_request + ) + assert developer_metadata is not None + active_metadata = CacheMetadata( + cache_name="cachedContents/developer-cache", + expire_time=time.time() + 1_800, + fingerprint=developer_metadata.fingerprint, + invocations_used=1, + contents_count=developer_metadata.contents_count, + created_at=time.time(), + ) + vertex_client = AsyncMock(spec=Client) + vertex_client.vertexai = True + vertex_client.aio.caches.delete = AsyncMock() + vertex_manager = GeminiContextCacheManager(vertex_client) + vertex_request = self.create_llm_request( + cache_metadata=active_metadata, contents_count=0 + ) + + vertex_metadata = await vertex_manager.handle_context_caching( + vertex_request + ) + + assert vertex_metadata is not None + assert vertex_metadata.cache_name is None + assert vertex_metadata.fingerprint != active_metadata.fingerprint + vertex_client.aio.caches.delete.assert_awaited_once_with( + name="cachedContents/developer-cache" + ) + async def test_is_cache_valid_fingerprint_mismatch(self): """Test cache validation with fingerprint mismatch.""" cache_metadata = self.create_cache_metadata() @@ -384,6 +447,38 @@ def test_generate_cache_fingerprint_different_requests(self): assert fingerprint1 != fingerprint2 + def test_generate_cache_fingerprint_canonicalizes_mapping_order(self): + """Equivalent argument mappings do not cause an avoidable cache miss.""" + first_request = self.create_llm_request(contents_count=0) + second_request = self.create_llm_request(contents_count=0) + first_request.contents = [ + types.ModelContent( + types.Part( + function_call=types.FunctionCall( + name="lookup", args={"first": 1, "second": 2} + ) + ) + ) + ] + second_request.contents = [ + types.ModelContent( + types.Part( + function_call=types.FunctionCall( + name="lookup", args={"second": 2, "first": 1} + ) + ) + ) + ] + + first_fingerprint = self.manager._generate_cache_fingerprint( + first_request, 1 + ) + second_fingerprint = self.manager._generate_cache_fingerprint( + second_request, 1 + ) + + assert first_fingerprint == second_fingerprint + def test_generate_cache_fingerprint_tool_config_variations(self): """Test that different tool configs generate different fingerprints.""" # Request with AUTO mode