feat(router): soft-floor adaptive mode for complexity router (#32947)

* feat(router): soft-floor adaptive mode for complexity router

Let complexity_router_config.adaptive=true Thompson-sample across the
union of tier pools with a tier-distance penalty, and wire the existing
adaptive post-call bandit so mis-tiered requests can still recover.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): reattach adaptive hooks for hybrid complexity

Finalize was wiping every AdaptiveRouterPostCallHook and only
re-registering standalone auto_router/adaptive_router deployments,
so complexity adaptive=true never received bandit updates.

Co-authored-by: Cursor <cursoragent@cursor.com>

* chore(router): drop unnecessary hybrid docstrings

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): attribute adaptive feedback

Credit user reactions to the model that produced the previous response while keeping current-response signals on the serving model

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): tune hybrid cold defaults

Use the cost-weighted policy that beat equal-pool complexity in the full bakeoff, and make the committed harness compare identical tier pools

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): preserve hybrid cold quality floor

Sample only unobserved models in the classified tier until feedback exists, then apply adaptive scoring without mis-penalizing models shared across tiers

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): bound feedback context cache

Cap retained session feedback so unique session IDs cannot exhaust router memory

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(router): preserve exhaustion signals

Include tool-result exhaustion in adaptive feedback and clear strict lint regressions blocking CI

Co-authored-by: Cursor <cursoragent@cursor.com>

* refactor(router): remove stale owner cache

Remove obsolete attribution state, tighten the embedded router type, and keep the test diff focused on adaptive behavior

Co-authored-by: Cursor <cursoragent@cursor.com>

* refactor(router): centralize hook cleanup

Use the callback manager to discover and remove adaptive hooks across every registered callback list

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Krrish Dholakia 2026-07-11 21:56:33 -07:00 • committed by GitHub
parent 85f9bdd412
commit 26ab730bfa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 1035 additions and 409 deletions

View file

@ -7556,7 +7556,11 @@ class Router:
if default_model is None and complexity_router_config:
tiers = complexity_router_config.get("tiers", {})
# Use MEDIUM tier as fallback default
default_model = tiers.get("MEDIUM") or tiers.get("SIMPLE")
medium = tiers.get("MEDIUM") or tiers.get("SIMPLE")
if isinstance(medium, list):
default_model = medium[0] if medium else None
else:
default_model = medium
if default_model is None:
raise ValueError(
@ -7593,15 +7597,6 @@ class Router:
AdaptiveRouterPostCallHook,
)
for _cb_list in (
litellm.callbacks,
litellm.success_callback,
litellm.failure_callback,
litellm._async_success_callback,
litellm._async_failure_callback,
):
litellm.logging_callback_manager.remove_callbacks_by_type(_cb_list, AdaptiveRouterPostCallHook)
for entry in self.model_list or []:
lp = entry.get("litellm_params") if isinstance(entry, dict) else entry.litellm_params
lp_model = (lp.get("model") if isinstance(lp, dict) else lp.model) if lp else None
@ -7619,6 +7614,20 @@ class Router:
)
self.init_adaptive_router_deployment(deployment=deployment)
for model_name, complexity_router in self.complexity_routers.items():
if not complexity_router.config.adaptive or model_name in self.adaptive_routers:
continue
adaptive_router = complexity_router._ensure_adaptive_router()
if adaptive_router is not None:
self.adaptive_routers[model_name] = adaptive_router
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook):
litellm.logging_callback_manager.remove_callback_from_all_lists(callback)
for adaptive_router in self.adaptive_routers.values():
litellm.logging_callback_manager.add_litellm_callback(
AdaptiveRouterPostCallHook(adaptive_router=adaptive_router)
)
def init_adaptive_router_deployment(self, deployment: Deployment) -> None:
"""
Build an AdaptiveRouter instance for this deployment and register its

View file

@ -56,11 +56,10 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
- **Per-request decision.** Sample once per eligible model, score with
`quality_weight·sample + cost_weight·normalized_cost`, pick the argmax.
Routing is stateless per-turn — no sticky lookup. Each call resamples.
- **Owner-cache attribution.** Post-call, the conversation's first picked
model claims an "owner slot" for `OWNER_CACHE_TTL_SECONDS` (24h). Later
turns of the same conversation only fire bandit/state updates if the
same model handled them — mismatches are dropped (no attribution) and
counted in `skipped_updates_total`. Conversation identity is the
- **Previous-response attribution.** Post-call, feedback from the current user
message is attributed to the model that produced the previous response, while
response signals are attributed to the current model. Contexts expire after
24 hours and the in-memory cache is capped at 1,024 sessions. Conversation identity is the
client-supplied `litellm_session_id` if present, otherwise a sha256 over
caller identity (api key hash, team, user, end-user) + the first message.
- **Per-turn updates.** `satisfaction → +α`. `misalignment, stagnation,
@ -76,12 +75,6 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
model can still be picked.
- **Hard sample cap at 200.** Once `α + β > 200`, deltas are silently dropped.
No rescaling — drift is a v1 concern.
- **24h owner-cache TTL.** No explicit eviction below TTL. The in-memory map
can grow if traffic patterns produce many one-shot sessions.
- **Owner-recovery skew.** If model A "owns" a conversation but is then
dethroned in the bandit, later turns served by model B are dropped — so
bandit updates for that conversation flatline until A's TTL expires.
Tracked via `skipped_updates_total`.
- **Signals are regex + tool-call only.** No LLM-judge, no embedding similarity,
no exemplar storage. Signals are best-effort and biased toward English.
- **One AdaptiveRouter per `Router`.** Multiple `adaptive_router/*` deployments

View file

@ -3,25 +3,21 @@ Main adaptive router strategy. See README.md for design overview.
One AdaptiveRouter instance per router_name. Holds in-memory caches:
- _cells: Beta(alpha, beta) bandit posteriors per (request_type, model)
- _owner_cache: session_key -> (owner_model, expires_at) — the first model
picked for a conversation owns its bandit-update slot
- _session_states: (session_key, model) -> SessionState for incremental signal updates
Owns the AdaptiveRouterUpdateQueue used by the proxy's flusher to persist
state and session snapshots back to Postgres.
Routing is stateless per-turn (Thompson sample fresh on every call). The
owner cache is consulted only at post-call time to decide whether a turn's
signals should fire a bandit update — turns served by a different model than
the conversation's owner are skipped to avoid cross-model misattribution.
Routing is stateless per-turn (Thompson sample fresh on every call).
"""
from __future__ import annotations
import asyncio
import time
from dataclasses import asdict
from typing import Any, Dict, List, Optional, Tuple, Union, cast
from collections import OrderedDict
from dataclasses import asdict, dataclass
from typing import Any, Union, cast
from litellm._logging import verbose_router_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -38,13 +34,18 @@ from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
MIN_QUALITY_TIER_HEADER,
MIN_QUALITY_TIER_METADATA_KEY,
MIN_TURNS_FOR_CLEAN_CREDIT,
OWNER_CACHE_TTL_SECONDS,
)
from litellm.router_strategy.adaptive_router.signals import (
SessionState,
SignalDelta,
Turn,
apply_turn,
advance_session_state,
apply_signal_delta,
detect_response_signals,
detect_user_feedback,
merge_signal_deltas,
)
from litellm.router_strategy.adaptive_router.update_queue import (
AdaptiveRouterUpdateQueue,
@ -53,8 +54,7 @@ from litellm.router_strategy.adaptive_router.update_queue import (
# Sweep session-state cache when it exceeds this many live entries. Expired
# entries are dropped in bulk; amortizes to O(1) per insert.
_SESSION_STATE_SWEEP_THRESHOLD: int = 1024
# Same pattern for the owner cache.
_OWNER_CACHE_SWEEP_THRESHOLD: int = 1024
_FEEDBACK_CONTEXT_MAX_ENTRIES: int = 1024
from litellm.repositories.table_repositories import AdaptiveRouterStateRepository
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import (
@ -70,6 +70,17 @@ def _default_prefs() -> AdaptiveRouterPreferences:
return AdaptiveRouterPreferences(quality_tier=2, strengths=[])
@dataclass(frozen=True, slots=True)
class _FeedbackContext:
model_name: str
request_type: RequestType
user_content: str | None
assistant_content: str | None
turn_count: int
clean_credit_awarded: bool
expires_at: float
class AdaptiveRouter:
"""One instance per router_name. Holds in-memory caches + the update queue."""
@ -77,8 +88,8 @@ class AdaptiveRouter:
self,
router_name: str,
config: AdaptiveRouterConfig,
model_to_prefs: Dict[str, AdaptiveRouterPreferences],
model_to_cost: Dict[str, float],
model_to_prefs: dict[str, AdaptiveRouterPreferences],
model_to_cost: dict[str, float],
) -> None:
self.router_name = router_name
self.config = config
@ -86,13 +97,14 @@ class AdaptiveRouter:
self.model_to_cost = model_to_cost
self.queue = AdaptiveRouterUpdateQueue()
self._cells: Dict[Tuple[RequestType, str], BanditCell] = {}
self._owner_cache: Dict[str, Tuple[str, float]] = {}
self._session_states: Dict[Tuple[str, str], SessionState] = {}
# Parallel expiry map for _session_states, same TTL as _owner_cache.
# Evicted opportunistically in `get_or_create_session_state`.
self._session_states_expiry: Dict[Tuple[str, str], float] = {}
self._skipped_updates_total: int = 0
self._cells: dict[tuple[RequestType, str], BanditCell] = {}
self._session_states: dict[tuple[str, str], SessionState] = {}
self._feedback_contexts: OrderedDict[str, _FeedbackContext] = OrderedDict()
self._session_states_expiry: dict[tuple[str, str], float] = {}
self._feedback_attributed_total: int = 0
self._feedback_without_context_total: int = 0
self._cross_model_feedback_total: int = 0
self._response_signal_updates_total: int = 0
# Set to True once the proxy flusher has loaded persisted priors from
# Postgres. Checked to support lazy-load on hot-reloaded routers.
self._state_loaded: bool = False
@ -145,11 +157,11 @@ class AdaptiveRouter:
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict[str, Any],
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:
request_kwargs: dict[str, Any],
messages: list[dict[str, Any]] | None = None,
input: Union[str, list] | None = None,
specific_deployment: bool | None = False,
) -> PreRoutingHookResponse | None:
"""
Plugin entry point invoked by `Router.async_pre_routing_hook` when the
inbound `model` matches this adaptive router's `router_name`.
@ -159,11 +171,9 @@ class AdaptiveRouter:
post-call hook can surface it as a response header.
Routing is stateless per-turn: every call Thompson-samples fresh,
regardless of any prior pick for the same session. Cross-turn
attribution is enforced post-call via the owner cache (see
`claim_or_check_owner`).
regardless of any prior pick for the same session.
"""
user_text = get_last_user_message(cast(List[AllMessageValues], messages or [])) or ""
user_text = get_last_user_message(cast(list[AllMessageValues], messages or [])) or ""
request_type = classify_prompt(user_text)
min_quality_tier = self._extract_min_quality_tier(request_kwargs)
@ -190,7 +200,7 @@ class AdaptiveRouter:
async def pick_model(
self,
request_type: RequestType,
min_quality_tier: Optional[int] = None,
min_quality_tier: int | None = None,
) -> str:
"""Thompson-sample across eligible models. Stateless per-turn."""
eligible = self._eligible_models(min_quality_tier)
@ -206,44 +216,7 @@ class AdaptiveRouter:
cost_weight=self.config.weights.cost,
)
def claim_or_check_owner(self, session_key: str, current_model: str) -> bool:
"""Resolve attribution for a turn under stateless routing.
Returns True iff this turn should fire a bandit/state update. The
first call for a `session_key` claims ownership for `current_model`
and returns True. Subsequent calls return True only if the owner is
still live AND matches `current_model`. Mismatches (a different
model handled this turn) and expired owners both increment
`_skipped_updates_total` and return False — no attribution.
"""
now = time.time()
existing = self._owner_cache.get(session_key)
if existing is not None and existing[1] > now:
owner_model, _ = existing
if owner_model == current_model:
return True
self._skipped_updates_total += 1
return False
# Opportunistic bulk sweep — sessions that never come back would
# otherwise pile up here forever. Same threshold pattern as the
# session-state cache.
if len(self._owner_cache) >= _OWNER_CACHE_SWEEP_THRESHOLD:
self._evict_expired_owner_cache(now)
# No live owner -> claim for current_model.
self._owner_cache[session_key] = (
current_model,
now + OWNER_CACHE_TTL_SECONDS,
)
return True
def _evict_expired_owner_cache(self, now: float) -> None:
expired = [k for k, (_, exp) in self._owner_cache.items() if exp <= now]
for k in expired:
self._owner_cache.pop(k, None)
async def get_state_snapshot(self) -> Dict[str, Any]:
async def get_state_snapshot(self) -> dict[str, Any]:
"""In-memory snapshot for the introspection endpoint. Cheap; no DB hit."""
cells = []
for (rt, model), cell in sorted(self._cells.items(), key=lambda kv: (kv[0][0].value, kv[0][1])):
@ -264,7 +237,7 @@ class AdaptiveRouter:
)
queue = await self.queue.queue_size()
now = time.time()
owner_cache_live = sum(1 for _, exp in self._owner_cache.values() if exp > now)
feedback_contexts_live = sum(1 for context in self._feedback_contexts.values() if context.expires_at > now)
return {
"router_name": self.router_name,
"available_models": list(self.config.available_models),
@ -274,15 +247,18 @@ class AdaptiveRouter:
},
"model_costs": dict(self.model_to_cost),
"cells": cells,
"owner_cache_live": owner_cache_live,
"skipped_updates_total": self._skipped_updates_total,
"feedback_contexts_live": feedback_contexts_live,
"feedback_attributed_total": self._feedback_attributed_total,
"feedback_without_context_total": self._feedback_without_context_total,
"cross_model_feedback_total": self._cross_model_feedback_total,
"response_signal_updates_total": self._response_signal_updates_total,
"queue": queue,
}
@staticmethod
def _extract_min_quality_tier(
request_kwargs: Dict[str, Any],
) -> Optional[int]:
request_kwargs: dict[str, Any],
) -> int | None:
"""Pull `min_quality_tier` from request headers or metadata.
Precedence: headers (`x-litellm-min-quality-tier`) over metadata
@ -310,7 +286,7 @@ class AdaptiveRouter:
return None
return None
def _eligible_models(self, min_quality_tier: Optional[int]) -> List[str]:
def _eligible_models(self, min_quality_tier: int | None) -> list[str]:
if min_quality_tier is None:
return list(self.config.available_models)
return [
@ -363,17 +339,131 @@ class AdaptiveRouter:
request_type: RequestType,
turn: Turn,
) -> SignalDelta:
"""Apply one turn, push session snapshot + bandit deltas to the queue."""
state = self.get_or_create_session_state(session_id, model_name, request_type)
delta = apply_turn(state, turn)
verbose_router_logger.debug("AdaptiveRouter[%s]: record_turn delta=%s", self.router_name, delta)
"""Attribute feedback to the previous response and response signals to the current model."""
async with self._lock:
now = time.time()
while self._feedback_contexts:
oldest_context = next(iter(self._feedback_contexts.values()))
if oldest_context.expires_at > now:
break
self._feedback_contexts.popitem(last=False)
previous = self._feedback_contexts.pop(session_id, None)
# Strip the raw conversation content before persisting. The
# last_user/assistant_content and tool_call_history fields are only
# needed in-memory for the next turn's incremental signal detection;
# writing user prompts and tool payloads to the DB would store PII
# for every adaptive-router conversation. Counts + bookkeeping is
# all the persisted row needs.
effective_request_type = (
previous.request_type if previous is not None and request_type == RequestType.GENERAL else request_type
)
current_state = self.get_or_create_session_state(
session_id,
model_name,
effective_request_type,
)
feedback_delta = detect_user_feedback(
previous.user_content if previous else None,
turn.user_content,
turn.tool_results,
allow_satisfaction=(
previous is not None
and not previous.clean_credit_awarded
and previous.turn_count + 1 >= MIN_TURNS_FOR_CLEAN_CREDIT
),
)
previous_assistant = previous.assistant_content if previous else None
response_delta = detect_response_signals(
previous_assistant,
turn.assistant_content,
current_state.tool_call_history,
turn.tool_calls,
turn.tool_results,
turn.response_status,
)
states_to_persist: dict[str, SessionState] = {model_name: current_state}
bandit_deltas: dict[tuple[RequestType, str], SignalDelta] = {}
if previous is not None:
feedback_state = self.get_or_create_session_state(
session_id,
previous.model_name,
previous.request_type,
)
apply_signal_delta(feedback_state, feedback_delta)
if feedback_delta.satisfaction:
feedback_state.clean_credit_awarded = True
states_to_persist[previous.model_name] = feedback_state
if feedback_delta.any_fired():
self._feedback_attributed_total += 1
if previous.model_name != model_name:
self._cross_model_feedback_total += 1
bandit_deltas[(previous.request_type, previous.model_name)] = feedback_delta
else:
if feedback_delta.any_fired():
self._feedback_without_context_total += 1
initial_failure = SignalDelta(failure=feedback_delta.failure)
apply_signal_delta(current_state, initial_failure)
bandit_deltas[(effective_request_type, model_name)] = initial_failure
apply_signal_delta(current_state, response_delta)
if self._compute_bandit_delta(response_delta) != (0.0, 0.0):
self._response_signal_updates_total += 1
current_key = (effective_request_type, model_name)
bandit_deltas[current_key] = merge_signal_deltas(
bandit_deltas.get(current_key, SignalDelta()),
response_delta,
)
advance_session_state(current_state, turn)
next_turn_count = (previous.turn_count if previous else 0) + 1
clean_credit_awarded = bool((previous and previous.clean_credit_awarded) or feedback_delta.satisfaction)
if len(self._feedback_contexts) >= _FEEDBACK_CONTEXT_MAX_ENTRIES:
self._feedback_contexts.popitem(last=False)
self._feedback_contexts[session_id] = _FeedbackContext(
model_name=model_name,
request_type=effective_request_type,
user_content=turn.user_content,
assistant_content=turn.assistant_content,
turn_count=next_turn_count,
clean_credit_awarded=clean_credit_awarded,
expires_at=now + OWNER_CACHE_TTL_SECONDS,
)
for state_model, state in states_to_persist.items():
await self.queue.add_session_state(
session_id,
self.router_name,
state_model,
self._persistable_session_snapshot(state),
)
combined_delta = SignalDelta()
for (attribution_type, target_model), delta in bandit_deltas.items():
combined_delta = merge_signal_deltas(combined_delta, delta)
d_alpha, d_beta = self._compute_bandit_delta(delta)
if d_alpha == 0 and d_beta == 0:
continue
cell_key = (attribution_type, target_model)
self._cells[cell_key] = apply_delta(
self._cells[cell_key],
d_alpha,
d_beta,
)
await self.queue.add_state_delta(
self.router_name,
attribution_type.value,
target_model,
d_alpha,
d_beta,
)
verbose_router_logger.debug(
"AdaptiveRouter[%s]: feedback_target=%s current_model=%s delta=%s",
self.router_name,
previous.model_name if previous else None,
model_name,
combined_delta,
)
return combined_delta
@staticmethod
def _persistable_session_snapshot(state: SessionState) -> dict[str, Any]:
snapshot = asdict(state)
for sensitive in (
"last_user_content",
@ -382,38 +472,10 @@ class AdaptiveRouter:
"pending_tool_calls",
):
snapshot.pop(sensitive, None)
await self.queue.add_session_state(session_id, self.router_name, model_name, snapshot)
d_alpha, d_beta = self._compute_bandit_delta(delta)
verbose_router_logger.debug(
"AdaptiveRouter[%s]: bandit delta alpha=%.2f beta=%.2f",
self.router_name,
d_alpha,
d_beta,
)
if d_alpha != 0 or d_beta != 0:
# For non-GENERAL turns, attribute to the current-turn classification
# so genuine mid-session topic shifts (e.g. code → math) update the
# correct cell. For GENERAL turns ("thanks!", "ok", "sounds good"), fall
# back to the session's original type so closing pleasantries don't
# misattribute the reward.
attribution_type = (
request_type if request_type != RequestType.GENERAL else RequestType(state.classified_type)
)
cell_key = (attribution_type, model_name)
self._cells[cell_key] = apply_delta(self._cells[cell_key], d_alpha, d_beta)
await self.queue.add_state_delta(
self.router_name,
attribution_type.value,
model_name,
d_alpha,
d_beta,
)
return delta
return snapshot
@staticmethod
def _compute_bandit_delta(delta: SignalDelta) -> Tuple[float, float]:
def _compute_bandit_delta(delta: SignalDelta) -> tuple[float, float]:
"""
Translate per-turn signal deltas into bandit-cell deltas.

View file

@ -214,10 +214,6 @@ class AdaptiveRouterPostCallHook(CustomLogger):
) -> None:
try:
messages = kwargs.get("messages") or []
if len(messages) < SIGNAL_GATE_MIN_MESSAGES:
# Too few turns for any signal to be meaningful — skip.
return
session_key = _resolve_session_key(kwargs)
if not session_key:
return
@ -233,10 +229,6 @@ class AdaptiveRouterPostCallHook(CustomLogger):
if not current_model:
return
if not self.adaptive_router.claim_or_check_owner(session_key, current_model):
# A different model owns this conversation — skip attribution.
return
user_text = _last_user_content(messages)
assistant_text, tool_calls = _assistant_content_and_tool_calls(response_obj)
tool_results = _recent_tool_results(messages)

View file

@ -14,7 +14,7 @@ from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Set
from typing import Any
from litellm.router_strategy.adaptive_router.config import (
LOOP_REPEAT_THRESHOLD,
@ -74,26 +74,26 @@ class SessionState:
loop_count: int = 0
exhaustion_count: int = 0
last_user_content: Optional[str] = None
last_assistant_content: Optional[str] = None
tool_call_history: List[str] = field(default_factory=list)
pending_tool_calls: Dict[str, str] = field(default_factory=dict)
last_user_content: str | None = None
last_assistant_content: str | None = None
tool_call_history: list[str] = field(default_factory=list)
pending_tool_calls: dict[str, str] = field(default_factory=dict)
turn_count: int = 0
last_processed_turn: int = -1
clean_credit_awarded: bool = False
terminal_status: Optional[int] = None
terminal_status: int | None = None
@dataclass
class Turn:
"""One turn of input. Caller assembles this from the request/response."""
user_content: Optional[str] = None
assistant_content: Optional[str] = None
tool_calls: List[Dict[str, Any]] = field(default_factory=list)
tool_results: List[Dict[str, Any]] = field(default_factory=list)
response_status: Optional[int] = None
user_content: str | None = None
assistant_content: str | None = None
tool_calls: list[dict[str, Any]] = field(default_factory=list)
tool_results: list[dict[str, Any]] = field(default_factory=list)
response_status: int | None = None
# ---- Detection helpers ----------------------------------------------------
@ -101,13 +101,13 @@ class Turn:
_TOKEN_RE = re.compile(r"[A-Za-z0-9]+")
def _tokens(text: Optional[str]) -> Set[str]:
def _tokens(text: str | None) -> set[str]:
if not text:
return set()
return {t.lower() for t in _TOKEN_RE.findall(text)}
def _jaccard(a: Set[str], b: Set[str]) -> float:
def _jaccard(a: set[str], b: set[str]) -> float:
union = a | b
if not union:
return 0.0
@ -130,7 +130,7 @@ _SATISFACTION_PATTERNS = [
]
def _detect_misalignment(prev_user: Optional[str], curr_user: Optional[str]) -> bool:
def _detect_misalignment(prev_user: str | None, curr_user: str | None) -> bool:
"""Fires when consecutive user messages share *some* topic (jaccard > 0)
but are sufficiently different (jaccard < threshold) — i.e. user is
rephrasing, not changing topic, not repeating."""
@ -140,7 +140,7 @@ def _detect_misalignment(prev_user: Optional[str], curr_user: Optional[str]) ->
return 0.0 < j < MISALIGNMENT_JACCARD_THRESHOLD
def _detect_stagnation(prev_asst: Optional[str], curr_asst: Optional[str]) -> bool:
def _detect_stagnation(prev_asst: str | None, curr_asst: str | None) -> bool:
"""Fires when consecutive assistant messages are near-duplicates."""
if not prev_asst or not curr_asst:
return False
@ -148,19 +148,19 @@ def _detect_stagnation(prev_asst: Optional[str], curr_asst: Optional[str]) -> bo
return j >= STAGNATION_JACCARD_NEAR_DUP
def _detect_disengagement(curr_user: Optional[str]) -> bool:
def _detect_disengagement(curr_user: str | None) -> bool:
if not curr_user:
return False
return any(p.search(curr_user) for p in _DISENGAGEMENT_PATTERNS)
def _detect_satisfaction(curr_user: Optional[str]) -> bool:
def _detect_satisfaction(curr_user: str | None) -> bool:
if not curr_user:
return False
return any(p.search(curr_user) for p in _SATISFACTION_PATTERNS)
def _detect_failure(tool_results: List[Dict[str, Any]]) -> bool:
def _detect_failure(tool_results: list[dict[str, Any]]) -> bool:
"""Any tool result explicitly flagged as an error.
We do NOT treat empty content as failure — many tools legitimately return
@ -173,7 +173,7 @@ def _detect_failure(tool_results: List[Dict[str, Any]]) -> bool:
return False
def _signature(call: Dict[str, Any]) -> str:
def _signature(call: dict[str, Any]) -> str:
"""Stable signature for loop detection: name + sorted JSON-ish args."""
name = call.get("name") or call.get("function", {}).get("name", "")
call_args = call.get("arguments")
@ -184,7 +184,7 @@ def _signature(call: Dict[str, Any]) -> str:
return f"{name}({call_args})"
def _detect_loop(history: List[str], new_calls: List[Dict[str, Any]]) -> bool:
def _detect_loop(history: list[str], new_calls: list[dict[str, Any]]) -> bool:
"""Fires if any new call's signature appears >= LOOP_REPEAT_THRESHOLD-1 times
in recent history (so this call would be the Nth)."""
if not new_calls:
@ -209,7 +209,7 @@ _EXHAUSTION_KEYWORDS = (
)
def _detect_exhaustion(status: Optional[int], tool_results: List[Dict[str, Any]]) -> bool:
def _detect_exhaustion(status: int | None, tool_results: list[dict[str, Any]]) -> bool:
if status is not None and status in _EXHAUSTION_STATUSES:
return True
for r in tool_results:
@ -219,39 +219,53 @@ def _detect_exhaustion(status: Optional[int], tool_results: List[Dict[str, Any]]
return False
# ---- Public entrypoint ----------------------------------------------------
def detect_user_feedback(
previous_user_content: str | None,
current_user_content: str | None,
tool_results: list[dict[str, Any]],
allow_satisfaction: bool,
) -> SignalDelta:
return SignalDelta(
misalignment=int(_detect_misalignment(previous_user_content, current_user_content)),
disengagement=int(_detect_disengagement(current_user_content)),
satisfaction=int(allow_satisfaction and _detect_satisfaction(current_user_content)),
failure=int(_detect_failure(tool_results)),
)
def apply_turn(state: SessionState, turn: Turn) -> SignalDelta:
"""
Detect signals on this turn, mutate state, return the delta.
def detect_response_signals(
previous_assistant_content: str | None,
current_assistant_content: str | None,
tool_call_history: list[str],
tool_calls: list[dict[str, Any]],
tool_results: list[dict[str, Any]],
response_status: int | None,
) -> SignalDelta:
return SignalDelta(
stagnation=int(
_detect_stagnation(
previous_assistant_content,
current_assistant_content,
)
),
loop=int(_detect_loop(tool_call_history, tool_calls)),
exhaustion=int(_detect_exhaustion(response_status, tool_results)),
)
O(1) per turn (no full-history rescan). Only inspects last_*, recent tool history
(which is bounded at TOOL_CALL_HISTORY_MAX), and the new turn payload.
"""
delta = SignalDelta()
if _detect_misalignment(state.last_user_content, turn.user_content):
delta.misalignment = 1
if _detect_stagnation(state.last_assistant_content, turn.assistant_content):
delta.stagnation = 1
if _detect_disengagement(turn.user_content):
delta.disengagement = 1
if _detect_satisfaction(turn.user_content):
# Gate: only award satisfaction credit once per session, and only
# after MIN_TURNS_FOR_CLEAN_CREDIT turns of context. Early "thanks"
# on turn 1-2 is noise, not a validated quality signal.
current_turn_index = state.turn_count + 1
if not state.clean_credit_awarded and current_turn_index >= MIN_TURNS_FOR_CLEAN_CREDIT:
delta.satisfaction = 1
state.clean_credit_awarded = True
if _detect_failure(turn.tool_results):
delta.failure = 1
if _detect_loop(state.tool_call_history, turn.tool_calls):
delta.loop = 1
if _detect_exhaustion(turn.response_status, turn.tool_results):
delta.exhaustion = 1
def merge_signal_deltas(*deltas: SignalDelta) -> SignalDelta:
return SignalDelta(
misalignment=sum(delta.misalignment for delta in deltas),
stagnation=sum(delta.stagnation for delta in deltas),
disengagement=sum(delta.disengagement for delta in deltas),
satisfaction=sum(delta.satisfaction for delta in deltas),
failure=sum(delta.failure for delta in deltas),
loop=sum(delta.loop for delta in deltas),
exhaustion=sum(delta.exhaustion for delta in deltas),
)
def apply_signal_delta(state: SessionState, delta: SignalDelta) -> None:
state.misalignment_count += delta.misalignment
state.stagnation_count += delta.stagnation
state.disengagement_count += delta.disengagement
@ -260,6 +274,8 @@ def apply_turn(state: SessionState, turn: Turn) -> SignalDelta:
state.loop_count += delta.loop
state.exhaustion_count += delta.exhaustion
def advance_session_state(state: SessionState, turn: Turn) -> None:
if turn.user_content:
state.last_user_content = turn.user_content
if turn.assistant_content:
@ -276,4 +292,38 @@ def apply_turn(state: SessionState, turn: Turn) -> SignalDelta:
state.turn_count += 1
state.last_processed_turn = state.turn_count
# ---- Public entrypoint ----------------------------------------------------
def apply_turn(state: SessionState, turn: Turn) -> SignalDelta:
"""
Detect signals on this turn, mutate state, return the delta.
O(1) per turn (no full-history rescan). Only inspects last_*, recent tool history
(which is bounded at TOOL_CALL_HISTORY_MAX), and the new turn payload.
"""
feedback_delta = detect_user_feedback(
state.last_user_content,
turn.user_content,
turn.tool_results,
allow_satisfaction=(not state.clean_credit_awarded and state.turn_count + 1 >= MIN_TURNS_FOR_CLEAN_CREDIT),
)
response_delta = detect_response_signals(
state.last_assistant_content,
turn.assistant_content,
state.tool_call_history,
turn.tool_calls,
turn.tool_results,
turn.response_status,
)
delta = merge_signal_deltas(
feedback_delta,
response_delta,
)
apply_signal_delta(state, delta)
if delta.satisfaction:
state.clean_credit_awarded = True
advance_session_state(state, turn)
return delta

View file

@ -13,10 +13,12 @@ evaluated before either classification strategy and force a tier outright when m
Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter
"""
from __future__ import annotations
import asyncio
import random
import re
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
from typing import TYPE_CHECKING, Any, Literal, Optional, Union, cast
from pydantic import BaseModel
@ -38,6 +40,7 @@ if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
from litellm.router import Router
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.types.router import PreRoutingHookResponse
else:
Router = Any
@ -63,7 +66,7 @@ Tiers:
{prompt}"""
def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[list[str]]) -> list[str]:
def _append_custom_keywords(base_keywords: list[str], custom_keywords: list[str] | None) -> list[str]:
if not custom_keywords:
return base_keywords
base_lowered = frozenset(keyword.lower() for keyword in base_keywords)
@ -95,7 +98,7 @@ def _sanitize_user_api_key_auth(auth: Any) -> Any:
return auth
def _classifier_call_metadata(metadata: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]:
def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] | None:
if not metadata:
return metadata
return {
@ -110,7 +113,7 @@ class DimensionScore:
__slots__ = ("name", "score", "signal")
def __init__(self, name: str, score: float, signal: Optional[str] = None):
def __init__(self, name: str, score: float, signal: str | None = None):
self.name = name
self.score = score
self.signal = signal
@ -134,9 +137,9 @@ class ComplexityRouter(CustomLogger):
def __init__(
self,
model_name: str,
litellm_router_instance: "Router",
complexity_router_config: Optional[Dict[str, Any]] = None,
default_model: Optional[str] = None,
litellm_router_instance: Router,
complexity_router_config: dict[str, Any] | None = None,
default_model: str | None = None,
):
"""
Initialize ComplexityRouter.
@ -173,7 +176,7 @@ class ComplexityRouter(CustomLogger):
# embeddings are static, only the prompt is embedded per request). The lock
# serializes the one-time build so concurrent cold-start requests don't each
# construct the index and fire duplicate embedding calls.
self._semantic_routelayer: Optional[SemanticRouter] = None
self._semantic_routelayer: SemanticRouter | None = None
self._semantic_routelayer_lock = asyncio.Lock()
# Pre-compile regex patterns for efficiency
@ -185,6 +188,10 @@ class ComplexityRouter(CustomLogger):
re.compile(r"[a-z]\)\s", re.IGNORECASE),
]
self.adaptive_router: AdaptiveRouter | None = None
self._model_tiers: dict[str, tuple[ComplexityTier, ...]] = {}
self._adaptive_init_attempted = False
verbose_router_logger.debug(f"ComplexityRouter initialized for {model_name} with tiers: {self.config.tiers}")
def _estimate_tokens(self, text: str) -> int:
@ -228,12 +235,12 @@ class ComplexityRouter(CustomLogger):
def _score_keyword_match(
self,
text: str,
keywords: List[str],
keywords: list[str],
name: str,
signal_label: str,
thresholds: Tuple[int, int], # (low, high)
scores: Tuple[float, float, float], # (none, low, high)
) -> Tuple[DimensionScore, int]:
thresholds: tuple[int, int], # (low, high)
scores: tuple[float, float, float], # (none, low, high)
) -> tuple[DimensionScore, int]:
"""Score based on keyword matches using word boundary matching.
Returns:
@ -271,7 +278,7 @@ class ComplexityRouter(CustomLogger):
return DimensionScore("questionComplexity", 0.5, f"{count} questions")
return DimensionScore("questionComplexity", 0, None)
def classify(self, prompt: str, system_prompt: Optional[str] = None) -> Tuple[ComplexityTier, float, List[str]]:
def classify(self, prompt: str, system_prompt: str | None = None) -> tuple[ComplexityTier, float, list[str]]:
"""
Classify a prompt by complexity.
@ -330,7 +337,7 @@ class ComplexityRouter(CustomLogger):
(0, -1.0, -1.0),
)
dimensions: List[DimensionScore] = [
dimensions: list[DimensionScore] = [
self._score_token_count(estimated_tokens),
code_score,
reasoning_score,
@ -372,8 +379,8 @@ class ComplexityRouter(CustomLogger):
async def aclassify(
self,
prompt: str,
system_prompt: Optional[str] = None,
request_kwargs: Optional[dict[str, Any]] = None,
system_prompt: str | None = None,
request_kwargs: dict[str, Any] | None = None,
) -> tuple[ComplexityTier, float, list[str]]:
"""
Classify a prompt by complexity, using the LLM classifier when configured.
@ -396,8 +403,8 @@ class ComplexityRouter(CustomLogger):
async def _classify_with_llm(
self,
prompt: str,
system_prompt: Optional[str] = None,
request_kwargs: Optional[dict[str, Any]] = None,
system_prompt: str | None = None,
request_kwargs: dict[str, Any] | None = None,
) -> ComplexityTier:
"""Call the configured classifier model and parse its structured tier response."""
llm_config = self.config.classifier_llm_config
@ -458,7 +465,176 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"Empty model pool for tier {tier_key}")
return random.choice(model)
def _lexical_tier_override(self, user_message: str) -> Optional[ComplexityTier]:
def _tier_pools(self) -> dict[str, list[str]]:
return {tier: (models if isinstance(models, list) else [models]) for tier, models in self.config.tiers.items()}
def _ensure_adaptive_router(self) -> Any | None:
if not self.config.adaptive:
return None
if self.adaptive_router is not None:
return self.adaptive_router
if self._adaptive_init_attempted:
return self.adaptive_router
self._adaptive_init_attempted = True
from litellm.router_strategy.adaptive_router.adaptive_router import (
AdaptiveRouter,
)
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
)
from litellm.types.router import (
AdaptiveRouterConfig,
AdaptiveRouterPreferences,
)
pools = self._tier_pools()
available_models = list(dict.fromkeys(model for models in pools.values() for model in models))
self._model_tiers = {
model: tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models)
for model in available_models
}
model_to_prefs: dict[str, AdaptiveRouterPreferences] = {}
model_to_cost: dict[str, float] = {}
model_list = getattr(self.litellm_router_instance, "model_list", None) or []
name_to_indices = getattr(self.litellm_router_instance, "model_name_to_deployment_indices", {}) or {}
for name in available_models:
indices = name_to_indices.get(name, [])
if not indices:
model_to_prefs[name] = AdaptiveRouterPreferences(quality_tier=2, strengths=[])
model_to_cost[name] = 0.0
continue
deployment = model_list[indices[0]]
mi = deployment.get("model_info") if isinstance(deployment, dict) else deployment.model_info
mi_dict: dict[str, Any] = mi if isinstance(mi, dict) else (mi.model_dump() if mi else {})
prefs_raw = mi_dict.get("adaptive_router_preferences")
if prefs_raw is not None:
model_to_prefs[name] = AdaptiveRouterPreferences(**prefs_raw)
else:
model_to_prefs[name] = AdaptiveRouterPreferences(quality_tier=2, strengths=[])
lp = deployment.get("litellm_params") if isinstance(deployment, dict) else deployment.litellm_params
lp_dict: dict[str, Any] = lp if isinstance(lp, dict) else (lp.model_dump() if lp else {})
cost = lp_dict.get("input_cost_per_token")
model_to_cost[name] = float(cost) if cost is not None else 0.0
self.adaptive_router = AdaptiveRouter(
router_name=self.model_name,
config=AdaptiveRouterConfig(
available_models=available_models,
weights=self.config.adaptive_weights,
),
model_to_prefs=model_to_prefs,
model_to_cost=model_to_cost,
)
self._adaptive_chosen_model_key = ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY
return self.adaptive_router
def _soft_floor_pick(
self,
classified_tier: ComplexityTier,
user_message: str,
request_kwargs: dict[str, Any] | None = None,
) -> str:
from litellm.router_strategy.adaptive_router.bandit import (
normalized_cost,
thompson_sample,
)
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
adaptive = self._ensure_adaptive_router()
if adaptive is None:
return self.get_model_for_tier(classified_tier)
request_type = classify_prompt(user_message)
classified_idx = TIER_SEVERITY_ORDER.index(classified_tier)
pools = self._tier_pools()
classified_candidates = tuple(pools.get(classified_tier.value, ()))
cold_start_candidates = tuple(
model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0
)
if cold_start_candidates:
chosen_model = random.choice(cold_start_candidates)
if request_kwargs is not None:
metadata = request_kwargs.setdefault("metadata", {})
if isinstance(metadata, dict):
metadata["adaptive_router_decision"] = {
"phase": "cold_start",
"classified_tier": classified_tier.value,
"request_type": request_type.value,
"eligible_mode": "classified_tier",
"quality_weight": self.config.adaptive_weights.quality,
"cost_weight": self.config.adaptive_weights.cost,
"tier_distance_penalty": self.config.tier_distance_penalty,
"chosen_model": chosen_model,
"candidates": [
{
"model": model,
"total_samples": adaptive._cells[(request_type, model)].total_samples,
}
for model in cold_start_candidates
],
}
return chosen_model
if self.config.adaptive_eligible == "classified_tier":
candidates = list(classified_candidates)
if not candidates:
return self.get_model_for_tier(classified_tier)
else:
candidates = list(adaptive.config.available_models)
all_costs = [adaptive.model_to_cost.get(m, 0.0) for m in candidates]
quality_weight = self.config.adaptive_weights.quality
cost_weight = self.config.adaptive_weights.cost
penalty_weight = self.config.tier_distance_penalty
best_model: str | None = None
best_score = float("-inf")
candidate_scores: list[dict[str, Any]] = []
for model in candidates:
cell = adaptive._cells[(request_type, model)]
quality_sample = thompson_sample(cell)
cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs)
if self.config.adaptive_eligible == "classified_tier":
distance = 0
else:
model_tiers = self._model_tiers.get(model, (classified_tier,))
distance = min(
abs(TIER_SEVERITY_ORDER.index(model_tier) - classified_idx) for model_tier in model_tiers
)
score = quality_weight * quality_sample + cost_weight * cost_score - penalty_weight * distance
candidate_scores.append(
{
"model": model,
"quality_sample": quality_sample,
"cost_score": cost_score,
"tier_distance": distance,
"score": score,
}
)
if score > best_score:
best_score = score
best_model = model
if best_model is None:
return self.get_model_for_tier(classified_tier)
if request_kwargs is not None:
metadata = request_kwargs.setdefault("metadata", {})
if isinstance(metadata, dict):
metadata["adaptive_router_decision"] = {
"phase": "adaptive",
"classified_tier": classified_tier.value,
"request_type": request_type.value,
"eligible_mode": self.config.adaptive_eligible,
"quality_weight": quality_weight,
"cost_weight": cost_weight,
"tier_distance_penalty": penalty_weight,
"chosen_model": best_model,
"candidates": candidate_scores,
}
return best_model
def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None:
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
Escalating to the highest tier (rather than the first rule in the list) keeps
@ -476,7 +652,7 @@ class ComplexityRouter(CustomLogger):
return None
return max(matched_tiers, key=TIER_SEVERITY_ORDER.index)
def _get_or_create_semantic_routelayer(self) -> "SemanticRouter":
def _get_or_create_semantic_routelayer(self) -> SemanticRouter:
"""Build (once) a SemanticRouter with one route per tier, utterances = that tier's keywords."""
if self._semantic_routelayer is not None:
return self._semantic_routelayer
@ -515,7 +691,7 @@ class ComplexityRouter(CustomLogger):
self._semantic_routelayer = routelayer
return routelayer
async def _ensure_semantic_routelayer(self) -> "SemanticRouter":
async def _ensure_semantic_routelayer(self) -> SemanticRouter:
"""Return the cached route layer, building it once under a lock if needed.
The build embeds the static route utterances via the encoder's synchronous path,
@ -531,7 +707,7 @@ class ComplexityRouter(CustomLogger):
routelayer = await asyncio.to_thread(self._get_or_create_semantic_routelayer)
return routelayer
async def _semantic_tier_override(self, user_message: str, request_kwargs: Dict) -> Optional[ComplexityTier]:
async def _semantic_tier_override(self, user_message: str, request_kwargs: dict) -> ComplexityTier | None:
"""Match the prompt against keyword_tier_rules by embedding similarity.
Embeds the query ourselves (instead of letting SemanticRouter.acall embed it
@ -571,7 +747,7 @@ class ComplexityRouter(CustomLogger):
except ValueError:
return None
async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: Dict) -> Optional[ComplexityTier]:
async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: dict) -> ComplexityTier | None:
"""Resolve a keyword_tier_rule override, semantically or lexically per config.
Returns None (no override -> fall through to the scorer) not only when no rule
@ -592,9 +768,9 @@ class ComplexityRouter(CustomLogger):
def _resolve_messages(
self,
messages: Optional[List[Dict[str, Any]]],
request_kwargs: Dict,
) -> Optional[List[Dict[str, Any]]]:
messages: list[dict[str, Any]] | None,
request_kwargs: dict,
) -> list[dict[str, Any]] | None:
"""
Resolve messages from the request, converting from other formats if needed.
@ -609,11 +785,11 @@ class ComplexityRouter(CustomLogger):
@staticmethod
def _extract_user_message_and_system_prompt(
messages: List[Dict[str, Any]],
) -> Tuple[Optional[str], Optional[str]]:
messages: list[dict[str, Any]],
) -> tuple[str | None, str | None]:
"""Extract the last user message text and last system prompt from messages."""
user_message: Optional[str] = None
system_prompt: Optional[str] = None
user_message: str | None = None
system_prompt: str | None = None
for msg in reversed(messages):
role = msg.get("role", "")
@ -636,11 +812,11 @@ class ComplexityRouter(CustomLogger):
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
request_kwargs: dict,
messages: list[dict[str, Any]] | None = None,
input: Union[str, list] | None = None,
specific_deployment: bool | None = False,
) -> Optional[PreRoutingHookResponse]:
"""
Pre-routing hook called before the routing decision.
@ -692,12 +868,25 @@ class ComplexityRouter(CustomLogger):
)
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
routed_model = self.get_model_for_tier(tier)
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause=complexity_scorer, tier={tier.value}, "
f"score={score:.3f}, signals={signals}, routed_model={routed_model}"
)
if self.config.adaptive:
routed_model = self._soft_floor_pick(tier, user_message, request_kwargs)
adaptive = self._ensure_adaptive_router()
if adaptive is not None:
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
chosen_key = getattr(self, "_adaptive_chosen_model_key", "adaptive_router_chosen_model")
kwargs_metadata[chosen_key] = routed_model
verbose_router_logger.info(
f"ComplexityRouter[adaptive]: routing decision cause=complexity_scorer, "
f"tier={tier.value}, score={score:.3f}, "
f"signals={signals}, routed_model={routed_model}"
)
else:
routed_model = self.get_model_for_tier(tier)
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause=complexity_scorer, tier={tier.value}, "
f"score={score:.3f}, signals={signals}, routed_model={routed_model}"
)
return PreRoutingHookResponse(
model=routed_model,

View file

@ -6,10 +6,12 @@ All values are configurable via proxy config.yaml.
"""
from enum import Enum
from typing import Dict, List, Literal, Optional
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from litellm.types.router import AdaptiveRouterWeights
class ComplexityTier(str, Enum):
"""Complexity tiers for routing decisions."""
@ -27,11 +29,13 @@ TIER_SEVERITY_ORDER: tuple[ComplexityTier, ...] = (
ComplexityTier.REASONING,
)
DEFAULT_TIER_DISTANCE_PENALTY: float = 0.5
class KeywordTierRule(BaseModel):
"""A deterministic override: if any keyword matches, route to this tier."""
keywords: List[str] = Field(
keywords: list[str] = Field(
min_length=1,
description="Keywords/phrases that trigger this rule (lexical or semantic match)",
)
@ -56,7 +60,7 @@ class KeywordTierRule(BaseModel):
# Note: Keywords should be full words/phrases to avoid substring false positives.
# The matching logic uses word boundary detection for single-word keywords.
DEFAULT_CODE_KEYWORDS: List[str] = [
DEFAULT_CODE_KEYWORDS: list[str] = [
"function",
"class",
"def",
@ -104,7 +108,7 @@ DEFAULT_CODE_KEYWORDS: List[str] = [
"pull request",
]
DEFAULT_REASONING_KEYWORDS: List[str] = [
DEFAULT_REASONING_KEYWORDS: list[str] = [
"step by step",
"think through",
"let's think",
@ -126,7 +130,7 @@ DEFAULT_REASONING_KEYWORDS: List[str] = [
"conclude",
]
DEFAULT_TECHNICAL_KEYWORDS: List[str] = [
DEFAULT_TECHNICAL_KEYWORDS: list[str] = [
"architecture",
"distributed",
"scalable",
@ -158,7 +162,7 @@ DEFAULT_TECHNICAL_KEYWORDS: List[str] = [
# Note: "async", "kubernetes", "docker" are in DEFAULT_CODE_KEYWORDS
]
DEFAULT_SIMPLE_KEYWORDS: List[str] = [
DEFAULT_SIMPLE_KEYWORDS: list[str] = [
"what is",
"what's",
"define",
@ -191,7 +195,7 @@ DEFAULT_SIMPLE_KEYWORDS: List[str] = [
# ─── Default Dimension Weights ───
DEFAULT_DIMENSION_WEIGHTS: Dict[str, float] = {
DEFAULT_DIMENSION_WEIGHTS: dict[str, float] = {
"tokenCount": 0.10, # Reduced - length is less important than content
"codePresence": 0.30, # High - code requests need capable models
"reasoningMarkers": 0.25, # High - explicit reasoning requests
@ -204,7 +208,7 @@ DEFAULT_DIMENSION_WEIGHTS: Dict[str, float] = {
# ─── Default Tier Boundaries ───
DEFAULT_TIER_BOUNDARIES: Dict[str, float] = {
DEFAULT_TIER_BOUNDARIES: dict[str, float] = {
"simple_medium": 0.15, # Lower threshold to catch more MEDIUM cases
"medium_complex": 0.35, # Lower threshold to catch technical COMPLEX cases
"complex_reasoning": 0.60, # Reasoning tier reserved for explicit reasoning markers
@ -213,7 +217,7 @@ DEFAULT_TIER_BOUNDARIES: Dict[str, float] = {
# ─── Default Token Thresholds ───
DEFAULT_TOKEN_THRESHOLDS: Dict[str, int] = {
DEFAULT_TOKEN_THRESHOLDS: dict[str, int] = {
"simple": 15, # Only very short prompts (<15 tokens) are penalized
"complex": 400, # Long prompts (>400 tokens) get complexity boost
}
@ -221,7 +225,7 @@ DEFAULT_TOKEN_THRESHOLDS: Dict[str, int] = {
# ─── Default Tier to Model Mapping ───
DEFAULT_TIER_MODELS: Dict[str, str] = {
DEFAULT_TIER_MODELS: dict[str, str] = {
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4-20250514",
@ -244,46 +248,47 @@ class ClassifierLLMConfig(BaseModel):
class ComplexityRouterConfig(BaseModel):
"""Configuration for the ComplexityRouter."""
# string = pin; list = random pick from the tier pool
# string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True
tiers: dict[str, str | list[str]] = Field(
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
description=(
"Mapping of complexity tiers to a model or model pool. A list is randomly picked from for that tier"
"Mapping of complexity tiers to a model or model pool. "
"A list is randomly picked from when adaptive=False, and used as a soft-floor home pool when adaptive=True"
),
)
# Tier boundaries (normalized scores)
tier_boundaries: Dict[str, float] = Field(
tier_boundaries: dict[str, float] = Field(
default_factory=lambda: DEFAULT_TIER_BOUNDARIES.copy(),
description="Score boundaries between tiers",
)
# Token count thresholds
token_thresholds: Dict[str, int] = Field(
token_thresholds: dict[str, int] = Field(
default_factory=lambda: DEFAULT_TOKEN_THRESHOLDS.copy(),
description="Token count thresholds for simple/complex classification",
)
# Dimension weights
dimension_weights: Dict[str, float] = Field(
dimension_weights: dict[str, float] = Field(
default_factory=lambda: DEFAULT_DIMENSION_WEIGHTS.copy(),
description="Weights for each scoring dimension",
)
# Keyword lists (overridable)
code_keywords: Optional[List[str]] = Field(
code_keywords: list[str] | None = Field(
default=None,
description="Keywords indicating code-related content",
)
reasoning_keywords: Optional[List[str]] = Field(
reasoning_keywords: list[str] | None = Field(
default=None,
description="Keywords indicating reasoning-required content",
)
technical_keywords: Optional[List[str]] = Field(
technical_keywords: list[str] | None = Field(
default=None,
description="Keywords indicating technical content",
)
custom_technical_keywords: Optional[list[str]] = Field(
custom_technical_keywords: list[str] | None = Field(
default=None,
description=(
"Domain-specific technical keywords appended to the effective base list "
@ -292,13 +297,13 @@ class ComplexityRouterConfig(BaseModel):
"the base list and within this list."
),
)
simple_keywords: Optional[List[str]] = Field(
simple_keywords: list[str] | None = Field(
default=None,
description="Keywords indicating simple/basic queries",
)
# Default model if scoring fails
default_model: Optional[str] = Field(
default_model: str | None = Field(
default=None,
description="Default model to use if tier cannot be determined",
)
@ -308,13 +313,34 @@ class ComplexityRouterConfig(BaseModel):
default="heuristic",
description="Classification strategy: local regex/keyword scoring, or an LLM call",
)
classifier_llm_config: Optional[ClassifierLLMConfig] = Field(
classifier_llm_config: ClassifierLLMConfig | None = Field(
default=None,
description="Configuration for the LLM classifier; required when classifier_type is 'llm'",
)
adaptive: bool = Field(
default=False,
description="Enable adaptive bandit selection with soft complexity floors",
)
adaptive_weights: AdaptiveRouterWeights = Field(
default_factory=lambda: AdaptiveRouterWeights(quality=0.3, cost=0.7),
description="Quality vs cost weights for adaptive selection (used when adaptive=True)",
)
tier_distance_penalty: float = Field(
default=DEFAULT_TIER_DISTANCE_PENALTY,
ge=0.0,
description="Score penalty per tier-step away from the classified tier when adaptive=True",
)
adaptive_eligible: Literal["all", "classified_tier"] = Field(
default="all",
description=(
"When adaptive=True: 'all' scores every pool model with a tier-distance penalty (soft floors); "
"'classified_tier' Thompson-samples only inside the classified tier's pool"
),
)
# Deterministic keyword -> tier overrides, evaluated before weighted scoring
keyword_tier_rules: Optional[List[KeywordTierRule]] = Field(
keyword_tier_rules: list[KeywordTierRule] | None = Field(
default=None,
description="Rules that force a specific tier when their keywords match the prompt",
)
@ -324,7 +350,7 @@ class ComplexityRouterConfig(BaseModel):
default=False,
description="Match keyword_tier_rules by embedding similarity instead of literal text",
)
embedding_model: Optional[str] = Field(
embedding_model: str | None = Field(
default=None,
description="Embedding model (LiteLLM model name) used when semantic_keyword_matching is enabled",
)
@ -358,6 +384,19 @@ class ComplexityRouterConfig(BaseModel):
raise ValueError("classifier_llm_config is required when classifier_type is 'llm'")
return self
@model_validator(mode="after")
def _validate_adaptive_pools(self) -> "ComplexityRouterConfig":
if not self.adaptive:
return self
normalized = {tier: (models if isinstance(models, list) else [models]) for tier, models in self.tiers.items()}
if not any(normalized.values()):
raise ValueError("adaptive=True requires at least one non-empty tier pool")
empty = [tier for tier, models in normalized.items() if not models]
if empty:
raise ValueError(f"adaptive=True tier pools must be non-empty; empty tiers: {empty}")
self.tiers = normalized
return self
@model_validator(mode="after")
def _validate_semantic_matching(self) -> "ComplexityRouterConfig":
if not self.semantic_keyword_matching:

View file

@ -306,7 +306,7 @@
"limit": 9
},
"TID251": {
"limit": 2710
"limit": 2701
},
"TRY002": {
"limit": 548
@ -324,7 +324,7 @@
"limit": 883
},
"UP006": {
"limit": 12869
"limit": 12792
},
"UP007": {
"limit": 2570
@ -354,7 +354,7 @@
"limit": 4
},
"UP035": {
"limit": 2295
"limit": 2284
},
"UP036": {
"limit": 4
@ -363,6 +363,6 @@
"limit": 105
},
"UP045": {
"limit": 18517
"limit": 18462
}
}

View file

@ -7,9 +7,6 @@ from litellm.router_strategy.adaptive_router import adaptive_router as ar_module
import pytest
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.router_strategy.adaptive_router.config import (
OWNER_CACHE_TTL_SECONDS,
)
from litellm.router_strategy.adaptive_router.signals import Turn
from litellm.types.router import (
AdaptiveRouterConfig,
@ -22,9 +19,7 @@ def _make_router() -> AdaptiveRouter:
cfg = AdaptiveRouterConfig(available_models=["fast", "smart"])
prefs = {
"fast": AdaptiveRouterPreferences(quality_tier=1, strengths=[]),
"smart": AdaptiveRouterPreferences(
quality_tier=3, strengths=[RequestType.CODE_GENERATION]
),
"smart": AdaptiveRouterPreferences(quality_tier=3, strengths=[RequestType.CODE_GENERATION]),
}
costs = {"fast": 0.0001, "smart": 0.001}
return AdaptiveRouter(
@ -58,85 +53,6 @@ async def test_pick_model_min_quality_tier_filter_raises_when_no_eligible():
await r.pick_model(RequestType.GENERAL, min_quality_tier=4)
@pytest.mark.asyncio
async def test_pick_model_is_stateless_no_owner_cache_writes():
"""pick_model must not touch the owner cache — that's gated post-call."""
r = _make_router()
for _ in range(5):
await r.pick_model(RequestType.GENERAL)
assert r._owner_cache == {}
# ---- claim_or_check_owner -----------------------------------------------
def test_claim_or_check_owner_first_call_claims_and_returns_true(monkeypatch):
r = _make_router()
monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0)
assert r.claim_or_check_owner("sess-A", "fast") is True
assert r._owner_cache["sess-A"] == ("fast", 1_000.0 + OWNER_CACHE_TTL_SECONDS)
assert r._skipped_updates_total == 0
def test_claim_or_check_owner_same_model_returns_true_without_extending_ttl(
monkeypatch,
):
r = _make_router()
monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0)
r.claim_or_check_owner("sess-A", "fast")
original_expiry = r._owner_cache["sess-A"][1]
monkeypatch.setattr(ar_module.time, "time", lambda: 1_500.0)
assert r.claim_or_check_owner("sess-A", "fast") is True
# No extension on hit — owner cache snapshots the first claim.
assert r._owner_cache["sess-A"][1] == original_expiry
def test_claim_or_check_owner_mismatch_skips_and_increments_counter(monkeypatch):
r = _make_router()
monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0)
r.claim_or_check_owner("sess-A", "fast")
assert r.claim_or_check_owner("sess-A", "smart") is False
assert r._skipped_updates_total == 1
# Owner unchanged.
assert r._owner_cache["sess-A"][0] == "fast"
def test_claim_or_check_owner_expired_owner_reclaims_for_new_model(monkeypatch):
r = _make_router()
monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0)
r.claim_or_check_owner("sess-A", "fast")
monkeypatch.setattr(
ar_module.time, "time", lambda: 1_000.0 + OWNER_CACHE_TTL_SECONDS + 1
)
assert r.claim_or_check_owner("sess-A", "smart") is True
assert r._owner_cache["sess-A"][0] == "smart"
# Reclaim isn't a skip.
assert r._skipped_updates_total == 0
def test_owner_cache_evicts_expired_entries_when_threshold_crossed(monkeypatch):
"""Past _OWNER_CACHE_SWEEP_THRESHOLD live entries, new claims sweep stale."""
r = _make_router()
monkeypatch.setattr(ar_module, "_OWNER_CACHE_SWEEP_THRESHOLD", 5)
monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0)
for i in range(5):
r.claim_or_check_owner(f"old-{i}", "fast")
assert len(r._owner_cache) == 5
# Jump past TTL so all "old-*" entries are now expired.
monkeypatch.setattr(
ar_module.time, "time", lambda: 1_000.0 + OWNER_CACHE_TTL_SECONDS + 1
)
r.claim_or_check_owner("new-1", "fast")
# Sweep ran -> only the new entry remains.
assert "new-1" in r._owner_cache
assert all(k.startswith("new-") for k in r._owner_cache)
# ---- record_turn --------------------------------------------------------
@ -185,9 +101,7 @@ async def test_record_turn_satisfaction_increments_alpha():
# Prime with 2 prior turns to clear the MIN_TURNS_FOR_CLEAN_CREDIT gate.
# Use distinct content to avoid incidentally firing stagnation/misalignment.
priming_turns = [
Turn(
user_content="alpha bravo charlie", assistant_content="delta echo foxtrot"
),
Turn(user_content="alpha bravo charlie", assistant_content="delta echo foxtrot"),
Turn(
user_content="golf hotel india juliet",
assistant_content="kilo lima mike november",
@ -232,6 +146,128 @@ async def test_record_turn_failure_increments_beta():
assert cell_after.alpha == pytest.approx(cell_before.alpha)
@pytest.mark.asyncio
async def test_record_turn_detects_exhaustion_in_tool_results():
r = _make_router()
delta = await r.record_turn(
session_id="exhausted",
model_name="smart",
request_type=RequestType.GENERAL,
turn=Turn(tool_results=[{"content": "rate limit exceeded"}]),
)
assert delta.exhaustion == 1
assert r._session_states[("exhausted", "smart")].exhaustion_count == 1
@pytest.mark.asyncio
async def test_record_turn_attributes_user_feedback_to_previous_response_model():
r = _make_router()
fast_before = r._cells[(RequestType.CODE_GENERATION, "fast")]
smart_before = r._cells[(RequestType.GENERAL, "smart")]
await r.record_turn(
session_id="feedback-switch",
model_name="fast",
request_type=RequestType.CODE_GENERATION,
turn=Turn(
user_content="fix this python retry bug",
assistant_content="clear the cache on every retry",
),
)
await r.record_turn(
session_id="feedback-switch",
model_name="smart",
request_type=RequestType.GENERAL,
turn=Turn(
user_content="the python fix is still broken",
assistant_content="keep successful cache entries",
),
)
fast_after = r._cells[(RequestType.CODE_GENERATION, "fast")]
smart_after = r._cells[(RequestType.GENERAL, "smart")]
assert fast_after.beta == pytest.approx(fast_before.beta + 1.0)
assert smart_after.beta == pytest.approx(smart_before.beta)
snapshot = await r.get_state_snapshot()
assert snapshot["feedback_attributed_total"] == 1
assert snapshot["cross_model_feedback_total"] == 1
assert snapshot["feedback_without_context_total"] == 0
@pytest.mark.asyncio
async def test_record_turn_attributes_satisfaction_to_previous_response_model():
r = _make_router()
await r.record_turn(
session_id="satisfaction-switch",
model_name="smart",
request_type=RequestType.CODE_GENERATION,
turn=Turn(
user_content="write a python retry helper",
assistant_content="first draft",
),
)
await r.record_turn(
session_id="satisfaction-switch",
model_name="fast",
request_type=RequestType.CODE_GENERATION,
turn=Turn(
user_content="add exponential backoff to the python helper",
assistant_content="updated draft",
),
)
fast_before = r._cells[(RequestType.CODE_GENERATION, "fast")]
smart_before = r._cells[(RequestType.GENERAL, "smart")]
await r.record_turn(
session_id="satisfaction-switch",
model_name="smart",
request_type=RequestType.GENERAL,
turn=Turn(
user_content="thanks, that worked",
assistant_content="glad to help",
),
)
fast_after = r._cells[(RequestType.CODE_GENERATION, "fast")]
smart_after = r._cells[(RequestType.GENERAL, "smart")]
assert fast_after.alpha == pytest.approx(fast_before.alpha + 1.0)
assert smart_after.alpha == pytest.approx(smart_before.alpha)
@pytest.mark.asyncio
async def test_record_turn_bounds_feedback_contexts_and_evicts_least_recent_session():
r = _make_router()
context_limit = ar_module._FEEDBACK_CONTEXT_MAX_ENTRIES
for index in range(context_limit):
await r.record_turn(
session_id=f"session-{index}",
model_name="fast",
request_type=RequestType.GENERAL,
turn=Turn(user_content="question", assistant_content="answer"),
)
await r.record_turn(
session_id="session-0",
model_name="fast",
request_type=RequestType.GENERAL,
turn=Turn(user_content="follow up", assistant_content="updated answer"),
)
await r.record_turn(
session_id="overflow",
model_name="fast",
request_type=RequestType.GENERAL,
turn=Turn(user_content="question", assistant_content="answer"),
)
assert len(r._feedback_contexts) == context_limit
assert "session-0" in r._feedback_contexts
assert "session-1" not in r._feedback_contexts
assert "overflow" in r._feedback_contexts
@pytest.mark.asyncio
async def test_load_state_from_db_overrides_cold_start():
r = _make_router()
@ -270,9 +306,7 @@ async def test_load_state_from_db_handles_unknown_request_type():
good_row.beta = 3.0
prisma = MagicMock()
prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(
return_value=[bad_row, good_row]
)
prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[bad_row, good_row])
await r.load_state_from_db(prisma)
# Unknown skipped; good applied.

View file

@ -110,23 +110,6 @@ async def test_pick_record_flush_full_cycle():
assert session_call.kwargs["data"]["create"]["model_name"] == chosen
@pytest.mark.asyncio
async def test_owner_cache_pins_attribution_to_first_picked_model():
"""First call claims ownership; matching model returns True, mismatch False."""
router = _make_router()
chosen = await router.pick_model(RequestType.GENERAL)
assert router.claim_or_check_owner("sess-own", chosen) is True
# Same model on later turns keeps attributing.
for _ in range(5):
assert router.claim_or_check_owner("sess-own", chosen) is True
# A different model on a later turn is rejected.
other = "gpt-4o" if chosen == "gpt-4o-mini" else "gpt-4o-mini"
assert router.claim_or_check_owner("sess-own", other) is False
assert router._skipped_updates_total == 1
@pytest.mark.asyncio
async def test_pick_model_returns_valid_models_without_error():
router = _make_router()

View file

@ -16,10 +16,9 @@ from litellm.router_strategy.adaptive_router.hooks import (
from litellm.router_strategy.adaptive_router.signals import Turn
def _make_hook(claim: bool = True) -> AdaptiveRouterPostCallHook:
def _make_hook() -> AdaptiveRouterPostCallHook:
fake_router = MagicMock()
fake_router.record_turn = AsyncMock()
fake_router.claim_or_check_owner = MagicMock(return_value=claim)
return AdaptiveRouterPostCallHook(adaptive_router=fake_router)
@ -151,7 +150,24 @@ async def test_hook_skips_when_below_signal_gate():
kwargs = _kwargs(messages=short)
await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0)
hook.adaptive_router.record_turn.assert_not_awaited()
hook.adaptive_router.claim_or_check_owner.assert_not_called()
@pytest.mark.asyncio
async def test_hook_tracks_short_conversation_with_explicit_session_id():
hook = _make_hook()
kwargs = _kwargs(
messages=[{"role": "user", "content": "hi"}],
extra_litellm_params={"litellm_session_id": "explicit-short"},
)
await hook.async_log_success_event(
kwargs,
_resp_with_content("hello"),
0.0,
1.0,
)
assert hook.adaptive_router.record_turn.await_args.kwargs["session_id"] == (
"explicit-short"
)
@pytest.mark.asyncio
@ -168,22 +184,19 @@ async def test_hook_skips_when_chosen_model_missing_from_metadata():
kwargs = _kwargs(chosen=None)
await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0)
hook.adaptive_router.record_turn.assert_not_awaited()
hook.adaptive_router.claim_or_check_owner.assert_not_called()
@pytest.mark.asyncio
async def test_hook_skips_when_owner_cache_mismatch():
"""A different model owns this conversation -> no attribution."""
hook = _make_hook(claim=False)
async def test_hook_records_when_model_changes():
hook = _make_hook()
kwargs = _kwargs(chosen="fast")
await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0)
hook.adaptive_router.claim_or_check_owner.assert_called_once()
hook.adaptive_router.record_turn.assert_not_awaited()
hook.adaptive_router.record_turn.assert_awaited_once()
@pytest.mark.asyncio
async def test_hook_records_turn_when_owner_claims():
hook = _make_hook(claim=True)
async def test_hook_records_turn():
hook = _make_hook()
kwargs = _kwargs(chosen="smart", messages=_long_messages("ask"))
await hook.async_log_success_event(
kwargs, _resp_with_content("answer here"), 0.0, 1.0
@ -205,8 +218,6 @@ async def test_hook_uses_explicit_session_id_when_provided():
extra_litellm_params={"litellm_session_id": "explicit-sess"},
)
await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0)
args, _ = hook.adaptive_router.claim_or_check_owner.call_args
assert args[0] == "explicit-sess"
assert hook.adaptive_router.record_turn.await_args.kwargs["session_id"] == (
"explicit-sess"
)

View file

@ -1,7 +1,6 @@
"""Tests for the GET /adaptive_router/state introspection endpoint and the
underlying `AdaptiveRouter.get_state_snapshot()` helper."""
import time
from unittest.mock import MagicMock
import pytest
@ -9,7 +8,7 @@ from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.router_strategy.adaptive_router.bandit import BanditCell, apply_delta
from litellm.router_strategy.adaptive_router.bandit import apply_delta
from litellm.types.router import (
AdaptiveRouterConfig,
AdaptiveRouterPreferences,
@ -47,8 +46,6 @@ async def test_get_state_snapshot_returns_cell_per_request_type_per_model():
assert snap["available_models"] == ["fast", "smart"]
assert snap["weights"] == {"quality": 0.7, "cost": 0.3}
assert snap["model_costs"] == {"fast": 0.0001, "smart": 0.001}
assert snap["owner_cache_live"] == 0
assert snap["skipped_updates_total"] == 0
assert set(snap["queue"].keys()) == {
"state_pending",
"session_pending",
@ -95,26 +92,6 @@ async def test_get_state_snapshot_quality_mean_matches_alpha_over_total():
assert cell["quality_mean"] == pytest.approx(expected_mean)
@pytest.mark.asyncio
async def test_get_state_snapshot_counts_only_live_owner_cache_entries():
r = _make_router()
now = time.time()
r._owner_cache["live-1"] = ("fast", now + 3600)
r._owner_cache["live-2"] = ("smart", now + 3600)
r._owner_cache["expired-1"] = ("fast", now - 1)
snap = await r.get_state_snapshot()
assert snap["owner_cache_live"] == 2
@pytest.mark.asyncio
async def test_get_state_snapshot_exposes_skipped_updates_total():
r = _make_router()
r._skipped_updates_total = 7
snap = await r.get_state_snapshot()
assert snap["skipped_updates_total"] == 7
# ---- endpoint --------------------------------------------------------

View file

@ -948,6 +948,60 @@ class TestRouterComplexityDeploymentMethods:
router.init_complexity_router_deployment(deployment)
assert "auto_router/complexity_router/test-router" in router.complexity_routers
def test_hybrid_initialization_waits_for_later_pool_deployments(self):
router = Router(
model_list=[
{
"model_name": "hybrid",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_default_model": "cheap",
"complexity_router_config": {
"adaptive": True,
"tiers": {
"SIMPLE": ["cheap"],
"MEDIUM": ["cheap", "premium"],
},
},
},
},
{
"model_name": "cheap",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"input_cost_per_token": 0.00000015,
},
"model_info": {
"adaptive_router_preferences": {
"quality_tier": 1,
"strengths": [],
}
},
},
{
"model_name": "premium",
"litellm_params": {
"model": "openai/gpt-4o",
"input_cost_per_token": 0.000005,
},
"model_info": {
"adaptive_router_preferences": {
"quality_tier": 3,
"strengths": [],
}
},
},
]
)
adaptive = router.adaptive_routers["hybrid"]
assert adaptive.model_to_cost == {
"cheap": pytest.approx(0.00000015),
"premium": pytest.approx(0.000005),
}
assert adaptive.model_to_prefs["cheap"].quality_tier == 1
assert adaptive.model_to_prefs["premium"].quality_tier == 3
class TestAsyncPreRoutingHookMultiFormat:
"""Test async_pre_routing_hook with multiple input formats."""
@ -1356,6 +1410,240 @@ class TestLLMClassifier:
assert call_kwargs["metadata"] == request_metadata
class TestAdaptiveSoftFloors:
def test_adaptive_defaults_use_cost_weighted_cold_policy(self):
config = ComplexityRouterConfig(
adaptive=True,
tiers={"SIMPLE": ["cheap"]},
)
assert config.adaptive_weights.quality == pytest.approx(0.3)
assert config.adaptive_weights.cost == pytest.approx(0.7)
assert config.tier_distance_penalty == pytest.approx(0.5)
@pytest.fixture
def adaptive_router_instance(self):
router = MagicMock()
router.model_list = [
{
"model_name": "cheap",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"input_cost_per_token": 0.00000015,
},
"model_info": {
"adaptive_router_preferences": {"quality_tier": 1, "strengths": []}
},
},
{
"model_name": "premium",
"litellm_params": {
"model": "openai/gpt-4o",
"input_cost_per_token": 0.000005,
},
"model_info": {
"adaptive_router_preferences": {"quality_tier": 3, "strengths": []}
},
},
]
router.model_name_to_deployment_indices = {"cheap": [0], "premium": [1]}
return router
@pytest.fixture
def hybrid_config(self) -> Dict:
return {
"adaptive": True,
"adaptive_weights": {"quality": 0.7, "cost": 0.3},
"tier_distance_penalty": 0.15,
"tiers": {
"SIMPLE": ["cheap"],
"MEDIUM": ["cheap"],
"COMPLEX": ["premium"],
"REASONING": ["premium"],
},
"default_model": "cheap",
}
def test_adaptive_config_requires_non_empty_pools(self):
with pytest.raises(ValidationError):
ComplexityRouterConfig(adaptive=True, tiers={"SIMPLE": []})
def test_cold_start_randomly_samples_unobserved_classified_tier_models(
self, adaptive_router_instance
):
cr = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=adaptive_router_instance,
complexity_router_config={
"adaptive": True,
"tiers": {
"SIMPLE": ["cheap", "premium"],
"MEDIUM": ["premium"],
},
},
)
request_kwargs: Dict = {"metadata": {}}
with patch(
"litellm.router_strategy.complexity_router.complexity_router.random.choice",
return_value="premium",
) as choice:
picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi", request_kwargs)
assert picked == "premium"
choice.assert_called_once_with(("cheap", "premium"))
decision = request_kwargs["metadata"]["adaptive_router_decision"]
assert decision["phase"] == "cold_start"
assert {candidate["model"] for candidate in decision["candidates"]} == {
"cheap",
"premium",
}
def test_get_model_for_tier_list_without_adaptive_random_choice(
self, mock_router_instance
):
router = ComplexityRouter(
model_name="test",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"adaptive": False,
"tiers": {"SIMPLE": ["cheap", "premium"], "MEDIUM": "mid"},
"default_model": "mid",
},
)
pool = ["cheap", "premium"]
with patch(
"litellm.router_strategy.complexity_router.complexity_router.random.choice",
return_value="premium",
) as choice:
assert router.get_model_for_tier(ComplexityTier.SIMPLE) == "premium"
choice.assert_called_once_with(pool)
assert router.get_model_for_tier(ComplexityTier.MEDIUM) == "mid"
def test_soft_floor_prefers_home_tier_when_posteriors_equal(
self, adaptive_router_instance, hybrid_config
):
from litellm.router_strategy.adaptive_router.bandit import BanditCell
from litellm.types.router import RequestType
cr = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=adaptive_router_instance,
complexity_router_config=hybrid_config,
)
adaptive = cr._ensure_adaptive_router()
assert adaptive is not None
for model in ("cheap", "premium"):
adaptive._cells[(RequestType.GENERAL, model)] = BanditCell(
alpha=5.0, beta=5.0
)
# Equal quality samples; home-tier penalty should favor cheap for SIMPLE.
with patch(
"litellm.router_strategy.adaptive_router.bandit.thompson_sample",
return_value=0.5,
):
picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi")
assert picked == "cheap"
def test_soft_floor_allows_cross_tier_when_posterior_dominates(
self, adaptive_router_instance, hybrid_config
):
from litellm.router_strategy.adaptive_router.bandit import BanditCell
from litellm.types.router import RequestType
cr = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=adaptive_router_instance,
complexity_router_config=hybrid_config,
)
adaptive = cr._ensure_adaptive_router()
assert adaptive is not None
adaptive._cells[(RequestType.GENERAL, "cheap")] = BanditCell(
alpha=1.0, beta=20.0
)
adaptive._cells[(RequestType.GENERAL, "premium")] = BanditCell(
alpha=20.0, beta=1.0
)
with patch(
"litellm.router_strategy.adaptive_router.bandit.thompson_sample",
side_effect=lambda cell, rng=None: cell.alpha / (cell.alpha + cell.beta),
):
picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi")
assert picked == "premium"
def test_reused_model_has_zero_distance_in_each_configured_tier(
self, adaptive_router_instance
):
from litellm.router_strategy.adaptive_router.bandit import BanditCell
from litellm.types.router import RequestType
cr = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=adaptive_router_instance,
complexity_router_config={
"adaptive": True,
"tiers": {
"SIMPLE": ["cheap"],
"MEDIUM": ["cheap", "premium"],
"COMPLEX": ["premium"],
},
},
)
adaptive = cr._ensure_adaptive_router()
assert adaptive is not None
for model in ("cheap", "premium"):
adaptive._cells[(RequestType.GENERAL, model)] = BanditCell(
alpha=6.0, beta=5.0
)
request_kwargs: Dict = {"metadata": {}}
with patch(
"litellm.router_strategy.adaptive_router.bandit.thompson_sample",
return_value=0.5,
):
cr._soft_floor_pick(ComplexityTier.MEDIUM, "hi", request_kwargs)
candidates = request_kwargs["metadata"]["adaptive_router_decision"][
"candidates"
]
assert {
candidate["model"]: candidate["tier_distance"] for candidate in candidates
} == {
"cheap": 0,
"premium": 0,
}
@pytest.mark.asyncio
async def test_pre_routing_hook_adaptive_stashes_chosen_model(
self, adaptive_router_instance, hybrid_config
):
cr = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=adaptive_router_instance,
complexity_router_config=hybrid_config,
)
request_kwargs: Dict = {"metadata": {}}
result = await cr.async_pre_routing_hook(
model="hybrid",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hi"}],
)
assert result is not None
assert result.model in {"cheap", "premium"}
assert (
request_kwargs["metadata"].get("adaptive_router_chosen_model")
== result.model
)
decision = request_kwargs["metadata"]["adaptive_router_decision"]
assert decision["phase"] == "cold_start"
assert decision["classified_tier"] == "SIMPLE"
assert decision["request_type"] == "general"
assert decision["eligible_mode"] == "classified_tier"
assert decision["chosen_model"] == result.model
assert {candidate["model"] for candidate in decision["candidates"]} == {"cheap"}
class TestLexicalKeywordTierRules:
"""Test deterministic (literal) keyword_tier_rules overrides."""

View file

@ -111,7 +111,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }
? (value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS })
: undefined,
});
};

View file

@ -43,8 +43,7 @@ export const getSemanticConfigError = ({
embeddingModel,
keywordTierRules,
}: Pick<BuildComplexityRouterConfigParams, "semanticMatchingEnabled" | "embeddingModel" | "keywordTierRules">):
| string
| null => {
string | null => {
if (!semanticMatchingEnabled) return null;
if (!embeddingModel) return "Select an embedding model to use semantic keyword matching";
if (keywordTierRules.length === 0) return "Add at least one keyword tier rule to use semantic keyword matching";