mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
feat: introduce canonical pipeline lifecycle contract
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
parent
5738339524
commit
cd9c2e1d01
11 changed files with 3477 additions and 2561 deletions
11
README.md
11
README.md
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
178
headroom/pipeline.py
Normal 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
|
||||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
187
tests/test_canonical_pipeline.py
Normal file
187
tests/test_canonical_pipeline.py
Normal 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,
|
||||
]
|
||||
Loading…
Add table
Add a link
Reference in a new issue