chore: import upstream snapshot with attribution
Security / Dependency audit (pip-audit) (push) Has been cancelled
Security / CodeQL (javascript-typescript) (push) Has been cancelled
Security / CodeQL (python) (push) Has been cancelled
Security / Secret scan (gitleaks) (push) Has been cancelled
rust / test (ubuntu) (push) Has been cancelled
rust / simulator e2e (macos-latest) (push) Has been cancelled
rust / simulator e2e (ubuntu-latest) (push) Has been cancelled
rust / simulator e2e (windows-latest) (push) Has been cancelled
rust / wheels (aarch64-apple-darwin) (push) Has been cancelled
rust / wheels (x86_64-unknown-linux-gnu) (push) Has been cancelled
rust / wheels (x86_64-apple-darwin) (push) Has been cancelled
rust / audit (push) Has been cancelled
rust / parity (nightly, allowed to fail during Phase 0) (push) Has been cancelled
CI / commitlint (push) Has been skipped
Dev Containers / validate (.devcontainer/devcontainer.json, default) (push) Failing after 0s
Dev Containers / validate (.devcontainer/memory-stack/devcontainer.json, memory-stack) (push) Failing after 0s
Dev Containers / validate-worktree (push) Failing after 0s
CI / changes (push) Failing after 4s
Deploy Documentation / validate (push) Has been skipped
Deploy Documentation / deploy (push) Failing after 1s
Init Native E2E / init-native (ubuntu-latest, claude) (push) Failing after 1s
Init Native E2E / init-native (ubuntu-latest, codex) (push) Failing after 1s
Install Native E2E / install-native (ubuntu-latest) (push) Failing after 1s
OpenCode Plugin / typecheck + build + test (push) Failing after 1s
Init Native E2E / init-native (ubuntu-latest, copilot) (push) Failing after 1s
Release Please / release-please (push) Failing after 1s
Wrap E2E / docker-wrap-e2e (push) Failing after 1s
Wrap Native E2E / wrap-native (ubuntu-latest) (push) Failing after 1s
Init E2E / docker-init-e2e (push) Failing after 4s
Merge Conflicts / merge-conflicts (push) Failing after 4s
CI / lint (push) Has been cancelled
CI / build-wheel (push) Has been cancelled
CI / build-wheel-windows (push) Has been cancelled
CI / prefetch-model (push) Has been cancelled
CI / test-dashboard-ui (push) Has been cancelled
CI / test (1) (push) Has been cancelled
CI / test (2) (push) Has been cancelled
CI / test (3) (push) Has been cancelled
CI / test (4) (push) Has been cancelled
CI / test-extras (push) Has been cancelled
CI / test-agno (push) Has been cancelled
CI / build (push) Has been cancelled
CI / workflow-validation (push) Has been cancelled
CI / docker-native-e2e (push) Has been cancelled
CI / windows-native-wrapper (push) Has been cancelled
CI / macos-native-wrapper (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-code-nonroot name:code-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-code-slim name:code-slim]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-code-slim-nonroot name:code-slim-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-nonroot name:nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-slim name:slim]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-slim-nonroot name:slim-nonroot]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime name:]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-code name:code]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-code-nonroot name:code-nonroot]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-code-slim name:code-slim]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-code-slim-nonroot name:code-slim-nonroot]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-nonroot name:nonroot]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-slim name:slim]) (push) Has been cancelled
Docker / docker-manifest (map[bake_target:runtime-slim-nonroot name:slim-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime name:]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-code name:code]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-code-nonroot name:code-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-code-slim name:code-slim]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-code-slim-nonroot name:code-slim-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-nonroot name:nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-slim name:slim]) (push) Has been cancelled
Docker / docker-build (map[name:amd64 platform:linux/amd64 runs_on:ubuntu-24.04], map[bake_target:runtime-slim-nonroot name:slim-nonroot]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime name:]) (push) Has been cancelled
Docker / docker-build (map[name:arm64 platform:linux/arm64 runs_on:ubuntu-24.04-arm], map[bake_target:runtime-code name:code]) (push) Has been cancelled
Docker / promote-latest (push) Has been cancelled
Init Native E2E / init-native (macos-latest, claude) (push) Has been cancelled
Init Native E2E / init-native (macos-latest, codex) (push) Has been cancelled
Init Native E2E / init-native (macos-latest, copilot) (push) Has been cancelled
Install Native E2E / install-native (macos-latest) (push) Has been cancelled
Wrap Native E2E / wrap-native (macos-latest) (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 12:03:20 +08:00
commit 0ef5fcb1c5
1951 changed files with 606278 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""Tests for the cache optimization module."""
+184
View File
@@ -0,0 +1,184 @@
"""Tests for AnthropicCacheOptimizer."""
import pytest
from headroom.cache import (
AnthropicCacheOptimizer,
CacheConfig,
OptimizationContext,
)
from headroom.cache.base import CacheStrategy
class TestAnthropicCacheOptimizer:
"""Test AnthropicCacheOptimizer functionality."""
@pytest.fixture
def optimizer(self):
"""Create optimizer instance."""
return AnthropicCacheOptimizer()
@pytest.fixture
def context(self):
"""Create optimization context."""
return OptimizationContext(
provider="anthropic",
model="claude-3-opus",
)
def test_optimizer_properties(self, optimizer):
"""Test optimizer properties."""
assert optimizer.name == "anthropic-cache-optimizer"
assert optimizer.provider == "anthropic"
assert optimizer.strategy == CacheStrategy.EXPLICIT_BREAKPOINTS
def test_enforces_minimum_tokens(self):
"""Test that Anthropic minimum is enforced."""
config = CacheConfig(min_cacheable_tokens=100)
optimizer = AnthropicCacheOptimizer(config)
assert optimizer.config.min_cacheable_tokens >= 1024
def test_enforces_maximum_breakpoints(self):
"""Test that Anthropic maximum breakpoints is enforced."""
config = CacheConfig(max_breakpoints=10)
optimizer = AnthropicCacheOptimizer(config)
assert optimizer.config.max_breakpoints <= 4
def test_optimize_simple_messages(self, optimizer, context):
"""Test optimizing simple messages."""
messages = [
{"role": "system", "content": "You are a helpful assistant. " * 500},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.messages is not None
assert len(result.messages) == 2
assert result.metrics.stable_prefix_hash != ""
def test_optimize_inserts_cache_control(self, optimizer, context):
"""Test that optimization inserts cache_control blocks."""
# Large system prompt to trigger caching
messages = [
{"role": "system", "content": "You are a helpful assistant. " * 500},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
# Check if cache_control was inserted
system_content = result.messages[0]["content"]
if isinstance(system_content, list):
has_cache_control = any(
"cache_control" in block for block in system_content if isinstance(block, dict)
)
assert has_cache_control
def test_optimize_with_dates(self, optimizer, context):
"""Test optimization extracts dates."""
messages = [
{
"role": "system",
"content": "Today is January 7, 2026. You are a helpful assistant. " * 300,
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
# Dates should be moved to end
assert (
"extracted_dates" in result.transforms_applied
or result.metrics.breakpoints_inserted >= 0
)
def test_optimize_disabled(self, context):
"""Test optimization when disabled."""
config = CacheConfig(enabled=False)
optimizer = AnthropicCacheOptimizer(config)
messages = [
{"role": "system", "content": "Test"},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.transforms_applied == []
def test_prefix_hash_tracking(self, optimizer, context):
"""Test that prefix hash is tracked between calls."""
messages = [
{"role": "system", "content": "You are a helpful assistant. " * 500},
{"role": "user", "content": "Hello!"},
]
result1 = optimizer.optimize(messages, context)
result2 = optimizer.optimize(messages, context)
# Second call should detect stable prefix
assert result2.metrics.previous_prefix_hash == result1.metrics.stable_prefix_hash
def test_estimate_savings(self, optimizer, context):
"""Test savings estimation."""
messages = [
{"role": "system", "content": "You are a helpful assistant. " * 500},
{"role": "user", "content": "Hello!"},
]
savings = optimizer.estimate_savings(messages, context)
assert savings >= 0.0
assert savings <= 100.0
def test_content_block_format(self, optimizer, context):
"""Test handling of content block format."""
messages = [
{
"role": "system",
"content": [{"type": "text", "text": "You are a helpful assistant. " * 500}],
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.messages is not None
def test_tools_are_cacheable(self, optimizer, context):
"""Test that tools are identified as cacheable."""
messages = [
{"role": "system", "content": "You are helpful. " * 300},
{
"role": "user",
"content": "Use tools",
"tools": [
{
"name": "search",
"description": "Search the web " * 200,
"input_schema": {"type": "object"},
}
],
},
]
result = optimizer.optimize(messages, context)
assert result.metrics.cacheable_tokens > 0
def test_metrics_history(self, optimizer, context):
"""Test that metrics are recorded."""
messages = [
{"role": "system", "content": "You are helpful. " * 500},
{"role": "user", "content": "Hello!"},
]
optimizer.optimize(messages, context)
metrics = optimizer.get_metrics()
assert metrics is not None
assert metrics.stable_prefix_hash != ""
def test_cache_constants(self, optimizer):
"""Test cache-related constants."""
assert optimizer.get_cache_write_cost_multiplier() == 1.25
assert optimizer.get_cache_read_cost_multiplier() == 0.10
assert optimizer.get_cache_ttl_seconds() == 300
+410
View File
@@ -0,0 +1,410 @@
"""Tests for CompressionStore storage backends.
These tests define the contract that all backends must fulfill.
Each backend implementation should pass all these tests.
"""
from __future__ import annotations
import threading
import time
from typing import TYPE_CHECKING
import pytest
from headroom.cache.backends import CompressionStoreBackend, InMemoryBackend
from headroom.cache.compression_store import CompressionEntry
if TYPE_CHECKING:
from collections.abc import Callable
def make_entry(
hash_key: str = "test_hash",
original: str = "original content",
compressed: str = "compressed",
original_tokens: int = 100,
compressed_tokens: int = 10,
) -> CompressionEntry:
"""Create a test CompressionEntry."""
return CompressionEntry(
hash=hash_key,
original_content=original,
compressed_content=compressed,
original_tokens=original_tokens,
compressed_tokens=compressed_tokens,
original_item_count=5,
compressed_item_count=2,
tool_name="test_tool",
tool_call_id="call_123",
query_context="test query",
created_at=time.time(),
ttl=300,
)
class TestCompressionStoreBackendProtocol:
"""Test that InMemoryBackend implements the protocol correctly."""
def test_inmemory_backend_implements_protocol(self) -> None:
"""InMemoryBackend should implement CompressionStoreBackend protocol."""
backend = InMemoryBackend()
assert isinstance(backend, CompressionStoreBackend)
def test_protocol_is_runtime_checkable(self) -> None:
"""Protocol should be runtime checkable."""
class NotABackend:
pass
assert not isinstance(NotABackend(), CompressionStoreBackend)
class TestInMemoryBackend:
"""Test suite for InMemoryBackend.
These tests define the contract for all backends.
"""
@pytest.fixture
def backend(self) -> InMemoryBackend:
"""Create a fresh backend for each test."""
return InMemoryBackend()
# --- Basic CRUD operations ---
def test_get_returns_none_for_missing_key(self, backend: InMemoryBackend) -> None:
"""get() should return None for keys that don't exist."""
assert backend.get("nonexistent") is None
def test_set_and_get_roundtrip(self, backend: InMemoryBackend) -> None:
"""set() followed by get() should return the same entry."""
entry = make_entry(hash_key="abc123")
backend.set("abc123", entry)
retrieved = backend.get("abc123")
assert retrieved is not None
assert retrieved.hash == "abc123"
assert retrieved.original_content == "original content"
assert retrieved.compressed_content == "compressed"
assert retrieved.original_tokens == 100
assert retrieved.compressed_tokens == 10
def test_set_overwrites_existing(self, backend: InMemoryBackend) -> None:
"""set() should overwrite existing entries with the same key."""
entry1 = make_entry(hash_key="abc123", original="first")
entry2 = make_entry(hash_key="abc123", original="second")
backend.set("abc123", entry1)
backend.set("abc123", entry2)
retrieved = backend.get("abc123")
assert retrieved is not None
assert retrieved.original_content == "second"
def test_delete_removes_entry(self, backend: InMemoryBackend) -> None:
"""delete() should remove the entry and return True."""
entry = make_entry(hash_key="abc123")
backend.set("abc123", entry)
result = backend.delete("abc123")
assert result is True
assert backend.get("abc123") is None
def test_delete_returns_false_for_missing(self, backend: InMemoryBackend) -> None:
"""delete() should return False for keys that don't exist."""
result = backend.delete("nonexistent")
assert result is False
def test_exists_returns_true_for_stored_entry(self, backend: InMemoryBackend) -> None:
"""exists() should return True for stored entries."""
entry = make_entry(hash_key="abc123")
backend.set("abc123", entry)
assert backend.exists("abc123") is True
def test_exists_returns_false_for_missing(self, backend: InMemoryBackend) -> None:
"""exists() should return False for missing entries."""
assert backend.exists("nonexistent") is False
def test_clear_removes_all_entries(self, backend: InMemoryBackend) -> None:
"""clear() should remove all entries."""
backend.set("key1", make_entry(hash_key="key1"))
backend.set("key2", make_entry(hash_key="key2"))
backend.set("key3", make_entry(hash_key="key3"))
backend.clear()
assert backend.count() == 0
assert backend.get("key1") is None
assert backend.get("key2") is None
assert backend.get("key3") is None
# --- Enumeration methods ---
def test_count_returns_zero_for_empty(self, backend: InMemoryBackend) -> None:
"""count() should return 0 for empty backend."""
assert backend.count() == 0
def test_count_returns_correct_count(self, backend: InMemoryBackend) -> None:
"""count() should return the number of entries."""
backend.set("key1", make_entry(hash_key="key1"))
backend.set("key2", make_entry(hash_key="key2"))
backend.set("key3", make_entry(hash_key="key3"))
assert backend.count() == 3
def test_keys_returns_empty_list_for_empty(self, backend: InMemoryBackend) -> None:
"""keys() should return empty list for empty backend."""
assert backend.keys() == []
def test_keys_returns_all_keys(self, backend: InMemoryBackend) -> None:
"""keys() should return all stored keys."""
backend.set("key1", make_entry(hash_key="key1"))
backend.set("key2", make_entry(hash_key="key2"))
backend.set("key3", make_entry(hash_key="key3"))
keys = backend.keys()
assert set(keys) == {"key1", "key2", "key3"}
def test_items_returns_empty_list_for_empty(self, backend: InMemoryBackend) -> None:
"""items() should return empty list for empty backend."""
assert backend.items() == []
def test_items_returns_all_entries(self, backend: InMemoryBackend) -> None:
"""items() should return all (key, entry) pairs."""
entry1 = make_entry(hash_key="key1", original="content1")
entry2 = make_entry(hash_key="key2", original="content2")
backend.set("key1", entry1)
backend.set("key2", entry2)
items = backend.items()
assert len(items) == 2
items_dict = dict(items)
assert items_dict["key1"].original_content == "content1"
assert items_dict["key2"].original_content == "content2"
# --- Statistics ---
def test_get_stats_returns_required_fields(self, backend: InMemoryBackend) -> None:
"""get_stats() should return required fields."""
stats = backend.get_stats()
assert "backend_type" in stats
assert "entry_count" in stats
assert stats["backend_type"] == "memory"
assert stats["entry_count"] == 0
def test_get_stats_entry_count_accurate(self, backend: InMemoryBackend) -> None:
"""get_stats() entry_count should match actual count."""
backend.set("key1", make_entry(hash_key="key1"))
backend.set("key2", make_entry(hash_key="key2"))
stats = backend.get_stats()
assert stats["entry_count"] == 2
def test_get_stats_bytes_used_increases(self, backend: InMemoryBackend) -> None:
"""get_stats() bytes_used should increase with entries."""
stats_empty = backend.get_stats()
backend.set(
"key1",
make_entry(hash_key="key1", original="x" * 1000),
)
stats_one = backend.get_stats()
backend.set(
"key2",
make_entry(hash_key="key2", original="y" * 1000),
)
stats_two = backend.get_stats()
assert stats_one["bytes_used"] > stats_empty["bytes_used"]
assert stats_two["bytes_used"] > stats_one["bytes_used"]
# --- Thread safety ---
def test_concurrent_set_operations(self, backend: InMemoryBackend) -> None:
"""Backend should handle concurrent set operations safely."""
num_threads = 10
entries_per_thread = 100
errors: list[Exception] = []
def worker(thread_id: int) -> None:
try:
for i in range(entries_per_thread):
key = f"thread{thread_id}_entry{i}"
entry = make_entry(hash_key=key)
backend.set(key, entry)
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=worker, args=(i,)) for i in range(num_threads)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(errors) == 0
assert backend.count() == num_threads * entries_per_thread
def test_concurrent_get_set_delete(self, backend: InMemoryBackend) -> None:
"""Backend should handle mixed concurrent operations safely."""
num_iterations = 100
errors: list[Exception] = []
# Pre-populate some entries
for i in range(50):
backend.set(f"key{i}", make_entry(hash_key=f"key{i}"))
def setter() -> None:
try:
for i in range(num_iterations):
backend.set(f"new_key{i}", make_entry(hash_key=f"new_key{i}"))
except Exception as e:
errors.append(e)
def getter() -> None:
try:
for i in range(num_iterations):
backend.get(f"key{i % 50}")
except Exception as e:
errors.append(e)
def deleter() -> None:
try:
for i in range(num_iterations):
backend.delete(f"key{i % 50}")
except Exception as e:
errors.append(e)
threads = [
threading.Thread(target=setter),
threading.Thread(target=setter),
threading.Thread(target=getter),
threading.Thread(target=getter),
threading.Thread(target=deleter),
]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(errors) == 0
# --- Edge cases ---
def test_empty_string_key(self, backend: InMemoryBackend) -> None:
"""Backend should handle empty string as key."""
entry = make_entry(hash_key="")
backend.set("", entry)
retrieved = backend.get("")
assert retrieved is not None
assert retrieved.hash == ""
def test_unicode_content(self, backend: InMemoryBackend) -> None:
"""Backend should handle unicode content correctly."""
entry = make_entry(
hash_key="unicode",
original="日本語テスト 🎉 émojis",
compressed="日本語",
)
backend.set("unicode", entry)
retrieved = backend.get("unicode")
assert retrieved is not None
assert retrieved.original_content == "日本語テスト 🎉 émojis"
assert retrieved.compressed_content == "日本語"
def test_large_content(self, backend: InMemoryBackend) -> None:
"""Backend should handle large content."""
large_content = "x" * 10_000_000 # 10MB
entry = make_entry(hash_key="large", original=large_content)
backend.set("large", entry)
retrieved = backend.get("large")
assert retrieved is not None
assert len(retrieved.original_content) == 10_000_000
# --- Parameterized tests for all backend implementations ---
def all_backends() -> list[Callable[[], CompressionStoreBackend]]:
"""Return factory functions for all backend implementations."""
return [
InMemoryBackend,
# Add more backends here as they're implemented:
# MongoDBBackend,
# RedisBackend,
]
@pytest.mark.parametrize("backend_factory", all_backends())
class TestBackendContract:
"""Contract tests that ALL backends must pass.
These tests are parameterized to run against every backend implementation.
Add new backends to all_backends() to include them in these tests.
"""
def test_implements_protocol(
self, backend_factory: Callable[[], CompressionStoreBackend]
) -> None:
"""All backends must implement CompressionStoreBackend protocol."""
backend = backend_factory()
assert isinstance(backend, CompressionStoreBackend)
def test_basic_crud_cycle(self, backend_factory: Callable[[], CompressionStoreBackend]) -> None:
"""All backends must support basic CRUD operations."""
backend = backend_factory()
# Create
entry = make_entry(hash_key="test")
backend.set("test", entry)
assert backend.exists("test")
# Read
retrieved = backend.get("test")
assert retrieved is not None
assert retrieved.original_content == entry.original_content
# Update (overwrite)
entry2 = make_entry(hash_key="test", original="updated")
backend.set("test", entry2)
retrieved2 = backend.get("test")
assert retrieved2 is not None
assert retrieved2.original_content == "updated"
# Delete
assert backend.delete("test") is True
assert backend.exists("test") is False
assert backend.get("test") is None
def test_clear_works(self, backend_factory: Callable[[], CompressionStoreBackend]) -> None:
"""All backends must support clear()."""
backend = backend_factory()
backend.set("key1", make_entry(hash_key="key1"))
backend.set("key2", make_entry(hash_key="key2"))
assert backend.count() == 2
backend.clear()
assert backend.count() == 0
def test_stats_has_required_fields(
self, backend_factory: Callable[[], CompressionStoreBackend]
) -> None:
"""All backends must return required stats fields."""
backend = backend_factory()
stats = backend.get_stats()
assert "backend_type" in stats
assert "entry_count" in stats
assert isinstance(stats["backend_type"], str)
assert isinstance(stats["entry_count"], int)
+140
View File
@@ -0,0 +1,140 @@
"""Tests for cache base types and interfaces."""
from headroom.cache.base import (
BreakpointLocation,
CacheBreakpoint,
CacheConfig,
CacheMetrics,
CacheResult,
CacheStrategy,
OptimizationContext,
)
class TestCacheStrategy:
"""Test CacheStrategy enum."""
def test_strategies_exist(self):
"""Test all expected strategies exist."""
assert CacheStrategy.PREFIX_STABILIZATION.value == "prefix_stabilization"
assert CacheStrategy.EXPLICIT_BREAKPOINTS.value == "explicit_breakpoints"
assert CacheStrategy.CACHED_CONTENT.value == "cached_content"
assert CacheStrategy.NONE.value == "none"
class TestCacheConfig:
"""Test CacheConfig dataclass."""
def test_default_values(self):
"""Test default configuration values."""
config = CacheConfig()
assert config.enabled is True
assert config.min_cacheable_tokens == 1024
assert config.max_breakpoints == 4
assert config.normalize_whitespace is True
assert config.collapse_blank_lines is True
def test_custom_values(self):
"""Test custom configuration."""
config = CacheConfig(
enabled=False,
min_cacheable_tokens=2048,
max_breakpoints=2,
)
assert config.enabled is False
assert config.min_cacheable_tokens == 2048
assert config.max_breakpoints == 2
def test_date_patterns(self):
"""Test date patterns are set."""
config = CacheConfig()
assert len(config.date_patterns) > 0
assert any("Today" in p for p in config.date_patterns)
class TestCacheMetrics:
"""Test CacheMetrics dataclass."""
def test_default_values(self):
"""Test default metrics values."""
metrics = CacheMetrics()
assert metrics.stable_prefix_tokens == 0
assert metrics.breakpoints_inserted == 0
assert metrics.estimated_cache_hit is False
assert metrics.estimated_savings_percent == 0.0
def test_custom_values(self):
"""Test custom metrics."""
metrics = CacheMetrics(
stable_prefix_tokens=5000,
breakpoints_inserted=2,
estimated_cache_hit=True,
estimated_savings_percent=90.0,
)
assert metrics.stable_prefix_tokens == 5000
assert metrics.breakpoints_inserted == 2
assert metrics.estimated_cache_hit is True
assert metrics.estimated_savings_percent == 90.0
class TestCacheBreakpoint:
"""Test CacheBreakpoint dataclass."""
def test_breakpoint_creation(self):
"""Test creating a breakpoint."""
bp = CacheBreakpoint(
message_index=0,
location=BreakpointLocation.AFTER_SYSTEM,
tokens_at_breakpoint=2000,
reason="System prompt is cacheable",
)
assert bp.message_index == 0
assert bp.location == BreakpointLocation.AFTER_SYSTEM
assert bp.tokens_at_breakpoint == 2000
assert bp.content_index is None
class TestCacheResult:
"""Test CacheResult dataclass."""
def test_result_creation(self):
"""Test creating a cache result."""
messages = [{"role": "system", "content": "Hello"}]
result = CacheResult(
messages=messages,
metrics=CacheMetrics(cacheable_tokens=1000),
transforms_applied=["normalized_whitespace"],
)
assert result.messages == messages
assert result.metrics.cacheable_tokens == 1000
assert "normalized_whitespace" in result.transforms_applied
def test_semantic_cache_hit(self):
"""Test semantic cache hit result."""
result = CacheResult(
messages=[],
semantic_cache_hit=True,
cached_response={"text": "cached response"},
)
assert result.semantic_cache_hit is True
assert result.cached_response["text"] == "cached response"
class TestOptimizationContext:
"""Test OptimizationContext dataclass."""
def test_context_creation(self):
"""Test creating optimization context."""
context = OptimizationContext(
provider="anthropic",
model="claude-3-opus",
request_id="req-123",
)
assert context.provider == "anthropic"
assert context.model == "claude-3-opus"
assert context.request_id == "req-123"
def test_default_timestamp(self):
"""Test default timestamp is set."""
context = OptimizationContext()
assert context.timestamp is not None
+636
View File
@@ -0,0 +1,636 @@
"""Tests for HeadroomClient cache optimizer integration."""
import os
import tempfile
from dataclasses import dataclass
from unittest.mock import MagicMock, patch
import pytest
from headroom import (
AnthropicCacheOptimizer,
HeadroomClient,
)
from headroom.cache.base import CacheMetrics, CacheResult
@pytest.fixture
def temp_db():
"""Create a temporary database file."""
fd, path = tempfile.mkstemp(suffix=".db")
os.close(fd)
yield f"sqlite:///{path}"
if os.path.exists(path):
os.unlink(path)
class MockTokenCounter:
"""Mock token counter for testing."""
def count_text(self, text: str) -> int:
"""Count tokens in text (required by Tokenizer interface)."""
return len(text) // 4
def count_tokens(self, text: str) -> int:
"""Alias for count_text."""
return self.count_text(text)
def count_message(self, message: dict) -> int:
"""Count tokens in a single message."""
content = message.get("content", "")
if isinstance(content, str):
return len(content) // 4
elif isinstance(content, list):
total = 0
for block in content:
if isinstance(block, dict):
total += len(block.get("text", "")) // 4
return total
return 0
def count_messages(self, messages: list) -> int:
"""Count tokens in messages."""
return sum(self.count_message(msg) for msg in messages)
class MockAnthropicProvider:
"""Mock Anthropic provider for testing."""
name = "anthropic"
def get_token_counter(self, model: str):
return MockTokenCounter()
def get_context_limit(self, model: str) -> int:
return 200000
class MockOpenAIProvider:
"""Mock OpenAI provider for testing."""
name = "openai"
def get_token_counter(self, model: str):
return MockTokenCounter()
def get_context_limit(self, model: str) -> int:
return 128000
# Mock response classes for testing (avoid MagicMock in sqlite)
@dataclass
class MockTextBlock:
"""Mock text block for Anthropic response."""
type: str = "text"
text: str = "Hello!"
@dataclass
class MockUsage:
"""Mock usage for Anthropic response."""
input_tokens: int = 100
output_tokens: int = 20
@dataclass
class MockAnthropicResponse:
"""Mock Anthropic API response."""
content: list = None
usage: MockUsage = None
model: str = "claude-sonnet-4-20250514"
id: str = "msg_123"
stop_reason: str = "end_turn"
def __post_init__(self):
if self.content is None:
self.content = [MockTextBlock()]
if self.usage is None:
self.usage = MockUsage()
class TestHeadroomClientCacheIntegration:
"""Test HeadroomClient cache optimizer integration."""
def test_auto_detect_anthropic_optimizer(self, temp_db):
"""Test that Anthropic optimizer is auto-detected."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=True,
)
assert client._cache_optimizer is not None
assert client._cache_optimizer.name == "anthropic-cache-optimizer"
def test_auto_detect_openai_optimizer(self, temp_db):
"""Test that OpenAI optimizer is auto-detected."""
mock_client = MagicMock()
provider = MockOpenAIProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=True,
)
assert client._cache_optimizer is not None
assert client._cache_optimizer.name == "openai-prefix-stabilizer"
def test_custom_optimizer(self, temp_db):
"""Test using a custom optimizer."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
custom_optimizer = AnthropicCacheOptimizer()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
cache_optimizer=custom_optimizer,
)
assert client._cache_optimizer is custom_optimizer
def test_disable_cache_optimizer(self, temp_db):
"""Test disabling cache optimizer."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=False,
)
assert client._cache_optimizer is None
def test_semantic_cache_layer_creation(self, temp_db):
"""Test semantic cache layer is created when enabled."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=True,
enable_semantic_cache=True,
)
assert client._semantic_cache_layer is not None
assert client._cache_optimizer is not None
def test_extract_query_from_string_content(self, temp_db):
"""Test query extraction from string content."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "What is 2+2?"},
]
query = client._extract_query(messages)
assert query == "What is 2+2?"
def test_extract_query_from_content_blocks(self, temp_db):
"""Test query extraction from content block format."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
messages = [
{"role": "system", "content": "You are helpful."},
{
"role": "user",
"content": [{"type": "text", "text": "What is 2+2?"}],
},
]
query = client._extract_query(messages)
assert query == "What is 2+2?"
def test_extract_query_last_user_message(self, temp_db):
"""Test that query extraction uses last user message."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
messages = [
{"role": "user", "content": "First question"},
{"role": "assistant", "content": "First answer"},
{"role": "user", "content": "Second question"},
]
query = client._extract_query(messages)
assert query == "Second question"
def test_config_propagation(self, temp_db):
"""Test that config is properly propagated."""
mock_client = MagicMock()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=True,
enable_semantic_cache=True,
)
assert client._config.cache_optimizer.enabled is True
assert client._config.cache_optimizer.enable_semantic_cache is True
class TestCacheOptimizerInvocation:
"""Test that cache optimizer is actually INVOKED during chat completion.
These tests catch bugs where the optimizer is assigned but never called
in the production code path.
"""
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_optimizer_optimize_is_called_during_chat(self, mock_save, temp_db):
"""CRITICAL: Verify optimizer.optimize() is called during chat completion.
This test catches the gap where tests verify assignment but not invocation.
Note: Cache optimizer is only invoked in OPTIMIZE mode, not AUDIT mode (the default).
"""
from headroom import HeadroomMode
# Use module-level mock classes to avoid sqlite issues with MagicMock
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
# Create a spy optimizer to track calls
real_optimizer = AnthropicCacheOptimizer()
spy_optimize = MagicMock(
return_value=CacheResult(
messages=[{"role": "user", "content": "test"}],
metrics=CacheMetrics(
cacheable_tokens=100,
breakpoints_inserted=1,
estimated_cache_hit=False,
estimated_savings_percent=0.0,
),
transforms_applied=["test_transform"],
)
)
real_optimizer.optimize = spy_optimize
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
cache_optimizer=real_optimizer,
)
# Make a chat completion call in OPTIMIZE mode (cache optimizer only runs in OPTIMIZE mode)
messages = [
{"role": "user", "content": "Hello, how are you?"},
]
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
# CRITICAL: Verify optimizer.optimize() was actually called
assert spy_optimize.called, (
"Cache optimizer.optimize() should be called during chat completion. "
"If this fails, the optimizer is assigned but never invoked."
)
# Verify it was called with the right arguments
call_args = spy_optimize.call_args
assert call_args is not None
optimized_messages, context = call_args[0]
assert len(optimized_messages) >= 1, "Should pass messages to optimizer"
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_optimizer_transforms_applied_in_response(self, mock_save, temp_db):
"""Verify optimizer transforms are reported in the response metadata."""
from headroom import HeadroomMode
# Use module-level mock classes to avoid sqlite issues with MagicMock
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
# Create optimizer that applies a transform
real_optimizer = AnthropicCacheOptimizer()
real_optimizer.optimize = MagicMock(
return_value=CacheResult(
messages=[{"role": "user", "content": "test"}],
metrics=CacheMetrics(
cacheable_tokens=500,
breakpoints_inserted=2,
estimated_cache_hit=True,
estimated_savings_percent=0.5,
),
transforms_applied=["add_cache_control"],
)
)
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
cache_optimizer=real_optimizer,
)
messages = [
{"role": "user", "content": "x" * 1000}, # Large message
]
# Use OPTIMIZE mode so cache optimizer is invoked
result = client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
# Verify the response includes cache optimizer info
assert hasattr(result, "headroom"), "Response should have headroom metadata"
headroom_meta = result.headroom
# Check that cache optimizer was reported
assert headroom_meta.cache_optimizer_used is not None or any(
"cache_optimizer" in t for t in (headroom_meta.transforms_applied or [])
), "Cache optimizer usage should be reported in metadata"
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_optimizer_not_called_in_audit_mode(self, mock_save, temp_db):
"""Verify optimizer is NOT called in AUDIT mode (observe only)."""
from headroom import HeadroomMode
# Use module-level mock classes to avoid sqlite issues with MagicMock
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
spy_optimize = MagicMock(
return_value=CacheResult(
messages=[{"role": "user", "content": "test"}],
metrics=CacheMetrics(),
)
)
real_optimizer = AnthropicCacheOptimizer()
real_optimizer.optimize = spy_optimize
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
cache_optimizer=real_optimizer,
)
messages = [{"role": "user", "content": "Hello"}]
# Make call in AUDIT mode (observe only, no modifications)
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.AUDIT,
)
# Optimizer should NOT be called in AUDIT mode
assert not spy_optimize.called, "Cache optimizer should NOT be called in AUDIT mode"
class TestSemanticCacheIntegration:
"""Test semantic cache integration with HeadroomClient.
These tests verify the full production code path for semantic caching,
including that cache hits actually return cached responses without calling
the underlying API.
"""
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_semantic_cache_hit_returns_cached_response_without_api_call(self, mock_save, temp_db):
"""CRITICAL: Verify semantic cache hit returns cached response without API call.
This test catches the gap where semantic cache is enabled but cached
responses are never actually returned (API is always called).
"""
from headroom import HeadroomMode
# Mock OpenAI-style response (chat.completions.create uses OpenAI API style)
mock_client = MagicMock()
mock_openai_response = MagicMock()
mock_openai_response.choices = [MagicMock(message=MagicMock(content="4"))]
mock_openai_response.usage = MagicMock(
prompt_tokens=10, completion_tokens=5, total_tokens=15
)
mock_openai_response.model = "claude-sonnet-4-20250514"
mock_openai_response.id = "chatcmpl-123"
mock_client.chat.completions.create.return_value = mock_openai_response
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
enable_cache_optimizer=True,
enable_semantic_cache=True,
)
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "What is 2+2?"},
]
# First call - should call API and potentially cache
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
first_call_count = mock_client.chat.completions.create.call_count
assert first_call_count == 1, "First call should hit API"
# Manually store response in semantic cache for test
if client._semantic_cache_layer is not None:
from headroom.cache import OptimizationContext
context = OptimizationContext(
provider="anthropic",
model="claude-sonnet-4-20250514",
query="What is 2+2?",
)
client._semantic_cache_layer.store_response(
messages,
{"text": "4", "role": "assistant"},
context,
)
# Second call with same messages - should hit cache, NOT call API
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
second_call_count = mock_client.chat.completions.create.call_count
# If semantic cache is working, API should NOT be called again
assert second_call_count == 1, (
f"Semantic cache hit should NOT call API. "
f"Expected 1 API call, got {second_call_count}. "
"If this fails, cached responses are not being returned."
)
class TestSessionStatsTracking:
"""Test session statistics tracking in HeadroomClient.
These tests verify that session stats are actually updated during
chat completion calls.
"""
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_session_stats_incremented_after_request(self, mock_save, temp_db):
"""CRITICAL: Verify session stats are incremented after requests."""
from headroom import HeadroomMode
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
# Get initial stats
initial_stats = client.get_stats()
initial_requests = initial_stats["session"]["requests_total"]
# Make a request in AUDIT mode
messages = [{"role": "user", "content": "Hello"}]
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.AUDIT,
)
# Verify stats were updated
after_stats = client.get_stats()
after_requests = after_stats["session"]["requests_total"]
assert after_requests == initial_requests + 1, (
f"requests_total should increment. Before: {initial_requests}, After: {after_requests}"
)
assert after_stats["session"]["requests_audit"] >= 1, (
"requests_audit should be at least 1 after AUDIT mode request"
)
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_session_stats_tracks_optimize_mode(self, mock_save, temp_db):
"""Verify session stats track OPTIMIZE mode requests separately."""
from headroom import HeadroomMode
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
messages = [{"role": "user", "content": "Hello"}]
# Make request in OPTIMIZE mode
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
stats = client.get_stats()
assert stats["session"]["requests_optimized"] >= 1, (
"requests_optimized should be at least 1 after OPTIMIZE mode request"
)
@patch("headroom.storage.sqlite.SQLiteStorage.save")
def test_session_stats_tracks_tokens_saved(self, mock_save, temp_db):
"""Verify session stats track tokens saved."""
from headroom import HeadroomMode
mock_client = MagicMock()
mock_client.messages.create.return_value = MockAnthropicResponse()
provider = MockAnthropicProvider()
client = HeadroomClient(
original_client=mock_client,
provider=provider,
store_url=temp_db,
)
# Create a conversation that will trigger some optimization
messages = [
{"role": "system", "content": "You are helpful. " * 100},
{"role": "user", "content": "Hello"},
]
client.chat.completions.create(
model="claude-sonnet-4-20250514",
messages=messages,
max_tokens=100,
headroom_mode=HeadroomMode.OPTIMIZE,
)
stats = client.get_stats()
# tokens_saved_total should be tracked (may be 0 if no compression)
assert "tokens_saved_total" in stats["session"], (
"Session stats should track tokens_saved_total"
)
+557
View File
@@ -0,0 +1,557 @@
"""Tests for the dynamic content detector."""
import pytest
from headroom.cache.dynamic_detector import (
DetectionResult,
DetectorConfig,
DynamicCategory,
DynamicContentDetector,
RegexDetector,
detect_dynamic_content,
)
class TestRegexDetector:
"""Test the Tier 1 regex detector."""
@pytest.fixture
def detector(self):
"""Create a regex detector."""
config = DetectorConfig(tiers=["regex"])
return RegexDetector(config)
def test_iso_date(self, detector):
"""Test ISO date detection."""
spans = detector.detect("The date is 2024-01-15.")
assert len(spans) == 1
assert spans[0].text == "2024-01-15"
assert spans[0].category == DynamicCategory.DATE
assert spans[0].tier == "regex"
def test_structural_detection(self, detector):
"""Test structural detection via 'Label: value' patterns."""
# New scalable approach: detect via structural "Today: value" pattern
spans = detector.detect("Date: 2024-01-15")
assert len(spans) == 1
assert spans[0].text == "2024-01-15"
assert spans[0].category == DynamicCategory.DATE
# Test user label detection
spans = detector.detect("User: john.doe@example.com")
user_spans = [s for s in spans if s.category == DynamicCategory.USER_DATA]
assert len(user_spans) == 1
def test_datetime_iso(self, detector):
"""Test ISO datetime detection."""
spans = detector.detect("Timestamp: 2024-01-15T10:30:00Z")
assert len(spans) == 1
assert spans[0].text == "2024-01-15T10:30:00Z"
assert spans[0].category == DynamicCategory.DATETIME
def test_uuid(self, detector):
"""Test UUID detection."""
spans = detector.detect("ID: 550e8400-e29b-41d4-a716-446655440000")
assert len(spans) == 1
assert spans[0].text == "550e8400-e29b-41d4-a716-446655440000"
assert spans[0].category == DynamicCategory.UUID
def test_request_id(self, detector):
"""Test request ID detection."""
spans = detector.detect("Request: req_abc123def456ghi789")
assert len(spans) == 1
assert "req_" in spans[0].text
assert spans[0].category == DynamicCategory.REQUEST_ID
def test_unix_timestamp(self, detector):
"""Test Unix timestamp detection."""
spans = detector.detect("Time: 1705312200")
assert len(spans) == 1
assert spans[0].text == "1705312200"
assert spans[0].category == DynamicCategory.TIMESTAMP
def test_time(self, detector):
"""Test time detection."""
spans = detector.detect("Meeting at 10:30 AM")
assert len(spans) == 1
assert spans[0].text == "10:30 AM"
assert spans[0].category == DynamicCategory.TIME
def test_version(self, detector):
"""Test version number detection."""
spans = detector.detect("Running v2.3.1-beta")
assert len(spans) == 1
assert spans[0].text == "v2.3.1-beta"
assert spans[0].category == DynamicCategory.VERSION
def test_date_prefix_pattern(self, detector):
"""Test full date prefix phrase detection."""
spans = detector.detect("Today is Monday, January 15, 2024. You are an assistant.")
assert len(spans) >= 1
# Should detect the full phrase
date_spans = [s for s in spans if s.category == DynamicCategory.DATE]
assert len(date_spans) >= 1
def test_multiple_dynamic_elements(self, detector):
"""Test detecting multiple dynamic elements."""
content = """
Date: 2024-01-15
Time: 10:30:00
Request ID: req_abc123def456ghi789xyz
UUID: 550e8400-e29b-41d4-a716-446655440000
"""
spans = detector.detect(content)
assert len(spans) == 4
categories = {s.category for s in spans}
assert DynamicCategory.DATE in categories
assert DynamicCategory.TIME in categories
assert DynamicCategory.REQUEST_ID in categories
assert DynamicCategory.UUID in categories
def test_no_false_positives_on_static(self, detector):
"""Test that static content doesn't trigger false positives."""
spans = detector.detect("You are a helpful assistant. Answer questions clearly.")
assert len(spans) == 0
def test_positions_are_correct(self, detector):
"""Test that span positions are correct."""
content = "Date: 2024-01-15"
spans = detector.detect(content)
assert len(spans) == 1
assert content[spans[0].start : spans[0].end] == spans[0].text
class TestDynamicContentDetector:
"""Test the unified dynamic content detector."""
def test_regex_only(self):
"""Test detector with regex tier only."""
config = DetectorConfig(tiers=["regex"])
detector = DynamicContentDetector(config)
result = detector.detect("Today is 2024-01-15. You are helpful.")
assert len(result.spans) == 1
assert result.spans[0].text == "2024-01-15"
assert "regex" in result.tiers_used
assert result.processing_time_ms < 10 # Should be very fast
def test_static_dynamic_split(self):
"""Test that content is properly split."""
config = DetectorConfig(tiers=["regex"])
detector = DynamicContentDetector(config)
result = detector.detect("Today is 2024-01-15. You are helpful.")
assert "2024-01-15" not in result.static_content
assert "2024-01-15" in result.dynamic_content
assert "You are helpful" in result.static_content
def test_complex_content(self):
"""Test with realistic system prompt."""
config = DetectorConfig(tiers=["regex"])
detector = DynamicContentDetector(config)
content = """You are a helpful AI assistant.
Today is January 15, 2024.
Current session: sess_abc123def456ghi789xyz
Instructions:
1. Be concise
2. Be accurate
3. Be helpful
Request ID: req_xyz789abc123def456ghi"""
result = detector.detect(content)
# Should find date, session ID, request ID
assert len(result.spans) >= 2
categories = {s.category for s in result.spans}
assert DynamicCategory.DATE in categories or DynamicCategory.REQUEST_ID in categories
def test_empty_content(self):
"""Test with empty content."""
detector = DynamicContentDetector()
result = detector.detect("")
assert len(result.spans) == 0
assert result.static_content == ""
assert result.dynamic_content == ""
def test_no_dynamic_content(self):
"""Test with fully static content."""
detector = DynamicContentDetector()
content = "You are a helpful assistant. Answer questions clearly and concisely."
result = detector.detect(content)
assert len(result.spans) == 0
assert result.static_content == content
assert result.dynamic_content == ""
def test_custom_patterns(self):
"""Test adding custom regex patterns."""
config = DetectorConfig(
tiers=["regex"],
custom_patterns=[
(r"CUSTOM_\d{4}", DynamicCategory.REQUEST_ID),
],
)
detector = DynamicContentDetector(config)
result = detector.detect("Code: CUSTOM_1234")
custom_spans = [s for s in result.spans if s.text == "CUSTOM_1234"]
assert len(custom_spans) == 1
def test_available_tiers(self):
"""Test that available_tiers reflects actual availability."""
config = DetectorConfig(tiers=["regex", "ner", "semantic"])
detector = DynamicContentDetector(config)
# Regex should always be available
assert "regex" in detector.available_tiers
# NER and semantic depend on optional dependencies
# They may or may not be available
def test_warnings_for_missing_dependencies(self):
"""Test that warnings are generated for missing dependencies."""
config = DetectorConfig(tiers=["regex", "ner", "semantic"])
detector = DynamicContentDetector(config)
detector.detect("Test content")
# If NER/semantic not installed, should have warnings
# (This test passes either way - it's informational)
# If deps ARE installed, no warnings. If not, warnings present.
class TestConvenienceFunction:
"""Test the detect_dynamic_content convenience function."""
def test_basic_usage(self):
"""Test basic convenience function usage."""
result = detect_dynamic_content("Date: 2024-01-15")
assert isinstance(result, DetectionResult)
assert len(result.spans) == 1
assert result.spans[0].text == "2024-01-15"
def test_with_tiers(self):
"""Test specifying tiers."""
result = detect_dynamic_content(
"Date: 2024-01-15",
tiers=["regex"],
)
assert "regex" in result.tiers_used
class TestEntropyDetection:
"""Test entropy-based detection for random IDs/tokens."""
def test_high_entropy_string(self):
"""Test that high-entropy strings are detected."""
from headroom.cache.dynamic_detector import calculate_entropy
# High entropy strings (random-looking)
assert calculate_entropy("abc123xyz789def") > 0.7
assert calculate_entropy("550e8400e29b41d4") > 0.7
# Low entropy strings (repetitive)
assert calculate_entropy("aaaaaaaaaa") < 0.3
assert calculate_entropy("abababab") < 0.6
def test_entropy_detection_finds_ids(self):
"""Test that entropy detection finds random IDs."""
detector = DynamicContentDetector()
# Random-looking ID that isn't covered by universal patterns
result = detector.detect("Auth: xK7mN2pQr9sT4vW")
# Should find the ID via entropy or structural detection
assert len(result.spans) >= 1
def test_entropy_skips_common_words(self):
"""Test that common words aren't flagged as high-entropy."""
detector = DynamicContentDetector()
# These words have mixed case/numbers but aren't IDs
result = detector.detect("Use username and password correctly.")
# "username" and "password" shouldn't be detected
flagged_words = [s.text for s in result.spans]
assert "username" not in flagged_words
assert "password" not in flagged_words
class TestEdgeCases:
"""Test edge cases and tricky inputs."""
def test_overlapping_patterns(self):
"""Test that overlapping patterns don't cause duplicates."""
detector = DynamicContentDetector()
# ISO datetime contains ISO date - shouldn't match both
result = detector.detect("Time: 2024-01-15T10:30:00Z")
# Should match datetime, not date separately
assert len(result.spans) == 1
assert result.spans[0].category == DynamicCategory.DATETIME
def test_adjacent_dynamic_content(self):
"""Test adjacent dynamic elements."""
detector = DynamicContentDetector()
result = detector.detect("2024-01-15 10:30:00")
# Should find both date and time
assert len(result.spans) == 2
def test_very_long_content(self):
"""Test with long content."""
detector = DynamicContentDetector()
# Create long content with some dynamic parts
static_parts = ["This is static text. "] * 100
content = "".join(static_parts) + "Date: 2024-01-15. " + "".join(static_parts)
result = detector.detect(content)
assert len(result.spans) == 1
assert result.processing_time_ms < 100 # Should still be fast
def test_special_characters(self):
"""Test content with special characters."""
detector = DynamicContentDetector()
content = "Date: 2024-01-15\nUUID: 550e8400-e29b-41d4-a716-446655440000\n\n---\n"
result = detector.detect(content)
assert len(result.spans) == 2
def test_unicode_content(self):
"""Test with Unicode content."""
detector = DynamicContentDetector()
content = "日期: 2024-01-15. Héllo wörld!"
result = detector.detect(content)
# Should still find the date
assert len(result.spans) == 1
assert result.spans[0].text == "2024-01-15"
class TestCacheAlignmentScenarios:
"""Test scenarios relevant to cache alignment."""
def test_system_prompt_dates(self):
"""Test extracting dates from system prompts."""
detector = DynamicContentDetector()
content = """You are Claude, an AI assistant by Anthropic.
Today is Monday, January 15, 2024.
Current time: 10:30 AM PST.
Your task is to help users with coding questions."""
result = detector.detect(content)
# Should extract date and time
assert len(result.spans) >= 1
# Static content should not have dates
assert "2024" not in result.static_content or "January" in result.static_content
# Dynamic content should have the dates
assert (
"January" in result.dynamic_content
or "2024-01-15" in result.dynamic_content
or "10:30" in result.dynamic_content
)
def test_request_metadata(self):
"""Test extracting request metadata."""
detector = DynamicContentDetector()
content = """Request ID: req_abc123xyz789
Trace ID: 550e8400-e29b-41d4-a716-446655440000
Timestamp: 1705312200
Process the following query:"""
result = detector.detect(content)
# Should find request ID, UUID, timestamp
{s.category for s in result.spans}
assert len(result.spans) >= 2
def test_mixed_static_dynamic(self):
"""Test content with interspersed static and dynamic parts."""
detector = DynamicContentDetector()
content = """You are helpful (static).
Today is 2024-01-15 (dynamic).
Always be accurate (static).
Session: sess_abc123xyz789 (dynamic).
Never lie (static)."""
result = detector.detect(content)
# Should find date and session ID
assert len(result.spans) >= 1
# Static content should preserve the static parts
assert "helpful" in result.static_content
assert "accurate" in result.static_content
class TestNERDetector:
"""Test Tier 2 NER detector (if spaCy available)."""
@pytest.fixture
def ner_detector(self):
"""Create detector with NER enabled."""
from headroom.cache.dynamic_detector import _SPACY_AVAILABLE, NERDetector
if not _SPACY_AVAILABLE:
pytest.skip("spaCy not installed")
config = DetectorConfig(tiers=["ner"])
detector = NERDetector(config)
if not detector.is_available:
pytest.skip("spaCy model not available")
return detector
def test_person_detection(self, ner_detector):
"""Test detecting person names."""
spans, _ = ner_detector.detect("John Smith sent the message.")
[s for s in spans if s.category == DynamicCategory.PERSON]
# NER might or might not detect "John Smith" depending on model
# This is more of an integration test
def test_money_detection(self, ner_detector):
"""Test detecting money amounts."""
spans, _ = ner_detector.detect("The total is $500.00")
[s for s in spans if s.category == DynamicCategory.MONEY]
# May or may not detect depending on spaCy model
class TestSemanticDetector:
"""Test Tier 3 semantic detector (if sentence-transformers available)."""
@pytest.fixture
def semantic_detector(self):
"""Create detector with semantic enabled."""
from headroom.cache.dynamic_detector import (
_SENTENCE_TRANSFORMERS_AVAILABLE,
SemanticDetector,
)
if not _SENTENCE_TRANSFORMERS_AVAILABLE:
pytest.skip("sentence-transformers not installed")
config = DetectorConfig(tiers=["semantic"])
detector = SemanticDetector(config)
if not detector.is_available:
pytest.skip("Embedding model not available")
return detector
def test_realtime_detection(self, semantic_detector):
"""Test detecting real-time/volatile content."""
content = "The current stock price is updated every minute."
spans, _ = semantic_detector.detect(content)
# Should detect this as volatile/realtime
# Depends on similarity threshold
def test_missing_exemplar_embeddings_returns_warning(self):
"""Semantic detector reports unavailable state when embeddings are missing."""
from headroom.cache.dynamic_detector import SemanticDetector
detector = object.__new__(SemanticDetector)
detector.config = DetectorConfig(tiers=["semantic"])
detector._model = object()
detector._exemplar_embeddings = None
detector._load_error = None
spans, warning = detector.detect("The current stock price changes every minute.")
assert spans == []
# Model present but exemplar matrix missing → the warning names the
# actual missing piece (matches TestSemanticDetectorGuards below).
assert warning == "exemplar embeddings not initialized"
class TestIntegrationWithAllTiers:
"""Integration tests using all available tiers."""
def test_all_tiers_together(self):
"""Test running all tiers on complex content."""
config = DetectorConfig(tiers=["regex", "ner", "semantic"])
detector = DynamicContentDetector(config)
content = """Today is January 15, 2024.
John paid $500 for the service.
Request ID: req_abc123xyz789.
The stock price updates in real-time.
Be helpful and accurate."""
result = detector.detect(content)
# Should find at least the regex matches
assert len(result.spans) >= 1
# Check processing time is reasonable
# NER + semantic might add 50-100ms
assert result.processing_time_ms < 5000 # Very generous timeout
# Should have used at least regex
assert "regex" in result.tiers_used
def test_tier_precedence(self):
"""Test that earlier tiers take precedence."""
config = DetectorConfig(tiers=["regex", "ner"])
detector = DynamicContentDetector(config)
# Date should be caught by regex, not NER
result = detector.detect("Date: 2024-01-15")
assert len(result.spans) == 1
assert result.spans[0].tier == "regex"
class TestSemanticDetectorGuards:
"""Defensive guards in SemanticDetector.detect()."""
def test_none_exemplars_early_return(self):
"""detect() must early-return, not crash, when exemplar embeddings
are unset while a model is present.
Regression for the `None.T` guard: `is_available` only checks
`_model`, so `_exemplar_embeddings` can be None at the `np.dot`
call. The guard returns the method's `(spans, warning)` contract.
"""
np = pytest.importorskip("numpy")
from unittest.mock import MagicMock
from headroom.cache.dynamic_detector import SemanticDetector
det = object.__new__(SemanticDetector)
det._model = MagicMock()
det._model.encode.return_value = np.zeros((1, 3))
det._exemplar_embeddings = None
det._load_error = None
spans, warning = det.detect("This is a sentence here. Here is another long one.")
assert spans == []
assert warning == "exemplar embeddings not initialized"
+272
View File
@@ -0,0 +1,272 @@
"""Tests for GoogleCacheOptimizer."""
from datetime import datetime, timedelta
import pytest
from headroom.cache import CacheConfig, GoogleCacheOptimizer, OptimizationContext
from headroom.cache.base import CacheStrategy
from headroom.cache.google import (
GOOGLE_CACHE_DISCOUNT,
GOOGLE_MIN_CACHE_TOKENS,
CacheabilityAnalysis,
CachedContentInfo,
)
class TestGoogleCacheOptimizer:
"""Test GoogleCacheOptimizer functionality."""
@pytest.fixture
def optimizer(self):
"""Create optimizer instance."""
return GoogleCacheOptimizer()
@pytest.fixture
def context(self):
"""Create optimization context."""
return OptimizationContext(
provider="google",
model="gemini-1.5-pro",
)
def test_optimizer_properties(self, optimizer):
"""Test optimizer properties."""
assert optimizer.name == "google-cached-content"
assert optimizer.provider == "google"
assert optimizer.strategy == CacheStrategy.CACHED_CONTENT
def test_enforces_minimum_tokens(self):
"""Test that Google's minimum is enforced."""
config = CacheConfig(min_cacheable_tokens=100)
optimizer = GoogleCacheOptimizer(config)
assert optimizer.config.min_cacheable_tokens >= GOOGLE_MIN_CACHE_TOKENS
def test_optimize_below_threshold(self, optimizer, context):
"""Test optimization with content below threshold."""
messages = [
{"role": "system", "content": "Short system prompt."},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.metrics.cacheable_tokens < GOOGLE_MIN_CACHE_TOKENS
assert any("32K" in w for w in result.warnings)
def test_optimize_above_threshold(self, optimizer, context):
"""Test optimization with content above threshold."""
# Create content above 32K tokens
messages = [
{"role": "system", "content": "You are helpful. " * 15000},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.metrics.cacheable_tokens > 0
if result.metrics.cacheable_tokens >= GOOGLE_MIN_CACHE_TOKENS:
assert result.metrics.estimated_savings_percent == GOOGLE_CACHE_DISCOUNT * 100
def test_analyze_cacheability(self, optimizer, context):
"""Test cacheability analysis."""
messages = [
{"role": "system", "content": "Short"},
{"role": "user", "content": "Hello!"},
]
analysis = optimizer.analyze_cacheability(messages, context)
assert isinstance(analysis, CacheabilityAnalysis)
assert analysis.total_tokens > 0
assert analysis.tokens_below_minimum > 0
assert analysis.is_cacheable is False
assert len(analysis.recommendations) > 0
def test_cache_registration(self, optimizer):
"""Test registering a cache."""
cache_info = optimizer.register_cache(
cache_id="test-cache-123",
content_hash="abc123",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
assert cache_info.cache_id == "test-cache-123"
assert cache_info.content_hash == "abc123"
assert cache_info.token_count == 50000
assert not cache_info.is_expired
assert cache_info.ttl_remaining_seconds > 0
def test_cache_lookup(self, optimizer):
"""Test looking up a registered cache."""
optimizer.register_cache(
cache_id="test-cache-456",
content_hash="def456",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
found = optimizer.get_reusable_cache("def456")
assert found is not None
assert found.cache_id == "test-cache-456"
def test_cache_lookup_expired(self, optimizer):
"""Test that expired caches are not returned."""
optimizer.register_cache(
cache_id="expired-cache",
content_hash="expired123",
token_count=50000,
expires_at=datetime.now() - timedelta(hours=1),
)
found = optimizer.get_reusable_cache("expired123")
assert found is None
def test_cache_lookup_insufficient_ttl(self, optimizer):
"""Test that caches with insufficient TTL are not returned."""
optimizer.register_cache(
cache_id="short-ttl-cache",
content_hash="shortttl123",
token_count=50000,
expires_at=datetime.now() + timedelta(seconds=30),
)
# Default min_ttl is 60 seconds
found = optimizer.get_reusable_cache("shortttl123")
assert found is None
def test_extend_cache_ttl(self, optimizer):
"""Test extending cache TTL."""
optimizer.register_cache(
cache_id="extend-cache",
content_hash="extend123",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
new_expires = datetime.now() + timedelta(hours=2)
updated = optimizer.extend_cache_ttl("extend-cache", new_expires)
assert updated is not None
assert updated.expires_at == new_expires
def test_remove_cache(self, optimizer):
"""Test removing a cache."""
optimizer.register_cache(
cache_id="remove-cache",
content_hash="remove123",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
removed = optimizer.remove_cache("remove-cache")
assert removed is True
found = optimizer.get_reusable_cache("remove123")
assert found is None
def test_cleanup_expired_caches(self, optimizer):
"""Test cleaning up expired caches."""
# Register expired cache
optimizer.register_cache(
cache_id="cleanup-expired",
content_hash="cleanup123",
token_count=50000,
expires_at=datetime.now() - timedelta(hours=1),
)
expired_ids = optimizer.cleanup_expired_caches()
assert "cleanup-expired" in expired_ids
def test_list_caches(self, optimizer):
"""Test listing caches."""
optimizer.register_cache(
cache_id="list-cache-1",
content_hash="list1",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
optimizer.register_cache(
cache_id="list-cache-2",
content_hash="list2",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=2),
)
caches = optimizer.list_caches()
assert len(caches) >= 2
def test_get_statistics(self, optimizer):
"""Test getting statistics."""
optimizer.register_cache(
cache_id="stats-cache",
content_hash="stats123",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
stats = optimizer.get_statistics()
assert "active_caches" in stats
assert "caches_created" in stats
assert stats["caches_created"] >= 1
def test_prepare_cache_creation(self, optimizer, context):
"""Test preparing cache creation parameters."""
messages = [
{"role": "system", "content": "You are helpful. " * 15000},
{"role": "user", "content": "Hello!"},
]
params = optimizer.prepare_cache_creation(messages, context)
if params is not None:
assert "contents" in params
assert "ttl" in params
assert "display_name" in params
def test_build_request_with_cache(self, optimizer):
"""Test building request with cache."""
messages = [
{"role": "system", "content": "System prompt"},
{"role": "user", "content": "User message"},
]
request = optimizer.build_request_with_cache(messages, "cache-123")
assert "cached_content" in request
assert request["cached_content"] == "cache-123"
assert "contents" in request
def test_export_import_registry(self, optimizer):
"""Test exporting and importing cache registry."""
optimizer.register_cache(
cache_id="export-cache",
content_hash="export123",
token_count=50000,
expires_at=datetime.now() + timedelta(hours=1),
)
exported = optimizer.export_cache_registry()
assert len(exported) >= 1
new_optimizer = GoogleCacheOptimizer()
imported = new_optimizer.import_cache_registry(exported)
assert imported >= 1
def test_cached_content_info_serialization(self):
"""Test CachedContentInfo serialization."""
info = CachedContentInfo(
cache_id="test",
content_hash="hash",
created_at=datetime.now(),
expires_at=datetime.now() + timedelta(hours=1),
token_count=50000,
)
data = info.to_dict()
restored = CachedContentInfo.from_dict(data)
assert restored.cache_id == info.cache_id
assert restored.content_hash == info.content_hash
assert restored.token_count == info.token_count
+197
View File
@@ -0,0 +1,197 @@
"""Tests for OpenAICacheOptimizer."""
import pytest
from headroom.cache import CacheConfig, OpenAICacheOptimizer, OptimizationContext
from headroom.cache.base import CacheStrategy
class TestOpenAICacheOptimizer:
"""Test OpenAICacheOptimizer functionality."""
@pytest.fixture
def optimizer(self):
"""Create optimizer instance."""
return OpenAICacheOptimizer()
@pytest.fixture
def context(self):
"""Create optimization context."""
return OptimizationContext(
provider="openai",
model="gpt-4",
)
def test_optimizer_properties(self, optimizer):
"""Test optimizer properties."""
assert optimizer.name == "openai-prefix-stabilizer"
assert optimizer.provider == "openai"
assert optimizer.strategy == CacheStrategy.PREFIX_STABILIZATION
def test_optimize_simple_messages(self, optimizer, context):
"""Test optimizing simple messages."""
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.messages is not None
assert len(result.messages) == 2
assert result.metrics.stable_prefix_hash != ""
def test_date_extraction(self, optimizer, context):
"""Test that dates are extracted from system prompt."""
messages = [
{
"role": "system",
"content": "Today is January 7, 2026. You are a helpful assistant.",
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
# Check that date was extracted and moved
system_content = result.messages[0]["content"]
# The date should be moved to a dynamic section at the end
assert "You are a helpful assistant" in system_content
def test_whitespace_normalization(self, optimizer, context):
"""Test whitespace normalization."""
messages = [
{
"role": "system",
"content": "You are a helpful assistant.",
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
# Whitespace should be normalized
system_content = result.messages[0]["content"]
assert " " not in system_content # Multiple spaces collapsed
def test_optimize_disabled(self, context):
"""Test optimization when disabled."""
config = CacheConfig(enabled=False)
optimizer = OpenAICacheOptimizer(config)
messages = [
{"role": "system", "content": "Test"},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.transforms_applied == []
def test_prefix_stability_tracking(self, optimizer, context):
"""Test that prefix stability is tracked."""
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
]
# First call
optimizer.optimize(messages, context)
# Second call with same messages
result2 = optimizer.optimize(messages, context)
# Second call should detect stable prefix
assert result2.metrics.estimated_cache_hit is True
assert result2.metrics.prefix_changed_from_previous is False
def test_prefix_change_detection(self, optimizer, context):
"""Test detection of prefix changes."""
messages1 = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
]
messages2 = [
{"role": "system", "content": "You are a different assistant."},
{"role": "user", "content": "Hello!"},
]
optimizer.optimize(messages1, context)
result2 = optimizer.optimize(messages2, context)
# Second call should detect prefix change
assert result2.metrics.prefix_changed_from_previous is True
def test_token_threshold_warning(self, optimizer, context):
"""Test warning when below token threshold."""
messages = [
{"role": "system", "content": "Short."},
{"role": "user", "content": "Hi"},
]
result = optimizer.optimize(messages, context)
# Should have warning about being below threshold
assert any("1024" in w for w in result.warnings)
def test_estimate_savings_below_threshold(self, optimizer, context):
"""Test savings estimation below threshold."""
messages = [
{"role": "system", "content": "Short system prompt."},
{"role": "user", "content": "Hello!"},
]
savings = optimizer.estimate_savings(messages, context)
assert savings == 0.0 # Below threshold
def test_estimate_savings_above_threshold(self, optimizer, context):
"""Test savings estimation above threshold."""
messages = [
{"role": "system", "content": "You are helpful. " * 500},
{"role": "user", "content": "Hello!"},
]
# First call to establish baseline
optimizer.optimize(messages, context)
# Second call should show savings
savings = optimizer.estimate_savings(messages, context)
assert savings > 0.0
def test_uuid_pattern_detection(self, optimizer, context):
"""Test detection of UUIDs in content."""
messages = [
{
"role": "system",
"content": "Request ID: 12345678-1234-1234-1234-123456789012. Be helpful.",
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
# UUID should be detected as dynamic content
assert result.metrics.stable_prefix_hash != ""
def test_content_block_format(self, optimizer, context):
"""Test handling of content block format."""
messages = [
{
"role": "system",
"content": [{"type": "text", "text": "You are helpful."}],
},
{"role": "user", "content": "Hello!"},
]
result = optimizer.optimize(messages, context)
assert result.messages is not None
def test_metrics_recording(self, optimizer, context):
"""Test that metrics are recorded."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello!"},
]
optimizer.optimize(messages, context)
metrics = optimizer.get_metrics()
assert metrics.stable_prefix_hash != ""
+645
View File
@@ -0,0 +1,645 @@
"""Tests for PrefixCacheTracker — cache-aware compression."""
import time
import pytest
from headroom.cache.prefix_tracker import (
MISS_COLD_START,
MISS_PREFIX_CHANGE,
MISS_TTL_EXPIRY,
MISS_UNKNOWN,
FreezeStats,
PrefixCacheTracker,
PrefixFreezeConfig,
SessionTrackerStore,
)
class TestPrefixCacheTracker:
"""Test PrefixCacheTracker core functionality."""
@pytest.fixture
def tracker(self):
return PrefixCacheTracker("anthropic")
@pytest.fixture
def openai_tracker(self):
return PrefixCacheTracker("openai")
def test_turn_0_no_freeze(self, tracker):
"""First turn should never freeze — no cache state yet."""
assert tracker.get_frozen_message_count() == 0
def test_turn_1_with_cache_hit_freezes(self, tracker):
"""After turn 1 with cache hits, turn 2 should freeze."""
messages = [
{"role": "system", "content": "You are a helpful assistant." * 100},
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"},
]
# Simulate: provider cached 2000 tokens (system + user)
token_counts = [1500, 50, 500]
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=2050,
messages=messages,
message_token_counts=token_counts,
)
# On turn 2, the first 2 messages (1500 + 50 = 1550 <= 2050) are frozen
assert tracker.get_frozen_message_count() == 3 # All 3 fit within 2050
def test_partial_freeze(self, tracker):
"""Only messages that fit within cached tokens are frozen."""
messages = [
{"role": "system", "content": "System prompt" * 50},
{"role": "user", "content": "First question" * 50},
{"role": "assistant", "content": "First answer" * 50},
{"role": "user", "content": "Second question"},
]
token_counts = [2000, 500, 500, 50]
tracker.update_from_response(
cache_read_tokens=2500,
cache_write_tokens=0,
messages=messages,
message_token_counts=token_counts,
)
# 2000 + 500 = 2500 <= 2500, but 2000 + 500 + 500 = 3000 > 2500
assert tracker.get_frozen_message_count() == 2
def test_cold_start_no_freeze(self, tracker):
"""If cache_read=0 and cache_write=0, don't freeze."""
messages = [{"role": "user", "content": "Hello"}]
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=0,
messages=messages,
)
assert tracker.get_frozen_message_count() == 0
def test_cache_write_freezes_next_turn(self, tracker):
"""Cache writes (new cache entries) should be frozen on the next turn."""
messages = [
{"role": "system", "content": "System" * 200},
{"role": "user", "content": "Hello"},
]
token_counts = [1500, 50]
# Turn 1: provider writes to cache (above min threshold)
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=1550,
messages=messages,
message_token_counts=token_counts,
)
# Turn 2: should freeze what was written
assert tracker.get_frozen_message_count() == 2
def test_min_cached_tokens_threshold(self):
"""Below min_cached_tokens, no freeze."""
config = PrefixFreezeConfig(min_cached_tokens=2000)
tracker = PrefixCacheTracker("anthropic", config)
messages = [{"role": "user", "content": "Hello"}]
# Turn 1: only 500 tokens cached — below threshold
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=500,
messages=messages,
message_token_counts=[500],
)
assert tracker.get_frozen_message_count() == 0
def test_disabled_config(self):
"""Disabled config always returns 0."""
config = PrefixFreezeConfig(enabled=False)
tracker = PrefixCacheTracker("anthropic", config)
messages = [{"role": "system", "content": "System" * 500}]
tracker.update_from_response(
cache_read_tokens=5000,
cache_write_tokens=0,
messages=messages,
message_token_counts=[5000],
)
assert tracker.get_frozen_message_count() == 0
def test_turn_number_increments(self, tracker):
"""Turn number should increment on each update."""
messages = [{"role": "user", "content": "Hello"}]
assert tracker._turn_number == 0
tracker.update_from_response(0, 0, messages)
assert tracker._turn_number == 1
tracker.update_from_response(0, 0, messages)
assert tracker._turn_number == 2
def test_stats_tracking(self, tracker):
"""Stats should reflect tracker state."""
stats = tracker.stats
assert isinstance(stats, FreezeStats)
assert stats.busts_avoided == 0
assert stats.tokens_preserved == 0
assert stats.turn_number == 0
def test_record_bust_avoided(self, tracker):
"""Recording bust avoided should update stats."""
tracker.record_bust_avoided(tokens_preserved=5000, compression_foregone=500)
tracker.record_bust_avoided(tokens_preserved=3000, compression_foregone=200)
stats = tracker.stats
assert stats.busts_avoided == 2
assert stats.tokens_preserved == 8000
assert stats.compression_foregone_tokens == 700
assert stats.net_benefit_tokens == 7300
def test_should_force_compress_outside_frozen(self, tracker):
"""Messages outside frozen prefix should always be compressed."""
tracker._cached_message_count = 3
assert tracker.should_force_compress(5, 1000, 200) is True
def test_should_force_compress_when_savings_exceed_discount(self, tracker):
"""For Anthropic (90% discount), compression must save >90% to be worth it."""
tracker._cached_message_count = 5
# 95% savings > 90% discount — should force compress
assert tracker.should_force_compress(2, 1000, 50) is True
# 50% savings < 90% discount — should NOT force compress
assert tracker.should_force_compress(2, 1000, 500) is False
def test_should_force_compress_openai(self, openai_tracker):
"""For OpenAI (50% discount), compression must save >50% to be worth it."""
openai_tracker._cached_message_count = 5
# 60% savings > 50% discount — should force compress
assert openai_tracker.should_force_compress(2, 1000, 400) is True
# 40% savings < 50% discount — should NOT force compress
assert openai_tracker.should_force_compress(2, 1000, 600) is False
def test_estimate_message_tokens(self):
"""Token estimation should roughly match character / 3.5."""
messages = [
{"role": "system", "content": "A" * 350}, # ~100 tokens
{"role": "user", "content": "B" * 70}, # ~20 tokens
]
counts = PrefixCacheTracker._estimate_message_tokens(messages)
assert len(counts) == 2
assert counts[0] > counts[1] # System should have more tokens
def test_estimate_content_blocks(self):
"""Token estimation should handle Anthropic content blocks."""
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "A" * 350},
{"type": "text", "text": "B" * 350},
],
},
]
counts = PrefixCacheTracker._estimate_message_tokens(messages)
assert len(counts) == 1
assert counts[0] > 100
def test_estimate_tool_result_content(self):
"""Token estimation should count tool_result content field."""
tool_content = "x" * 3500 # ~1000 tokens
messages = [
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "t1",
"content": tool_content,
}
],
},
]
counts = PrefixCacheTracker._estimate_message_tokens(messages)
assert len(counts) == 1
# Should be ~1000 tokens, definitely > 100
assert counts[0] > 100
def test_estimate_tool_use_input(self):
"""Token estimation should count tool_use input field."""
messages = [
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "t1",
"name": "Read",
"input": {"file_path": "/very/long/path/" + "x" * 700},
}
],
},
]
counts = PrefixCacheTracker._estimate_message_tokens(messages)
assert len(counts) == 1
# Should count the serialized input dict
assert counts[0] > 50
def test_estimate_tool_result_nested_blocks(self):
"""Token estimation should handle nested content blocks in tool_result."""
messages = [
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "t1",
"content": [
{"type": "text", "text": "A" * 3500},
],
}
],
},
]
counts = PrefixCacheTracker._estimate_message_tokens(messages)
assert len(counts) == 1
assert counts[0] > 100
def test_session_ttl_expiry(self):
"""Tracker should report as expired after TTL."""
config = PrefixFreezeConfig(session_ttl_seconds=1)
tracker = PrefixCacheTracker("anthropic", config)
assert tracker.is_expired is False
# Simulate time passing
tracker._last_activity = time.time() - 2
assert tracker.is_expired is True
class TestSessionTrackerStore:
"""Test SessionTrackerStore management."""
@pytest.fixture
def store(self):
return SessionTrackerStore()
def test_get_or_create_new(self, store):
"""Should create a new tracker for unknown session."""
tracker = store.get_or_create("session-1", "anthropic")
assert isinstance(tracker, PrefixCacheTracker)
assert tracker.provider == "anthropic"
def test_get_or_create_existing(self, store):
"""Should return the same tracker for the same session."""
tracker1 = store.get_or_create("session-1", "anthropic")
tracker2 = store.get_or_create("session-1", "anthropic")
assert tracker1 is tracker2
def test_different_sessions(self, store):
"""Different sessions should get different trackers."""
tracker1 = store.get_or_create("session-1", "anthropic")
tracker2 = store.get_or_create("session-2", "openai")
assert tracker1 is not tracker2
assert tracker1.provider == "anthropic"
assert tracker2.provider == "openai"
def test_active_sessions_count(self, store):
"""Should track the number of active sessions."""
assert store.active_sessions == 0
store.get_or_create("s1", "anthropic")
assert store.active_sessions == 1
store.get_or_create("s2", "openai")
assert store.active_sessions == 2
def test_cleanup_expired(self, store):
"""Should remove expired sessions on cleanup."""
config = PrefixFreezeConfig(session_ttl_seconds=1)
store = SessionTrackerStore(default_config=config)
tracker = store.get_or_create("expired-session", "anthropic")
tracker._last_activity = time.time() - 2
# Force cleanup
store._last_cleanup = 0
store._maybe_cleanup()
assert store.active_sessions == 0
def test_compute_session_id_from_header(self, store):
"""Should use x-headroom-session-id header if present."""
class MockRequest:
headers = {"x-headroom-session-id": "explicit-id-123"}
session_id = store.compute_session_id(
MockRequest(), "claude-3", [{"role": "user", "content": "Hi"}]
)
assert session_id == "explicit-id-123"
def test_compute_session_id_from_hash(self, store):
"""Should hash model + system prompt as fallback."""
class MockRequest:
headers = {}
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hi"},
]
id1 = store.compute_session_id(MockRequest(), "claude-3", messages)
id2 = store.compute_session_id(MockRequest(), "claude-3", messages)
assert id1 == id2 # Stable hash
assert len(id1) == 16
# Different model = different session
id3 = store.compute_session_id(MockRequest(), "gpt-4", messages)
assert id3 != id1
def test_compute_session_id_uses_all_system_messages(self, store):
"""Different dynamic system messages should not collide."""
class MockRequest:
headers = {}
static_prompt = "framework prompt " * 80
conv_a = [
{"role": "system", "content": [{"type": "text", "text": static_prompt}]},
{"role": "system", "content": [{"type": "text", "text": "context: session A"}]},
{"role": "user", "content": "hello"},
]
conv_b = [
{"role": "system", "content": [{"type": "text", "text": static_prompt}]},
{"role": "system", "content": [{"type": "text", "text": "context: session B"}]},
{"role": "user", "content": "hello"},
]
id_a = store.compute_session_id(MockRequest(), "claude-3", conv_a)
id_b = store.compute_session_id(MockRequest(), "claude-3", conv_b)
assert id_a != id_b
def test_compute_session_id_distinguishes_top_level_system(self, store):
"""Anthropic carries the system prompt as a top-level field (not a
role:'system' message). The handler folds it in as a synthetic system
message so two conversations with the same model and turns but different
system prompts get distinct ids — otherwise they share one tracker and
their sticky state cross-contaminates. This exercises that mechanism."""
class MockRequest:
headers = {}
turns = [{"role": "user", "content": "hello"}]
def with_system(system):
# Mirror what handlers/anthropic.py does for the top-level system.
return [{"role": "system", "content": system}, *turns]
id_a = store.compute_session_id(
MockRequest(), "claude-3", with_system("You are a Python expert.")
)
id_b = store.compute_session_id(
MockRequest(), "claude-3", with_system("You are a Rust expert.")
)
assert id_a != id_b
# A list-of-text-blocks system folds the same text as the string form.
id_a_list = store.compute_session_id(
MockRequest(),
"claude-3",
with_system([{"type": "text", "text": "You are a Python expert."}]),
)
assert id_a_list == id_a
def test_compute_session_id_is_stable_when_only_non_system_turns_change(self, store):
"""Appending non-system turns should keep the same fallback session id."""
class MockRequest:
headers = {}
base_messages = [
{"role": "system", "content": [{"type": "text", "text": "framework prompt"}]},
{"role": "system", "content": [{"type": "text", "text": "context: session A"}]},
{"role": "user", "content": "hello"},
]
extended_messages = base_messages + [{"role": "assistant", "content": "hi there"}]
id1 = store.compute_session_id(MockRequest(), "claude-3", base_messages)
id2 = store.compute_session_id(MockRequest(), "claude-3", extended_messages)
assert id1 == id2
def test_compute_session_id_no_system(self, store):
"""Should work without system messages."""
class MockRequest:
headers = {}
messages = [{"role": "user", "content": "Hi"}]
session_id = store.compute_session_id(MockRequest(), "claude-3", messages)
assert isinstance(session_id, str)
assert len(session_id) == 16
class TestMultiTurnScenario:
"""Integration-style tests simulating multi-turn conversations."""
def test_five_turn_conversation(self):
"""Simulate a 5-turn conversation with growing prefix."""
tracker = PrefixCacheTracker("anthropic")
# Turn 1: System + User (cold start, no cache)
messages_t1 = [
{"role": "system", "content": "System prompt" * 200},
{"role": "user", "content": "Question 1"},
]
token_counts_t1 = [2000, 50]
assert tracker.get_frozen_message_count() == 0 # No freeze on turn 1
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=2050,
messages=messages_t1,
message_token_counts=token_counts_t1,
)
# Turn 2: Previous messages cached, new user message added
messages_t2 = messages_t1 + [
{"role": "assistant", "content": "Answer 1"},
{"role": "user", "content": "Question 2"},
]
token_counts_t2 = [2000, 50, 200, 50]
frozen = tracker.get_frozen_message_count()
assert frozen == 2 # System + User1 frozen
tracker.update_from_response(
cache_read_tokens=2050,
cache_write_tokens=250,
messages=messages_t2,
message_token_counts=token_counts_t2,
)
# Turn 3: Even more cached
messages_t3 = messages_t2 + [
{"role": "assistant", "content": "Answer 2"},
{"role": "user", "content": "Question 3"},
]
token_counts_t3 = [2000, 50, 200, 50, 200, 50]
frozen = tracker.get_frozen_message_count()
assert frozen == 4 # System + User1 + Asst1 + User2 frozen
tracker.update_from_response(
cache_read_tokens=2300,
cache_write_tokens=250,
messages=messages_t3,
message_token_counts=token_counts_t3,
)
# Verify turn count
assert tracker._turn_number == 3
def test_cache_bust_resets_freeze(self):
"""If cache is busted (0 read, 0 write), freeze should reset."""
tracker = PrefixCacheTracker("anthropic")
messages = [
{"role": "system", "content": "System" * 200},
{"role": "user", "content": "Hello"},
]
# Turn 1: Cache established
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=2000,
messages=messages,
message_token_counts=[1500, 500],
)
assert tracker.get_frozen_message_count() == 2 # Both fit within 2000
# Turn 2: Cache bust (0 reads, system prompt changed)
tracker.update_from_response(
cache_read_tokens=0,
cache_write_tokens=0,
messages=messages,
message_token_counts=[1500, 500],
)
# After a bust with 0 total, freeze should reset
assert tracker.get_frozen_message_count() == 0
class TestClassifyCacheMiss:
"""Cache-miss attribution (#1313): TTL lapse vs prefix change vs unknown."""
BASE = [
{"role": "system", "content": "x" * 4000},
{"role": "user", "content": "hello"},
]
CHANGED = [
{"role": "system", "content": "DIFFERENT" * 400},
{"role": "user", "content": "hello"},
]
def _warm(self, tracker, messages, read=500, write=500):
"""Simulate a turn that left `messages` cached."""
tracker.update_from_response(
cache_read_tokens=read, cache_write_tokens=write, messages=messages
)
def test_cold_start_is_not_a_miss(self):
"""No prior cached prefix → cold start, is_miss False."""
tracker = PrefixCacheTracker("anthropic")
result = tracker.classify_cache_miss(0, self.BASE)
assert result.is_miss is False
assert result.reason == MISS_COLD_START
def test_cache_read_is_a_hit(self):
"""A non-zero read on an expected-cached prefix is a hit, not a miss."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
result = tracker.classify_cache_miss(800, self.BASE)
assert result.is_miss is False
assert result.reason == "hit"
def test_ttl_expiry_when_idle_exceeds_ttl(self):
"""Idle longer than the cache TTL → ttl_expiry."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
result = tracker.classify_cache_miss(0, self.BASE, idle_seconds=400)
assert result.is_miss is True
assert result.reason == MISS_TTL_EXPIRY
assert result.ttl_exceeded is True
assert result.cache_ttl_seconds == 300
def test_ttl_wins_tie_when_prefix_also_changed(self):
"""When idle past TTL AND prefix changed, TTL expiry wins (docstring)."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
result = tracker.classify_cache_miss(0, self.CHANGED, idle_seconds=400)
assert result.reason == MISS_TTL_EXPIRY
assert result.ttl_exceeded is True
assert result.prefix_changed is True
def test_prefix_change_within_ttl(self):
"""Within TTL but the forwarded prefix differs → prefix_change."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
result = tracker.classify_cache_miss(0, self.CHANGED, idle_seconds=10)
assert result.is_miss is True
assert result.reason == MISS_PREFIX_CHANGE
assert result.prefix_changed is True
assert result.ttl_exceeded is False
def test_unknown_when_stable_prefix_within_ttl(self):
"""Within TTL, prefix unchanged, but still no read → unknown."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
result = tracker.classify_cache_miss(0, self.BASE, idle_seconds=10)
assert result.is_miss is True
assert result.reason == MISS_UNKNOWN
def test_growing_prefix_is_stable(self):
"""A turn that appends to last turn's forwarded prefix is not a change."""
tracker = PrefixCacheTracker("anthropic")
self._warm(tracker, self.BASE)
grown = self.BASE + [{"role": "assistant", "content": "hi back"}]
result = tracker.classify_cache_miss(0, grown, idle_seconds=10)
# Prefix preserved (only appended) → not a prefix_change.
assert result.prefix_changed is False
assert result.reason == MISS_UNKNOWN
def test_one_hour_ttl_override(self):
"""cache_ttl_seconds override widens the TTL window (1h breakpoint)."""
tracker = PrefixCacheTracker("anthropic", PrefixFreezeConfig(cache_ttl_seconds=3600))
self._warm(tracker, self.BASE)
# 400s idle is past the 300s default but within 3600s → not TTL expiry.
result = tracker.classify_cache_miss(0, self.BASE, idle_seconds=400)
assert result.cache_ttl_seconds == 3600
assert result.ttl_exceeded is False
assert result.reason == MISS_UNKNOWN
def test_resolved_ttl_falls_back_to_provider_default(self):
assert PrefixCacheTracker("anthropic").resolved_cache_ttl_seconds() == 300
assert (
PrefixCacheTracker(
"anthropic", PrefixFreezeConfig(cache_ttl_seconds=3600)
).resolved_cache_ttl_seconds()
== 3600
)
+136
View File
@@ -0,0 +1,136 @@
"""Tests for CacheOptimizerRegistry."""
import pytest
from headroom.cache import (
AnthropicCacheOptimizer,
CacheConfig,
CacheOptimizerRegistry,
GoogleCacheOptimizer,
OpenAICacheOptimizer,
)
from headroom.cache.base import BaseCacheOptimizer, CacheResult, CacheStrategy
class MockOptimizer(BaseCacheOptimizer):
"""Mock optimizer for testing."""
@property
def name(self) -> str:
return "mock-optimizer"
@property
def provider(self) -> str:
return "mock"
@property
def strategy(self) -> CacheStrategy:
return CacheStrategy.NONE
def optimize(self, messages, context, config=None):
return CacheResult(messages=messages)
class TestCacheOptimizerRegistry:
"""Test CacheOptimizerRegistry functionality."""
def test_default_providers_registered(self):
"""Test that default providers are registered on import."""
providers = CacheOptimizerRegistry.list_all()
assert "anthropic" in providers
assert "openai" in providers
assert "google" in providers
def test_get_anthropic(self):
"""Test getting Anthropic optimizer."""
optimizer = CacheOptimizerRegistry.get("anthropic")
assert isinstance(optimizer, AnthropicCacheOptimizer)
assert optimizer.provider == "anthropic"
assert optimizer.strategy == CacheStrategy.EXPLICIT_BREAKPOINTS
def test_get_openai(self):
"""Test getting OpenAI optimizer."""
optimizer = CacheOptimizerRegistry.get("openai")
assert isinstance(optimizer, OpenAICacheOptimizer)
assert optimizer.provider == "openai"
assert optimizer.strategy == CacheStrategy.PREFIX_STABILIZATION
def test_get_google(self):
"""Test getting Google optimizer."""
optimizer = CacheOptimizerRegistry.get("google")
assert isinstance(optimizer, GoogleCacheOptimizer)
assert optimizer.provider == "google"
assert optimizer.strategy == CacheStrategy.CACHED_CONTENT
def test_get_with_config(self):
"""Test getting optimizer with custom config."""
config = CacheConfig(min_cacheable_tokens=2048)
optimizer = CacheOptimizerRegistry.get("anthropic", config=config, cached=False)
assert optimizer.config.min_cacheable_tokens >= 1024 # Anthropic enforces minimum
def test_register_custom_optimizer(self):
"""Test registering a custom optimizer."""
CacheOptimizerRegistry.register("mock", MockOptimizer)
try:
optimizer = CacheOptimizerRegistry.get("mock")
assert isinstance(optimizer, MockOptimizer)
finally:
CacheOptimizerRegistry.unregister("mock")
def test_register_duplicate_raises(self):
"""Test that registering duplicate without override raises."""
CacheOptimizerRegistry.register("test-dup", MockOptimizer)
try:
with pytest.raises(ValueError):
CacheOptimizerRegistry.register("test-dup", MockOptimizer)
finally:
CacheOptimizerRegistry.unregister("test-dup")
def test_register_with_override(self):
"""Test registering with override."""
CacheOptimizerRegistry.register("test-override", MockOptimizer)
try:
CacheOptimizerRegistry.register("test-override", MockOptimizer, override=True)
optimizer = CacheOptimizerRegistry.get("test-override")
assert isinstance(optimizer, MockOptimizer)
finally:
CacheOptimizerRegistry.unregister("test-override")
def test_get_unknown_provider_raises(self):
"""Test getting unknown provider raises KeyError."""
with pytest.raises(KeyError):
CacheOptimizerRegistry.get("unknown-provider")
def test_list_providers(self):
"""Test listing providers."""
providers = CacheOptimizerRegistry.list_providers()
assert "anthropic" in providers
assert "openai" in providers
assert "google" in providers
def test_is_registered(self):
"""Test is_registered check."""
assert CacheOptimizerRegistry.is_registered("anthropic")
assert not CacheOptimizerRegistry.is_registered("nonexistent")
def test_cached_instances(self):
"""Test that cached instances are reused."""
opt1 = CacheOptimizerRegistry.get("anthropic", cached=True)
opt2 = CacheOptimizerRegistry.get("anthropic", cached=True)
assert opt1 is opt2
def test_uncached_instances(self):
"""Test that uncached instances are not reused."""
opt1 = CacheOptimizerRegistry.get("anthropic", cached=False)
opt2 = CacheOptimizerRegistry.get("anthropic", cached=False)
assert opt1 is not opt2
def test_tier_based_selection(self):
"""Test tier-based optimizer selection."""
# OSS tier should work
oss_opt = CacheOptimizerRegistry.get("anthropic", tier="oss")
assert oss_opt is not None
# Enterprise tier falls back to OSS if not registered
ent_opt = CacheOptimizerRegistry.get("anthropic", tier="enterprise")
assert ent_opt is not None
+320
View File
@@ -0,0 +1,320 @@
"""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