diff --git a/headroom/proxy/ccr_golden_policy.py b/headroom/proxy/ccr_golden_policy.py new file mode 100644 index 000000000..4d54fddb7 --- /dev/null +++ b/headroom/proxy/ccr_golden_policy.py @@ -0,0 +1,52 @@ +"""Policy helpers for replaying CCR golden tool definitions.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Literal, cast + +from headroom.ccr.tool_injection import create_ccr_tool_definition + + +@dataclass(frozen=True) +class CcrToolDefinitionReplay: + """CCR tool definition selected for sticky replay or fresh injection.""" + + tool_definition: dict[str, Any] + canonical_bytes: bytes + used_golden_bytes: bool + + +def serialize_ccr_tool_definition_canonical(tool_definition: dict[str, Any]) -> bytes: + """Return stable canonical bytes for a CCR tool definition.""" + + return json.dumps( + tool_definition, + ensure_ascii=False, + separators=(",", ":"), + ).encode("utf-8") + + +def replay_golden_ccr_tool_definition(golden_tool_bytes: bytes) -> CcrToolDefinitionReplay: + """Decode a stored CCR tool definition and preserve its original bytes.""" + + tool_definition = json.loads(golden_tool_bytes.decode("utf-8")) + return CcrToolDefinitionReplay( + tool_definition=cast(dict[str, Any], tool_definition), + canonical_bytes=golden_tool_bytes, + used_golden_bytes=True, + ) + + +def create_fresh_ccr_tool_definition( + provider: Literal["anthropic", "openai", "google"], +) -> CcrToolDefinitionReplay: + """Create and canonicalize a fresh CCR tool definition for ``provider``.""" + + tool_definition = create_ccr_tool_definition(provider) + return CcrToolDefinitionReplay( + tool_definition=tool_definition, + canonical_bytes=serialize_ccr_tool_definition_canonical(tool_definition), + used_golden_bytes=False, + ) diff --git a/headroom/proxy/helpers.py b/headroom/proxy/helpers.py index f3c95e2a8..1e299cbe1 100644 --- a/headroom/proxy/helpers.py +++ b/headroom/proxy/helpers.py @@ -38,6 +38,10 @@ from headroom.proxy.body_forwarding import ( prepare_outbound_body_bytes as prepare_outbound_body_bytes, # noqa: F401 - compatibility export ) from headroom.proxy.body_forwarding import serialize_body_canonical +from headroom.proxy.ccr_golden_policy import ( + create_fresh_ccr_tool_definition, + replay_golden_ccr_tool_definition, +) from headroom.proxy.ccr_session_tracker import SessionCcrTracker as _SessionCcrTracker from headroom.proxy.internal_header_policy import ( INTERNAL_HEADER_PREFIX, @@ -2303,7 +2307,7 @@ def apply_session_sticky_ccr_tool( Returns ``(updated_tools, was_injected)``. ``updated_tools`` is a fresh list (caller-safe). """ - from headroom.ccr.tool_injection import CCR_TOOL_NAME, create_ccr_tool_definition + from headroom.ccr.tool_injection import CCR_TOOL_NAME if provider not in ("anthropic", "openai", "google"): raise ValueError(f"unsupported provider: {provider!r}") @@ -2337,14 +2341,13 @@ def apply_session_sticky_ccr_tool( request_id=request_id, ) return tools_out, False - tool_def = create_ccr_tool_definition(provider) - canonical = serialize_tool_definition_canonical(tool_def) - tools_out.append(tool_def) + replay = create_fresh_ccr_tool_definition(provider) + tools_out.append(replay.tool_definition) log_tool_injection_decision( provider=provider, session_id=None, decision="inject_first_time", - tool_definition_bytes_count=len(canonical), + tool_definition_bytes_count=len(replay.canonical_bytes), request_id=request_id, ) return tools_out, True @@ -2361,13 +2364,13 @@ def apply_session_sticky_ccr_tool( golden = tracker.get_golden_tool_bytes(provider, session_id) if golden is not None: try: - tool_def = json.loads(golden.decode("utf-8")) - tools_out.append(tool_def) + replay = replay_golden_ccr_tool_definition(golden) + tools_out.append(replay.tool_definition) log_tool_injection_decision( provider=provider, session_id=session_id, decision="inject_sticky_replay", - tool_definition_bytes_count=len(golden), + tool_definition_bytes_count=len(replay.canonical_bytes), request_id=request_id, ) return tools_out, True @@ -2381,15 +2384,14 @@ def apply_session_sticky_ccr_tool( # Fall through to fresh creation below # Tracker says "done CCR" but has no golden bytes (or they were corrupt). Pin # them now so future turns are stable. - tool_def = create_ccr_tool_definition(provider) - canonical = serialize_tool_definition_canonical(tool_def) - tracker.record_ccr_done(provider, session_id, canonical) - tools_out.append(tool_def) + replay = create_fresh_ccr_tool_definition(provider) + tracker.record_ccr_done(provider, session_id, replay.canonical_bytes) + tools_out.append(replay.tool_definition) log_tool_injection_decision( provider=provider, session_id=session_id, decision="inject_sticky_replay", - tool_definition_bytes_count=len(canonical), + tool_definition_bytes_count=len(replay.canonical_bytes), request_id=request_id, ) return tools_out, True @@ -2405,15 +2407,14 @@ def apply_session_sticky_ccr_tool( ) return tools_out, False - tool_def = create_ccr_tool_definition(provider) - canonical = serialize_tool_definition_canonical(tool_def) - tracker.record_ccr_done(provider, session_id, canonical) - tools_out.append(tool_def) + replay = create_fresh_ccr_tool_definition(provider) + tracker.record_ccr_done(provider, session_id, replay.canonical_bytes) + tools_out.append(replay.tool_definition) log_tool_injection_decision( provider=provider, session_id=session_id, decision="inject_first_time", - tool_definition_bytes_count=len(canonical), + tool_definition_bytes_count=len(replay.canonical_bytes), request_id=request_id, ) return tools_out, True diff --git a/tests/test_ccr_golden_policy.py b/tests/test_ccr_golden_policy.py new file mode 100644 index 000000000..1604f7936 --- /dev/null +++ b/tests/test_ccr_golden_policy.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +import pytest + +from headroom.ccr.tool_injection import CCR_TOOL_NAME, create_ccr_tool_definition +from headroom.proxy.ccr_golden_policy import ( + create_fresh_ccr_tool_definition, + replay_golden_ccr_tool_definition, + serialize_ccr_tool_definition_canonical, +) + + +def test_replays_golden_definition_without_reserializing() -> None: + golden = b'{ "name" : "headroom_retrieve" , "description" : "client bytes" }' + + replay = replay_golden_ccr_tool_definition(golden) + + assert replay.tool_definition["name"] == CCR_TOOL_NAME + assert replay.canonical_bytes == golden + assert replay.used_golden_bytes is True + + +def test_rejects_invalid_golden_json() -> None: + with pytest.raises(ValueError): + replay_golden_ccr_tool_definition(b"not-json") + + +def test_rejects_non_utf8_golden_bytes() -> None: + with pytest.raises(UnicodeDecodeError): + replay_golden_ccr_tool_definition(b"\x80\x81") + + +def test_fresh_definition_uses_canonical_bytes() -> None: + replay = create_fresh_ccr_tool_definition("anthropic") + + assert replay.tool_definition == create_ccr_tool_definition("anthropic") + assert replay.canonical_bytes == serialize_ccr_tool_definition_canonical(replay.tool_definition) + assert replay.used_golden_bytes is False