headroom/tests/test_cache/test_semantic.py
Abhay Singh d8783ab89b
fix(cache/semantic): key entries by context hash, not query text (#2022)
## Description

`SemanticCache` (`headroom/cache/semantic.py`) derives each entry's key
from the **query text
only** — where `query` is just the trailing user message — and its
exact-match lookup returns
the slot without checking the stored entry's `messages_hash`:

```python
# put()
key = self._generate_key(query)          # sha256(query)[:16]
self._cache[key] = entry
if messages_hash:
    self._hash_index[messages_hash] = key

# get() — exact-match branch
key = self._hash_index.get(messages_hash)
if key and key in self._cache:
    entry = self._cache[key]
    ...
    return entry                         # never checks entry.messages_hash
```

So two requests that share a trailing user message but differ in earlier
context map to the
**same** key. The second `put` overwrites the first, and the first
request's `messages_hash`
still points at that (now overwritten) slot — so it is served the
**other conversation's**
cached response.

Trailing messages like `"continue"`, `"yes"`, `"fix it"`, `"run the
tests"` are extremely
common in agentic/coding sessions, so this collides constantly. It's
independent of the
proxy-level `_compute_key` fix (that's about what goes *into*
`messages_hash`; here the entry
is stored under a query-only key regardless of how good the hash is).
This `SemanticCache` is
the one used by the SDK client's `enable_semantic_cache` path.

Concretely:
1. `put("run the tests", A, messages_hash=HA)` → key `K = sha256("run
the tests")`; `_cache[K]=A`.
2. `put("run the tests", B, messages_hash=HB)` → same `K`; `_cache[K]`
overwritten with `B`.
3. `get("run the tests", HA)` → `_hash_index[HA]=K`, `K in _cache` →
returns **B**.

Closes: no issue filed — found while auditing the cache key derivation.

## Fix

1. Key entries by the full-context `messages_hash` when present, falling
back to the query hash
   only when no hash is supplied:
   ```python
   key = messages_hash or self._generate_key(query)
   ```
2. Defensively verify `entry.messages_hash == messages_hash` in the
exact-match branch of `get`,
   so any residual stale mapping becomes a miss rather than wrong data.

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)

## Changes Made

- `headroom/cache/semantic.py`: key `put` entries by `messages_hash`
when present; verify `entry.messages_hash` in the `get` exact-match
branch.
- `tests/test_cache/test_semantic.py`: add
`test_same_query_different_context_does_not_collide` and
`test_exact_match_verifies_messages_hash`.

## Testing

- [x] New regression tests added (`tests/test_cache/test_semantic.py`)
- [x] Linting/formatting clean — run with the CI-pinned `ruff==0.15.17`
- [ ] Full `pytest` deferred to CI (local-OOM reason below).

```text
$ uvx ruff@0.15.17 check headroom/cache/semantic.py tests/test_cache/test_semantic.py
All checks passed!
```

## Real Behavior Proof

- Environment: Windows 11, Python 3.10, headroom from this branch.
Importing `headroom` pulls in the torch/transformers stack and a full
`pytest` gets OOM-killed on this box, so I verified the `put`/`get`
logic with a dependency-free script and left the full pytest to CI.
- Exact command / steps: stored responses A and B under the same query
`"run the tests"` with different `messages_hash`, then read each hash
back — through both the old (query-keyed) and new (hash-keyed) logic.
- Observed result: the old logic serves B's response to request A; the
new logic isolates them:

```text
OLD: A->RESPONSE_B  B->RESPONSE_B
NEW: A->RESPONSE_A  B->RESPONSE_B
SEMANTIC CACHE COLLISION FIX VERIFIED (OLD served B to A; NEW isolates)
```

- Not tested: the full SDK `HeadroomClient` round-trip with
`enable_semantic_cache=True` (needs the heavy stack). The fix is
confined to `SemanticCache.put`/`get` and the new tests drive them
directly. Full local `pytest` deferred to CI (OOM, per above).

## Review Readiness

- [x] I have performed a self-review
- [x] This PR is ready for human review

## Checklist

- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [x] I have added tests that prove my fix is effective or that my
feature works
- [ ] New and existing unit tests pass locally with my changes — ran
lint + a standalone logic check; full pytest deferred to CI (local OOM,
disclosed above)
- [x] I have updated the CHANGELOG.md if applicable

## Additional Notes

- Small, contained fix — the key derivation plus a verification guard,
no new dependencies.
- @JerrettDavis tagging you — this one can serve one conversation's
cached response to another when the last message matches, so it seemed
worth surfacing. Thanks!
2026-07-11 10:11:09 -05:00

320 lines
11 KiB
Python

"""Tests for SemanticCache and SemanticCacheLayer."""
import time
import pytest
from headroom.cache import (
AnthropicCacheOptimizer,
OptimizationContext,
SemanticCache,
SemanticCacheLayer,
)
from headroom.cache.semantic import SemanticCacheConfig
class TestSemanticCacheConfig:
"""Test SemanticCacheConfig."""
def test_default_values(self):
"""Test default configuration values."""
config = SemanticCacheConfig()
assert config.similarity_threshold == 0.95
assert config.max_entries == 1000
assert config.ttl_seconds == 300
assert config.use_exact_matching is True
class TestSemanticCache:
"""Test SemanticCache functionality."""
@pytest.fixture
def cache(self):
"""Create cache instance."""
config = SemanticCacheConfig(
max_entries=10,
ttl_seconds=60,
)
return SemanticCache(config)
def test_put_and_get_exact_match(self, cache):
"""Test storing and retrieving with exact hash matching."""
response = {"text": "Hello, how can I help?"}
cache.put("What is the weather?", response, messages_hash="hash123")
entry = cache.get("What is the weather?", messages_hash="hash123")
assert entry is not None
assert entry.response == response
def test_get_miss(self, cache):
"""Test cache miss."""
entry = cache.get("Unknown query", messages_hash="unknown")
assert entry is None
def test_same_query_different_context_does_not_collide(self, cache):
"""Two requests that share a trailing user message but differ in earlier
context (distinct messages_hash) must not overwrite each other. Before the
fix both were keyed by sha256(query), so the second clobbered the first and
the first's hash resolved to the second's response."""
cache.put("run the tests", {"text": "response A"}, messages_hash="ctxA")
cache.put("run the tests", {"text": "response B"}, messages_hash="ctxB")
got_a = cache.get("run the tests", messages_hash="ctxA")
got_b = cache.get("run the tests", messages_hash="ctxB")
assert got_a is not None and got_a.response == {"text": "response A"}
assert got_b is not None and got_b.response == {"text": "response B"}
def test_exact_match_verifies_messages_hash(self, cache):
"""A stored entry is only returned when its messages_hash matches the
looked-up hash — never another conversation's cached response."""
cache.put("continue", {"text": "A"}, messages_hash="hA")
# A lookup for a hash that isn't stored is a miss, not a wrong hit.
assert cache.get("continue", messages_hash="hB") is None
def test_lru_eviction(self):
"""Test LRU eviction when at capacity."""
config = SemanticCacheConfig(max_entries=3)
cache = SemanticCache(config)
# Fill cache
cache.put("query1", "response1", messages_hash="h1")
cache.put("query2", "response2", messages_hash="h2")
cache.put("query3", "response3", messages_hash="h3")
# Access query1 to make it recently used
cache.get("query1", messages_hash="h1")
# Add new entry, should evict query2 (oldest unused)
cache.put("query4", "response4", messages_hash="h4")
# query1 should still be there (recently accessed)
assert cache.get("query1", messages_hash="h1") is not None
# query2 should be evicted
assert cache.get("query2", messages_hash="h2") is None
# query3 and query4 should be there
assert cache.get("query3", messages_hash="h3") is not None
assert cache.get("query4", messages_hash="h4") is not None
def test_ttl_expiration(self):
"""Test TTL expiration."""
config = SemanticCacheConfig(ttl_seconds=1)
cache = SemanticCache(config)
cache.put("expiring query", "response", messages_hash="exp1")
# Should be available immediately
assert cache.get("expiring query", messages_hash="exp1") is not None
# Wait for TTL
time.sleep(1.1)
# Should be expired
assert cache.get("expiring query", messages_hash="exp1") is None
def test_invalidate(self, cache):
"""Test invalidating an entry."""
key = cache.put("query", "response", messages_hash="inv1")
assert cache.get("query", messages_hash="inv1") is not None
cache.invalidate(key)
assert cache.get("query", messages_hash="inv1") is None
def test_clear(self, cache):
"""Test clearing cache."""
cache.put("query1", "response1", messages_hash="c1")
cache.put("query2", "response2", messages_hash="c2")
cache.clear()
stats = cache.get_stats()
assert stats["entries"] == 0
def test_stats(self, cache):
"""Test statistics."""
cache.put("query", "response", messages_hash="s1")
cache.get("query", messages_hash="s1") # hit
cache.get("unknown", messages_hash="unknown") # miss
stats = cache.get_stats()
assert stats["entries"] == 1
assert stats["hits"] == 1
assert stats["misses"] == 1
assert stats["hit_rate"] == 0.5
def test_access_count(self, cache):
"""Test that access count is tracked."""
cache.put("query", "response", messages_hash="ac1")
# Access multiple times
for _ in range(5):
entry = cache.get("query", messages_hash="ac1")
# Initial count is 1, plus 5 accesses = 6
assert entry.access_count == 6
def test_semantic_similarity_with_embedding_fn(self):
"""Test semantic similarity with custom embedding function."""
def mock_embedding(text: str) -> list[float]:
# Simple mock: return consistent embedding for similar queries
if "weather" in text.lower():
return [1.0, 0.0, 0.0]
elif "time" in text.lower():
return [0.0, 1.0, 0.0]
else:
return [0.0, 0.0, 1.0]
config = SemanticCacheConfig(similarity_threshold=0.9)
cache = SemanticCache(config, embedding_fn=mock_embedding)
# Store a weather query
cache.put("What is the weather today?", "It's sunny", messages_hash="w1")
# Similar weather query should hit
entry = cache.get("How is the weather?")
assert entry is not None
assert entry.response == "It's sunny"
# Different query should miss
entry = cache.get("What time is it?")
assert entry is None
class TestSemanticCacheLayer:
"""Test SemanticCacheLayer functionality."""
@pytest.fixture
def layer(self):
"""Create cache layer with Anthropic optimizer."""
optimizer = AnthropicCacheOptimizer()
return SemanticCacheLayer(
optimizer,
similarity_threshold=0.95,
max_entries=100,
ttl_seconds=60,
)
@pytest.fixture
def context(self):
"""Create optimization context."""
return OptimizationContext(
provider="anthropic",
model="claude-3-opus",
)
def test_process_no_cache_hit(self, layer, context):
"""Test processing with no cache hit."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello!"},
]
result = layer.process(messages, context)
assert result.semantic_cache_hit is False
assert result.cached_response is None
def test_process_with_cache_hit(self, layer, context):
"""Test processing with cache hit."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "What is 2+2?"},
]
# First, store a response
layer.store_response(messages, {"text": "4"}, context)
# Now process same messages
result = layer.process(messages, context)
assert result.semantic_cache_hit is True
assert result.cached_response == {"text": "4"}
def test_store_response(self, layer, context):
"""Test storing a response."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Tell me a joke"},
]
key = layer.store_response(messages, {"text": "Why did..."}, context)
assert key is not None
assert len(key) > 0
def test_get_stats(self, layer, context):
"""Test getting statistics."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
]
layer.process(messages, context)
stats = layer.get_stats()
assert "semantic_cache" in stats
assert "provider_optimizer" in stats
assert stats["provider_optimizer"] == "anthropic-cache-optimizer"
def test_query_extraction(self, layer, context):
"""Test query extraction from messages."""
messages = [
{"role": "system", "content": "System"},
{"role": "user", "content": "First question"},
{"role": "assistant", "content": "Answer"},
{"role": "user", "content": "Second question"},
]
# Store response
layer.store_response(messages, {"text": "Response"}, context)
# The query should be the last user message
result = layer.process(messages, context)
assert result.semantic_cache_hit is True
def test_query_from_context(self, layer):
"""Test using query from context."""
messages = [
{"role": "user", "content": "Some message"},
]
context = OptimizationContext(
query="Specific query for caching",
)
layer.store_response(messages, {"text": "Response"}, context)
result = layer.process(messages, context)
assert result.semantic_cache_hit is True
def test_provider_optimizer_fallback(self, layer, context):
"""Test that provider optimizer is used on cache miss."""
messages = [
{"role": "system", "content": "You are helpful. " * 500},
{"role": "user", "content": "New uncached question"},
]
result = layer.process(messages, context)
# Should have used provider optimizer
assert result.semantic_cache_hit is False
# Provider optimizer should have processed
assert result.metrics.stable_prefix_hash != ""
def test_content_block_query_extraction(self, layer, context):
"""Test query extraction from content block format."""
messages = [
{"role": "system", "content": "System"},
{
"role": "user",
"content": [{"type": "text", "text": "Block format question"}],
},
]
layer.store_response(messages, {"text": "Response"}, context)
result = layer.process(messages, context)
assert result.semantic_cache_hit is True