Add Compression Hooks — extension points for SaaS and advanced customization

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
This commit is contained in:
chopratejas 2026-02-19 08:17:30 -08:00
parent 78d2847398
commit 7c2c55abc0
4 changed files with 369 additions and 2 deletions

137
headroom/hooks.py Normal file
View file

@ -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

View file

@ -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:

View file

@ -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:

151
tests/test_hooks.py Normal file
View file

@ -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