diff --git a/litellm/router.py b/litellm/router.py index 633f060f208..a2aec6847a3 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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.