From 26ab730bfaac36b6d96af68d5fe5e7eb867af2ca Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 11 Jul 2026 21:56:33 -0700 Subject: [PATCH] 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 * 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 * chore(router): drop unnecessary hybrid docstrings Co-authored-by: Cursor * 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 * 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 * 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 * fix(router): bound feedback context cache Cap retained session feedback so unique session IDs cannot exhaust router memory Co-authored-by: Cursor * fix(router): preserve exhaustion signals Include tool-result exhaustion in adaptive feedback and clear strict lint regressions blocking CI Co-authored-by: Cursor * 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 * refactor(router): centralize hook cleanup Use the callback manager to discover and remove adaptive hooks across every registered callback list Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/router.py | 29 +- .../router_strategy/adaptive_router/README.md | 15 +- .../adaptive_router/adaptive_router.py | 290 +++++++++++------- .../router_strategy/adaptive_router/hooks.py | 8 - .../adaptive_router/signals.py | 148 ++++++--- .../complexity_router/complexity_router.py | 271 +++++++++++++--- .../complexity_router/config.py | 87 ++++-- ruff-strict-budget.json | 8 +- .../adaptive_router/test_adaptive_router.py | 216 +++++++------ .../test_e2e_adaptive_router.py | 17 - .../adaptive_router/test_hooks.py | 37 ++- .../adaptive_router/test_state_endpoint.py | 25 +- .../router_strategy/test_complexity_router.py | 288 +++++++++++++++++ .../add_model/ComplexityRouterConfig.tsx | 2 +- .../build_complexity_router_config.ts | 3 +- 15 files changed, 1035 insertions(+), 409 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 245a50545e7..6539d3c0c43 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_strategy/adaptive_router/README.md b/litellm/router_strategy/adaptive_router/README.md index 7f5d7aa21d0..09420a8dd9d 100644 --- a/litellm/router_strategy/adaptive_router/README.md +++ b/litellm/router_strategy/adaptive_router/README.md @@ -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 diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 69d6a019e68..ec84eb1decf 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -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. diff --git a/litellm/router_strategy/adaptive_router/hooks.py b/litellm/router_strategy/adaptive_router/hooks.py index c3e3f8ca74a..89ae28be227 100644 --- a/litellm/router_strategy/adaptive_router/hooks.py +++ b/litellm/router_strategy/adaptive_router/hooks.py @@ -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) diff --git a/litellm/router_strategy/adaptive_router/signals.py b/litellm/router_strategy/adaptive_router/signals.py index 2fd1d24fbbe..74fa8936098 100644 --- a/litellm/router_strategy/adaptive_router/signals.py +++ b/litellm/router_strategy/adaptive_router/signals.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 74644f01be8..bebdbba90ef 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 8c8e5acb51f..df699d1a059 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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: diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 7750ac6628a..dcde6fd1641 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -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 } } diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py index 93c4db90dad..cbf5635a5ae 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py @@ -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. diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py index 9786832b4ae..3071f916ef1 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py @@ -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() diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py b/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py index a2b85f2ce53..ad61f43c5a0 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py @@ -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" ) diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py b/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py index 753a449791b..d6d89c8e811 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py @@ -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 -------------------------------------------------------- diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index e1133620a57..da02b774e41 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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.""" diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 50168056eaa..c31ee41a6ec 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -111,7 +111,7 @@ const ComplexityRouterConfig: React.FC = ({ 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, }); }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 3eddca8c35b..a4a8ee6b074 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -43,8 +43,7 @@ export const getSemanticConfigError = ({ embeddingModel, keywordTierRules, }: Pick): - | 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";