fix(router): avoid typed decision expansion errors

This commit is contained in:
mikemikimike 2026-08-30 09:44:16 +08:00
parent dd6a7dfd9d
commit 31d5460171
2 changed files with 43 additions and 40 deletions

View file

@ -16,7 +16,6 @@ from __future__ import annotations
import asyncio
import time
from collections import OrderedDict
from collections.abc import Mapping
from dataclasses import asdict, dataclass
from typing import TYPE_CHECKING, Any, Final, cast
@ -204,27 +203,30 @@ class AdaptiveRouter:
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = chosen_model
routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.router_name,
router_type="adaptive",
routed_model=chosen_model,
cause="bandit",
request_type=request_type.value,
conversation_continuing=conversation_is_continuing(messages),
)
routing_decision.update(self._savings_baseline_fields())
return PreRoutingHookResponse(
model=chosen_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.router_name,
router_type="adaptive",
routed_model=chosen_model,
cause="bandit",
request_type=request_type.value,
conversation_continuing=conversation_is_continuing(messages),
**self._savings_baseline_fields(),
),
routing_decision=routing_decision,
)
def _savings_baseline_fields(self) -> Mapping[str, str]:
def _savings_baseline_fields(self) -> StandardLoggingRoutingDecision:
if self.litellm_router_instance is None:
return {} # mutable-ok: immutable empty result for absent router
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {} # mutable-ok: immutable empty result for unresolved baseline
fields = {"savings_baseline_model": baseline.model} # mutable-ok: assemble TypedDict kwargs
fields: StandardLoggingRoutingDecision = {
"savings_baseline_model": baseline.model,
} # mutable-ok: assemble TypedDict kwargs
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields

View file

@ -16,7 +16,6 @@ then cheapest `model_info.input_cost_per_token`).
"""
import math
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Optional
from litellm._logging import verbose_router_logger
@ -323,11 +322,13 @@ class QualityRouter(CustomLogger):
if isinstance(metadata, dict):
metadata["quality_router_decision"] = decision
def _savings_fields(self) -> Mapping[str, str]:
def _savings_fields(self) -> StandardLoggingRoutingDecision:
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {} # mutable-ok: immutable empty result for unresolved baseline
fields = {"savings_baseline_model": baseline.model} # mutable-ok: assemble TypedDict kwargs
fields: StandardLoggingRoutingDecision = {
"savings_baseline_model": baseline.model,
} # mutable-ok: assemble TypedDict kwargs
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields
@ -370,17 +371,18 @@ class QualityRouter(CustomLogger):
verbose_router_logger.debug("QualityRouter: No user message found, routing to default model")
if not self.config.default_model:
raise ValueError("QualityRouter: no user message and no default_model configured")
default_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
conversation_continuing=conversation_is_continuing(messages),
)
default_routing_decision.update(self._savings_fields())
return PreRoutingHookResponse(
model=self.config.default_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
),
routing_decision=default_routing_decision,
)
# Try keyword override first — it short-circuits complexity classification.
@ -404,25 +406,24 @@ class QualityRouter(CustomLogger):
"quality_tier": self._model_quality.get(routed_model),
"complexity_tier": None,
"conversation_continuing": conversation_is_continuing(messages),
**self._savings_fields(),
},
)
routing_decision: Final = StandardLoggingRoutingDecision(
keyword_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="keyword",
matched_keyword=matched_keyword,
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
)
keyword_routing_decision.update(self._savings_fields())
keyword_quality_tier: Final = self._model_quality.get(routed_model)
if keyword_quality_tier is not None:
routing_decision["tier"] = str(keyword_quality_tier)
keyword_routing_decision["tier"] = str(keyword_quality_tier)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
routing_decision=routing_decision,
routing_decision=keyword_routing_decision,
)
# No keyword match → complexity classification flow.
@ -454,22 +455,22 @@ class QualityRouter(CustomLogger):
"quality_tier": int(quality_tier),
"complexity_tier": complexity_name,
"conversation_continuing": conversation_is_continuing(messages),
**self._savings_fields(),
},
)
quality_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="quality_tier",
tier=str(int(quality_tier)),
score=score,
signals=list(signals),
conversation_continuing=conversation_is_continuing(messages),
)
quality_routing_decision.update(self._savings_fields())
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="quality_tier",
tier=str(int(quality_tier)),
score=score,
signals=list(signals),
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
),
routing_decision=quality_routing_decision,
)