feat(router): treat_finish_reason_as_failure config surface and helpers

This commit is contained in:
Paolo Antinori 2026-09-18 15:27:18 +02:00
parent 8fc9c46d1a
commit f6ec7dfeb9
No known key found for this signature in database

View file

@ -752,6 +752,7 @@ class Router:
fallbacks: list = [],
context_window_fallbacks: list = [],
content_policy_fallbacks: list = [],
treat_finish_reason_as_failure: dict[str, str] | None = None,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = {},
enable_pre_call_checks: bool = False,
enable_tag_filtering: bool = False,
@ -1072,6 +1073,30 @@ class Router:
_content_policy_fallbacks: Final = content_policy_fallbacks or litellm.content_policy_fallbacks
self.validate_fallbacks(fallback_param=_content_policy_fallbacks)
self.content_policy_fallbacks = _content_policy_fallbacks
## treat_finish_reason_as_failure: map a terminal finish/stop reason on a 200 response to a
## router-understood exception class, so the mapped reason engages allowed_fails/cooldowns/
## fallbacks like any failure. Values must name one of: RateLimitError, APIError,
## BadRequestError, Timeout, ServiceUnavailableError, InternalServerError (resolved from
## litellm at use time). Reason strings are matched exactly.
_finish_reason_failure_exception_names: Final = frozenset(
{
"RateLimitError",
"APIError",
"BadRequestError",
"Timeout",
"ServiceUnavailableError",
"InternalServerError",
}
)
if treat_finish_reason_as_failure is not None:
for exception_name in treat_finish_reason_as_failure.values():
if exception_name not in _finish_reason_failure_exception_names:
raise ValueError(
f"treat_finish_reason_as_failure values must be one of {sorted(_finish_reason_failure_exception_names)}, got {exception_name}"
)
self.treat_finish_reason_as_failure = treat_finish_reason_as_failure
self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model
self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model
self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model
@ -8591,6 +8616,93 @@ class Router:
)
return resolved is not None
def _get_mapped_finish_reason(self, response: ModelResponse) -> str | None:
"""
The finish reason configured in treat_finish_reason_as_failure that this response carries,
or None. Checks both the mapped finish_reason and the pre-mapping value stashed in
provider_specific_fields["native_finish_reason"]. Streaming detection is a follow-up
modeled on _aanthropic_messages_streaming_iterator.
"""
if not self.treat_finish_reason_as_failure:
return None
if not (response.choices and len(response.choices) > 0):
return None
choice: Final = response.choices[0]
if choice.finish_reason in self.treat_finish_reason_as_failure:
return choice.finish_reason
native_reason: Final = (choice.provider_specific_fields or {}).get("native_finish_reason")
if native_reason in self.treat_finish_reason_as_failure:
return native_reason
return None
def _finish_reason_failure_fallback_available(self, model_group: str, kwargs: Mapping[str, Any]) -> bool:
"""
Whether a generic fallback can serve the retry after a mapped finish-reason failure.
Mirrors the tail of _refusal_fallback_available without the content-policy branch: the
dispatcher falls through to the generic fallbacks lookup, so the gate arms on default
fallbacks or a resolving generic chain.
"""
if fallbacks_disabled_for_request(kwargs):
return False
if self._has_default_fallbacks():
return True
fallbacks: Final = kwargs.get("fallbacks", self.fallbacks)
if fallbacks is None:
return False
resolved, _ = get_fallback_model_group_for_lookup_groups(
fallbacks=fallbacks,
lookup_groups=fallback_lookup_groups(kwargs, model_group),
)
return resolved is not None
def _should_raise_mapped_finish_reason_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool:
"""
True when the response carries a reason from treat_finish_reason_as_failure and a generic
fallback can serve the retry. When a reason is mapped but no fallback is available the
caller must still account for the failure via _account_mapped_finish_reason_failure.
"""
if self._get_mapped_finish_reason(response) is None:
return False
return self._finish_reason_failure_fallback_available(model, kwargs)
def _finish_reason_failure_error(self, model: str, reason: str) -> Exception:
"""Build the exception instance configured for a mapped finish reason."""
exception_name: Final = self.treat_finish_reason_as_failure[reason]
exception_cls: Final = getattr(litellm, exception_name)
message: Final = f"Response finished with reason '{reason}' (treat_finish_reason_as_failure)."
if exception_name == "APIError":
return exception_cls(status_code=500, message=message, llm_provider="", model=model)
return exception_cls(message=message, llm_provider="", model=model)
def _account_mapped_finish_reason_failure(
self, model: str, deployment: dict, response: ModelResponse, kwargs: dict
) -> None:
"""
Count and park a mapped finish-reason failure when nothing is raised (no fallback can
serve the retry): increment the per-minute failure counter and set the cooldown the way a
raised exception would. In the raising branch the normal exception flow does the accounting.
"""
reason: Final = self._get_mapped_finish_reason(response)
if reason is None:
return
model_info: Final = deployment.get("model_info") or {}
deployment_id: Final = model_info.get("id")
if deployment_id is None:
return
exception: Final = self._finish_reason_failure_error(model=model, reason=reason)
increment_deployment_failures_for_current_minute(
litellm_router_instance=self,
deployment_id=deployment_id,
)
_set_cooldown_deployments(
litellm_router_instance=self,
exception_status=exception.status_code,
original_exception=exception,
deployment=deployment_id,
time_to_cooldown=self.cooldown_time,
requested_model_group=(get_litellm_metadata_from_kwargs(kwargs) or {}).get("model_group"),
)
def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool:
"""
Determines if a content policy error should be raised.