From f8b0615ae9c74c70b1d4abb34d17cc41c86afcea Mon Sep 17 00:00:00 2001 From: Jason Matthew Suhari Date: Fri, 20 Mar 2026 16:50:32 +0800 Subject: [PATCH] fix: make all LiteLLM exception classes pickle-safe MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Standard Python exception pickling calls cls(*self.args) to reconstruct the object. LiteLLM exceptions require several positional args (message, llm_provider, model, …) that are not stored in self.args, so the default protocol raises TypeError. Additionally, response/request attributes hold httpx.Response/httpx.Request instances which are not picklable. Add _LiteLLMPickleMixin with __reduce__ and __getstate__ that: - Uses Exception.__new__(cls) + dict restoration to bypass __init__ - Serialises httpx.Response status codes separately and recreates minimal response/request objects on unpickle Apply the mixin to all 19 exception classes that directly inherit from openai exceptions or Exception, covering all 28 exception types via inheritance. Fixes #24136 --- litellm/exceptions.py | 114 +++++-- tests/test_litellm/test_exceptions_pickle.py | 297 +++++++++++++++++++ 2 files changed, 391 insertions(+), 20 deletions(-) create mode 100644 tests/test_litellm/test_exceptions_pickle.py diff --git a/litellm/exceptions.py b/litellm/exceptions.py index abdba09dd8d..cfcd631fb24 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -19,6 +19,80 @@ 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: + state[k] = 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 +104,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 +150,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 +195,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 +270,7 @@ class ImageFetchError(BadRequestError): ) -class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore +class UnprocessableEntityError(_LiteLLMPickleMixin, openai.UnprocessableEntityError): # type: ignore def __init__( self, message, @@ -235,7 +309,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 +355,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 +394,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 +576,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 +626,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 +676,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 +727,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 +768,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 +807,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 +858,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 +915,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 +930,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 +947,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 +985,7 @@ class LiteLLMUnknownProvider(BadRequestError): return self.message -class GuardrailRaisedException(Exception): +class GuardrailRaisedException(_LiteLLMPickleMixin, Exception): def __init__( self, guardrail_name: Optional[str] = None, @@ -924,7 +998,7 @@ class GuardrailRaisedException(Exception): super().__init__(self.message) -class BlockedPiiEntityError(Exception): +class BlockedPiiEntityError(_LiteLLMPickleMixin, Exception): def __init__( self, entity_type: str, @@ -1017,7 +1091,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 diff --git a/tests/test_litellm/test_exceptions_pickle.py b/tests/test_litellm/test_exceptions_pickle.py new file mode 100644 index 00000000000..9bce5ace072 --- /dev/null +++ b/tests/test_litellm/test_exceptions_pickle.py @@ -0,0 +1,297 @@ +""" +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_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