feat: introduce canonical pipeline lifecycle contract

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
JerrettDavis 2026-04-21 21:04:04 -05:00
parent 5738339524
commit cd9c2e1d01
11 changed files with 3477 additions and 2561 deletions

View file

@ -112,6 +112,17 @@ Bundles the [RTK](https://github.com/rtk-ai/rtk) binary for shell-output rewriti
→ [Architecture](https://headroom-docs.vercel.app/docs/architecture) · [CCR reversible compression](https://headroom-docs.vercel.app/docs/ccr) · [Kompress-base model card](https://huggingface.co/chopratejas/kompress-base)
### Canonical pipeline lifecycle
Headroom now exposes one stable request lifecycle across `compress()`, the SDK, and the proxy:
`Setup``Pre-Start``Post-Start``Input Received``Input Cached``Input Routed``Input Compressed``Input Remembered``Pre-Send``Post-Send``Response Received`
- **Transforms** still do the work: CacheAligner, ContentRouter, SmartCrusher, CodeCompressor, Kompress-base, IntelligentContext / RollingWindow.
- **Pipeline extensions** observe or customize those lifecycle stages via `on_pipeline_event(...)`.
- **Compression hooks** still work and now sit alongside the canonical lifecycle instead of being the only extension seam.
- **Proxy extensions** remain the server/app integration seam for ASGI middleware, routes, and startup policy.
---
## Proof

View file

@ -175,6 +175,11 @@ __all__ = [
"CompressionHooks",
"CompressContext",
"CompressEvent",
# Canonical pipeline
"PipelineStage",
"PipelineEvent",
"PipelineExtensionManager",
"CANONICAL_PIPELINE_STAGES",
# Shared context for multi-agent workflows
"SharedContext",
]
@ -268,6 +273,11 @@ _LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"CompressionHooks": ("headroom.hooks", "CompressionHooks"),
"CompressContext": ("headroom.hooks", "CompressContext"),
"CompressEvent": ("headroom.hooks", "CompressEvent"),
# Canonical pipeline
"PipelineStage": ("headroom.pipeline", "PipelineStage"),
"PipelineEvent": ("headroom.pipeline", "PipelineEvent"),
"PipelineExtensionManager": ("headroom.pipeline", "PipelineExtensionManager"),
"CANONICAL_PIPELINE_STAGES": ("headroom.pipeline", "CANONICAL_PIPELINE_STAGES"),
# Shared context
"SharedContext": ("headroom.shared_context", "SharedContext"),
}

View file

@ -20,6 +20,7 @@ from .config import (
SimulationResult,
)
from .parser import parse_messages
from .pipeline import PipelineExtensionManager, PipelineStage, summarize_routing_markers
from .providers.base import Provider
from .storage import create_storage
from .tokenizer import Tokenizer
@ -324,6 +325,10 @@ class HeadroomClient:
# Initialize transform pipeline
self._pipeline = TransformPipeline(self._config, provider=self._provider)
self._pipeline_extensions = PipelineExtensionManager(
extensions=self._config.pipeline_extensions,
discover=self._config.discover_pipeline_extensions,
)
# Initialize cache optimizer
self._cache_optimizer: BaseCacheOptimizer | None = None
@ -357,6 +362,16 @@ class HeadroomClient:
self.chat = type("Chat", (), {"completions": ChatCompletions(self)})()
# Public API - Anthropic style
self.messages = Messages(self)
self._pipeline_extensions.emit(
PipelineStage.SETUP,
operation="sdk.setup",
provider=self._provider.name.lower(),
metadata={
"default_mode": self._default_mode.value,
"cache_optimizer_enabled": enable_cache_optimizer,
"semantic_cache_enabled": enable_semantic_cache,
},
)
def _get_tokenizer(self, model: str) -> Tokenizer:
"""Get tokenizer for model using provider."""
@ -391,6 +406,18 @@ class HeadroomClient:
timestamp = datetime.now(timezone.utc).replace(tzinfo=None)
mode = HeadroomMode(headroom_mode) if headroom_mode else self._default_mode
input_event = self._pipeline_extensions.emit(
PipelineStage.INPUT_RECEIVED,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
messages=messages,
metadata={"api_style": api_style, "stream": stream, "mode": mode.value},
)
if input_event.messages is not None:
messages = input_event.messages
tokenizer = self._get_tokenizer(model)
# Analyze original messages
@ -433,6 +460,23 @@ class HeadroomClient:
tokens_after = result.tokens_after
transforms_applied = result.transforms_applied
routing_markers = summarize_routing_markers(transforms_applied)
if routing_markers:
routed_event = self._pipeline_extensions.emit(
PipelineStage.INPUT_ROUTED,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
messages=optimized_messages,
metadata={
"routing_markers": routing_markers,
"transforms_applied": transforms_applied,
},
)
if routed_event.messages is not None:
optimized_messages = routed_event.messages
# Apply provider-specific cache optimization
if self._cache_optimizer is not None or self._semantic_cache_layer is not None:
cache_context = OptimizationContext(
@ -481,6 +525,23 @@ class HeadroomClient:
f"cache_optimizer:{t}" for t in (cache_result.transforms_applied or [])
)
compressed_event = self._pipeline_extensions.emit(
PipelineStage.INPUT_COMPRESSED,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
messages=optimized_messages,
metadata={
"tokens_before": tokens_before,
"tokens_after": tokens_after,
"transforms_applied": transforms_applied,
},
)
if compressed_event.messages is not None:
optimized_messages = compressed_event.messages
tokens_after = tokenizer.count_messages(optimized_messages)
# Recalculate prefix hash after optimization
stable_prefix_hash = compute_prefix_hash(optimized_messages)
else:
@ -489,6 +550,20 @@ class HeadroomClient:
tokens_after = tokens_before
transforms_applied = []
presend_event = self._pipeline_extensions.emit(
PipelineStage.PRE_SEND,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
messages=optimized_messages,
metadata={"api_style": api_style, "stream": stream},
)
if presend_event.messages is not None:
optimized_messages = presend_event.messages
tokens_after = tokenizer.count_messages(optimized_messages)
stable_prefix_hash = compute_prefix_hash(optimized_messages)
# Create metrics
metrics = RequestMetrics(
request_id=request_id,
@ -530,7 +605,7 @@ class HeadroomClient:
# Call underlying client based on API style
try:
if api_style == "anthropic":
return self._call_anthropic(
response = self._call_anthropic(
model=model,
messages=optimized_messages,
stream=stream,
@ -538,7 +613,7 @@ class HeadroomClient:
**kwargs,
)
else:
return self._call_openai(
response = self._call_openai(
model=model,
messages=optimized_messages,
stream=stream,
@ -546,6 +621,27 @@ class HeadroomClient:
**kwargs,
)
self._pipeline_extensions.emit(
PipelineStage.POST_SEND,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
messages=optimized_messages,
response=response,
metadata={"api_style": api_style, "stream": stream},
)
self._pipeline_extensions.emit(
PipelineStage.RESPONSE_RECEIVED,
operation="sdk.request",
request_id=request_id,
provider=self._provider.name.lower(),
model=model,
response=response,
metadata={"api_style": api_style, "stream": stream},
)
return response
except Exception as e:
metrics.error = str(e)
self._storage.save(metrics)

View file

@ -636,6 +636,10 @@ class HeadroomConfig:
# Debugging - opt-in diff artifact generation
generate_diff_artifact: bool = False # Enable to get detailed transform diffs
# Canonical pipeline lifecycle extensions
pipeline_extensions: list[Any] = field(default_factory=list)
discover_pipeline_extensions: bool = True
def get_context_limit(self, model: str) -> int | None:
"""
Get context limit for a model from user overrides.

View file

@ -1,4 +1,4 @@
"""Compression Hooks — extension points for customizing compression behavior.
"""Compression hooks and pipeline lifecycle events.
Three hooks at well-defined pipeline stages:
@ -6,6 +6,10 @@ Three hooks at well-defined pipeline stages:
2. compute_biases: set per-message compression aggressiveness (position-aware, phase-aware)
3. post_compress: observe results after compression (learning, analytics, logging)
The canonical pipeline also emits lifecycle events through ``on_pipeline_event``.
That gives extensions one stable contract across SDK, ``compress()``, and proxy
request flow without replacing the existing compression hooks.
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).
@ -30,6 +34,8 @@ from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from .pipeline import PipelineEvent
@dataclass
class CompressContext:
@ -135,3 +141,11 @@ class CompressionHooks:
event: Full compression event with before/after metrics.
"""
pass
def on_pipeline_event(self, event: PipelineEvent) -> PipelineEvent | None:
"""Observe canonical pipeline lifecycle events.
Override when the integration needs stable lifecycle notifications beyond
the three legacy compression-specific hooks.
"""
return None

178
headroom/pipeline.py Normal file
View file

@ -0,0 +1,178 @@
"""Canonical Headroom pipeline lifecycle and extension contracts."""
from __future__ import annotations
import importlib.metadata
import logging
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Protocol
log = logging.getLogger(__name__)
ENTRY_POINT_GROUP = "headroom.pipeline_extension"
class PipelineStage(str, Enum):
"""Stable lifecycle stages for the canonical Headroom pipeline."""
SETUP = "setup"
PRE_START = "pre_start"
POST_START = "post_start"
INPUT_RECEIVED = "input_received"
INPUT_CACHED = "input_cached"
INPUT_ROUTED = "input_routed"
INPUT_COMPRESSED = "input_compressed"
INPUT_REMEMBERED = "input_remembered"
PRE_SEND = "pre_send"
POST_SEND = "post_send"
RESPONSE_RECEIVED = "response_received"
CANONICAL_PIPELINE_STAGES: tuple[PipelineStage, ...] = (
PipelineStage.SETUP,
PipelineStage.PRE_START,
PipelineStage.POST_START,
PipelineStage.INPUT_RECEIVED,
PipelineStage.INPUT_CACHED,
PipelineStage.INPUT_ROUTED,
PipelineStage.INPUT_COMPRESSED,
PipelineStage.INPUT_REMEMBERED,
PipelineStage.PRE_SEND,
PipelineStage.POST_SEND,
PipelineStage.RESPONSE_RECEIVED,
)
@dataclass
class PipelineEvent:
"""Event emitted at a canonical pipeline stage.
Extensions may mutate ``messages``, ``tools``, ``headers``, or ``metadata`` in
place, or return a replacement ``PipelineEvent`` from ``on_pipeline_event``.
"""
stage: PipelineStage
operation: str
request_id: str = ""
provider: str = ""
model: str = ""
messages: list[dict[str, Any]] | None = None
tools: list[dict[str, Any]] | None = None
headers: dict[str, str] | None = None
response: Any = None
metadata: dict[str, Any] = field(default_factory=dict)
class PipelineExtension(Protocol):
"""Request lifecycle extension contract for the canonical pipeline."""
def on_pipeline_event(self, event: PipelineEvent) -> PipelineEvent | None:
"""Handle a canonical pipeline event."""
def discover_pipeline_extensions() -> list[PipelineExtension]:
"""Load registered pipeline extensions from Python entry points."""
discovered: list[PipelineExtension] = []
try:
entries = importlib.metadata.entry_points(group=ENTRY_POINT_GROUP)
except Exception as exc: # noqa: BLE001 - importlib metadata varies by runtime
log.debug("pipeline extensions: entry-point enumeration failed: %s", exc)
return discovered
for entry in entries:
try:
extension = entry.load()
except Exception as exc: # noqa: BLE001 - third-party load failures are isolated
log.warning("pipeline extension %r failed to load: %s", entry.name, exc)
continue
if isinstance(extension, type):
try:
extension = extension()
except Exception as exc: # noqa: BLE001
log.warning("pipeline extension %r failed to initialize: %s", entry.name, exc)
continue
discovered.append(extension)
return discovered
def summarize_routing_markers(transforms_applied: list[str]) -> list[str]:
"""Return the routed transform markers emitted by ContentRouter."""
return [item for item in transforms_applied if item.startswith("router:")]
class PipelineExtensionManager:
"""Dispatch canonical pipeline events to configured extensions."""
def __init__(
self,
*,
hooks: Any = None,
extensions: list[Any] | None = None,
discover: bool = True,
) -> None:
resolved: list[Any] = []
if hooks is not None and callable(getattr(hooks, "on_pipeline_event", None)):
resolved.append(hooks)
if extensions:
resolved.extend(extensions)
if discover:
resolved.extend(discover_pipeline_extensions())
self._extensions = resolved
@property
def enabled(self) -> bool:
return bool(self._extensions)
def emit(
self,
stage: PipelineStage,
*,
operation: str,
request_id: str = "",
provider: str = "",
model: str = "",
messages: list[dict[str, Any]] | None = None,
tools: list[dict[str, Any]] | None = None,
headers: dict[str, str] | None = None,
response: Any = None,
metadata: dict[str, Any] | None = None,
) -> PipelineEvent:
"""Emit a canonical lifecycle event and return the final event state."""
event = PipelineEvent(
stage=stage,
operation=operation,
request_id=request_id,
provider=provider,
model=model,
messages=messages,
tools=tools,
headers=headers,
response=response,
metadata=metadata or {},
)
for extension in self._extensions:
handler = getattr(extension, "on_pipeline_event", None)
if not callable(handler):
continue
try:
updated = handler(event)
except Exception as exc: # noqa: BLE001 - preserve hook fail-open behavior
log.warning(
"pipeline extension %r failed during %s: %s",
type(extension).__name__,
stage.value,
exc,
)
continue
if isinstance(updated, PipelineEvent):
event = updated
return event

View file

@ -23,6 +23,8 @@ if TYPE_CHECKING:
import httpx
from headroom.pipeline import PipelineStage, summarize_routing_markers
logger = logging.getLogger("headroom.proxy")
@ -297,6 +299,11 @@ class AnthropicHandlerMixin:
request: Request,
) -> Response | StreamingResponse:
"""Handle Anthropic /v1/messages endpoint."""
if not hasattr(self, "pipeline_extensions"):
from headroom.pipeline import PipelineExtensionManager
self.pipeline_extensions = PipelineExtensionManager(discover=False)
from fastapi import HTTPException
from fastapi.responses import JSONResponse, Response
@ -466,6 +473,22 @@ class AnthropicHandlerMixin:
messages = body.get("messages", [])
with stage_timer.measure("deep_copy"):
original_client_messages = copy.deepcopy(messages)
input_event = self.pipeline_extensions.emit(
PipelineStage.INPUT_RECEIVED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=messages,
tools=body.get("tools"),
metadata={"path": "/v1/messages", "stream": body.get("stream", False)},
)
if input_event.messages is not None:
messages = input_event.messages
with stage_timer.measure("deep_copy"):
original_client_messages = copy.deepcopy(messages)
if input_event.tools is not None:
body["tools"] = input_event.tools
# Validate message array size
if len(messages) > MAX_MESSAGE_ARRAY_LENGTH:
@ -570,6 +593,15 @@ class AnthropicHandlerMixin:
cached = await self.cache.get(messages, model)
if cached:
cache_hit = True
self.pipeline_extensions.emit(
PipelineStage.INPUT_CACHED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=messages,
metadata={"cache_hit": True, "path": "/v1/messages"},
)
optimization_latency = (time.time() - start_time) * 1000
await self.metrics.record_request(
@ -857,6 +889,43 @@ class AnthropicHandlerMixin:
tokens_saved = max(0, original_tokens - optimized_tokens)
optimization_latency = (time.time() - start_time) * 1000
routing_markers = summarize_routing_markers(transforms_applied)
if routing_markers:
routed_event = self.pipeline_extensions.emit(
PipelineStage.INPUT_ROUTED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=optimized_messages,
metadata={
"routing_markers": routing_markers,
"transforms_applied": transforms_applied,
},
)
if routed_event.messages is not None:
optimized_messages = routed_event.messages
optimized_tokens = tokenizer.count_messages(optimized_messages)
tokens_saved = max(0, original_tokens - optimized_tokens)
compressed_event = self.pipeline_extensions.emit(
PipelineStage.INPUT_COMPRESSED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=optimized_messages,
metadata={
"tokens_before": original_tokens,
"tokens_after": optimized_tokens,
"transforms_applied": transforms_applied,
},
)
if compressed_event.messages is not None:
optimized_messages = compressed_event.messages
optimized_tokens = tokenizer.count_messages(optimized_messages)
tokens_saved = max(0, original_tokens - optimized_tokens)
# Hook: post_compress — let hooks observe compression results
if self.config.hooks and tokens_saved > 0:
from headroom.hooks import CompressEvent
@ -1005,6 +1074,8 @@ class AnthropicHandlerMixin:
logger.debug(f"[{request_id}] Traffic learner: {e}")
# Memory: Inject context and tools
memory_context_injected = False
memory_tools_injected = False
if self.memory_handler and memory_user_id:
# Search and inject memory context
if self.memory_handler.config.inject_context:
@ -1040,6 +1111,7 @@ class AnthropicHandlerMixin:
frozen_message_count=frozen_message_count,
)
)
memory_context_injected = True
logger.info(
f"[{request_id}] Memory: Appended {len(memory_context)} chars "
f"to latest non-frozen user turn (prefix cache-safe)"
@ -1048,6 +1120,7 @@ class AnthropicHandlerMixin:
optimized_messages = self._inject_system_context(
optimized_messages, memory_context, body=body
)
memory_context_injected = True
logger.info(
f"[{request_id}] Memory: Injected {len(memory_context)} chars of context"
)
@ -1058,6 +1131,7 @@ class AnthropicHandlerMixin:
if self.memory_handler.config.inject_tools:
tools, mem_tools_injected = self.memory_handler.inject_tools(tools, "anthropic")
if mem_tools_injected:
memory_tools_injected = True
tool_names = [
t.get("name") or t.get("type", "")
for t in tools
@ -1080,6 +1154,28 @@ class AnthropicHandlerMixin:
f"[{request_id}] Memory: Added beta header: {key}={headers[key]}"
)
if memory_context_injected or memory_tools_injected:
remembered_event = self.pipeline_extensions.emit(
PipelineStage.INPUT_REMEMBERED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=optimized_messages,
tools=tools,
headers=headers,
metadata={
"memory_context_injected": memory_context_injected,
"memory_tools_injected": memory_tools_injected,
},
)
if remembered_event.messages is not None:
optimized_messages = remembered_event.messages
if remembered_event.tools is not None:
tools = remembered_event.tools
if remembered_event.headers is not None:
headers = remembered_event.headers
# Query Echo: disabled — hurts prefix caching in long conversations.
# The echo changes every turn, invalidating the cached prefix.
# To re-enable, uncomment and set query_echo_enabled on ProxyConfig.
@ -1090,6 +1186,28 @@ class AnthropicHandlerMixin:
tools = self._sort_tools_deterministically(tools)
body["tools"] = tools
presend_event = self.pipeline_extensions.emit(
PipelineStage.PRE_SEND,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=optimized_messages,
tools=tools,
headers=headers,
metadata={"path": "/v1/messages", "stream": stream},
)
if presend_event.messages is not None:
optimized_messages = presend_event.messages
body["messages"] = optimized_messages
if presend_event.tools is not None:
tools = self._sort_tools_deterministically(presend_event.tools)
body["tools"] = tools
if presend_event.headers is not None:
headers = presend_event.headers
optimized_tokens = tokenizer.count_messages(body["messages"])
tokens_saved = max(0, original_tokens - optimized_tokens)
# Unit 2: mark end of pre-upstream phase. Everything after this
# point is upstream I/O or post-response bookkeeping.
stage_timer.record(
@ -1102,6 +1220,16 @@ class AnthropicHandlerMixin:
# Route through Bedrock backend
try:
if stream:
self.pipeline_extensions.emit(
PipelineStage.POST_SEND,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=body["messages"],
tools=tools,
metadata={"path": "/v1/messages", "stream": True},
)
await _finalize_pre_upstream()
return await self._stream_response_bedrock(
body,
@ -1122,6 +1250,34 @@ class AnthropicHandlerMixin:
backend_response = await self.anthropic_backend.send_message(
body, headers
)
self.pipeline_extensions.emit(
PipelineStage.POST_SEND,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=body["messages"],
tools=tools,
response=backend_response.body,
metadata={
"path": "/v1/messages",
"stream": False,
"status_code": backend_response.status_code,
},
)
self.pipeline_extensions.emit(
PipelineStage.RESPONSE_RECEIVED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
response=backend_response.body,
metadata={
"path": "/v1/messages",
"stream": False,
"status_code": backend_response.status_code,
},
)
# Non-stream: first-byte and connect are effectively
# the same horizon — ``send_message`` awaits until
# the response body is fully buffered.
@ -1210,6 +1366,16 @@ class AnthropicHandlerMixin:
try:
if stream:
self.pipeline_extensions.emit(
PipelineStage.POST_SEND,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=body["messages"],
tools=tools,
metadata={"path": "/v1/messages", "stream": True},
)
await _finalize_pre_upstream()
return await self._stream_response(
url,
@ -1232,6 +1398,34 @@ class AnthropicHandlerMixin:
else:
async with stage_timer.measure("upstream_connect"):
response = await self._retry_request("POST", url, headers, body)
self.pipeline_extensions.emit(
PipelineStage.POST_SEND,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
messages=body["messages"],
tools=tools,
response=response,
metadata={
"path": "/v1/messages",
"stream": False,
"status_code": response.status_code,
},
)
self.pipeline_extensions.emit(
PipelineStage.RESPONSE_RECEIVED,
operation="proxy.request",
request_id=request_id,
provider="anthropic",
model=model,
response=response,
metadata={
"path": "/v1/messages",
"stream": False,
"status_code": response.status_code,
},
)
if (
"upstream_first_byte" not in stage_timer
and "upstream_connect" in stage_timer

File diff suppressed because it is too large Load diff

View file

@ -211,6 +211,8 @@ class ProxyConfig:
# Compression Hooks
hooks: Any = None
pipeline_extensions: list[Any] = field(default_factory=list)
discover_pipeline_extensions: bool = True
# Subscription Window Tracking (Anthropic OAuth accounts)
subscription_tracking_enabled: bool = True

View file

@ -87,6 +87,7 @@ from headroom.observability import (
shutdown_headroom_tracing,
shutdown_otel_metrics,
)
from headroom.pipeline import PipelineExtensionManager, PipelineStage
from headroom.providers.anthropic import AnthropicProvider
from headroom.providers.openai import OpenAIProvider
@ -211,6 +212,11 @@ class HeadroomProxy(
def __init__(self, config: ProxyConfig):
self.config = config
self.config.mode = normalize_proxy_mode(self.config.mode)
self.pipeline_extensions = PipelineExtensionManager(
hooks=config.hooks,
extensions=config.pipeline_extensions,
discover=config.discover_pipeline_extensions,
)
# Reset per-instance API targets first so test runs and multiple app instances
# do not leak class-level overrides across each other.
@ -590,6 +596,17 @@ class HeadroomProxy(
else:
self.code_graph_watcher = None
self.pipeline_extensions.emit(
PipelineStage.SETUP,
operation="proxy.setup",
metadata={
"mode": self.config.mode,
"optimize": self.config.optimize,
"backend": self.config.backend,
"memory_enabled": self.config.memory_enabled,
},
)
def _get_compression_cache(self, session_id: str) -> CompressionCache:
"""Get or create a CompressionCache for a session."""
if session_id not in self._compression_caches:
@ -645,6 +662,11 @@ class HeadroomProxy(
async def startup(self):
"""Initialize async resources."""
self.pipeline_extensions.emit(
PipelineStage.PRE_START,
operation="proxy.startup",
metadata={"port": self.config.port, "host": self.config.host},
)
self.http_client = httpx.AsyncClient(
timeout=httpx.Timeout(
connect=self.config.connect_timeout_seconds,
@ -863,6 +885,16 @@ class HeadroomProxy(
else:
logger.info("Anonymous telemetry: DISABLED")
self.pipeline_extensions.emit(
PipelineStage.POST_START,
operation="proxy.startup",
metadata={
"port": self.config.port,
"host": self.config.host,
"warmup": self.warmup.to_dict(),
},
)
async def shutdown(self):
"""Cleanup async resources."""
if self.http_client:

View file

@ -0,0 +1,187 @@
from __future__ import annotations
import importlib
from types import SimpleNamespace
from typing import Any
from headroom.client import HeadroomClient
from headroom.compress import compress
from headroom.config import HeadroomConfig, HeadroomMode, TransformResult
from headroom.hooks import CompressionHooks
from headroom.pipeline import (
CANONICAL_PIPELINE_STAGES,
PipelineExtensionManager,
PipelineStage,
summarize_routing_markers,
)
from headroom.providers.base import Provider, TokenCounter
class RecordingExtension:
def __init__(self) -> None:
self.stages: list[PipelineStage] = []
def on_pipeline_event(self, event):
self.stages.append(event.stage)
return None
class MutatingExtension:
def on_pipeline_event(self, event):
if event.stage == PipelineStage.INPUT_RECEIVED:
event.messages = [{"role": "user", "content": "mutated"}]
return event
class RecordingHooks(CompressionHooks):
def __init__(self) -> None:
self.stages: list[PipelineStage] = []
self.post_event = None
def pre_compress(self, messages, ctx):
return messages
def compute_biases(self, messages, ctx):
return {}
def post_compress(self, event):
self.post_event = event
def on_pipeline_event(self, event):
self.stages.append(event.stage)
return None
class StubPipeline:
def apply(self, messages, model, **kwargs):
return TransformResult(
messages=messages,
tokens_before=20,
tokens_after=8,
transforms_applied=["router:text:kompress", "kompress:user:0.40"],
)
def _get_tokenizer(self, model):
return StubTokenCounter()
class StubTokenCounter(TokenCounter):
def count_text(self, text: str) -> int:
return len(text.split())
def count_message(self, message: dict[str, Any]) -> int:
content = message.get("content", "")
if isinstance(content, str):
return len(content.split())
return 1
def count_messages(self, messages: list[dict[str, Any]]) -> int:
return sum(self.count_message(message) for message in messages)
class StubProvider(Provider):
@property
def name(self) -> str:
return "openai"
def get_token_counter(self, model: str) -> TokenCounter:
return StubTokenCounter()
def get_context_limit(self, model: str) -> int:
return 128000
def supports_model(self, model: str) -> bool:
return True
class DummyCompletions:
def __init__(self) -> None:
self.calls: list[dict[str, Any]] = []
def create(self, **kwargs: Any) -> dict[str, Any]:
self.calls.append(kwargs)
return {"id": "resp_123", "messages": kwargs["messages"]}
class DummyOriginalClient:
def __init__(self) -> None:
self.chat = SimpleNamespace(completions=DummyCompletions())
def test_pipeline_extension_manager_uses_canonical_stage_contract():
recorder = RecordingExtension()
manager = PipelineExtensionManager(
extensions=[recorder, MutatingExtension()],
discover=False,
)
event = manager.emit(
PipelineStage.INPUT_RECEIVED,
operation="test",
model="gpt-4o",
messages=[{"role": "user", "content": "hello"}],
)
assert list(CANONICAL_PIPELINE_STAGES)[0] is PipelineStage.SETUP
assert summarize_routing_markers(["router:text:kompress", "smart:kept=3"]) == [
"router:text:kompress"
]
assert recorder.stages == [PipelineStage.INPUT_RECEIVED]
assert event.messages == [{"role": "user", "content": "mutated"}]
def test_compress_emits_canonical_pipeline_events(monkeypatch):
hooks = RecordingHooks()
compress_module = importlib.import_module("headroom.compress")
monkeypatch.setattr(compress_module, "_get_pipeline", lambda: StubPipeline())
result = compress(
[{"role": "user", "content": "hello world"}],
model="gpt-4o",
hooks=hooks,
)
assert result.tokens_before == 20
assert result.tokens_after == 2
assert hooks.post_event is not None
assert hooks.post_event.tokens_saved == 18
assert hooks.stages == [
PipelineStage.INPUT_RECEIVED,
PipelineStage.INPUT_ROUTED,
PipelineStage.INPUT_COMPRESSED,
]
def test_headroom_client_emits_canonical_pipeline_events(tmp_path):
recorder = RecordingExtension()
original = DummyOriginalClient()
config = HeadroomConfig(
store_url=f"jsonl://{tmp_path / 'headroom.jsonl'}",
default_mode=HeadroomMode.OPTIMIZE,
pipeline_extensions=[recorder],
discover_pipeline_extensions=False,
)
client = HeadroomClient(
original_client=original,
provider=StubProvider(),
store_url=f"jsonl://{tmp_path / 'headroom-client.jsonl'}",
enable_cache_optimizer=False,
config=config,
)
client._pipeline = StubPipeline()
response = client.chat.completions.create(
model="gpt-4o",
messages=[{"role": "user", "content": "hello world"}],
)
assert response["id"] == "resp_123"
assert recorder.stages == [
PipelineStage.SETUP,
PipelineStage.INPUT_RECEIVED,
PipelineStage.INPUT_ROUTED,
PipelineStage.INPUT_COMPRESSED,
PipelineStage.PRE_SEND,
PipelineStage.POST_SEND,
PipelineStage.RESPONSE_RECEIVED,
]