mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
## Description Adds an optional, configuration driven model router (closes #1706). With `HEADROOM_MODEL_ROUTER_ENABLED` set, ordered rules in `HEADROOM_MODEL_ROUTES` rewrite the upstream model by estimated input size and tool presence, complementary to content compression, for example sending small, tool-free requests to a cheaper model. First matching rule wins and every decision is logged with a reason. Off by default so behavior is unchanged, skipped under `x-headroom-bypass`/passthrough, and wired on the Anthropic `/v1/messages` path. Malformed rules fail open, so a bad rule is skipped rather than silently widened. Closes #1706 ## Type of Change - [ ] Bug fix (non-breaking change that fixes an issue) - [x] New feature (non-breaking change that adds functionality) - [ ] Breaking change (fix or feature that would cause existing functionality to change) - [x] Documentation update - [ ] Performance improvement - [ ] Code refactoring (no functional changes) ## Changes Made - `headroom/proxy/model_router.py`: new `ModelRouter` component (ordered rules, first-match decision with reason, fail-open env parsing, tokenizer-free input estimate). - `headroom/proxy/models.py` + `headroom/proxy/server.py`: `ProxyConfig.model_router` field, env loader (`HEADROOM_MODEL_ROUTER_ENABLED` / `HEADROOM_MODEL_ROUTES`), and proxy wiring. - `headroom/proxy/handlers/anthropic.py`: apply routing on `/v1/messages` after the bypass gate, tracked as a body mutation. - Tests, docs (`configuration.mdx`), and a CHANGELOG entry. ## Testing - [x] Unit tests pass (`pytest`) - [x] Linting passes (`ruff check .`) - [x] Type checking passes (`mypy headroom`) - [x] New tests added for new functionality - [x] Manual testing performed ### Test Output ```text $ pytest -q tests/test_proxy/test_model_router.py tests/test_proxy/test_model_router_wiring.py 36 passed, 1 warning $ ruff check . All checks passed! $ mypy headroom --ignore-missing-imports Success: no issues found in 477 source files ``` ## Real Behavior Proof - Environment: local, macOS, Python 3.12, headroom `.venv`, upstream mocked (no live provider call). - Exact command / steps: enable the router via `ProxyConfig(model_router=...)`, POST `/v1/messages` through `TestClient` with a rule routing low-risk requests to a cheaper model; repeat with header `x-headroom-bypass: true`. - Observed result: the forwarded upstream body model is rewritten from `claude-sonnet-4-6` to `claude-haiku-4-5` when the router is enabled, and is left unchanged under bypass (see `tests/test_proxy/test_model_router_wiring.py`). - Not tested: the OpenAI and Gemini handler paths (this PR wires the Anthropic path only); no live provider request (upstream is mocked). ## Review Readiness - [x] I have performed a self-review - [x] This PR is ready for human review ## Checklist - [x] My code follows the project's style guidelines - [x] I have performed a self-review of my code - [x] I have commented my code, particularly in hard-to-understand areas - [x] I have made corresponding changes to the documentation - [x] My changes generate no new warnings - [x] I have added tests that prove my fix is effective or that my feature works - [x] New and existing unit tests pass locally with my changes - [x] I have updated the CHANGELOG.md if applicable ## Additional Notes Happy to adjust the interface or scope (for example OpenAI and Gemini parity) if you'd prefer a different shape. --------- Co-authored-by: JerrettDavis <mxjerrett@gmail.com> Co-authored-by: reneleonhardt <65483435+reneleonhardt@users.noreply.github.com>
This commit is contained in:
parent
aa4515cf7a
commit
57e8dcb425
9 changed files with 904 additions and 0 deletions
|
|
@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||
## Unreleased
|
||||
|
||||
### Features
|
||||
- **proxy:** opt-in cost-aware model routing ([#1706](https://github.com/headroomlabs-ai/headroom/issues/1706)). Set `HEADROOM_MODEL_ROUTER_ENABLED=1` and `HEADROOM_MODEL_ROUTES` (a JSON array of ordered rules) to rewrite the upstream model based on estimated input size and tool presence, complementary to content compression, e.g. send small, tool-free requests to a cheaper model. First matching rule wins, and each decision is logged with a reason so routing stays observable. Malformed rules fail open (the rule is skipped, never silently widened). Disabled by default so behavior is unchanged, skipped under `x-headroom-bypass`/passthrough, and currently applied on the Anthropic `/v1/messages` path.
|
||||
- **install:** `headroom install apply` now accepts `--code-aware/--no-code-aware`, `--intercept-tool-results`, `--protect-tool-results`, and `--bedrock-profile`, mirroring the equivalent flags already on `headroom proxy`. Previously the only way to run a persistent deployment with these settings was to hand-edit `manifest.json` after the fact, which silently reverts on the next `install apply`.
|
||||
- **install:** `headroom install apply --env KEY=VALUE` (repeatable) passes environment variables into supervised runners (macOS launchd, Linux systemd/cron, Windows services/tasks). These runners previously started with a bare environment and did not inherit the interactive shell's exports — e.g. a custom `HEADROOM_WORKSPACE_DIR` never reached the supervised process, so `headroom install agent run` looked for its manifest in the wrong location and failed outright even though `install apply` itself succeeded. `--env` values are merged into `DeploymentManifest.base_env` last, so they can override auto-derived defaults, and are threaded into the generated `run-headroom.sh`/`ensure-headroom.sh` (and Windows equivalents) as `export`/`$env:` lines before the `exec`.
|
||||
|
||||
|
|
|
|||
|
|
@ -312,9 +312,41 @@ headroom proxy --learn --min-evidence 3
|
|||
| `HEADROOM_REQUEST_TIMEOUT` | Request timeout in seconds | `300` |
|
||||
| `HEADROOM_BETA_HEADER_STICKY` | Controls per-session `anthropic-beta` / `OpenAI-Beta` re-echo. `enabled` (default): the proxy unions beta tokens across turns within a session — if the client sends a token in turn N and omits it in turn N+1, the proxy re-injects it to preserve prefix-cache stability. `disabled`: the client's value is forwarded verbatim with no accumulation. Any other value raises at request time. See [Session Beta Header Tracking](/docs/configuration#session-beta-header-tracking). | `enabled` |
|
||||
| `HEADROOM_BETA_TRACKER_MAX_SESSIONS` | LRU capacity of the in-memory session beta tracker. Once full, the oldest session entry is evicted. | `1000` |
|
||||
| `HEADROOM_MODEL_ROUTER_ENABLED` | Enable cost-aware model routing. `1`/`true`/`yes`/`on`/`enabled` turns it on and requires `HEADROOM_MODEL_ROUTES`. See [Cost-aware model routing](/docs/configuration#cost-aware-model-routing). | `off` |
|
||||
| `HEADROOM_MODEL_ROUTES` | JSON array of ordered routing rules for cost-aware model routing (schema below). | -- |
|
||||
|
||||
For provider-only proxying, prefer `HEADROOM_HTTP_PROXY` over process-wide variables such as `HTTP_PROXY`, `HTTPS_PROXY`, `ALL_PROXY`, or `NO_PROXY`. HTTPX reads those global variables, but Headroom also passes them through to tool executions.
|
||||
|
||||
### Cost-aware model routing
|
||||
|
||||
Complementary to content compression, Headroom can rewrite the upstream model per request to stretch quota and control spend, for example by sending small, tool-free requests to a cheaper model. Routing is opt-in and disabled by default, so behavior is unchanged unless you configure it.
|
||||
|
||||
Enable it with `HEADROOM_MODEL_ROUTER_ENABLED=1` and declare ordered rules in `HEADROOM_MODEL_ROUTES` (a JSON array). The router evaluates rules top to bottom and the first rule whose conditions all match wins; every decision is logged with a reason so routing stays observable. Each rule object supports:
|
||||
|
||||
| Field | Type | Meaning |
|
||||
|-------|------|---------|
|
||||
| `to_model` | string (required) | Model to route to when the rule matches. |
|
||||
| `max_input_tokens` | integer | Match only when the estimated input size is at or below this. |
|
||||
| `min_input_tokens` | integer | Match only when the estimated input size is at or above this. |
|
||||
| `require_no_tools` | boolean | Match only when the request declares no tools (a proxy for low-risk work). |
|
||||
| `from_models` | list of strings | Restrict the rule to these source models. Omit for any source model. |
|
||||
| `name` | string | Label surfaced in the decision log. |
|
||||
|
||||
```bash
|
||||
export HEADROOM_MODEL_ROUTER_ENABLED=1
|
||||
export HEADROOM_MODEL_ROUTES='[
|
||||
{"name": "small-no-tools", "max_input_tokens": 4000, "require_no_tools": true,
|
||||
"from_models": ["claude-sonnet-4-6"], "to_model": "claude-haiku-4-5"}
|
||||
]'
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- Input size is a fast, tokenizer-free estimate over the messages, tools, and top-level `system` prompt, meant for tier selection rather than exact accounting.
|
||||
- A malformed rule fails open: it is skipped (never silently widened), and the rest of the rules still apply.
|
||||
- Routing is skipped for byte-faithful passthrough requests (`x-headroom-bypass: true` or `x-headroom-mode: passthrough`), so those are never model-rewritten.
|
||||
- Routing currently applies on the Anthropic `/v1/messages` path.
|
||||
|
||||
### Session Beta Header Tracking
|
||||
|
||||
When running as a proxy, Headroom maintains a per-session union of `anthropic-beta` (and `OpenAI-Beta`) tokens via `SessionBetaTracker`. The session key is derived from the `x-headroom-session-id` header if present, otherwise from `md5(model + system_prompt[:500])[:16]` — stable across turns of the same conversation.
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ from headroom.proxy.helpers import extract_tags
|
|||
from headroom.proxy.image_isolation import run_image_compression_isolated
|
||||
from headroom.proxy.memory_decision import MemoryDecision
|
||||
from headroom.proxy.memory_query import MemoryQuery
|
||||
from headroom.proxy.model_router import estimate_input_tokens
|
||||
from headroom.proxy.outcome import RequestOutcome
|
||||
|
||||
logger = logging.getLogger("headroom.proxy")
|
||||
|
|
@ -514,6 +515,37 @@ class AnthropicHandlerMixin:
|
|||
"content": copy.deepcopy(resp_json.get("content", "")),
|
||||
}
|
||||
|
||||
def _maybe_route_model(
|
||||
self,
|
||||
model: str,
|
||||
messages: object,
|
||||
body: dict[str, Any],
|
||||
body_mutation_tracker: Any,
|
||||
bypass: bool,
|
||||
) -> str:
|
||||
"""Apply cost-aware model routing (#1706), returning the model to forward.
|
||||
|
||||
Fails closed to disabled when no ``model_router`` is present: alternate
|
||||
mixin hosts and test doubles that do not run ``HeadroomProxy.__init__``
|
||||
never set the attribute, and reading it unconditionally would crash them
|
||||
even when routing is off. Also skipped under bypass/passthrough so a
|
||||
byte-faithful request is never model-rewritten.
|
||||
"""
|
||||
router = getattr(self, "model_router", None)
|
||||
if router is None or not router.enabled or bypass:
|
||||
return model
|
||||
decision = router.select(
|
||||
model=model,
|
||||
input_tokens=estimate_input_tokens(messages, body.get("tools"), body.get("system")),
|
||||
has_tools=bool(body.get("tools")),
|
||||
)
|
||||
logger.info("model routing decision: %s", decision.reason)
|
||||
if not decision.changed:
|
||||
return model
|
||||
body["model"] = decision.routed_model
|
||||
body_mutation_tracker.mark_mutated("model_router")
|
||||
return decision.routed_model
|
||||
|
||||
async def handle_anthropic_messages(
|
||||
self,
|
||||
request: Request,
|
||||
|
|
@ -770,6 +802,15 @@ class AnthropicHandlerMixin:
|
|||
if _bypass:
|
||||
logger.info(f"[{request_id}] Bypass: skipping compression (header)")
|
||||
|
||||
# Cost-aware model routing (#1706). Opt-in and disabled by default;
|
||||
# fail-closed and bypass handling live in the helper. A model override
|
||||
# comes from a provider URL (for example Vertex rawPredict), where
|
||||
# rewriting body["model"] would not change the upstream model.
|
||||
if model_override is None:
|
||||
model = self._maybe_route_model(
|
||||
model, messages, body, body_mutation_tracker, _bypass
|
||||
)
|
||||
|
||||
# NOTE: Upstream temporarily disabled broad image compression due to
|
||||
# token-counting inaccuracies. We only compress the latest non-frozen
|
||||
# user turn later in this handler to preserve Anthropic prefix caching.
|
||||
|
|
|
|||
289
headroom/proxy/model_router.py
Normal file
289
headroom/proxy/model_router.py
Normal file
|
|
@ -0,0 +1,289 @@
|
|||
"""Cost-aware model routing (issue #1706).
|
||||
|
||||
Complementary to content compression: route a request to a cheaper (or more
|
||||
capable) model based on request characteristics, so callers can stretch quota
|
||||
and control spend without changing their client.
|
||||
|
||||
This is an opt-in, config-driven **mechanism**, not an opinionated built-in
|
||||
policy. The operator declares an ordered list of rules mapping request
|
||||
characteristics to a target model; the router picks the first rule whose
|
||||
conditions all match and records the decision, with a human-readable reason,
|
||||
so routing is observable and never a black box. When disabled (the default) or
|
||||
when no rule matches, the original model is returned unchanged, so behavior is
|
||||
identical to today.
|
||||
|
||||
The router is a pure component: no I/O, no global state, fully unit-testable.
|
||||
Wiring into the request path (reading the decision, rewriting the outgoing
|
||||
model, logging, and metrics) lives in the handlers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelRoute:
|
||||
"""One ordered routing rule.
|
||||
|
||||
A rule matches when every condition that is set is satisfied (logical AND).
|
||||
Conditions left as ``None``/empty are ignored. Rules are evaluated in order
|
||||
and the first match wins.
|
||||
"""
|
||||
|
||||
to_model: str
|
||||
"""Model to route to when this rule matches."""
|
||||
|
||||
max_input_tokens: int | None = None
|
||||
"""Match only when estimated input tokens are <= this (cheap for small requests)."""
|
||||
|
||||
min_input_tokens: int | None = None
|
||||
"""Match only when estimated input tokens are >= this."""
|
||||
|
||||
require_no_tools: bool = False
|
||||
"""Match only when the request declares no tools (a proxy for low-risk work)."""
|
||||
|
||||
from_models: tuple[str, ...] = ()
|
||||
"""Restrict this rule to these source models. Empty = any source model."""
|
||||
|
||||
name: str = ""
|
||||
"""Human-readable label surfaced in decision logs."""
|
||||
|
||||
def matches(self, *, model: str, input_tokens: int, has_tools: bool) -> bool:
|
||||
"""True when every set condition is satisfied for this request."""
|
||||
if self.from_models and model not in self.from_models:
|
||||
return False
|
||||
if self.require_no_tools and has_tools:
|
||||
return False
|
||||
if self.max_input_tokens is not None and input_tokens > self.max_input_tokens:
|
||||
return False
|
||||
if self.min_input_tokens is not None and input_tokens < self.min_input_tokens:
|
||||
return False
|
||||
# A rule whose ``to_model`` equals the current model still MATCHES (strict
|
||||
# first-match-wins): it is a no-op (``changed`` is False) that short-circuits
|
||||
# later rules, which lets an operator write an explicit exemption rule.
|
||||
return True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelRouterConfig:
|
||||
"""Configuration for :class:`ModelRouter`. Disabled by default."""
|
||||
|
||||
enabled: bool = False
|
||||
routes: tuple[ModelRoute, ...] = ()
|
||||
|
||||
@classmethod
|
||||
def from_env(cls, enabled_raw: str | None, routes_raw: str | None) -> ModelRouterConfig:
|
||||
"""Build config from env-style strings, failing open to disabled.
|
||||
|
||||
``routes_raw`` is a JSON array of rule objects, e.g.::
|
||||
|
||||
[{"name": "small->mini", "max_input_tokens": 4000,
|
||||
"require_no_tools": true, "to_model": "gpt-5.4-mini"}]
|
||||
|
||||
A malformed value logs a warning and disables routing rather than
|
||||
raising, so a bad config can never take the proxy down.
|
||||
"""
|
||||
enabled = _truthy(enabled_raw)
|
||||
routes = _parse_routes(routes_raw)
|
||||
if enabled and not routes:
|
||||
logger.warning("model router enabled but no valid routes configured; disabling")
|
||||
return cls(enabled=False, routes=())
|
||||
return cls(enabled=enabled and bool(routes), routes=routes)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelDecision:
|
||||
"""The outcome of a routing evaluation for one request."""
|
||||
|
||||
original_model: str
|
||||
routed_model: str
|
||||
matched: bool
|
||||
reason: str
|
||||
rule_name: str = ""
|
||||
|
||||
@property
|
||||
def changed(self) -> bool:
|
||||
"""True when the caller should rewrite the outgoing model."""
|
||||
return self.matched and self.routed_model != self.original_model
|
||||
|
||||
|
||||
class ModelRouter:
|
||||
"""Selects an outgoing model from ordered, config-driven rules."""
|
||||
|
||||
def __init__(self, config: ModelRouterConfig | None) -> None:
|
||||
self._config = config or ModelRouterConfig()
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._config.enabled and bool(self._config.routes)
|
||||
|
||||
def select(self, *, model: str, input_tokens: int, has_tools: bool) -> ModelDecision:
|
||||
"""Return the routing decision for a request.
|
||||
|
||||
Never raises: on a disabled router or no matching rule, returns a
|
||||
non-matching decision that leaves the original model in place.
|
||||
"""
|
||||
if not self.enabled:
|
||||
return ModelDecision(model, model, matched=False, reason="router disabled")
|
||||
if not isinstance(model, str) or not model:
|
||||
return ModelDecision(model, model, matched=False, reason="no source model")
|
||||
|
||||
for route in self._config.routes:
|
||||
if route.matches(model=model, input_tokens=input_tokens, has_tools=has_tools):
|
||||
reason = (
|
||||
f"matched rule {route.name or route.to_model!r}: "
|
||||
f"{model} -> {route.to_model} "
|
||||
f"(input_tokens={input_tokens}, has_tools={has_tools})"
|
||||
)
|
||||
return ModelDecision(
|
||||
original_model=model,
|
||||
routed_model=route.to_model,
|
||||
matched=True,
|
||||
reason=reason,
|
||||
rule_name=route.name,
|
||||
)
|
||||
return ModelDecision(model, model, matched=False, reason="no rule matched")
|
||||
|
||||
|
||||
def estimate_input_tokens(messages: object, tools: object = None, system: object = None) -> int:
|
||||
"""Cheap, tokenizer-free estimate of request input size, for routing only.
|
||||
|
||||
Uses a ~4-chars-per-token heuristic over the serialized message, tool, and
|
||||
system content. ``system`` covers Anthropic's top-level ``system`` field
|
||||
(string or content-block list), which is not part of ``messages`` but can
|
||||
dominate request size, so omitting it would let a large system prompt route
|
||||
as if the request were tiny. This is deliberately approximate: it runs on
|
||||
the hot path purely to pick a route tier, so it must not pay for a real
|
||||
tokenizer. It never raises.
|
||||
"""
|
||||
try:
|
||||
chars = 0
|
||||
if isinstance(messages, list):
|
||||
for msg in messages:
|
||||
chars += (
|
||||
len(str(msg.get("content", ""))) if isinstance(msg, dict) else len(str(msg))
|
||||
)
|
||||
if tools:
|
||||
chars += len(str(tools))
|
||||
if system:
|
||||
chars += len(str(system))
|
||||
return chars // 4
|
||||
except Exception: # noqa: BLE001 — estimation must never break the request path
|
||||
return 0
|
||||
|
||||
|
||||
def _truthy(value: str | None) -> bool:
|
||||
return (value or "").strip().lower() in {"1", "true", "yes", "on", "enable", "enabled"}
|
||||
|
||||
|
||||
def _parse_routes(routes_raw: str | None) -> tuple[ModelRoute, ...]:
|
||||
if not routes_raw or not routes_raw.strip():
|
||||
return ()
|
||||
try:
|
||||
parsed = json.loads(routes_raw)
|
||||
except (ValueError, TypeError) as exc:
|
||||
logger.warning("invalid HEADROOM_MODEL_ROUTES JSON; ignoring: %s", exc)
|
||||
return ()
|
||||
if not isinstance(parsed, list):
|
||||
logger.warning("HEADROOM_MODEL_ROUTES must be a JSON array; ignoring")
|
||||
return ()
|
||||
|
||||
routes: list[ModelRoute] = []
|
||||
for i, entry in enumerate(parsed):
|
||||
route = _route_from_entry(entry, i)
|
||||
if route is not None:
|
||||
routes.append(route)
|
||||
return tuple(routes)
|
||||
|
||||
|
||||
_INVALID = object()
|
||||
"""Sentinel: a route field was present but malformed (fail open, skip the route)."""
|
||||
|
||||
_ALLOWED_ROUTE_KEYS = frozenset(
|
||||
{"to_model", "max_input_tokens", "min_input_tokens", "require_no_tools", "from_models", "name"}
|
||||
)
|
||||
|
||||
|
||||
def _route_from_entry(entry: object, index: int) -> ModelRoute | None:
|
||||
"""Parse one route object, failing open (skip) on any malformed condition.
|
||||
|
||||
A silently-broadened rule (e.g. an unparseable ``max_input_tokens`` treated
|
||||
as "no cap") could route far more traffic than the operator intended, so an
|
||||
invalid condition disables just that rule rather than widening it.
|
||||
"""
|
||||
if not isinstance(entry, dict):
|
||||
logger.warning("model route #%d is not an object; skipping", index)
|
||||
return None
|
||||
unknown_keys = set(entry) - _ALLOWED_ROUTE_KEYS
|
||||
if unknown_keys:
|
||||
# A misspelled condition (e.g. "max_input_token") would otherwise be
|
||||
# ignored, silently widening the rule. Reject unknown keys instead.
|
||||
logger.warning(
|
||||
"model route #%d has unknown key(s) %s; skipping route", index, sorted(unknown_keys)
|
||||
)
|
||||
return None
|
||||
to_model = entry.get("to_model")
|
||||
if not isinstance(to_model, str) or not to_model:
|
||||
logger.warning("model route #%d missing string 'to_model'; skipping", index)
|
||||
return None
|
||||
|
||||
max_tokens = _strict_opt_int(entry, "max_input_tokens", index)
|
||||
min_tokens = _strict_opt_int(entry, "min_input_tokens", index)
|
||||
if max_tokens is _INVALID or min_tokens is _INVALID:
|
||||
return None
|
||||
|
||||
require_no_tools = entry.get("require_no_tools", False)
|
||||
if not isinstance(require_no_tools, bool):
|
||||
logger.warning(
|
||||
"model route #%d 'require_no_tools' must be a boolean; skipping route", index
|
||||
)
|
||||
return None
|
||||
|
||||
from_models_raw = entry.get("from_models", [])
|
||||
if not isinstance(from_models_raw, list) or not all(
|
||||
isinstance(m, str) for m in from_models_raw
|
||||
):
|
||||
logger.warning(
|
||||
"model route #%d 'from_models' must be a list of strings; skipping route", index
|
||||
)
|
||||
return None
|
||||
|
||||
return ModelRoute(
|
||||
to_model=to_model,
|
||||
max_input_tokens=max_tokens, # type: ignore[arg-type]
|
||||
min_input_tokens=min_tokens, # type: ignore[arg-type]
|
||||
require_no_tools=require_no_tools,
|
||||
from_models=tuple(from_models_raw),
|
||||
name=str(entry.get("name", "")),
|
||||
)
|
||||
|
||||
|
||||
def _strict_opt_int(entry: dict, key: str, index: int) -> int | None | object:
|
||||
"""Return the int at ``key``, ``None`` if absent, or ``_INVALID`` if malformed.
|
||||
|
||||
Accepts JSON integers and digit strings; rejects booleans, floats, and
|
||||
non-numeric values so a typo cannot silently remove a token bound.
|
||||
"""
|
||||
if key not in entry or entry[key] is None:
|
||||
return None
|
||||
value = entry[key]
|
||||
if isinstance(value, bool):
|
||||
logger.warning("model route #%d '%s' must be an integer, not a boolean", index, key)
|
||||
return _INVALID
|
||||
if isinstance(value, int):
|
||||
parsed = value
|
||||
else:
|
||||
try:
|
||||
parsed = int(str(value))
|
||||
except (ValueError, TypeError):
|
||||
logger.warning("model route #%d '%s' is not a valid integer", index, key)
|
||||
return _INVALID
|
||||
if parsed < 0:
|
||||
logger.warning("model route #%d '%s' must be non-negative", index, key)
|
||||
return _INVALID
|
||||
return parsed
|
||||
|
|
@ -13,6 +13,7 @@ from typing import Any, Literal
|
|||
|
||||
from headroom.memory import qdrant_env
|
||||
from headroom.providers.registry import ProviderApiOverrides
|
||||
from headroom.proxy.model_router import ModelRouterConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -163,6 +164,11 @@ class ProxyConfig:
|
|||
smart_crusher_with_compaction: bool | None = None
|
||||
keep_last_turns: int = 4
|
||||
|
||||
# Cost-aware model routing (issue #1706). Opt-in and disabled by default;
|
||||
# when configured, an ordered rule set can rewrite the outgoing model based
|
||||
# on request size / tool presence. None keeps behavior unchanged.
|
||||
model_router: ModelRouterConfig | None = None
|
||||
|
||||
# CCR Tool Injection
|
||||
ccr_inject_tool: bool = True
|
||||
ccr_inject_system_instructions: bool = False
|
||||
|
|
|
|||
|
|
@ -149,6 +149,7 @@ from headroom.proxy.loopback_guard import is_loopback_host
|
|||
from headroom.proxy.memory_handler import MemoryConfig, MemoryHandler
|
||||
|
||||
# Data models (extracted to headroom/proxy/models.py for maintainability)
|
||||
from headroom.proxy.model_router import ModelRouter, ModelRouterConfig
|
||||
from headroom.proxy.models import CacheEntry, ProxyConfig, RateLimitState, RequestLog # noqa: F401
|
||||
from headroom.proxy.modes import (
|
||||
PROXY_MODE_CACHE,
|
||||
|
|
@ -708,6 +709,10 @@ class HeadroomProxy(
|
|||
)
|
||||
self.metrics = PrometheusMetrics(cost_tracker=self.cost_tracker, stateless=config.stateless)
|
||||
|
||||
# Cost-aware model routing (issue #1706). Disabled unless configured, so
|
||||
# the default request path is unchanged.
|
||||
self.model_router = ModelRouter(config.model_router)
|
||||
|
||||
# Initialize transforms based on routing mode.
|
||||
#
|
||||
# Phase B PR-B1 retired the IntelligentContextManager / RollingWindow
|
||||
|
|
@ -4693,6 +4698,10 @@ def _proxy_config_from_env() -> ProxyConfig:
|
|||
read_maturation_min_size_bytes=_get_env_int(
|
||||
"HEADROOM_READ_MATURATION_MIN_SIZE_BYTES", 2048
|
||||
),
|
||||
model_router=ModelRouterConfig.from_env(
|
||||
os.environ.get("HEADROOM_MODEL_ROUTER_ENABLED"),
|
||||
os.environ.get("HEADROOM_MODEL_ROUTES"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
252
tests/test_proxy/test_model_router.py
Normal file
252
tests/test_proxy/test_model_router.py
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
"""Tests for cost-aware model routing (issue #1706)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from headroom.proxy.model_router import (
|
||||
ModelDecision,
|
||||
ModelRoute,
|
||||
ModelRouter,
|
||||
ModelRouterConfig,
|
||||
estimate_input_tokens,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelRoute.matches
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_route_matches_on_max_tokens_and_no_tools() -> None:
|
||||
route = ModelRoute(to_model="cheap", max_input_tokens=4000, require_no_tools=True)
|
||||
assert route.matches(model="strong", input_tokens=1000, has_tools=False)
|
||||
# too many tokens
|
||||
assert not route.matches(model="strong", input_tokens=5000, has_tools=False)
|
||||
# tools present
|
||||
assert not route.matches(model="strong", input_tokens=1000, has_tools=True)
|
||||
|
||||
|
||||
def test_route_min_tokens() -> None:
|
||||
route = ModelRoute(to_model="strong", min_input_tokens=10000)
|
||||
assert route.matches(model="cheap", input_tokens=20000, has_tools=True)
|
||||
assert not route.matches(model="cheap", input_tokens=5000, has_tools=True)
|
||||
|
||||
|
||||
def test_route_from_models_restriction() -> None:
|
||||
route = ModelRoute(to_model="cheap", from_models=("gpt-5.5", "gpt-5.4"))
|
||||
assert route.matches(model="gpt-5.5", input_tokens=1, has_tools=False)
|
||||
assert not route.matches(model="claude-sonnet-4-6", input_tokens=1, has_tools=False)
|
||||
|
||||
|
||||
def test_route_matches_even_for_same_model() -> None:
|
||||
# A same-model rule still MATCHES (strict first-match-wins); it is a no-op
|
||||
# that short-circuits later rules, enabling explicit exemption rules.
|
||||
route = ModelRoute(to_model="cheap")
|
||||
assert route.matches(model="cheap", input_tokens=1, has_tools=False)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelRouter.select
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _router(*routes: ModelRoute, enabled: bool = True) -> ModelRouter:
|
||||
return ModelRouter(ModelRouterConfig(enabled=enabled, routes=tuple(routes)))
|
||||
|
||||
|
||||
def test_disabled_router_is_passthrough() -> None:
|
||||
router = _router(ModelRoute(to_model="cheap", max_input_tokens=10_000), enabled=False)
|
||||
d = router.select(model="strong", input_tokens=10, has_tools=False)
|
||||
assert not d.matched and not d.changed
|
||||
assert d.routed_model == "strong"
|
||||
|
||||
|
||||
def test_first_matching_rule_wins() -> None:
|
||||
router = _router(
|
||||
ModelRoute(to_model="nano", max_input_tokens=2000, name="tiny"),
|
||||
ModelRoute(to_model="mini", max_input_tokens=8000, name="small"),
|
||||
)
|
||||
d = router.select(model="gpt-5.5", input_tokens=1500, has_tools=False)
|
||||
assert d.changed and d.routed_model == "nano" and d.rule_name == "tiny"
|
||||
|
||||
d2 = router.select(model="gpt-5.5", input_tokens=5000, has_tools=False)
|
||||
assert d2.changed and d2.routed_model == "mini" and d2.rule_name == "small"
|
||||
|
||||
|
||||
def test_exemption_rule_short_circuits_later_rules() -> None:
|
||||
# An explicit same-model rule wins first and stops a later downgrade rule.
|
||||
router = _router(
|
||||
ModelRoute(to_model="keep", from_models=("keep",), name="exempt"),
|
||||
ModelRoute(to_model="cheap", max_input_tokens=10_000, name="downgrade"),
|
||||
)
|
||||
d = router.select(model="keep", input_tokens=100, has_tools=False)
|
||||
assert d.matched and not d.changed
|
||||
assert d.routed_model == "keep" and d.rule_name == "exempt"
|
||||
|
||||
|
||||
def test_no_rule_matches_is_passthrough() -> None:
|
||||
router = _router(ModelRoute(to_model="mini", max_input_tokens=1000))
|
||||
d = router.select(model="gpt-5.5", input_tokens=50_000, has_tools=True)
|
||||
assert not d.matched and not d.changed and d.routed_model == "gpt-5.5"
|
||||
assert d.reason == "no rule matched"
|
||||
|
||||
|
||||
def test_empty_source_model_is_passthrough() -> None:
|
||||
router = _router(ModelRoute(to_model="mini"))
|
||||
d = router.select(model="", input_tokens=10, has_tools=False)
|
||||
assert not d.matched and d.routed_model == ""
|
||||
|
||||
|
||||
def test_enabled_requires_routes() -> None:
|
||||
assert not ModelRouter(ModelRouterConfig(enabled=True, routes=())).enabled
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelDecision
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_decision_changed_only_when_model_differs() -> None:
|
||||
assert ModelDecision("a", "b", matched=True, reason="x").changed
|
||||
assert not ModelDecision("a", "a", matched=True, reason="x").changed
|
||||
assert not ModelDecision("a", "b", matched=False, reason="x").changed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelRouterConfig.from_env (fail-open parsing)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_from_env_disabled_by_default() -> None:
|
||||
cfg = ModelRouterConfig.from_env(None, None)
|
||||
assert not cfg.enabled and cfg.routes == ()
|
||||
|
||||
|
||||
def test_from_env_parses_routes() -> None:
|
||||
routes = (
|
||||
'[{"name":"small","max_input_tokens":4000,"require_no_tools":true,'
|
||||
'"to_model":"gpt-5.4-mini","from_models":["gpt-5.5"]}]'
|
||||
)
|
||||
cfg = ModelRouterConfig.from_env("true", routes)
|
||||
assert cfg.enabled
|
||||
assert len(cfg.routes) == 1
|
||||
r = cfg.routes[0]
|
||||
assert r.to_model == "gpt-5.4-mini"
|
||||
assert r.max_input_tokens == 4000
|
||||
assert r.require_no_tools is True
|
||||
assert r.from_models == ("gpt-5.5",)
|
||||
|
||||
|
||||
def test_from_env_enabled_but_no_routes_disables() -> None:
|
||||
cfg = ModelRouterConfig.from_env("true", None)
|
||||
assert not cfg.enabled
|
||||
|
||||
|
||||
def test_from_env_malformed_json_fails_open() -> None:
|
||||
cfg = ModelRouterConfig.from_env("true", "{not json")
|
||||
assert not cfg.enabled and cfg.routes == ()
|
||||
|
||||
|
||||
def test_from_env_non_array_json_ignored() -> None:
|
||||
cfg = ModelRouterConfig.from_env("true", '{"to_model":"x"}')
|
||||
assert cfg.routes == ()
|
||||
|
||||
|
||||
def test_from_env_skips_bad_entries_keeps_good() -> None:
|
||||
routes = '[{"no_to_model":true}, {"to_model":"mini","max_input_tokens":"3000"}]'
|
||||
cfg = ModelRouterConfig.from_env("1", routes)
|
||||
assert len(cfg.routes) == 1
|
||||
assert cfg.routes[0].to_model == "mini"
|
||||
# numeric string coerced
|
||||
assert cfg.routes[0].max_input_tokens == 3000
|
||||
|
||||
|
||||
def test_from_env_malformed_int_skips_route() -> None:
|
||||
# A bool or non-numeric token bound must fail open (skip the route), never
|
||||
# silently widen to "no cap".
|
||||
assert (
|
||||
ModelRouterConfig.from_env("yes", '[{"to_model":"m","max_input_tokens":true}]').routes == ()
|
||||
)
|
||||
assert (
|
||||
ModelRouterConfig.from_env("yes", '[{"to_model":"m","min_input_tokens":"abc"}]').routes
|
||||
== ()
|
||||
)
|
||||
|
||||
|
||||
def test_from_env_malformed_require_no_tools_skips_route() -> None:
|
||||
# A string "false" must not be coerced to True.
|
||||
cfg = ModelRouterConfig.from_env("yes", '[{"to_model":"m","require_no_tools":"false"}]')
|
||||
assert cfg.routes == ()
|
||||
|
||||
|
||||
def test_from_env_malformed_from_models_skips_route() -> None:
|
||||
assert (
|
||||
ModelRouterConfig.from_env("yes", '[{"to_model":"m","from_models":"gpt-5.5"}]').routes == ()
|
||||
)
|
||||
assert ModelRouterConfig.from_env("yes", '[{"to_model":"m","from_models":[1,2]}]').routes == ()
|
||||
|
||||
|
||||
def test_from_env_negative_bound_skips_route() -> None:
|
||||
# A negative bound would match everything; it must fail open (skip the route).
|
||||
assert (
|
||||
ModelRouterConfig.from_env("yes", '[{"to_model":"m","min_input_tokens":-1}]').routes == ()
|
||||
)
|
||||
assert (
|
||||
ModelRouterConfig.from_env("yes", '[{"to_model":"m","max_input_tokens":-5}]').routes == ()
|
||||
)
|
||||
|
||||
|
||||
def test_from_env_unknown_key_skips_route() -> None:
|
||||
# A misspelled condition key must not be silently ignored (which would widen
|
||||
# the rule to match everything).
|
||||
assert ModelRouterConfig.from_env("yes", '[{"to_model":"m","max_input_token":5}]').routes == ()
|
||||
assert ModelRouterConfig.from_env("yes", '[{"to_model":"m","typo":true}]').routes == ()
|
||||
|
||||
|
||||
def test_from_env_valid_bool_and_ints_kept() -> None:
|
||||
cfg = ModelRouterConfig.from_env(
|
||||
"yes",
|
||||
'[{"to_model":"m","require_no_tools":false,"max_input_tokens":10,"min_input_tokens":0}]',
|
||||
)
|
||||
assert len(cfg.routes) == 1
|
||||
r = cfg.routes[0]
|
||||
assert r.require_no_tools is False and r.max_input_tokens == 10 and r.min_input_tokens == 0
|
||||
|
||||
|
||||
def test_from_env_various_truthy_values() -> None:
|
||||
for v in ("1", "true", "YES", "on", "enabled"):
|
||||
assert ModelRouterConfig.from_env(v, '[{"to_model":"m"}]').enabled, v
|
||||
for v in ("0", "false", "", "off", None):
|
||||
assert not ModelRouterConfig.from_env(v, '[{"to_model":"m"}]').enabled
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# estimate_input_tokens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_estimate_input_tokens_basic() -> None:
|
||||
messages = [{"role": "user", "content": "a" * 400}]
|
||||
assert estimate_input_tokens(messages) == 100
|
||||
|
||||
|
||||
def test_estimate_input_tokens_includes_tools() -> None:
|
||||
with_tools = estimate_input_tokens([{"content": "x" * 40}], tools=[{"name": "y" * 40}])
|
||||
without = estimate_input_tokens([{"content": "x" * 40}])
|
||||
assert with_tools > without
|
||||
|
||||
|
||||
def test_estimate_input_tokens_never_raises() -> None:
|
||||
assert estimate_input_tokens(None) == 0
|
||||
assert estimate_input_tokens("not a list") == 0
|
||||
assert estimate_input_tokens([123, {"content": "ok"}]) >= 0
|
||||
|
||||
|
||||
def test_estimate_input_tokens_counts_system_string() -> None:
|
||||
# A large top-level system prompt must not be ignored.
|
||||
small = estimate_input_tokens([{"content": "hi"}])
|
||||
with_system = estimate_input_tokens([{"content": "hi"}], system="s" * 4000)
|
||||
assert with_system >= small + 900
|
||||
|
||||
|
||||
def test_estimate_input_tokens_counts_system_blocks() -> None:
|
||||
blocks = [{"type": "text", "text": "x" * 4000}]
|
||||
assert estimate_input_tokens([{"content": "hi"}], system=blocks) > 100
|
||||
266
tests/test_proxy/test_model_router_wiring.py
Normal file
266
tests/test_proxy/test_model_router_wiring.py
Normal file
|
|
@ -0,0 +1,266 @@
|
|||
"""Wiring tests for cost-aware model routing (issue #1706).
|
||||
|
||||
Covers env -> ProxyConfig, ProxyConfig -> live proxy, and the presence of the
|
||||
routing block in the Anthropic request handler.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from headroom.proxy.handlers.anthropic import AnthropicHandlerMixin
|
||||
from headroom.proxy.model_router import ModelRoute, ModelRouter, ModelRouterConfig
|
||||
from headroom.proxy.server import ProxyConfig, _proxy_config_from_env, create_app
|
||||
|
||||
MESSAGES = "/v1/messages"
|
||||
|
||||
|
||||
def _install_fake_client(proxy) -> MagicMock:
|
||||
"""Replace proxy.http_client so forwarding never touches the network.
|
||||
|
||||
The buffered ``/v1/messages`` path forwards via ``http_client.post(content=...)``;
|
||||
the other forward shapes are stubbed too so the mock is robust to path choice.
|
||||
"""
|
||||
response = httpx.Response(
|
||||
200, json={"ok": True}, request=httpx.Request("POST", "http://upstream/v1/messages")
|
||||
)
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
client.request = AsyncMock(return_value=response)
|
||||
client.send = AsyncMock(return_value=response)
|
||||
client.build_request = MagicMock(
|
||||
return_value=httpx.Request("POST", "http://upstream/v1/messages", content=b"{}")
|
||||
)
|
||||
client.aclose = AsyncMock()
|
||||
proxy.http_client = client
|
||||
return client
|
||||
|
||||
|
||||
def _forwarded_model(client: MagicMock) -> str:
|
||||
"""Parse the outgoing model from the content forwarded upstream."""
|
||||
return _forwarded_body(client)["model"]
|
||||
|
||||
|
||||
def _forwarded_body(client: MagicMock) -> dict:
|
||||
"""Parse the JSON body forwarded upstream."""
|
||||
content = client.post.call_args.kwargs["content"]
|
||||
return json.loads(content)
|
||||
|
||||
|
||||
def _router_config() -> ModelRouterConfig:
|
||||
return ModelRouterConfig(
|
||||
enabled=True,
|
||||
routes=(
|
||||
ModelRoute(
|
||||
to_model="claude-haiku-4-5",
|
||||
max_input_tokens=100_000,
|
||||
require_no_tools=True,
|
||||
name="low-risk",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_proxy_config_from_env_reads_router(monkeypatch) -> None:
|
||||
monkeypatch.setenv("HEADROOM_MODEL_ROUTER_ENABLED", "true")
|
||||
monkeypatch.setenv(
|
||||
"HEADROOM_MODEL_ROUTES",
|
||||
'[{"name":"small","max_input_tokens":4000,"require_no_tools":true,'
|
||||
'"to_model":"claude-haiku-4-5"}]',
|
||||
)
|
||||
config = _proxy_config_from_env()
|
||||
assert config.model_router is not None
|
||||
assert config.model_router.enabled
|
||||
assert config.model_router.routes[0].to_model == "claude-haiku-4-5"
|
||||
|
||||
|
||||
def test_proxy_config_from_env_router_disabled_by_default(monkeypatch) -> None:
|
||||
monkeypatch.delenv("HEADROOM_MODEL_ROUTER_ENABLED", raising=False)
|
||||
monkeypatch.delenv("HEADROOM_MODEL_ROUTES", raising=False)
|
||||
config = _proxy_config_from_env()
|
||||
assert config.model_router is not None
|
||||
assert not config.model_router.enabled
|
||||
|
||||
|
||||
def test_create_app_wires_model_router() -> None:
|
||||
config = ProxyConfig(
|
||||
optimize=False,
|
||||
image_optimize=False,
|
||||
cache_enabled=False,
|
||||
rate_limit_enabled=False,
|
||||
cost_tracking_enabled=False,
|
||||
ccr_inject_tool=False,
|
||||
ccr_handle_responses=False,
|
||||
ccr_context_tracking=False,
|
||||
model_router=ModelRouterConfig(
|
||||
enabled=True,
|
||||
routes=(ModelRoute(to_model="cheap", max_input_tokens=10_000, name="small"),),
|
||||
),
|
||||
)
|
||||
app = create_app(config)
|
||||
with TestClient(app) as client:
|
||||
router = client.app.state.proxy.model_router
|
||||
assert router.enabled
|
||||
decision = router.select(model="strong", input_tokens=500, has_tools=False)
|
||||
assert decision.changed and decision.routed_model == "cheap"
|
||||
|
||||
|
||||
def test_create_app_router_disabled_when_unset() -> None:
|
||||
app = create_app(ProxyConfig(optimize=False, cost_tracking_enabled=False))
|
||||
with TestClient(app) as client:
|
||||
assert not client.app.state.proxy.model_router.enabled
|
||||
|
||||
|
||||
def test_handler_delegates_to_maybe_route_model() -> None:
|
||||
src = inspect.getsource(AnthropicHandlerMixin.handle_anthropic_messages)
|
||||
assert "_maybe_route_model(" in src, "handler must apply model routing"
|
||||
|
||||
|
||||
class _RouterHost(AnthropicHandlerMixin):
|
||||
"""Minimal mixin host (like a handler test double) for routing-only tests."""
|
||||
|
||||
|
||||
def test_maybe_route_model_fails_closed_without_router() -> None:
|
||||
# A host that never set model_router (test doubles, alternate mixin hosts that
|
||||
# do not run HeadroomProxy.__init__) must not crash when routing is off.
|
||||
host = _RouterHost()
|
||||
tracker = MagicMock()
|
||||
out = host._maybe_route_model(
|
||||
"claude-sonnet-4-6", [{"content": "hi"}], {"model": "claude-sonnet-4-6"}, tracker, False
|
||||
)
|
||||
assert out == "claude-sonnet-4-6"
|
||||
tracker.mark_mutated.assert_not_called()
|
||||
|
||||
|
||||
def test_maybe_route_model_routes_when_enabled() -> None:
|
||||
host = _RouterHost()
|
||||
host.model_router = ModelRouter(
|
||||
ModelRouterConfig(
|
||||
enabled=True,
|
||||
routes=(
|
||||
ModelRoute(
|
||||
to_model="claude-haiku-4-5", max_input_tokens=100_000, require_no_tools=True
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
tracker = MagicMock()
|
||||
body = {"model": "claude-sonnet-4-6"}
|
||||
out = host._maybe_route_model("claude-sonnet-4-6", [{"content": "hi"}], body, tracker, False)
|
||||
assert out == "claude-haiku-4-5"
|
||||
assert body["model"] == "claude-haiku-4-5"
|
||||
tracker.mark_mutated.assert_called_once_with("model_router")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("routes", "expected_reason"),
|
||||
[
|
||||
((ModelRoute(to_model="keep", from_models=("keep",), name="exempt"),), "exempt"),
|
||||
((ModelRoute(to_model="cheap", from_models=("other",)),), "no rule matched"),
|
||||
],
|
||||
)
|
||||
def test_maybe_route_model_logs_unchanged_decision(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
routes: tuple[ModelRoute, ...],
|
||||
expected_reason: str,
|
||||
) -> None:
|
||||
host = _RouterHost()
|
||||
host.model_router = ModelRouter(ModelRouterConfig(enabled=True, routes=routes))
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="headroom.proxy"):
|
||||
out = host._maybe_route_model(
|
||||
"keep", [{"content": "hi"}], {"model": "keep"}, MagicMock(), False
|
||||
)
|
||||
|
||||
assert out == "keep"
|
||||
decisions = [
|
||||
record.message for record in caplog.records if "model routing decision" in record.message
|
||||
]
|
||||
assert len(decisions) == 1
|
||||
assert expected_reason in decisions[0]
|
||||
|
||||
|
||||
def test_maybe_route_model_skips_on_bypass() -> None:
|
||||
host = _RouterHost()
|
||||
host.model_router = ModelRouter(
|
||||
ModelRouterConfig(enabled=True, routes=(ModelRoute(to_model="cheap"),))
|
||||
)
|
||||
tracker = MagicMock()
|
||||
out = host._maybe_route_model("keep", [{"content": "hi"}], {"model": "keep"}, tracker, True)
|
||||
assert out == "keep"
|
||||
tracker.mark_mutated.assert_not_called()
|
||||
|
||||
|
||||
def _messages_config() -> ProxyConfig:
|
||||
return ProxyConfig(
|
||||
optimize=False,
|
||||
cache_enabled=False,
|
||||
rate_limit_enabled=False,
|
||||
cost_tracking_enabled=False,
|
||||
ccr_inject_tool=False,
|
||||
ccr_handle_responses=False,
|
||||
ccr_context_tracking=False,
|
||||
mode="token",
|
||||
model_router=_router_config(),
|
||||
)
|
||||
|
||||
|
||||
def test_messages_request_gets_model_rewritten_when_enabled() -> None:
|
||||
app = create_app(_messages_config())
|
||||
with TestClient(app) as client:
|
||||
http = _install_fake_client(client.app.state.proxy)
|
||||
resp = client.post(
|
||||
MESSAGES,
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
# A low-risk request routes to the cheaper model on the forwarded body.
|
||||
assert _forwarded_model(http) == "claude-haiku-4-5"
|
||||
|
||||
|
||||
def test_bypass_request_is_never_model_rewritten() -> None:
|
||||
app = create_app(_messages_config())
|
||||
with TestClient(app) as client:
|
||||
http = _install_fake_client(client.app.state.proxy)
|
||||
resp = client.post(
|
||||
MESSAGES,
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
headers={"x-headroom-bypass": "true"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
# Byte-faithful passthrough must keep the client's original model.
|
||||
assert _forwarded_model(http) == "claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_vertex_raw_predict_model_is_not_rewritten_in_body() -> None:
|
||||
# When the model comes from the provider URL (Vertex rawPredict), the upstream
|
||||
# model is set by the path, so routing must not rewrite body["model"].
|
||||
app = create_app(_messages_config())
|
||||
with TestClient(app) as client:
|
||||
http = _install_fake_client(client.app.state.proxy)
|
||||
resp = client.post(
|
||||
"/v1/projects/p/locations/us-central1/publishers/anthropic/models/"
|
||||
"claude-sonnet-4-6:rawPredict",
|
||||
json={
|
||||
"anthropic_version": "vertex-2023-10-16",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert "model" not in _forwarded_body(http)
|
||||
|
|
@ -48,6 +48,14 @@ from headroom.proxy.server import ProxyConfig, create_app
|
|||
|
||||
pytest.importorskip("fastapi")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _disable_output_shaper(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
# Isolate this suite from the opt-in HEADROOM_OUTPUT_SHAPER a developer shell
|
||||
# may export, which otherwise perturbs the byte-faithful assertions.
|
||||
monkeypatch.delenv("HEADROOM_OUTPUT_SHAPER", raising=False)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit tests for serializer + tracker
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue