From 57debdb461d866477051d7c304c37112658ba032 Mon Sep 17 00:00:00 2001 From: wzj1228516103 <154067345+wzj1228516103@users.noreply.github.com> Date: Thu, 10 Sep 2026 18:37:30 +0800 Subject: [PATCH] fix(streaming): preserve provider status for string error codes --- litellm/exceptions.py | 20 +++++++++++--- .../test_exception_header_preservation.py | 26 +++++++++++++++++++ 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/litellm/exceptions.py b/litellm/exceptions.py index f9215267bf3..b3fc6866529 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -1086,6 +1086,15 @@ class BlockedPiiEntityError(Exception): super().__init__(self.message) +def _numeric_status_code(value: object) -> int | None: + if value is None: + return None + try: + return int(value) + except (TypeError, ValueError): + return None + + class MidStreamFallbackError(ServiceUnavailableError): def __init__( self, @@ -1100,8 +1109,13 @@ class MidStreamFallbackError(ServiceUnavailableError): generated_content: str = "", is_pre_first_chunk: bool = False, ): - original_status: Final = getattr(original_exception, "status_code", None) - self.status_code = int(original_status) if original_status is not None else 503 + original_status: Final = _numeric_status_code(getattr(original_exception, "status_code", None)) + original_response: Final = getattr(original_exception, "response", None) + response_status: Final = _numeric_status_code( + getattr(response, "status_code", getattr(original_response, "status_code", None)) + ) + status_code: Final = original_status or response_status or 503 + self.status_code = status_code self.message = f"litellm.MidStreamFallbackError: {message}" self.model = model self.llm_provider = llm_provider @@ -1143,7 +1157,7 @@ class MidStreamFallbackError(ServiceUnavailableError): ) # Restore the propagated status and original response/request objects - self.status_code = int(original_status) if original_status is not None else 503 + self.status_code = status_code self.response = _saved_response self.request = _saved_request self.message = _saved_message diff --git a/tests/test_litellm/test_exception_header_preservation.py b/tests/test_litellm/test_exception_header_preservation.py index dd142d9d40b..3c9a29b09c2 100644 --- a/tests/test_litellm/test_exception_header_preservation.py +++ b/tests/test_litellm/test_exception_header_preservation.py @@ -22,6 +22,13 @@ from litellm.exceptions import ( ) +class ProviderToolSchemaError(Exception): + def __init__(self, response: httpx.Response) -> None: + self.status_code = "tool_use_failed" + self.response = response + super().__init__("tool call validation failed") + + class TestExceptionHeaderPreservation: """Test that exception classes preserve headers from provider responses.""" @@ -255,6 +262,25 @@ class TestExceptionAttributes: assert midstream_fallback.response.status_code == 503 assert str(midstream_fallback.response.request.url) == "https://openai.com/v1/" + def test_midstream_fallback_error_accepts_non_numeric_provider_status(self): + """Provider error codes can be strings while the response still has an HTTP status.""" + original_response = httpx.Response( + status_code=400, + request=httpx.Request("POST", "https://api.groq.com/openai/v1/chat/completions"), + ) + provider_error = ProviderToolSchemaError(original_response) + + midstream_error = MidStreamFallbackError( + message=str(provider_error), + model="openai/gpt-oss-120b", + llm_provider="groq", + original_exception=provider_error, + ) + + assert midstream_error.status_code == 400 + assert midstream_error.response.status_code == 400 + assert "tool call validation failed" in midstream_error.message + class TestProxyHeaderExtraction: """Test that proxy correctly extracts headers from exceptions."""