mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
f9a1359ca2
commit
2210f948ab
3 changed files with 72 additions and 25 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue