This commit is contained in:
Jason Matthew Suhari 2026-04-17 07:48:24 +00:00 • committed by GitHub
commit 0c26246916
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 423 additions and 20 deletions

View file

@ -9,6 +9,7 @@
## LiteLLM versions of the OpenAI Exception Types
import pickle
from typing import Optional
import httpx
@ -19,6 +20,84 @@ from litellm.types.utils import LiteLLMCommonStrings
_MINIMAL_ERROR_RESPONSE: Optional[httpx.Response] = None
def _restore_litellm_exception(cls, state):
"""Module-level callable used by pickle to reconstruct LiteLLM exceptions.
Standard Python exception pickling calls ``cls(*self.args)`` to rebuild the
object. LiteLLM exceptions require several positional arguments
(``message``, ``llm_provider``, ``model``, …) that are *not* stored in
``self.args``, so the default protocol fails with ``TypeError``.
Additionally, ``response`` / ``request`` attributes hold
``httpx.Response`` / ``httpx.Request`` instances which are not themselves
picklable.
This function reconstructs the exception via ``object.__new__`` +
``__dict__`` restoration, recreating minimal httpx objects for any
``response`` / ``request`` attributes that were serialised separately.
"""
obj = Exception.__new__(cls)
# Restore Exception.args so str(obj) and repr(obj) work correctly.
Exception.__init__(obj, state.get("message", ""))
http_attrs = state.get("_pickled_http_attrs", [])
clean_state = {k: v for k, v in state.items() if not k.startswith("_pickled_")}
obj.__dict__.update(clean_state)
for attr in http_attrs:
status_key = f"_pickled_{attr}_status"
if status_key in state:
setattr(
obj,
attr,
httpx.Response(
status_code=state[status_key],
request=httpx.Request(method="GET", url="https://litellm.ai"),
),
)
else:
setattr(obj, attr, httpx.Request(method="GET", url="https://litellm.ai"))
return obj
class _LiteLLMPickleMixin:
"""Mixin that makes LiteLLM exception classes safe to pickle and unpickle.
Mix this in as the *first* base class of every LiteLLM exception that
inherits from an openai exception or from ``Exception`` directly with
non-trivial ``__init__`` arguments::
class AuthenticationError(_LiteLLMPickleMixin, openai.AuthenticationError):
...
Subclasses that already inherit from a LiteLLM exception (e.g.
``ContextWindowExceededError(BadRequestError)``) do **not** need to list
this mixin explicitly — they inherit the behaviour from their parent.
"""
def __reduce__(self):
return _restore_litellm_exception, (type(self), self.__getstate__())
def __getstate__(self):
state = {}
http_attrs = []
for k, v in self.__dict__.items():
if isinstance(v, httpx.Response):
http_attrs.append(k)
state[f"_pickled_{k}_status"] = v.status_code
elif isinstance(v, httpx.Request):
http_attrs.append(k)
else:
try:
pickle.loads(pickle.dumps(v))
state[k] = v
except Exception:
state[k] = str(v)
if http_attrs:
state["_pickled_http_attrs"] = http_attrs
return state
def _get_minimal_error_response() -> httpx.Response:
"""Get a cached minimal httpx.Response object for error cases."""
global _MINIMAL_ERROR_RESPONSE
@ -30,7 +109,7 @@ def _get_minimal_error_response() -> httpx.Response:
return _MINIMAL_ERROR_RESPONSE
class AuthenticationError(openai.AuthenticationError): # type: ignore
class AuthenticationError(_LiteLLMPickleMixin, openai.AuthenticationError): # type: ignore
def __init__(
self,
message,
@ -76,7 +155,7 @@ class AuthenticationError(openai.AuthenticationError): # type: ignore
# raise when invalid models passed, example gpt-8
class NotFoundError(openai.NotFoundError): # type: ignore
class NotFoundError(_LiteLLMPickleMixin, openai.NotFoundError): # type: ignore
def __init__(
self,
message,
@ -121,7 +200,7 @@ class NotFoundError(openai.NotFoundError): # type: ignore
return _message
class BadRequestError(openai.BadRequestError): # type: ignore
class BadRequestError(_LiteLLMPickleMixin, openai.BadRequestError): # type: ignore
def __init__(
self,
message,
@ -196,7 +275,7 @@ class ImageFetchError(BadRequestError):
)
class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore
class UnprocessableEntityError(_LiteLLMPickleMixin, openai.UnprocessableEntityError): # type: ignore
def __init__(
self,
message,
@ -235,7 +314,7 @@ class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore
return _message
class Timeout(openai.APITimeoutError): # type: ignore
class Timeout(_LiteLLMPickleMixin, openai.APITimeoutError): # type: ignore
def __init__(
self,
message,
@ -281,7 +360,7 @@ class Timeout(openai.APITimeoutError): # type: ignore
return _message
class PermissionDeniedError(openai.PermissionDeniedError): # type:ignore
class PermissionDeniedError(_LiteLLMPickleMixin, openai.PermissionDeniedError): # type:ignore
def __init__(
self,
message,
@ -320,7 +399,7 @@ class PermissionDeniedError(openai.PermissionDeniedError): # type:ignore
return _message
class RateLimitError(openai.RateLimitError): # type: ignore
class RateLimitError(_LiteLLMPickleMixin, openai.RateLimitError): # type: ignore
def __init__(
self,
message,
@ -502,7 +581,7 @@ class ContentPolicyViolationError(BadRequestError): # type: ignore
return _message
class ServiceUnavailableError(openai.APIStatusError): # type: ignore
class ServiceUnavailableError(_LiteLLMPickleMixin, openai.APIStatusError): # type: ignore
def __init__(
self,
message,
@ -552,7 +631,7 @@ class ServiceUnavailableError(openai.APIStatusError): # type: ignore
return _message
class BadGatewayError(openai.APIStatusError): # type: ignore
class BadGatewayError(_LiteLLMPickleMixin, openai.APIStatusError): # type: ignore
def __init__(
self,
message,
@ -602,7 +681,7 @@ class BadGatewayError(openai.APIStatusError): # type: ignore
return _message
class InternalServerError(openai.InternalServerError): # type: ignore
class InternalServerError(_LiteLLMPickleMixin, openai.InternalServerError): # type: ignore
def __init__(
self,
message,
@ -653,7 +732,7 @@ class InternalServerError(openai.InternalServerError): # type: ignore
# raise this when the API returns an invalid response object - https://github.com/openai/openai-python/blob/1be14ee34a0f8e42d3f9aa5451aa4cb161f1781f/openai/api_requestor.py#L401
class APIError(openai.APIError): # type: ignore
class APIError(_LiteLLMPickleMixin, openai.APIError): # type: ignore
def __init__(
self,
status_code: int,
@ -694,7 +773,7 @@ class APIError(openai.APIError): # type: ignore
# raised if an invalid request (not get, delete, put, post) is made
class APIConnectionError(openai.APIConnectionError): # type: ignore
class APIConnectionError(_LiteLLMPickleMixin, openai.APIConnectionError): # type: ignore
def __init__(
self,
message,
@ -733,7 +812,7 @@ class APIConnectionError(openai.APIConnectionError): # type: ignore
# raised if an invalid request (not get, delete, put, post) is made
class APIResponseValidationError(openai.APIResponseValidationError): # type: ignore
class APIResponseValidationError(_LiteLLMPickleMixin, openai.APIResponseValidationError): # type: ignore
def __init__(
self,
message,
@ -784,7 +863,7 @@ class JSONSchemaValidationError(APIResponseValidationError):
super().__init__(model=model, message=message, llm_provider=llm_provider)
class OpenAIError(openai.OpenAIError): # type: ignore
class OpenAIError(_LiteLLMPickleMixin, openai.OpenAIError): # type: ignore
def __init__(self, original_exception=None):
super().__init__()
self.llm_provider = "openai"
@ -841,7 +920,7 @@ LITELLM_EXCEPTION_TYPES = [
]
class BudgetExceededError(Exception):
class BudgetExceededError(_LiteLLMPickleMixin, Exception):
def __init__(
self, current_cost: float, max_budget: float, message: Optional[str] = None
):
@ -856,7 +935,7 @@ class BudgetExceededError(Exception):
## DEPRECATED ##
class InvalidRequestError(openai.BadRequestError): # type: ignore
class InvalidRequestError(_LiteLLMPickleMixin, openai.BadRequestError): # type: ignore
def __init__(self, message, model, llm_provider):
self.status_code = 400
self.message = message
@ -873,7 +952,7 @@ class InvalidRequestError(openai.BadRequestError): # type: ignore
) # Call the base class constructor with the parameters it needs
class MockException(openai.APIError):
class MockException(_LiteLLMPickleMixin, openai.APIError):
# used for testing
def __init__(
self,
@ -911,7 +990,7 @@ class LiteLLMUnknownProvider(BadRequestError):
return self.message
class GuardrailRaisedException(Exception):
class GuardrailRaisedException(_LiteLLMPickleMixin, Exception):
def __init__(
self,
guardrail_name: Optional[str] = None,
@ -924,7 +1003,7 @@ class GuardrailRaisedException(Exception):
super().__init__(self.message)
class BlockedPiiEntityError(Exception):
class BlockedPiiEntityError(_LiteLLMPickleMixin, Exception):
def __init__(
self,
entity_type: str,
@ -1017,7 +1096,7 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore
class GuardrailInterventionNormalStringError(
Exception
_LiteLLMPickleMixin, Exception
): # custom exception to raise when a guardrail intervenes, but we want to return a normal string to the user
def __init__(self, message: str):
self.message = message

View file

@ -0,0 +1,324 @@
"""
Tests that all LiteLLM exception classes survive a pickle round-trip.
Relevant issue: https://github.com/BerriAI/litellm/issues/24136
Without the fix, exceptions fail with TypeError when pickle tries to
reconstruct them via cls(*self.args), because the required positional
arguments (message, llm_provider, model, …) are not stored in self.args.
Additionally, httpx.Response / httpx.Request attributes are not picklable
by default.
"""
import pickle
import httpx
import pytest
import litellm.exceptions as exc
def _roundtrip(obj):
"""Pickle and unpickle an object, returning the reconstructed copy."""
return pickle.loads(pickle.dumps(obj))
def _make_response(status: int = 400) -> httpx.Response:
return httpx.Response(
status_code=status,
request=httpx.Request(method="GET", url="https://litellm.ai"),
)
# ---------------------------------------------------------------------------
# Fixtures — one instance per exception class
# ---------------------------------------------------------------------------
CASES = [
exc.AuthenticationError(
message="bad key",
llm_provider="openai",
model="gpt-4",
),
exc.NotFoundError(
message="model not found",
model="gpt-99",
llm_provider="openai",
),
exc.BadRequestError(
message="invalid param",
model="gpt-4",
llm_provider="openai",
),
exc.ImageFetchError(
message="cannot fetch image",
model="gpt-4",
llm_provider="openai",
),
exc.UnprocessableEntityError(
message="unprocessable",
model="gpt-4",
llm_provider="openai",
response=_make_response(422),
),
exc.Timeout(
message="timed out",
model="gpt-4",
llm_provider="openai",
),
exc.PermissionDeniedError(
message="denied",
llm_provider="openai",
model="gpt-4",
response=_make_response(403),
),
exc.RateLimitError(
message="rate limited",
llm_provider="openai",
model="gpt-4",
),
exc.ContextWindowExceededError(
message="context too long",
model="gpt-4",
llm_provider="openai",
),
exc.RejectedRequestError(
message="guardrail rejected",
model="gpt-4",
llm_provider="openai",
request_data={"messages": []},
),
exc.ContentPolicyViolationError(
message="content violation",
model="gpt-4",
llm_provider="openai",
),
exc.ServiceUnavailableError(
message="service down",
llm_provider="openai",
model="gpt-4",
),
exc.BadGatewayError(
message="bad gateway",
llm_provider="openai",
model="gpt-4",
),
exc.InternalServerError(
message="internal error",
llm_provider="openai",
model="gpt-4",
),
exc.APIError(
status_code=500,
message="api error",
llm_provider="openai",
model="gpt-4",
),
exc.APIConnectionError(
message="connection failed",
llm_provider="openai",
model="gpt-4",
),
exc.APIResponseValidationError(
message="bad response",
llm_provider="openai",
model="gpt-4",
),
exc.JSONSchemaValidationError(
model="gpt-4",
llm_provider="openai",
raw_response='{"bad": true}',
schema='{"type": "object"}',
),
exc.OpenAIError(),
exc.UnsupportedParamsError(
message="unsupported param",
llm_provider="openai",
model="gpt-4",
),
exc.BudgetExceededError(current_cost=1.5, max_budget=1.0),
exc.InvalidRequestError(
message="invalid request",
model="gpt-4",
llm_provider="openai",
),
exc.MockException(
status_code=500,
message="mock error",
llm_provider="openai",
model="gpt-4",
),
exc.LiteLLMUnknownProvider(model="unknown/model"),
exc.GuardrailRaisedException(guardrail_name="my-guard", message="blocked"),
exc.BlockedPiiEntityError(entity_type="email", guardrail_name="pii-guard"),
exc.MidStreamFallbackError(
message="mid-stream fail",
model="gpt-4",
llm_provider="openai",
),
exc.GuardrailInterventionNormalStringError(message="intervention"),
]
@pytest.mark.parametrize("original", CASES, ids=lambda e: type(e).__name__)
def test_pickle_roundtrip(original):
"""Each exception must survive pickle.dumps → pickle.loads without error."""
restored = _roundtrip(original)
assert type(restored) is type(original)
@pytest.mark.parametrize("original", CASES, ids=lambda e: type(e).__name__)
def test_pickle_preserves_message(original):
"""The message attribute must be identical after the round-trip."""
if not hasattr(original, "message"):
pytest.skip("exception has no message attribute")
restored = _roundtrip(original)
assert restored.message == original.message
@pytest.mark.parametrize("original", CASES, ids=lambda e: type(e).__name__)
def test_pickle_preserves_llm_provider(original):
"""llm_provider must be preserved when present."""
if not hasattr(original, "llm_provider"):
pytest.skip("exception has no llm_provider attribute")
restored = _roundtrip(original)
assert restored.llm_provider == original.llm_provider
@pytest.mark.parametrize("original", CASES, ids=lambda e: type(e).__name__)
def test_pickle_preserves_model(original):
"""model must be preserved when present."""
if not hasattr(original, "model"):
pytest.skip("exception has no model attribute")
restored = _roundtrip(original)
assert restored.model == original.model
@pytest.mark.parametrize("original", CASES, ids=lambda e: type(e).__name__)
def test_pickle_isinstance_checks_still_work(original):
"""After round-trip, isinstance checks against the original class must pass."""
restored = _roundtrip(original)
assert isinstance(restored, type(original))
assert isinstance(restored, Exception)
def test_pickle_rejected_request_preserves_request_data():
original = exc.RejectedRequestError(
message="blocked",
model="gpt-4",
llm_provider="openai",
request_data={"messages": [{"role": "user", "content": "hello"}]},
)
restored = _roundtrip(original)
assert restored.request_data == original.request_data
def test_pickle_budget_exceeded_preserves_costs():
original = exc.BudgetExceededError(current_cost=2.5, max_budget=1.0)
restored = _roundtrip(original)
assert restored.current_cost == original.current_cost
assert restored.max_budget == original.max_budget
def test_pickle_json_schema_error_preserves_raw_response():
original = exc.JSONSchemaValidationError(
model="gpt-4",
llm_provider="openai",
raw_response='{"unexpected": "field"}',
schema='{"required": ["name"]}',
)
restored = _roundtrip(original)
assert restored.raw_response == original.raw_response
assert restored.schema == original.schema
def test_pickle_guardrail_raised_preserves_guardrail_name():
original = exc.GuardrailRaisedException(
guardrail_name="content-filter", message="profanity detected"
)
restored = _roundtrip(original)
assert restored.guardrail_name == original.guardrail_name
def test_pickle_blocked_pii_preserves_entity_type():
original = exc.BlockedPiiEntityError(
entity_type="credit_card", guardrail_name="pii-guard"
)
restored = _roundtrip(original)
assert restored.entity_type == original.entity_type
assert restored.guardrail_name == original.guardrail_name
def test_pickle_midstream_preserves_generated_content():
original = exc.MidStreamFallbackError(
message="stream interrupted",
model="gpt-4",
llm_provider="openai",
generated_content="partial response text",
is_pre_first_chunk=False,
)
restored = _roundtrip(original)
assert restored.generated_content == original.generated_content
assert restored.is_pre_first_chunk == original.is_pre_first_chunk
def test_pickle_midstream_with_non_picklable_original_exception():
"""MidStreamFallbackError must survive pickle even when original_exception
is a third-party exception that cannot be reconstructed by pickle.loads
(e.g. openai.RateLimitError which requires response and body kwargs).
The original_exception should be preserved as its string representation."""
import openai
request = _make_response(429).request
response = _make_response(429)
original_exc = openai.RateLimitError("rate limited", response=response, body={})
original = exc.MidStreamFallbackError(
message="fallback triggered mid-stream",
model="gpt-4",
llm_provider="openai",
original_exception=original_exc,
generated_content="partial text",
)
restored = _roundtrip(original)
assert restored.generated_content == original.generated_content
assert restored.model == original.model
assert restored.llm_provider == original.llm_provider
# original_exception falls back to str() when it cannot be round-tripped
assert restored.original_exception == str(original_exc)
assert isinstance(restored, exc.MidStreamFallbackError)
def test_pickle_with_retries_info():
original = exc.RateLimitError(
message="too many requests",
llm_provider="openai",
model="gpt-4",
max_retries=3,
num_retries=3,
)
restored = _roundtrip(original)
assert restored.max_retries == 3
assert restored.num_retries == 3
def test_pickle_exception_is_raiseable():
"""Unpickled exceptions must still be raise-able and catchable."""
original = exc.AuthenticationError(
message="invalid api key", llm_provider="openai", model="gpt-4"
)
restored = _roundtrip(original)
with pytest.raises(exc.AuthenticationError):
raise restored
def test_pickle_multiple_roundtrips():
"""Exceptions must survive multiple sequential pickle round-trips."""
original = exc.InternalServerError(
message="server error", llm_provider="anthropic", model="claude-3"
)
result = original
for _ in range(3):
result = _roundtrip(result)
assert result.message == original.message
assert result.llm_provider == original.llm_provider