mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
refactor(proxy): extract ccr session tracker (#2003)
## Description Extracts the sticky CCR session tracker from `headroom.proxy.helpers` into a focused state module. `helpers.SessionCcrTracker` remains as an env-aware compatibility wrapper so existing CCR tool injection and singleton call sites keep the same API. Closes # ## Type of Change - [ ] Bug fix (non-breaking change that fixes an issue) - [ ] New feature (non-breaking change that adds functionality) - [ ] Breaking change (fix or feature that would cause existing functionality to change) - [ ] Documentation update - [ ] Performance improvement - [x] Code refactoring (no functional changes) ## Changes Made - Added `headroom.proxy.ccr_session_tracker.SessionCcrTracker` as the pure bounded LRU CCR state holder. - Replaced the in-helper CCR tracker implementation with a small env-aware wrapper. - Added direct tracker tests for unknown sessions, monotonic done state, first-write golden bytes, provider isolation, LRU eviction, reset, and input validation. ## Testing - [x] Unit tests pass (`pytest`) - [x] Linting passes (`ruff check .`) - [x] Type checking passes (`mypy headroom`) - [x] New tests added for new functionality - [ ] Manual testing performed ### Test Output ```text python -m pytest tests/test_ccr_session_tracker.py tests/test_ccr_tool_always_on.py tests/test_corrupt_golden_bytes_recovery.py tests/test_issue_728_empty_tools_injection.py 35 passed in 0.69s python -m ruff check . All checks passed! python -m ruff format --check . 1069 files already formatted python -m mypy headroom --ignore-missing-imports Success: no issues found in 410 source files gitleaks protect --staged --no-banner --redact no leaks found ``` ## Real Behavior Proof - Environment: Windows, Python 3.13.13 - Exact command / steps: Ran direct CCR tracker tests, CCR always-on tests, corrupt golden byte recovery tests, empty tools injection regression tests, full ruff, format check, mypy, and staged gitleaks scan. - Observed result: Existing sticky CCR tool behavior and recovery behavior remain green while the CCR session state domain is directly covered. - Not tested: Full repository pytest suite locally; CI covers the broader matrix. ## Review Readiness - [x] I have performed a self-review - [x] This PR is ready for human review ## Checklist - [x] My code follows the project's style guidelines - [x] I have performed a self-review of my code - [x] I have commented my code, particularly in hard-to-understand areas - [ ] I have made corresponding changes to the documentation - [x] My changes generate no new warnings - [x] I have added tests that prove my fix is effective or that my feature works - [x] New and existing unit tests pass locally with my changes - [ ] I have updated the CHANGELOG.md if applicable ## Screenshots (if applicable) N/A ## Additional Notes Documentation and changelog updates are not applicable for this internal refactor. The default-branch Dependabot alerts reported during push are pre-existing and unrelated to this PR.
This commit is contained in:
parent
4e19bcf6ce
commit
e92c253977
3 changed files with 160 additions and 95 deletions
87
headroom/proxy/ccr_session_tracker.py
Normal file
87
headroom/proxy/ccr_session_tracker.py
Normal file
|
|
@ -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()
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
69
tests/test_ccr_session_tracker.py
Normal file
69
tests/test_ccr_session_tracker.py
Normal file
|
|
@ -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"")
|
||||
Loading…
Add table
Add a link
Reference in a new issue