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
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:
@@ -0,0 +1 @@
|
||||
"""Tests for the cache optimization module."""
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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 != ""
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user