diff --git a/headroom/proxy/ccr_session_tracker.py b/headroom/proxy/ccr_session_tracker.py new file mode 100644 index 000000000..6ddc14856 --- /dev/null +++ b/headroom/proxy/ccr_session_tracker.py @@ -0,0 +1,87 @@ +"""Session-scoped state for sticky CCR retrieval tool injection.""" + +from __future__ import annotations + +import threading +from collections import OrderedDict + + +class SessionCcrTracker: + """Bounded LRU tracker recording per-provider/session CCR state.""" + + def __init__(self, max_sessions: int) -> None: + if max_sessions <= 0: + raise ValueError("max_sessions must be > 0") + self._max_sessions = max_sessions + self._lock = threading.RLock() + self._sessions: OrderedDict[tuple[str, str], tuple[bool, bytes | None]] = OrderedDict() + + @property + def active_sessions(self) -> int: + with self._lock: + return len(self._sessions) + + def _key(self, provider: str, session_id: str) -> tuple[str, str]: + return (provider, session_id) + + def has_done_ccr(self, provider: str, session_id: str) -> bool: + """Return True when this session has previously performed CCR.""" + + if not provider: + raise ValueError("provider must be non-empty") + if not session_id: + raise ValueError("session_id must be non-empty") + key = self._key(provider, session_id) + with self._lock: + entry = self._sessions.get(key) + if entry is None: + return False + self._sessions.move_to_end(key) + return entry[0] + + def get_golden_tool_bytes(self, provider: str, session_id: str) -> bytes | None: + """Return recorded golden CCR tool-definition bytes, if any.""" + + if not provider: + raise ValueError("provider must be non-empty") + if not session_id: + raise ValueError("session_id must be non-empty") + key = self._key(provider, session_id) + with self._lock: + entry = self._sessions.get(key) + if entry is None: + return None + self._sessions.move_to_end(key) + return entry[1] + + def record_ccr_done( + self, + provider: str, + session_id: str, + golden_tool_bytes: bytes, + ) -> None: + """Mark the session as having performed CCR and pin golden tool bytes.""" + + if not provider: + raise ValueError("provider must be non-empty") + if not session_id: + raise ValueError("session_id must be non-empty") + if not golden_tool_bytes: + raise ValueError("golden_tool_bytes must be non-empty") + key = self._key(provider, session_id) + with self._lock: + existing = self._sessions.get(key) + if existing is None: + self._sessions[key] = (True, golden_tool_bytes) + else: + pinned = existing[1] if existing[1] is not None else golden_tool_bytes + self._sessions[key] = (True, pinned) + self._sessions.move_to_end(key) + while len(self._sessions) > self._max_sessions: + self._sessions.popitem(last=False) + + def reset(self) -> None: + """Clear all session state.""" + + with self._lock: + self._sessions.clear() diff --git a/headroom/proxy/helpers.py b/headroom/proxy/helpers.py index 68d1ae7f4..8e48aac03 100644 --- a/headroom/proxy/helpers.py +++ b/headroom/proxy/helpers.py @@ -38,6 +38,7 @@ 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_session_tracker import SessionCcrTracker as _SessionCcrTracker from headroom.proxy.tool_injection_config import ( ToolInjectionStickyMode, ) @@ -2375,105 +2376,13 @@ def apply_session_sticky_memory_tools( # tool list bytes stay byte-stable across turns once injected. -class SessionCcrTracker: - """Bounded LRU tracker recording per-(provider, session_id) CCR state. - - Two pieces of state per session: - - * ``has_done_ccr``: True once the proxy observed any CCR - compression marker in the messages of a request. Once True, it - never flips back to False (the prompt cache anchored on the - previous turn's tool list demands the tool stays present). - * ``golden_tool_bytes``: canonical serialization of the - ``headroom_retrieve`` tool definition recorded the first time - the tracker injected it. Subsequent turns replay these bytes - verbatim. - - Bounded by ``max_sessions`` via ``OrderedDict`` LRU. Mirrors - :class:`SessionToolTracker` semantics so the operator's mental model - is one tracker pattern, not two. - """ +class SessionCcrTracker(_SessionCcrTracker): + """Env-aware compatibility wrapper for the pure CCR session tracker.""" def __init__(self, max_sessions: int | None = None) -> None: if max_sessions is None: max_sessions = get_tool_tracker_max_sessions() - if max_sessions <= 0: - raise ValueError("max_sessions must be > 0") - self._max_sessions = max_sessions - self._lock = threading.RLock() - # Value is (has_done_ccr, golden_tool_bytes_or_none). - self._sessions: OrderedDict[tuple[str, str], tuple[bool, bytes | None]] = OrderedDict() - - @property - def active_sessions(self) -> int: - with self._lock: - return len(self._sessions) - - def _key(self, provider: str, session_id: str) -> tuple[str, str]: - return (provider, session_id) - - def has_done_ccr(self, provider: str, session_id: str) -> bool: - """Return True iff this session has previously performed CCR.""" - if not provider: - raise ValueError("provider must be non-empty") - if not session_id: - raise ValueError("session_id must be non-empty") - with self._lock: - entry = self._sessions.get(self._key(provider, session_id)) - if entry is None: - return False - self._sessions.move_to_end(self._key(provider, session_id)) - return entry[0] - - def get_golden_tool_bytes(self, provider: str, session_id: str) -> bytes | None: - """Return the recorded golden tool-definition bytes, or None.""" - if not provider: - raise ValueError("provider must be non-empty") - if not session_id: - raise ValueError("session_id must be non-empty") - with self._lock: - entry = self._sessions.get(self._key(provider, session_id)) - if entry is None: - return None - self._sessions.move_to_end(self._key(provider, session_id)) - return entry[1] - - def record_ccr_done( - self, - provider: str, - session_id: str, - golden_tool_bytes: bytes, - ) -> None: - """Mark the session as having performed CCR and pin the golden bytes. - - First-write wins for ``golden_tool_bytes`` (subsequent calls - with the same session keep the original bytes — prevents drift - if the canonical serialization changed mid-session). The - ``has_done_ccr`` flag is monotonic: once True, never False. - """ - if not provider: - raise ValueError("provider must be non-empty") - if not session_id: - raise ValueError("session_id must be non-empty") - if not golden_tool_bytes: - raise ValueError("golden_tool_bytes must be non-empty") - key = self._key(provider, session_id) - with self._lock: - existing = self._sessions.get(key) - if existing is None: - self._sessions[key] = (True, golden_tool_bytes) - else: - # Preserve original golden bytes; just promote the flag. - pinned = existing[1] if existing[1] is not None else golden_tool_bytes - self._sessions[key] = (True, pinned) - self._sessions.move_to_end(key) - while len(self._sessions) > self._max_sessions: - self._sessions.popitem(last=False) - - def reset(self) -> None: - """Clear all session state (test helper).""" - with self._lock: - self._sessions.clear() + super().__init__(max_sessions=max_sessions) # Process-wide singleton. diff --git a/tests/test_ccr_session_tracker.py b/tests/test_ccr_session_tracker.py new file mode 100644 index 000000000..884e7e644 --- /dev/null +++ b/tests/test_ccr_session_tracker.py @@ -0,0 +1,69 @@ +from __future__ import annotations + +import pytest + +from headroom.proxy.ccr_session_tracker import SessionCcrTracker + + +def test_tracker_reports_unknown_session_as_not_done() -> None: + tracker = SessionCcrTracker(max_sessions=10) + + assert tracker.has_done_ccr("anthropic", "s-1") is False + assert tracker.get_golden_tool_bytes("anthropic", "s-1") is None + + +def test_tracker_records_monotonic_done_state_and_golden_bytes() -> None: + tracker = SessionCcrTracker(max_sessions=10) + + tracker.record_ccr_done("anthropic", "s-1", b"first") + tracker.record_ccr_done("anthropic", "s-1", b"second") + + assert tracker.has_done_ccr("anthropic", "s-1") is True + assert tracker.get_golden_tool_bytes("anthropic", "s-1") == b"first" + + +def test_tracker_keeps_provider_namespaces_independent() -> None: + tracker = SessionCcrTracker(max_sessions=10) + + tracker.record_ccr_done("anthropic", "shared", b"anthropic") + tracker.record_ccr_done("openai", "shared", b"openai") + + assert tracker.get_golden_tool_bytes("anthropic", "shared") == b"anthropic" + assert tracker.get_golden_tool_bytes("openai", "shared") == b"openai" + + +def test_tracker_evicts_least_recently_used_session() -> None: + tracker = SessionCcrTracker(max_sessions=2) + tracker.record_ccr_done("anthropic", "s-1", b"a") + tracker.record_ccr_done("anthropic", "s-2", b"b") + + assert tracker.has_done_ccr("anthropic", "s-1") is True + tracker.record_ccr_done("anthropic", "s-3", b"c") + + assert tracker.active_sessions == 2 + assert tracker.has_done_ccr("anthropic", "s-1") is True + assert tracker.has_done_ccr("anthropic", "s-2") is False + assert tracker.has_done_ccr("anthropic", "s-3") is True + + +def test_tracker_reset_clears_state() -> None: + tracker = SessionCcrTracker(max_sessions=10) + tracker.record_ccr_done("anthropic", "s-1", b"bytes") + + tracker.reset() + + assert tracker.active_sessions == 0 + assert tracker.has_done_ccr("anthropic", "s-1") is False + + +def test_tracker_validates_inputs() -> None: + with pytest.raises(ValueError, match="max_sessions"): + SessionCcrTracker(max_sessions=0) + + tracker = SessionCcrTracker(max_sessions=10) + with pytest.raises(ValueError, match="provider"): + tracker.has_done_ccr("", "s-1") + with pytest.raises(ValueError, match="session_id"): + tracker.get_golden_tool_bytes("anthropic", "") + with pytest.raises(ValueError, match="golden_tool_bytes"): + tracker.record_ccr_done("anthropic", "s-1", b"")