mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(router): treat_finish_reason_as_failure config surface and helpers
This commit is contained in:
parent
8fc9c46d1a
commit
f6ec7dfeb9
1 changed files with 112 additions and 0 deletions
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue