From 7c2c55abc0250259ed0260f0db5257b1308659d5 Mon Sep 17 00:00:00 2001 From: chopratejas Date: Thu, 19 Feb 2026 08:17:30 -0800 Subject: [PATCH] =?UTF-8?q?Add=20Compression=20Hooks=20=E2=80=94=20extensi?= =?UTF-8?q?on=20points=20for=20SaaS=20and=20advanced=20customization?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three hooks at well-defined pipeline stages: 1. pre_compress(messages, ctx) → messages Modify messages before compression: cross-turn dedup, memory injection. 2. compute_biases(messages, ctx) → dict[int, float] Per-message compression bias: position-aware, phase-aware, learned. 3. post_compress(event) → None Observe results: failure-driven learning, analytics, A/B testing. - headroom/hooks.py: CompressionHooks, CompressContext, CompressEvent - ProxyConfig.hooks: optional, default None (zero overhead) - Wired into Anthropic and OpenAI handlers - ContentRouter reads hook biases, multiplies with tool bias - 11 tests --- headroom/hooks.py | 137 +++++++++++++++++++++++ headroom/proxy/server.py | 78 ++++++++++++- headroom/transforms/content_router.py | 5 +- tests/test_hooks.py | 151 ++++++++++++++++++++++++++ 4 files changed, 369 insertions(+), 2 deletions(-) create mode 100644 headroom/hooks.py create mode 100644 tests/test_hooks.py diff --git a/headroom/hooks.py b/headroom/hooks.py new file mode 100644 index 000000000..dc1f7454b --- /dev/null +++ b/headroom/hooks.py @@ -0,0 +1,137 @@ +"""Compression Hooks — extension points for customizing compression behavior. + +Three hooks at well-defined pipeline stages: + +1. pre_compress: modify messages before compression (dedup, filter, inject) +2. compute_biases: set per-message compression aggressiveness (position-aware, phase-aware) +3. post_compress: observe results after compression (learning, analytics, logging) + +Default implementation is no-op — OSS behavior unchanged. Override these +in a subclass to customize (e.g., Headroom SaaS implements position-aware +compression and cross-turn deduplication via these hooks). + +Usage: + from headroom.hooks import CompressionHooks, CompressContext + + class MyHooks(CompressionHooks): + def compute_biases(self, messages, ctx): + # Position-aware: keep more in the middle (attention is weakest there) + biases = {} + for i in range(len(messages)): + pos = i / max(len(messages) - 1, 1) + biases[i] = 1.0 + 0.5 * (1.0 - abs(2 * pos - 1)) + return biases + + config = ProxyConfig(hooks=MyHooks()) +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass +class CompressContext: + """Context passed to pre_compress and compute_biases hooks. + + Provides enough information for hooks to make decisions without + needing to understand the proxy's internals. + """ + + model: str = "" + user_query: str = "" + turn_number: int = 0 + tool_calls: list[str] = field(default_factory=list) + provider: str = "" # "anthropic", "openai", "gemini" + + +@dataclass +class CompressEvent: + """Data passed to post_compress hook after compression completes. + + Contains before/after state and full metrics for learning and analytics. + """ + + tokens_before: int = 0 + tokens_after: int = 0 + tokens_saved: int = 0 + compression_ratio: float = 0.0 + transforms_applied: list[str] = field(default_factory=list) + ccr_hashes: list[str] = field(default_factory=list) + model: str = "" + user_query: str = "" + provider: str = "" + + +class CompressionHooks: + """Base class for compression hooks. Override methods to customize. + + All methods have no-op defaults — OSS behavior is unchanged unless + a subclass is provided via ProxyConfig(hooks=MyHooks()). + """ + + def pre_compress( + self, + messages: list[dict[str, Any]], + ctx: CompressContext, + ) -> list[dict[str, Any]]: + """Called before the compression pipeline runs. + + Modify and return the messages list. Use for: + - Cross-turn deduplication (compare against recent CCR entries) + - Memory injection (add relevant context from external sources) + - Pre-filtering (remove messages irrelevant to the user's query) + - Phase detection (reorder/prioritize based on task phase) + + Args: + messages: The full message list (will be compressed next). + ctx: Compression context (model, query, turn, tool calls). + + Returns: + Modified (or unmodified) messages list. + """ + return messages + + def compute_biases( + self, + messages: list[dict[str, Any]], + ctx: CompressContext, + ) -> dict[int, float]: + """Compute per-message compression bias. + + Return a dict mapping message index to compression bias: + - 1.0 = default compression + - >1.0 = keep more (compress less aggressively) + - <1.0 = compress more aggressively + - Missing indices get 1.0 + + Use for: + - Position-aware compression (middle messages get higher bias + because LLM attention is weakest there) + - Phase-aware budgets (old exploration messages get lower bias, + recent execution messages get higher bias) + - Per-tool learned biases (from TOIN analysis) + + Args: + messages: The full message list. + ctx: Compression context. + + Returns: + Dict of {message_index: bias_float}. Empty dict = all default. + """ + return {} + + def post_compress(self, event: CompressEvent) -> None: + """Called after compression completes. Observational only. + + Use for: + - Failure-driven learning (log events, analyze offline) + - Per-org analytics and dashboards + - A/B testing of compression strategies + - Anomaly detection (alert on sudden ratio changes) + + Args: + event: Full compression event with before/after metrics. + """ + pass diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index 9db32ab83..83dbbb2ea 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -313,6 +313,9 @@ class ProxyConfig: memory_neo4j_user: str = "neo4j" memory_neo4j_password: str = "password" + # Compression Hooks (for SaaS and advanced customization) + hooks: Any = None # CompressionHooks instance, or None for default behavior + # ============================================================================= # Caching @@ -1625,6 +1628,21 @@ class HeadroomProxy: tokenizer = get_tokenizer(model) original_tokens = tokenizer.count_messages(messages) + # Hook: pre_compress — let hooks modify messages before compression + if self.config.hooks: + from headroom.hooks import CompressContext + from headroom.transforms.query_echo import extract_user_query + + _hook_ctx = CompressContext( + model=model, + user_query=extract_user_query(messages), + provider="anthropic", + ) + try: + messages = self.config.hooks.pre_compress(messages, _hook_ctx) + except Exception as e: + logger.debug(f"[{request_id}] pre_compress hook error: {e}") + # Apply optimization transforms_applied = [] optimized_messages = messages @@ -1637,6 +1655,9 @@ class HeadroomProxy: messages=messages, model=model, model_limit=context_limit, + biases=self.config.hooks.compute_biases(messages, _hook_ctx) + if self.config.hooks + else None, ) if result.messages != messages: @@ -1651,6 +1672,28 @@ class HeadroomProxy: tokens_saved = original_tokens - optimized_tokens optimization_latency = (time.time() - start_time) * 1000 + # Hook: post_compress — let hooks observe compression results + if self.config.hooks and tokens_saved > 0: + from headroom.hooks import CompressEvent + + try: + self.config.hooks.post_compress( + CompressEvent( + tokens_before=original_tokens, + tokens_after=optimized_tokens, + tokens_saved=tokens_saved, + compression_ratio=tokens_saved / original_tokens + if original_tokens > 0 + else 0, + transforms_applied=transforms_applied, + model=model, + user_query=_hook_ctx.user_query if self.config.hooks else "", + provider="anthropic", + ) + ) + except Exception as e: + logger.debug(f"[{request_id}] post_compress hook error: {e}") + # CCR Tool Injection: Inject retrieval tool if compression occurred tools = body.get("tools") if self.config.ccr_inject_tool or self.config.ccr_inject_system_instructions: @@ -4065,6 +4108,18 @@ class HeadroomProxy: tokenizer = get_tokenizer(model) original_tokens = tokenizer.count_messages(messages) + # Hook: pre_compress + _hook_biases = None + if self.config.hooks: + from headroom.hooks import CompressContext + + _hook_ctx = CompressContext(model=model, provider="openai") + try: + messages = self.config.hooks.pre_compress(messages, _hook_ctx) + _hook_biases = self.config.hooks.compute_biases(messages, _hook_ctx) + except Exception as e: + logger.debug(f"[{request_id}] Hook error: {e}") + # Optimization transforms_applied = [] optimized_messages = messages @@ -4077,11 +4132,11 @@ class HeadroomProxy: messages=messages, model=model, model_limit=context_limit, + biases=_hook_biases, ) if result.messages != messages: optimized_messages = result.messages transforms_applied = result.transforms_applied - # Use pipeline's token counts for consistency with pipeline logs original_tokens = result.tokens_before optimized_tokens = result.tokens_after except Exception as e: @@ -4090,6 +4145,27 @@ class HeadroomProxy: tokens_saved = original_tokens - optimized_tokens optimization_latency = (time.time() - start_time) * 1000 + # Hook: post_compress + if self.config.hooks and tokens_saved > 0: + from headroom.hooks import CompressEvent + + try: + self.config.hooks.post_compress( + CompressEvent( + tokens_before=original_tokens, + tokens_after=optimized_tokens, + tokens_saved=tokens_saved, + compression_ratio=tokens_saved / original_tokens + if original_tokens > 0 + else 0, + transforms_applied=transforms_applied, + model=model, + provider="openai", + ) + ) + except Exception as e: + logger.debug(f"[{request_id}] post_compress hook error: {e}") + # CCR Tool Injection: Inject retrieval tool if compression occurred tools = body.get("tools") if self.config.ccr_inject_tool or self.config.ccr_inject_system_instructions: diff --git a/headroom/transforms/content_router.py b/headroom/transforms/content_router.py index 852a11822..93a1ad6ea 100644 --- a/headroom/transforms/content_router.py +++ b/headroom/transforms/content_router.py @@ -1173,6 +1173,7 @@ class ContentRouter(Transform): """ tokens_before = sum(tokenizer.count_text(str(m.get("content", ""))) for m in messages) context = kwargs.get("context", "") + hook_biases: dict[int, float] = kwargs.get("biases") or {} # Build tool name map for exclusion checking tool_name_map = self._build_tool_name_map(messages) @@ -1265,8 +1266,10 @@ class ContentRouter(Transform): continue # Route and compress based on content detection - # Use tool-specific bias for tool messages, default 1.0 for others + # Merge tool-specific bias with hook-provided bias (multiplicative) msg_bias = bias if role == "tool" else 1.0 + if i in hook_biases: + msg_bias *= hook_biases[i] result = self.compress(content, context=context, bias=msg_bias) if result.compression_ratio < 0.9: diff --git a/tests/test_hooks.py b/tests/test_hooks.py new file mode 100644 index 000000000..8a2231005 --- /dev/null +++ b/tests/test_hooks.py @@ -0,0 +1,151 @@ +"""Tests for Compression Hooks interface.""" + +from headroom.hooks import CompressContext, CompressEvent, CompressionHooks + + +class TestCompressionHooksDefaults: + """Default (no-op) hooks don't modify anything.""" + + def test_pre_compress_returns_messages_unchanged(self): + hooks = CompressionHooks() + messages = [{"role": "user", "content": "hello"}] + ctx = CompressContext(model="test") + result = hooks.pre_compress(messages, ctx) + assert result is messages + + def test_compute_biases_returns_empty(self): + hooks = CompressionHooks() + messages = [{"role": "user", "content": "hello"}] + ctx = CompressContext(model="test") + result = hooks.compute_biases(messages, ctx) + assert result == {} + + def test_post_compress_is_noop(self): + hooks = CompressionHooks() + event = CompressEvent(tokens_before=100, tokens_after=50) + hooks.post_compress(event) # Should not raise + + +class TestCustomHooks: + """Custom hook implementations work correctly.""" + + def test_pre_compress_can_modify_messages(self): + class FilterHooks(CompressionHooks): + def pre_compress(self, messages, ctx): + return [m for m in messages if m.get("role") != "system"] + + hooks = FilterHooks() + messages = [ + {"role": "system", "content": "you are helpful"}, + {"role": "user", "content": "hello"}, + ] + result = hooks.pre_compress(messages, CompressContext()) + assert len(result) == 1 + assert result[0]["role"] == "user" + + def test_compute_biases_position_aware(self): + class PositionAwareHooks(CompressionHooks): + def compute_biases(self, messages, ctx): + biases = {} + n = len(messages) + for i in range(n): + pos = i / max(n - 1, 1) + # U-curve: middle gets higher bias + biases[i] = 1.0 + 0.5 * (1.0 - abs(2 * pos - 1)) + return biases + + hooks = PositionAwareHooks() + messages = [{"role": "user"}] * 10 + biases = hooks.compute_biases(messages, CompressContext()) + + # Edges should have lower bias, middle should have higher + assert biases[0] < biases[5] # start < middle + assert biases[9] < biases[5] # end < middle + assert biases[5] > 1.0 # middle is above default + + def test_post_compress_records_event(self): + events = [] + + class LoggingHooks(CompressionHooks): + def post_compress(self, event): + events.append(event) + + hooks = LoggingHooks() + event = CompressEvent( + tokens_before=1000, + tokens_after=300, + tokens_saved=700, + compression_ratio=0.7, + model="claude-sonnet", + provider="anthropic", + ) + hooks.post_compress(event) + assert len(events) == 1 + assert events[0].tokens_saved == 700 + + def test_hooks_receive_correct_context(self): + received_ctx = [] + + class ContextCapture(CompressionHooks): + def pre_compress(self, messages, ctx): + received_ctx.append(ctx) + return messages + + hooks = ContextCapture() + ctx = CompressContext( + model="gpt-4o", + user_query="find errors", + provider="openai", + turn_number=5, + tool_calls=["read_file", "grep"], + ) + hooks.pre_compress([], ctx) + + assert received_ctx[0].model == "gpt-4o" + assert received_ctx[0].user_query == "find errors" + assert received_ctx[0].provider == "openai" + assert received_ctx[0].turn_number == 5 + assert "read_file" in received_ctx[0].tool_calls + + +class TestCompressEvent: + def test_event_fields(self): + event = CompressEvent( + tokens_before=1000, + tokens_after=200, + tokens_saved=800, + compression_ratio=0.8, + transforms_applied=["smart:relevance(500->20)", "router:code_aware:0.45"], + ccr_hashes=["abc123", "def456"], + model="claude-sonnet-4-5-20250929", + user_query="What are the test failures?", + provider="anthropic", + ) + assert event.compression_ratio == 0.8 + assert len(event.transforms_applied) == 2 + assert len(event.ccr_hashes) == 2 + + def test_event_defaults(self): + event = CompressEvent() + assert event.tokens_before == 0 + assert event.transforms_applied == [] + assert event.provider == "" + + +class TestCompressContext: + def test_context_defaults(self): + ctx = CompressContext() + assert ctx.model == "" + assert ctx.tool_calls == [] + assert ctx.turn_number == 0 + + def test_context_with_values(self): + ctx = CompressContext( + model="gpt-4o", + user_query="find the bug", + turn_number=3, + tool_calls=["read_file", "bash"], + provider="openai", + ) + assert ctx.model == "gpt-4o" + assert len(ctx.tool_calls) == 2