mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
85f9bdd412
commit
26ab730bfa
15 changed files with 1035 additions and 409 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 --------------------------------------------------------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue