mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
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:
parent
78d2847398
commit
7c2c55abc0
4 changed files with 369 additions and 2 deletions
137
headroom/hooks.py
Normal file
137
headroom/hooks.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
151
tests/test_hooks.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue