mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(streaming): preserve provider status for string error codes
This commit is contained in:
parent
e5da59336d
commit
57debdb461
2 changed files with 43 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue