diff --git a/headroom/proxy/semantic_cache.py b/headroom/proxy/semantic_cache.py index 1ce8baf48..998e8f3c4 100644 --- a/headroom/proxy/semantic_cache.py +++ b/headroom/proxy/semantic_cache.py @@ -17,7 +17,7 @@ if TYPE_CHECKING: from ..memory.tracker import ComponentStats from headroom.proxy.models import CacheEntry -from headroom.proxy.semantic_cache_key import compute_semantic_cache_key, strip_cache_control +from headroom.proxy.semantic_cache_key_policy import compute_semantic_cache_key, strip_cache_control _strip_cache_control = strip_cache_control diff --git a/headroom/proxy/semantic_cache_key_policy.py b/headroom/proxy/semantic_cache_key_policy.py new file mode 100644 index 000000000..1fef49a7c --- /dev/null +++ b/headroom/proxy/semantic_cache_key_policy.py @@ -0,0 +1,33 @@ +"""Pure key policy for proxy semantic response cache.""" + +from __future__ import annotations + +import hashlib +import json +from typing import Any + + +def strip_cache_control(obj: Any) -> Any: + """Recursively drop ``cache_control`` annotations before hashing.""" + if isinstance(obj, dict): + return {k: strip_cache_control(v) for k, v in obj.items() if k != "cache_control"} + if isinstance(obj, list): + return [strip_cache_control(item) for item in obj] + return obj + + +def compute_semantic_cache_key( + messages: list[dict], + model: str, + **key_fields: Any, +) -> str: + """Compute a stable cache key from request content and shaping fields.""" + normalized = json.dumps( + { + "model": model, + "messages": messages, + **{k: strip_cache_control(v) for k, v in key_fields.items()}, + }, + sort_keys=True, + ) + return hashlib.sha256(normalized.encode()).hexdigest()[:32] diff --git a/tests/test_semantic_cache_key_policy.py b/tests/test_semantic_cache_key_policy.py new file mode 100644 index 000000000..44797b857 --- /dev/null +++ b/tests/test_semantic_cache_key_policy.py @@ -0,0 +1,71 @@ +"""Tests for pure proxy semantic cache key policy.""" + +from __future__ import annotations + +from headroom.proxy.semantic_cache import SemanticCache +from headroom.proxy.semantic_cache_key_policy import ( + compute_semantic_cache_key, + strip_cache_control, +) + +MESSAGES = [{"role": "user", "content": "hello"}] +MODEL = "claude-haiku-4-5" + + +def test_strip_cache_control_recurses_through_dicts_and_lists() -> None: + payload = { + "system": [ + {"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}, + {"nested": {"cache_control": "drop", "value": 1}}, + ], + "cache_control": "drop-root", + } + assert strip_cache_control(payload) == { + "system": [ + {"type": "text", "text": "sys"}, + {"nested": {"value": 1}}, + ] + } + + +def test_compute_semantic_cache_key_is_stable_for_identical_inputs() -> None: + kwargs = {"system": "sys", "tools": [{"name": "read"}], "temperature": 0.2} + assert compute_semantic_cache_key(MESSAGES, MODEL, **kwargs) == compute_semantic_cache_key( + MESSAGES, + MODEL, + **kwargs, + ) + + +def test_compute_semantic_cache_key_distinguishes_response_shaping_fields() -> None: + assert compute_semantic_cache_key( + MESSAGES, MODEL, temperature=0.0 + ) != compute_semantic_cache_key( + MESSAGES, + MODEL, + temperature=1.0, + ) + + +def test_compute_semantic_cache_key_ignores_moved_cache_control_breakpoints() -> None: + with_breakpoint = [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}] + without_breakpoint = [{"type": "text", "text": "sys"}] + assert compute_semantic_cache_key( + MESSAGES, + MODEL, + system=with_breakpoint, + ) == compute_semantic_cache_key( + MESSAGES, + MODEL, + system=without_breakpoint, + ) + + +def test_semantic_cache_private_key_wrapper_delegates_to_policy() -> None: + cache = SemanticCache() + kwargs = {"system": "sys", "tools": [{"name": "read"}], "temperature": 0.2} + assert cache._compute_key(MESSAGES, MODEL, **kwargs) == compute_semantic_cache_key( + MESSAGES, + MODEL, + **kwargs, + )