fix(router): balance the type-discipline budget and cover the knob helpers by name

- Mapping annotations and Final locals keep every LIT rule at or below the base count
- the stash helper builds under a new name; its mutable copy carries a reason
- direct helper calls in the test satisfy the router code-coverage name check
This commit is contained in:
Paolo Antinori 2026-09-18 17:45:15 +02:00
parent f9a1359ca2
commit 2210f948ab
No known key found for this signature in database
3 changed files with 72 additions and 25 deletions

View file

@ -721,6 +721,10 @@ _FINISH_REASON_FAILURE_EXCEPTION_NAMES: Final = frozenset(
}
)
## Healthy terminal reasons in the mapped OpenAI set: keys of treat_finish_reason_as_failure that
## name one of these would fail successful responses, so construction warns about them.
_HEALTHY_TERMINAL_FINISH_REASONS: Final = frozenset(("stop", "length", "tool_calls", "function_call"))
class Router:
model_names: set = set()
@ -766,7 +770,7 @@ class Router:
fallbacks: list = [],
context_window_fallbacks: list = [],
content_policy_fallbacks: list = [],
treat_finish_reason_as_failure: dict[str, str] | None = None,
treat_finish_reason_as_failure: Mapping[str, str] | None = None,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = {},
enable_pre_call_checks: bool = False,
enable_tag_filtering: bool = False,
@ -1102,12 +1106,7 @@ class Router:
verbose_router_logger.warning(
"treat_finish_reason_as_failure applies to non-streaming responses only; a streamed 200 with the mapped stop reason is delivered unchanged."
)
healthy_terminal_keys: Final = treat_finish_reason_as_failure.keys() & {
"stop",
"length",
"tool_calls",
"function_call",
}
healthy_terminal_keys: Final = treat_finish_reason_as_failure.keys() & _HEALTHY_TERMINAL_FINISH_REASONS
if healthy_terminal_keys:
verbose_router_logger.warning(
"treat_finish_reason_as_failure keys %s are healthy terminal reasons in the mapped OpenAI set; mapping them fails successful responses. Keys are matched against provider-native stop reasons.",
@ -2593,7 +2592,7 @@ class Router:
## CHECK MAPPED FINISH REASON ERROR ##
if isinstance(response, ModelResponse):
_mapped_reason = self._get_mapped_finish_reason(response)
_mapped_reason: Final = self._get_mapped_finish_reason(response)
if _mapped_reason is not None:
self._handle_mapped_finish_reason_failure(
model=model, deployment=deployment, reason=_mapped_reason, kwargs=kwargs
@ -3725,7 +3724,7 @@ class Router:
## CHECK MAPPED FINISH REASON ERROR ##
if isinstance(response, ModelResponse):
_mapped_reason = self._get_mapped_finish_reason(response)
_mapped_reason: Final = self._get_mapped_finish_reason(response)
if _mapped_reason is not None:
self._handle_mapped_finish_reason_failure(
model=model, deployment=deployment, reason=_mapped_reason, kwargs=kwargs
@ -8665,7 +8664,10 @@ class Router:
choice: Final = response.choices[0]
if choice.finish_reason in self.treat_finish_reason_as_failure:
return choice.finish_reason
native_reason: Final = (getattr(choice, "provider_specific_fields", None) or {}).get("native_finish_reason")
_provider_specific_fields: Final = getattr(choice, "provider_specific_fields", None)
native_reason: Final = (
_provider_specific_fields.get("native_finish_reason") if _provider_specific_fields else None
)
if native_reason in self.treat_finish_reason_as_failure:
return native_reason
return None
@ -8688,7 +8690,9 @@ class Router:
)
return resolved is not None
def _handle_mapped_finish_reason_failure(self, model: str, deployment: dict, reason: str, kwargs: dict) -> None:
def _handle_mapped_finish_reason_failure(
self, model: str, deployment: Mapping[str, Any], reason: str, kwargs: Mapping[str, Any]
) -> None:
"""
Account for a mapped finish-reason failure, then raise the configured exception into the
fallback chain when a generic fallback can serve. Accounting happens before the gate:
@ -8712,7 +8716,7 @@ class Router:
return exception_cls(message=message, llm_provider="", model=model)
def _account_mapped_finish_reason_failure(
self, model: str, deployment: dict, reason: str, kwargs: dict
self, model: str, deployment: Mapping[str, Any], reason: str, kwargs: Mapping[str, Any]
) -> Exception | None:
"""
Count and park a mapped finish-reason failure: increment the per-minute failure counter
@ -8722,18 +8726,20 @@ class Router:
exception so the caller can raise the same instance it accounted for, or None when the
deployment has no id to account against.
"""
model_info: Final = deployment.get("model_info") or {}
deployment_id: Final = model_info.get("id") if isinstance(model_info, dict) else None
raw_model_info: Final = deployment.get("model_info")
model_info: Final = raw_model_info if isinstance(raw_model_info, dict) else None
deployment_id: Final = model_info.get("id") if model_info is not None else None
if deployment_id is None:
return None
litellm_params: Final = deployment.get("litellm_params") or {}
deployment_cooldown: Final = _first_present(
model_info if isinstance(model_info, dict) else None, litellm_params, key="cooldown_time"
)
raw_litellm_params: Final = deployment.get("litellm_params")
litellm_params: Final = raw_litellm_params if isinstance(raw_litellm_params, dict) else None
deployment_cooldown: Final = _first_present(model_info, litellm_params, key="cooldown_time")
time_to_cooldown: Final = (
deployment_cooldown if deployment_cooldown is not None and deployment_cooldown >= 0 else self.cooldown_time
)
exception: Final = self._finish_reason_failure_error(model=model, reason=reason)
litellm_metadata: Final = get_litellm_metadata_from_kwargs(kwargs)
requested_model_group: Final = litellm_metadata.get("model_group") if litellm_metadata else None
increment_deployment_failures_for_current_minute(
litellm_router_instance=self,
deployment_id=deployment_id,
@ -8744,7 +8750,7 @@ class Router:
original_exception=exception,
deployment=deployment_id,
time_to_cooldown=time_to_cooldown,
requested_model_group=(get_litellm_metadata_from_kwargs(kwargs) or {}).get("model_group"),
requested_model_group=requested_model_group,
)
return exception

View file

@ -1568,16 +1568,17 @@ class Delta(SafeAttributeModel, OpenAIObject):
def map_finish_reason_and_stash_native(
finish_reason: str, provider_specific_fields: dict[str, Any] | None
) -> tuple[OpenAIChatCompletionFinishReason, dict[str, Any] | None]:
finish_reason: str, provider_specific_fields: Mapping[str, Any] | None
) -> tuple[OpenAIChatCompletionFinishReason, dict[str, Any] | None]: # mutable-ok: callers extend the returned stash
"""Map a provider-native finish reason to the OpenAI set; when the native value differs
from the mapped one, preserve it under provider_specific_fields["native_finish_reason"]
so downstream consumers can still see what the provider actually sent."""
mapped: Final = map_finish_reason(finish_reason)
if finish_reason != mapped:
provider_specific_fields = dict(provider_specific_fields) if provider_specific_fields else {}
provider_specific_fields["native_finish_reason"] = finish_reason
return mapped, provider_specific_fields
if finish_reason == mapped:
return mapped, provider_specific_fields
stash: Final = dict(provider_specific_fields or ()) # mutable-ok: the stash must stay a plain extensible dict
stash["native_finish_reason"] = finish_reason
return mapped, stash
class Choices(SafeAttributeModel, OpenAIObject):

View file

@ -16,6 +16,7 @@ import httpx
import pytest
from pytest import MonkeyPatch
import litellm
from litellm import Router
from litellm.router_utils.cooldown_handlers import _get_cooldown_deployments
@ -172,3 +173,42 @@ async def test_knob_unset_ignores_terminal_stop_reason(monkeypatch: MonkeyPatch)
def test_unknown_exception_name_raises_at_construction():
with pytest.raises(ValueError, match="NotAnException"):
Router(model_list=[], treat_finish_reason_as_failure={"x": "NotAnException"})
@pytest.mark.asyncio
async def test_mapped_finish_reason_helpers_direct(monkeypatch: MonkeyPatch):
"""Direct coverage of the knob helpers (the router code-coverage check matches by name)."""
fake = FakeAnthropicUpstream()
router = Router(
model_list=[FABLE_TIER, OPUS_TARGET],
treat_finish_reason_as_failure=_knob(),
default_fallbacks=["opus-target"],
num_retries=0,
allowed_fails=0,
cooldown_time=10,
)
fake.install(monkeypatch)
ok = await router.acompletion(model="opus-target", max_tokens=16, messages=[{"role": "user", "content": "hi"}])
assert router._get_mapped_finish_reason(ok) is None
assert router._generic_fallback_available("fable-tier", {}) is True
error = router._finish_reason_failure_error(model="fable-tier", reason="model_context_window_exceeded")
assert error.status_code == 429
deployment = router.model_list[0]
accounted = router._account_mapped_finish_reason_failure(
model="fable-tier",
deployment=deployment,
reason="model_context_window_exceeded",
kwargs={},
)
assert accounted is not None
with pytest.raises(litellm.RateLimitError):
router._handle_mapped_finish_reason_failure(
model="fable-tier",
deployment=deployment,
reason="model_context_window_exceeded",
kwargs={},
)