Skip to content

Commit 96f3989

Browse files
committed
refactor(proxy): isolate semantic cache key policy
1 parent 41af39d commit 96f3989

4 files changed

Lines changed: 107 additions & 28 deletions

File tree

headroom/integrations/litellm_callback.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -104,6 +104,7 @@ async def async_pre_call_hook(
104104
data, call_type = cache, data
105105
if data is None:
106106
return None
107+
107108
if call_type not in ("completion", "acompletion"):
108109
return data
109110

headroom/proxy/semantic_cache.py

Lines changed: 2 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -8,8 +8,6 @@
88
from __future__ import annotations
99

1010
import asyncio
11-
import hashlib
12-
import json
1311
import sys
1412
from collections import OrderedDict
1513
from datetime import datetime
@@ -19,23 +17,7 @@
1917
from ..memory.tracker import ComponentStats
2018

2119
from headroom.proxy.models import CacheEntry
22-
23-
24-
def _strip_cache_control(obj: Any) -> Any:
25-
"""Recursively drop ``cache_control`` annotations before hashing.
26-
27-
Clients (notably Claude Code) move the ``cache_control`` cache breakpoint to
28-
the newest content on each call, so the same logical ``system``/``tools``
29-
payload carries the marker on one call and not the next. Stripping it keeps
30-
the cache key stable across that movement. Mirrors
31-
``helpers._strip_per_call_annotations`` but kept local so the cache module
32-
stays free of the heavier proxy-helpers import chain.
33-
"""
34-
if isinstance(obj, dict):
35-
return {k: _strip_cache_control(v) for k, v in obj.items() if k != "cache_control"}
36-
if isinstance(obj, list):
37-
return [_strip_cache_control(item) for item in obj]
38-
return obj
20+
from headroom.proxy.semantic_cache_key_policy import compute_semantic_cache_key
3921

4022

4123
class SemanticCache:
@@ -66,15 +48,7 @@ def _compute_key(self, messages: list[dict], model: str, **key_fields: Any) -> s
6648
untouched). Absent fields don't contribute, so truly-identical requests
6749
still hit.
6850
"""
69-
normalized = json.dumps(
70-
{
71-
"model": model,
72-
"messages": messages,
73-
**{k: _strip_cache_control(v) for k, v in key_fields.items()},
74-
},
75-
sort_keys=True,
76-
)
77-
return hashlib.sha256(normalized.encode()).hexdigest()[:32]
51+
return compute_semantic_cache_key(messages, model, **key_fields)
7852

7953
async def get(self, messages: list[dict], model: str, **key_fields: Any) -> CacheEntry | None:
8054
"""Get cached response if exists and not expired."""
Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
"""Pure key policy for proxy semantic response cache."""
2+
3+
from __future__ import annotations
4+
5+
import hashlib
6+
import json
7+
from typing import Any
8+
9+
10+
def strip_cache_control(obj: Any) -> Any:
11+
"""Recursively drop ``cache_control`` annotations before hashing."""
12+
if isinstance(obj, dict):
13+
return {k: strip_cache_control(v) for k, v in obj.items() if k != "cache_control"}
14+
if isinstance(obj, list):
15+
return [strip_cache_control(item) for item in obj]
16+
return obj
17+
18+
19+
def compute_semantic_cache_key(
20+
messages: list[dict],
21+
model: str,
22+
**key_fields: Any,
23+
) -> str:
24+
"""Compute a stable cache key from request content and shaping fields."""
25+
normalized = json.dumps(
26+
{
27+
"model": model,
28+
"messages": messages,
29+
**{k: strip_cache_control(v) for k, v in key_fields.items()},
30+
},
31+
sort_keys=True,
32+
)
33+
return hashlib.sha256(normalized.encode()).hexdigest()[:32]
Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,71 @@
1+
"""Tests for pure proxy semantic cache key policy."""
2+
3+
from __future__ import annotations
4+
5+
from headroom.proxy.semantic_cache import SemanticCache
6+
from headroom.proxy.semantic_cache_key_policy import (
7+
compute_semantic_cache_key,
8+
strip_cache_control,
9+
)
10+
11+
MESSAGES = [{"role": "user", "content": "hello"}]
12+
MODEL = "claude-haiku-4-5"
13+
14+
15+
def test_strip_cache_control_recurses_through_dicts_and_lists() -> None:
16+
payload = {
17+
"system": [
18+
{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}},
19+
{"nested": {"cache_control": "drop", "value": 1}},
20+
],
21+
"cache_control": "drop-root",
22+
}
23+
assert strip_cache_control(payload) == {
24+
"system": [
25+
{"type": "text", "text": "sys"},
26+
{"nested": {"value": 1}},
27+
]
28+
}
29+
30+
31+
def test_compute_semantic_cache_key_is_stable_for_identical_inputs() -> None:
32+
kwargs = {"system": "sys", "tools": [{"name": "read"}], "temperature": 0.2}
33+
assert compute_semantic_cache_key(MESSAGES, MODEL, **kwargs) == compute_semantic_cache_key(
34+
MESSAGES,
35+
MODEL,
36+
**kwargs,
37+
)
38+
39+
40+
def test_compute_semantic_cache_key_distinguishes_response_shaping_fields() -> None:
41+
assert compute_semantic_cache_key(
42+
MESSAGES, MODEL, temperature=0.0
43+
) != compute_semantic_cache_key(
44+
MESSAGES,
45+
MODEL,
46+
temperature=1.0,
47+
)
48+
49+
50+
def test_compute_semantic_cache_key_ignores_moved_cache_control_breakpoints() -> None:
51+
with_breakpoint = [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]
52+
without_breakpoint = [{"type": "text", "text": "sys"}]
53+
assert compute_semantic_cache_key(
54+
MESSAGES,
55+
MODEL,
56+
system=with_breakpoint,
57+
) == compute_semantic_cache_key(
58+
MESSAGES,
59+
MODEL,
60+
system=without_breakpoint,
61+
)
62+
63+
64+
def test_semantic_cache_private_key_wrapper_delegates_to_policy() -> None:
65+
cache = SemanticCache()
66+
kwargs = {"system": "sys", "tools": [{"name": "read"}], "temperature": 0.2}
67+
assert cache._compute_key(MESSAGES, MODEL, **kwargs) == compute_semantic_cache_key(
68+
MESSAGES,
69+
MODEL,
70+
**kwargs,
71+
)

0 commit comments

Comments
 (0)