diff --git a/CHANGELOG.md b/CHANGELOG.md index 9a2db3b3f..ac1af9e49 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,8 +32,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 * **proxy:** force Responses API `store=true` when Headroom injects memory tools so `previous_response_id` continuations work after memory tool calls from clients that requested `store=false` ([#1103](https://github.com/chopratejas/headroom/pull/1103)). * **proxy:** build SSL contexts for custom CA bundles so enterprise/private PKI roots work with Python/OpenSSL strict verification. -* **tokenizers:** bound token-counting of oversized tool-content blobs instead of running `count_text` over the whole serialized string. `count_messages` runs on the proxy request path; serializing is cheap, but `count_text` over a multi-megabyte `tool_result` / `tool_use` string took seconds and could freeze `/health` and in-flight requests. For payloads over ~50KB serialized, `count_text` now runs on an even-spread sample of the string and scales by length — model-accurate (tracks the active tokenizer), bounded for any blob shape, and biased to under-count (the safe direction). Smaller payloads stay exact. +* **tokenizers:** bound token-counting of oversized tool-content blobs instead of running `count_text` over the whole serialized string. `count_messages` runs on the proxy request path; serializing is cheap, but `count_text` over a multi-megabyte `tool_result` / `tool_use` string took seconds and could freeze `/health` and in-flight requests. For payloads over ~50KB serialized, `count_text` now runs on an even-spread sample of the string and scales by length; it stays model-accurate, bounded for any blob shape, and biased to under-count. Smaller payloads stay exact. * **codex:** stop persisting a project-specific `--db` path in the global `headroom_memory` MCP config, so `headroom wrap codex --memory` falls back to the active cwd's `.headroom/memory.db` at runtime while keeping the current project's local bootstrap work scoped correctly ([#1147](https://github.com/chopratejas/headroom/issues/1147)). +* **ccr:** stop emitting Anthropic request-side retrieval markers on frozen-prefix turns when `headroom_retrieve` injection is deferred, so cache-preserving requests forward original content instead of irrecoverable marker-only payloads ([#1006](https://github.com/chopratejas/headroom/issues/1006)). * **proxy:** route Codex OAuth image generation and edit requests through the ChatGPT Codex image backend, while preserving OpenAI API-key image passthrough ([#1215](https://github.com/chopratejas/headroom/pull/1215)). * **wrap (codex):** keep RTK guidance in the global Codex `AGENTS.md` instead of modifying the shared project `AGENTS.md` ([#1235](https://github.com/chopratejas/headroom/issues/1235)). * **proxy:** enable SSO credential resolution in the native Bedrock route via the `aws-config` `sso` feature flag, making the credential chain match what `docs/bedrock.md` already documented ([#999](https://github.com/chopratejas/headroom/pull/999)). diff --git a/headroom/proxy/handlers/anthropic.py b/headroom/proxy/handlers/anthropic.py index a5761b421..fd9297162 100644 --- a/headroom/proxy/handlers/anthropic.py +++ b/headroom/proxy/handlers/anthropic.py @@ -1029,13 +1029,27 @@ class AnthropicHandlerMixin: from headroom.transforms.compression_policy import resolve_policy compression_policy = resolve_policy(getattr(request.state, "auth_mode", None)) + from headroom.ccr.tool_injection import CCR_TOOL_NAME + + existing_tool_names = { + tool.get("name") or tool.get("function", {}).get("name") + for tool in (body.get("tools") or []) + if isinstance(tool, dict) + } + + def should_skip_ccr_request_compression( + current_frozen_message_count: int, + ) -> bool: + # If the tool is already present, CCR stays reversible even on frozen turns. + return ( + self.config.ccr_inject_tool + and current_frozen_message_count > 0 + and CCR_TOOL_NAME not in existing_tool_names + ) if is_token_mode(self.config.mode): comp_cache = self._get_compression_cache(session_id) - # Zone 1: Swap cached compressed versions into working copy - working_messages = comp_cache.apply_cached(messages) - # Re-freeze boundary: consecutive stable messages from start. # Safety: never freeze beyond provider-confirmed cached prefix. # `prefix_tracker.frozen_message_count` (set above) is the @@ -1064,64 +1078,27 @@ class AnthropicHandlerMixin: # Record all tool_results in the verified frozen prefix as stable comp_cache.mark_stable_from_messages(messages, frozen_message_count) - # Phase 3 (#1171): off-path deferral gate. On a cold- - # start-large request (frozen=0 + large live zone) the - # synchronous kompress run would blow the 30s budget and - # leak a non-preemptible worker. Forward the cache- - # swapped messages uncompressed NOW and compress off the - # request path; the result lands in the SAME - # CompressionCache and apply_cached swaps it in (live - # zone) on a later turn. Byte-identity holds — see the - # frozen/live invariant; one-time upstream cache miss - # when the compressed form first lands, then stable. - if ( - getattr(self, "_background_compression_enabled", False) - and frozen_message_count == 0 - and original_tokens >= self._background_compression_min_tokens - ): - # Snapshot refs for the async job. The handler must - # NOT mutate these lists/dicts in-place after this - # point -- the background job reads them on a later - # turn. It doesn't today; keep it that way. - _bg_messages = messages - _bg_working = working_messages - _bg_frozen = frozen_message_count - # Dedup key: the gate only fires at frozen==0 (the - # first in-flight deferral of a session episode), so - # session_id alone is the right granularity -- one - # background job per session in flight -- and avoids - # JSON-serializing the large message list for a key. - accepted = self._background_compressor.enqueue( - session_id, - lambda: self.anthropic_pipeline.apply( - messages=_bg_working, - model=model, - model_limit=context_limit, - context=extract_user_query(_bg_working), - frozen_message_count=_bg_frozen, - biases=biases, - request_id=request_id, - compression_policy=compression_policy, - **proxy_pipeline_kwargs(self.config), - ), - lambda result: comp_cache.update_from_result( - _bg_messages, result.messages - ), + skip_ccr_request_compression = should_skip_ccr_request_compression( + frozen_message_count + ) + if skip_ccr_request_compression: + logger.info( + f"[{request_id}] CCR: skipping request-side compression " + f"(frozen prefix={frozen_message_count}) because tool injection is deferred" ) - # Forward uncompressed either way (the request can't - # wait); only CLAIM deferral when the job was actually - # queued. A full-queue drop is visible in telemetry as - # "deferred:dropped" and self-heals on a later turn. - optimized_messages = working_messages - transforms_applied = [ - "deferred:background_compression" - if accepted - else "deferred:dropped" - ] - pipeline_timing = {} + if skip_ccr_request_compression: + optimized_messages = messages + optimized_tokens = tokenizer.count_messages(optimized_messages) else: - async with stage_timer.measure("compression_first_stage"): - result = await self._run_compression_in_executor( + # Zone 1: Swap cached compressed versions into working copy + working_messages = comp_cache.apply_cached(messages) + if ( + getattr(self, "_background_compression_enabled", False) + and frozen_message_count == 0 + and original_tokens >= self._background_compression_min_tokens + ): + accepted = self._background_compressor.enqueue( + session_id, lambda: self.anthropic_pipeline.apply( messages=working_messages, model=model, @@ -1133,57 +1110,103 @@ class AnthropicHandlerMixin: compression_policy=compression_policy, **proxy_pipeline_kwargs(self.config), ), - timeout=COMPRESSION_TIMEOUT_SECONDS, + lambda bg_result: comp_cache.update_from_result( + messages, bg_result.messages + ), ) + class _DeferredCompressionResult: + messages = working_messages + transforms_applied = [ + "deferred:background_compression" + if accepted + else "deferred:dropped" + ] + timing = {} + + result = _DeferredCompressionResult() + else: + async with stage_timer.measure("compression_first_stage"): + result = await self._run_compression_in_executor( + lambda: self.anthropic_pipeline.apply( + messages=working_messages, + model=model, + model_limit=context_limit, + context=extract_user_query(working_messages), + frozen_message_count=frozen_message_count, + biases=biases, + request_id=request_id, + compression_policy=compression_policy, + **proxy_pipeline_kwargs(self.config), + ), + timeout=COMPRESSION_TIMEOUT_SECONDS, + ) + # Cache newly compressed messages (index-aligned diff) if result.messages != working_messages: comp_cache.update_from_result(messages, result.messages) - # Always use pipeline result — Zone 1 swaps applied + # Always use pipeline result — Zone 1 swaps are already applied optimized_messages = result.messages transforms_applied = result.transforms_applied pipeline_timing = result.timing - # Issue #327 / Bug 3: pipeline.apply uses the provider- - # side tokenizer (AnthropicProvider tiktoken estimator), - # which counts ~25% higher than the proxy-side - # EstimatingTokenCounter used to set `original_tokens` - # at line 634. Reusing `result.tokens_after` here - # produced an apples-vs-oranges comparison against - # `original_tokens` in the inflation guard below - # (line ~901): even after a real 12% compression the - # provider-tokenizer figure was higher than the proxy- - # tokenizer baseline, triggering a spurious revert. - # Recount optimized_messages with the proxy tokenizer - # so original_tokens vs optimized_tokens is self- - # consistent. The recount cost (~ms on a 50K-token - # request) is paid once per request and is dwarfed by - # the upstream call latency. - optimized_tokens = tokenizer.count_messages(optimized_messages) + # Issue #327 / Bug 3: pipeline.apply uses the provider- + # side tokenizer (AnthropicProvider tiktoken estimator), + # which counts ~25% higher than the proxy-side + # EstimatingTokenCounter used to set `original_tokens` + # at line 634. Reusing `result.tokens_after` here + # produced an apples-vs-oranges comparison against + # `original_tokens` in the inflation guard below + # (line ~901): even after a real 12% compression the + # provider-tokenizer figure was higher than the proxy- + # tokenizer baseline, triggering a spurious revert. + # Recount optimized_messages with the proxy tokenizer + # so original_tokens vs optimized_tokens is self- + # consistent. The recount cost (~ms on a 50K-token + # request) is paid once per request and is dwarfed by + # the upstream call latency. + optimized_tokens = tokenizer.count_messages(optimized_messages) elif not is_cache_mode(self.config.mode): - async with stage_timer.measure("compression_first_stage"): - result = await self._run_compression_in_executor( - lambda: self.anthropic_pipeline.apply( - messages=messages, - model=model, - model_limit=context_limit, - context=extract_user_query(messages), - frozen_message_count=frozen_message_count, - biases=biases, - request_id=request_id, - compression_policy=compression_policy, - **proxy_pipeline_kwargs(self.config), - ), - timeout=COMPRESSION_TIMEOUT_SECONDS, + skip_ccr_request_compression = should_skip_ccr_request_compression( + frozen_message_count + ) + if skip_ccr_request_compression: + logger.info( + f"[{request_id}] CCR: skipping request-side compression " + f"(frozen prefix={frozen_message_count}) because tool injection is deferred" ) + if not skip_ccr_request_compression: + async with stage_timer.measure("compression_first_stage"): + result = await self._run_compression_in_executor( + lambda: self.anthropic_pipeline.apply( + messages=messages, + model=model, + model_limit=context_limit, + context=extract_user_query(messages), + frozen_message_count=frozen_message_count, + biases=biases, + request_id=request_id, + compression_policy=compression_policy, + **proxy_pipeline_kwargs(self.config), + ), + timeout=COMPRESSION_TIMEOUT_SECONDS, + ) - if result.messages != messages: - optimized_messages = result.messages - transforms_applied = result.transforms_applied - pipeline_timing = result.timing - original_tokens = result.tokens_before - optimized_tokens = result.tokens_after + if result.messages != messages: + optimized_messages = result.messages + transforms_applied = result.transforms_applied + pipeline_timing = result.timing + original_tokens = result.tokens_before + optimized_tokens = result.tokens_after else: + skip_ccr_request_compression = should_skip_ccr_request_compression( + frozen_message_count + ) + if skip_ccr_request_compression: + logger.info( + f"[{request_id}] CCR: skipping request-side compression " + f"(frozen prefix={frozen_message_count}) because tool injection is deferred" + ) previous_original_messages = prefix_tracker.get_last_original_messages() previous_forwarded_messages = prefix_tracker.get_last_forwarded_messages() delta = self._extract_cache_stable_delta( @@ -1194,27 +1217,35 @@ class AnthropicHandlerMixin: if delta is not None: stable_forwarded_prefix, delta_messages = delta if delta_messages: - result = await self._run_compression_in_executor( - lambda: self.anthropic_pipeline.apply( - messages=delta_messages, - model=model, - model_limit=context_limit, - context=extract_user_query(delta_messages), - frozen_message_count=0, - biases=biases, - request_id=request_id, - compression_policy=compression_policy, - **proxy_pipeline_kwargs(self.config), - ), - timeout=COMPRESSION_TIMEOUT_SECONDS, - ) - optimized_messages = stable_forwarded_prefix + result.messages - transforms_applied = result.transforms_applied - pipeline_timing = result.timing - optimized_tokens = tokenizer.count_messages(optimized_messages) + if skip_ccr_request_compression: + optimized_messages = messages + optimized_tokens = tokenizer.count_messages(optimized_messages) + else: + result = await self._run_compression_in_executor( + lambda: self.anthropic_pipeline.apply( + messages=delta_messages, + model=model, + model_limit=context_limit, + context=extract_user_query(delta_messages), + frozen_message_count=0, + biases=biases, + request_id=request_id, + compression_policy=compression_policy, + **proxy_pipeline_kwargs(self.config), + ), + timeout=COMPRESSION_TIMEOUT_SECONDS, + ) + optimized_messages = stable_forwarded_prefix + result.messages + transforms_applied = result.transforms_applied + pipeline_timing = result.timing + optimized_tokens = tokenizer.count_messages(optimized_messages) else: - optimized_messages = stable_forwarded_prefix - optimized_tokens = tokenizer.count_messages(optimized_messages) + if skip_ccr_request_compression: + optimized_messages = messages + optimized_tokens = tokenizer.count_messages(optimized_messages) + else: + optimized_messages = stable_forwarded_prefix + optimized_tokens = tokenizer.count_messages(optimized_messages) else: # Conservative rule for cache mode: # only replay exact stable message-prefix extensions. diff --git a/tests/_skip_helpers.py b/tests/_skip_helpers.py new file mode 100644 index 000000000..ce8fb9715 --- /dev/null +++ b/tests/_skip_helpers.py @@ -0,0 +1,45 @@ +"""Helpers for skipping tests when external model dependencies are unavailable.""" + +from __future__ import annotations + +from collections.abc import Iterator + + +def _iter_exception_chain(exc: BaseException) -> Iterator[BaseException]: + """Yield an exception and its direct cause/context chain.""" + seen: set[int] = set() + current: BaseException | None = exc + while current is not None and id(current) not in seen: + yield current + seen.add(id(current)) + current = current.__cause__ or current.__context__ + + +def external_model_skip_reason(exc: BaseException) -> str | None: + """Return a pytest skip reason for transient or offline model dependency errors.""" + try: + import httpx + except ImportError: # pragma: no cover + httpx = None + + try: + from huggingface_hub.errors import LocalEntryNotFoundError + except ImportError: # pragma: no cover + LocalEntryNotFoundError = None + + for candidate in _iter_exception_chain(exc): + if httpx is not None and isinstance(candidate, httpx.ReadTimeout): + return "Skipped due to network timeout (flaky CI)" + + if LocalEntryNotFoundError is not None and isinstance(candidate, LocalEntryNotFoundError): + return "Skipped because required Hugging Face model files are unavailable offline" + + if isinstance(candidate, OSError): + message = str(candidate) + if ( + "couldn't connect to 'https://huggingface.co'" in message + and "couldn't find them in the cached files" in message + ): + return "Skipped because required Hugging Face model files are unavailable offline" + + return None diff --git a/tests/conftest.py b/tests/conftest.py index 892c20c27..62519e4c5 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,14 +14,7 @@ from unittest.mock import Mock import pytest -# Import httpx for timeout handling (will be available since it's a dependency) -try: - import httpx - - HTTPX_AVAILABLE = True -except ImportError: - HTTPX_AVAILABLE = False - +from tests._skip_helpers import external_model_skip_reason # ============================================================================= # Global test hooks @@ -30,19 +23,21 @@ except ImportError: @pytest.hookimpl(hookwrapper=True) def pytest_runtest_call(item): - """Wrap test execution to catch httpx.ReadTimeout and skip instead of fail. + """Wrap test execution to skip transient or offline external model failures. - This handles flaky network timeouts that occur when: + This handles model-loading failures that occur when: - HuggingFace Hub is slow during model downloads (sentence-transformers) + - Required HuggingFace model files were not restored into the offline CI cache - External embedding APIs timeout - Network connectivity issues in CI """ outcome = yield - if HTTPX_AVAILABLE and outcome.excinfo is not None: + if outcome.excinfo is not None: exc_type, exc_value, exc_tb = outcome.excinfo - if isinstance(exc_value, httpx.ReadTimeout): - pytest.skip("Skipped due to network timeout (flaky CI)") + reason = external_model_skip_reason(exc_value) + if reason is not None: + pytest.skip(reason) @pytest.fixture(autouse=True) diff --git a/tests/test_memory/conftest.py b/tests/test_memory/conftest.py index 3dfe47bda..01ba184e7 100644 --- a/tests/test_memory/conftest.py +++ b/tests/test_memory/conftest.py @@ -9,27 +9,23 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false" import pytest -# Import httpx for type checking (will be available since it's a dependency) -try: - import httpx - - HTTPX_AVAILABLE = True -except ImportError: - HTTPX_AVAILABLE = False +from tests._skip_helpers import external_model_skip_reason @pytest.hookimpl(hookwrapper=True) def pytest_runtest_call(item): - """Wrap test execution to catch httpx.ReadTimeout and skip instead of fail. + """Wrap test execution to skip transient or offline external model failures. - This handles flaky network timeouts that occur when: + This handles model-loading failures that occur when: - HuggingFace Hub is slow during model downloads (sentence-transformers) + - Required HuggingFace model files were not restored into the offline CI cache - External embedding APIs timeout - Network connectivity issues in CI """ outcome = yield - if HTTPX_AVAILABLE and outcome.excinfo is not None: + if outcome.excinfo is not None: exc_type, exc_value, exc_tb = outcome.excinfo - if isinstance(exc_value, httpx.ReadTimeout): - pytest.skip("Skipped due to network timeout (flaky CI)") + reason = external_model_skip_reason(exc_value) + if reason is not None: + pytest.skip(reason) diff --git a/tests/test_memory/test_skip_helpers.py b/tests/test_memory/test_skip_helpers.py new file mode 100644 index 000000000..1e5a8ba01 --- /dev/null +++ b/tests/test_memory/test_skip_helpers.py @@ -0,0 +1,29 @@ +from __future__ import annotations + +import httpx +from huggingface_hub.errors import LocalEntryNotFoundError + +from tests._skip_helpers import external_model_skip_reason + + +def test_external_model_skip_reason_handles_httpx_timeout() -> None: + reason = external_model_skip_reason(httpx.ReadTimeout("slow network")) + assert reason == "Skipped due to network timeout (flaky CI)" + + +def test_external_model_skip_reason_handles_hf_local_cache_miss() -> None: + reason = external_model_skip_reason(LocalEntryNotFoundError("not cached")) + assert reason == "Skipped because required Hugging Face model files are unavailable offline" + + +def test_external_model_skip_reason_handles_transformers_offline_oserror() -> None: + exc = OSError( + "We couldn't connect to 'https://huggingface.co' to load the files, " + "and couldn't find them in the cached files." + ) + reason = external_model_skip_reason(exc) + assert reason == "Skipped because required Hugging Face model files are unavailable offline" + + +def test_external_model_skip_reason_ignores_unrelated_errors() -> None: + assert external_model_skip_reason(RuntimeError("boom")) is None diff --git a/tests/test_memory_bridge.py b/tests/test_memory_bridge.py index 3d698eccb..789e1ced0 100644 --- a/tests/test_memory_bridge.py +++ b/tests/test_memory_bridge.py @@ -8,7 +8,9 @@ Run with: pytest tests/test_memory_bridge.py -v from __future__ import annotations +import functools import json +import os import uuid import pytest @@ -24,6 +26,7 @@ from headroom.memory.bridge_parsers import ( parse_generic_markdown, parse_markdown, ) +from tests._skip_helpers import external_model_skip_reason # Sample content for testing CLAUDE_CODE_MEMORY = """\ @@ -63,6 +66,40 @@ The system uses FastAPI for the proxy layer. """ +def skip_offline_model_failures(func): + """Skip bridge integration tests when the local embedder cannot start offline.""" + + @functools.wraps(func) + async def wrapper(*args, **kwargs): + try: + return await func(*args, **kwargs) + except Exception as exc: + reason = external_model_skip_reason(exc) + if reason is not None: + pytest.skip(reason) + raise + + return wrapper + + +def decorate_async_test_methods(cls): + """Wrap every async test method on a class with the offline-model skip helper.""" + for name, value in vars(cls).items(): + if name.startswith("test_"): + setattr(cls, name, skip_offline_model_failures(value)) + return cls + + +def skip_if_offline_bridge_import_failed(stats) -> None: + """Skip bridge assertions when every section failed only because the model cache is offline.""" + if ( + os.environ.get("TRANSFORMERS_OFFLINE") == "1" + and stats.sections_imported == 0 + and stats.sections_failed > 0 + ): + pytest.skip("Skipped because required Hugging Face model files are unavailable offline") + + # ============================================================================= # Parser Tests (pure functions, no backend) # ============================================================================= @@ -273,6 +310,7 @@ def bridge(bridge_config, backend): return MemoryBridge(bridge_config, backend) +@decorate_async_test_methods class TestMemoryBridgeImport: @pytest.mark.asyncio async def test_import_claude_code_memory(self, bridge, tmp_dir, backend): @@ -281,6 +319,7 @@ class TestMemoryBridgeImport: md_path.write_text(CLAUDE_CODE_MEMORY, encoding="utf-8") stats = await bridge.import_from_markdown(paths=[md_path], user_id="test_user") + skip_if_offline_bridge_import_failed(stats) assert stats.files_processed == 1 assert stats.sections_imported > 0 @@ -297,6 +336,7 @@ class TestMemoryBridgeImport: md_path.write_text(CLAUDE_CODE_MEMORY, encoding="utf-8") stats1 = await bridge.import_from_markdown(paths=[md_path], user_id="test_user") + skip_if_offline_bridge_import_failed(stats1) assert stats1.sections_imported > 0 stats2 = await bridge.import_from_markdown(paths=[md_path], user_id="test_user") @@ -316,6 +356,7 @@ class TestMemoryBridgeImport: md_path.write_text(modified, encoding="utf-8") stats = await bridge.import_from_markdown(paths=[md_path], user_id="test_user") + skip_if_offline_bridge_import_failed(stats) assert stats.files_processed == 1 assert stats.sections_imported >= 1 # At least the new section @@ -339,6 +380,7 @@ class TestMemoryBridgeImport: bridge._config.md_format = MarkdownFormat.CHATGPT stats = await bridge.import_from_markdown(paths=[md_path], user_id="test_user") + skip_if_offline_bridge_import_failed(stats) assert stats.sections_imported > 0 @pytest.mark.asyncio @@ -366,6 +408,7 @@ class TestMemoryBridgeImport: assert "source_file" in metadata +@decorate_async_test_methods class TestMemoryBridgeExport: @pytest.mark.asyncio async def test_export_claude_code_style(self, bridge, tmp_dir, backend): @@ -422,6 +465,7 @@ class TestMemoryBridgeExport: assert "No memories" in markdown +@decorate_async_test_methods class TestMemoryBridgeSync: @pytest.mark.asyncio async def test_sync_imports_and_exports(self, bridge, tmp_dir, backend): @@ -432,6 +476,7 @@ class TestMemoryBridgeSync: # First sync: imports from file stats = await bridge.sync(user_id="test_user") + skip_if_offline_bridge_import_failed(stats.import_stats) assert stats.import_stats.sections_imported > 0 # Add an organic memory (not from bridge) @@ -465,6 +510,7 @@ class TestMemoryBridgeSync: assert stats.memories_exported == 0 +@decorate_async_test_methods class TestSyncStatePersistence: @pytest.mark.asyncio async def test_state_saved_and_loaded(self, tmp_dir, backend): @@ -496,6 +542,7 @@ class TestSyncStatePersistence: assert stats.files_skipped_unchanged == 1 +@decorate_async_test_methods class TestRoundTrip: @pytest.mark.asyncio async def test_import_export_preserves_facts(self, bridge, tmp_dir, backend): diff --git a/tests/test_memory_handler_concurrent_init.py b/tests/test_memory_handler_concurrent_init.py index d26279c40..418f6b83e 100644 --- a/tests/test_memory_handler_concurrent_init.py +++ b/tests/test_memory_handler_concurrent_init.py @@ -12,6 +12,8 @@ Covers: from __future__ import annotations import asyncio +import functools +import os from typing import Any from unittest.mock import patch @@ -22,6 +24,24 @@ from headroom.proxy.memory_handler import ( MemoryConfig, MemoryHandler, ) +from tests._skip_helpers import external_model_skip_reason + + +def skip_offline_model_failures(func): + """Skip real-backend smoke tests when the local embedder cannot start offline.""" + + @functools.wraps(func) + async def wrapper(*args, **kwargs): + try: + return await func(*args, **kwargs) + except Exception as exc: + reason = external_model_skip_reason(exc) + if reason is not None: + pytest.skip(reason) + raise + + return wrapper + # ------------------------------------------------------------------- # Singleflight under concurrent callers @@ -224,6 +244,7 @@ async def test_ensure_initialized_cancellation_propagates_and_resets_state(tmp_p @pytest.mark.asyncio +@skip_offline_model_failures async def test_real_localbackend_initializes_via_public_entrypoint(tmp_path): """End-to-end sanity check: the public ``ensure_initialized`` path works against a real LocalBackend. This catches regressions where the new @@ -253,6 +274,8 @@ async def test_real_localbackend_initializes_via_public_entrypoint(tmp_path): # warmup_embedder is best-effort; on a real backend it should succeed. warmed = await handler.warmup_embedder() + if not warmed and os.environ.get("TRANSFORMERS_OFFLINE") == "1": + pytest.skip("Skipped because required Hugging Face model files are unavailable offline") assert warmed is True await handler.close() diff --git a/tests/test_proxy/test_anthropic_ccr_deferred_injection.py b/tests/test_proxy/test_anthropic_ccr_deferred_injection.py new file mode 100644 index 000000000..ee9adea26 --- /dev/null +++ b/tests/test_proxy/test_anthropic_ccr_deferred_injection.py @@ -0,0 +1,1256 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import httpx +import pytest + +pytest.importorskip("fastapi") + +from fastapi.testclient import TestClient + +from headroom.proxy.server import ProxyConfig, create_app + +_RAW_TRANSCRIPT = "\n".join(f"row {idx}: payload payload payload" for idx in range(80)) + + +class _FakePrefixTracker: + def __init__(self, frozen_count: int): + self._frozen_count = frozen_count + self._cached_token_count = 0 + self._last_original_messages: list[dict] = [] + self._last_forwarded_messages: list[dict] = [] + + def get_frozen_message_count(self) -> int: + return self._frozen_count + + def get_last_original_messages(self): # noqa: ANN201 + return self._last_original_messages.copy() + + def get_last_forwarded_messages(self): # noqa: ANN201 + return self._last_forwarded_messages.copy() + + def update_from_response(self, **kwargs): # noqa: ANN003 + self._cached_token_count = kwargs.get("cache_read_tokens", 0) + kwargs.get( + "cache_write_tokens", 0 + ) + self._last_original_messages = kwargs.get( + "original_messages", kwargs.get("messages", []) + ).copy() + self._last_forwarded_messages = kwargs.get("messages", []).copy() + return None + + +class _FakeCompressionCache: + def __init__( + self, + frozen_count: int, + cached_messages: list[dict] | None = None, + ): + self._frozen_count = frozen_count + self._cached_messages = cached_messages + + def apply_cached(self, messages): # noqa: ANN201 + if self._cached_messages is not None: + return self._cached_messages + return messages + + def compute_frozen_count(self, messages) -> int: # noqa: ARG002 + return self._frozen_count + + def mark_stable_from_messages(self, messages, frozen_count) -> None: # noqa: ARG002 + return None + + def update_from_result(self, originals, compressed) -> None: # noqa: ARG002 + return None + + +def _make_proxy_client() -> TestClient: + config = ProxyConfig( + optimize=False, + cache_enabled=False, + rate_limit_enabled=False, + cost_tracking_enabled=False, + log_requests=False, + ccr_inject_tool=False, + ccr_handle_responses=False, + ccr_context_tracking=False, + image_optimize=False, + ) + app = create_app(config) + return TestClient(app) + + +def _force_compression(monkeypatch) -> None: # noqa: ANN001 + decision = SimpleNamespace(should_compress=True, passthrough_reason=None) + decision.apply_to_tags = lambda tags: None + monkeypatch.setattr( + "headroom.proxy.handlers.anthropic.CompressionDecision.decide", + lambda **kwargs: decision, + ) + + +def _disable_pipeline_extensions(proxy) -> None: # noqa: ANN001 + proxy.pipeline_extensions.emit = lambda *args, **kwargs: SimpleNamespace( + messages=kwargs.get("messages"), + tools=kwargs.get("tools"), + headers=kwargs.get("headers"), + metadata=kwargs.get("metadata"), + ) + + +def test_frozen_prefix_skips_marker_emission_when_tool_injection_is_deferred(monkeypatch) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + cached_marker_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + ] + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + proxy._get_compression_cache = lambda session_id: _FakeCompressionCache( + frozen_count=1, + cached_messages=cached_marker_messages, + ) + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=abc123def456abc123def456]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_frozen", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert "tools" not in forwarded + + +def test_unfrozen_prefix_keeps_reversible_ccr_path(monkeypatch) -> None: + captured: dict[str, object] = {} + marker_message = { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=0) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[marker_message], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_unfrozen", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": _RAW_TRANSCRIPT}], + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + forwarded = captured["body"] + assert forwarded["messages"] == [marker_message] + assert any(tool.get("name") == "headroom_retrieve" for tool in forwarded["tools"]) + + +def test_token_mode_reclamp_keeps_reversible_ccr_path_when_effective_prefix_drops_to_zero( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + marker_message = { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "token" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + proxy._get_compression_cache = lambda session_id: _FakeCompressionCache(frozen_count=0) + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + captured["frozen_message_count"] = kwargs["frozen_message_count"] + return SimpleNamespace( + messages=[marker_message], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_reclamp", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": _RAW_TRANSCRIPT}], + }, + ) + + assert response.status_code == 200 + assert captured.get("frozen_message_count") == 0 + assert len(captured.get("compression_calls", [])) == 1 + forwarded = captured["body"] + assert forwarded["messages"] == [marker_message] + assert any(tool.get("name") == "headroom_retrieve" for tool in forwarded["tools"]) + + +def test_existing_retrieve_tool_keeps_reversible_ccr_path_when_prefix_is_frozen( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + marker_message = { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[marker_message], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_frozen_existing_tool", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": [{"role": "user", "content": _RAW_TRANSCRIPT}], + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + forwarded = captured["body"] + assert forwarded["messages"] == [marker_message] + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_cache_mode_skip_forwards_original_prefix_when_tool_injection_is_deferred( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + original_messages = [ + {"role": "user", "content": "prefix raw content"}, + {"role": "user", "content": _RAW_TRANSCRIPT}, + ] + previous_forwarded_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + ] + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "cache" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + fake_tracker._last_original_messages = [original_messages[0]] + fake_tracker._last_forwarded_messages = previous_forwarded_messages + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=abc123def456abc123def456]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_cache_mode_skip", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert "tools" not in forwarded + + +def test_cache_mode_exact_prefix_replay_forwards_original_messages_when_tool_injection_is_deferred( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + previous_forwarded_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + ] + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "cache" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + fake_tracker._last_original_messages = original_messages.copy() + fake_tracker._last_forwarded_messages = previous_forwarded_messages + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=abc123def456abc123def456]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_cache_mode_exact_prefix", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert "tools" not in forwarded + + +def test_token_mode_cached_messages_skip_cache_update_when_pipeline_result_is_unchanged( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + marker_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + ] + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=0) + cache = _FakeCompressionCache(frozen_count=0, cached_messages=marker_messages) + cache_updates: list[tuple[list[dict], list[dict]]] = [] + cache.update_from_result = lambda originals, compressed: cache_updates.append( # type: ignore[method-assign] + (originals, compressed) + ) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + proxy._get_compression_cache = lambda session_id: cache + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=marker_messages, + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_token_cache_hit", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": _RAW_TRANSCRIPT}], + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + assert cache_updates == [] + forwarded = captured["body"] + assert forwarded["messages"] == marker_messages + + +def test_non_token_non_cache_mode_still_skips_marker_emission_when_tool_is_unavailable( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + _force_compression(monkeypatch) + monkeypatch.setattr("headroom.proxy.modes.is_token_mode", lambda mode: False) + monkeypatch.setattr("headroom.proxy.modes.is_cache_mode", lambda mode: False) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=abc123def456abc123def456]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_non_token_skip", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert "tools" not in forwarded + + +def test_non_token_non_cache_mode_keeps_reversible_path_and_records_waste_signals( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + marker_message = { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + monkeypatch.setattr("headroom.proxy.modes.is_token_mode", lambda mode: False) + monkeypatch.setattr("headroom.proxy.modes.is_cache_mode", lambda mode: False) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + class _FakeWasteSignals: + def to_dict(self) -> dict[str, bool]: + return {"oversized_tool_result": True} + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[marker_message], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=_FakeWasteSignals(), + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_non_token_reversible", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": [{"role": "user", "content": _RAW_TRANSCRIPT}], + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + forwarded = captured["body"] + assert forwarded["messages"] == [marker_message] + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_cache_mode_existing_retrieve_tool_keeps_exact_prefix_replay(monkeypatch) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + previous_forwarded_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=abc123def456abc123def456]", + } + ] + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "cache" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + fake_tracker._last_original_messages = original_messages.copy() + fake_tracker._last_forwarded_messages = previous_forwarded_messages + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[previous_forwarded_messages[0]], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_cache_mode_exact_prefix_existing_tool", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == previous_forwarded_messages + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_cache_mode_existing_retrieve_tool_compresses_only_the_unfrozen_delta( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + original_messages = [ + {"role": "user", "content": "prefix raw content"}, + {"role": "user", "content": _RAW_TRANSCRIPT}, + ] + previous_forwarded_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=prefixprefixprefixprefix]", + } + ] + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "cache" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + fake_tracker._last_original_messages = [original_messages[0]] + fake_tracker._last_forwarded_messages = previous_forwarded_messages + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=deltaabcd1234deltaabcd1234]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_cache_mode_delta_existing_tool", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + assert captured["compression_calls"][0] == [original_messages[1]] + forwarded = captured["body"] + assert forwarded["messages"] == [ + previous_forwarded_messages[0], + { + "role": "user", + "content": ( + "[100 items compressed to 10. Retrieve more: hash=deltaabcd1234deltaabcd1234]" + ), + }, + ] + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_non_token_non_cache_mode_preserves_original_messages_when_result_is_unchanged( + monkeypatch, +) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + monkeypatch.setattr("headroom.proxy.modes.is_token_mode", lambda mode: False) + monkeypatch.setattr("headroom.proxy.modes.is_cache_mode", lambda mode: False) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=original_messages, + transforms_applied=[], + timing={}, + tokens_before=40, + tokens_after=40, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_non_token_unchanged", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert len(captured.get("compression_calls", [])) == 1 + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_non_token_non_cache_mode_recovers_from_compression_errors(monkeypatch) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + monkeypatch.setattr("headroom.proxy.modes.is_token_mode", lambda mode: False) + monkeypatch.setattr("headroom.proxy.modes.is_cache_mode", lambda mode: False) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + proxy.anthropic_pipeline.apply = lambda **kwargs: (_ for _ in ()).throw( + RuntimeError("synthetic compression failure") + ) + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_non_token_error_recovery", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"] + + +def test_cache_mode_without_stable_delta_keeps_original_messages(monkeypatch) -> None: + captured: dict[str, object] = {} + original_messages = [{"role": "user", "content": _RAW_TRANSCRIPT}] + existing_tool = { + "name": "headroom_retrieve", + "description": "Retrieve compressed content", + "input_schema": {"type": "object", "properties": {}}, + } + _force_compression(monkeypatch) + + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.image_optimize = False + proxy.config.ccr_inject_tool = True + proxy.config.mode = "cache" + _disable_pipeline_extensions(proxy) + + fake_tracker = _FakePrefixTracker(frozen_count=1) + fake_tracker._last_original_messages = [{"role": "user", "content": "different prefix"}] + fake_tracker._last_forwarded_messages = [ + { + "role": "user", + "content": "[100 items compressed to 10. Retrieve more: hash=unrelatedhashunrelatedhash]", + } + ] + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured.setdefault("compression_calls", []).append(kwargs["messages"]) + return SimpleNamespace( + messages=[ + { + "role": "user", + "content": ( + "[100 items compressed to 10. " + "Retrieve more: hash=deltaabcd1234deltaabcd1234]" + ), + } + ], + transforms_applied=["fake:ccr"], + timing={}, + tokens_before=40, + tokens_after=10, + waste_signals=None, + ) + + proxy.anthropic_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "msg_ccr_cache_mode_no_delta", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 20, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/messages", + headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"}, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "tools": [existing_tool], + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + assert captured.get("compression_calls", []) == [] + forwarded = captured["body"] + assert forwarded["messages"] == original_messages + assert [tool["name"] for tool in forwarded["tools"]] == ["headroom_retrieve"]