mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge da00342b24 into 40ff01b987
This commit is contained in:
commit
23bc453e63
2 changed files with 420 additions and 0 deletions
194
litellm/error_categories.py
Normal file
194
litellm/error_categories.py
Normal file
|
|
@ -0,0 +1,194 @@
|
|||
"""
|
||||
Protocol-level error categorization for provider-agnostic retry/circuit-breaker logic.
|
||||
|
||||
Each provider adapter maps its native errors into exactly four canonical categories:
|
||||
auth — authentication/authorization failure (401, 403)
|
||||
rate_limit — rate limit exceeded (429)
|
||||
server — upstream server error (5xx)
|
||||
client — invalid request or client-side error (4xx, excluding 401/403/429)
|
||||
|
||||
This module is the Python equivalent of zeshim's `ParsedError` type and the
|
||||
`parseError` protocol function — adapted to LiteLLM's existing exception hierarchy.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class ErrorCategory(str, Enum):
|
||||
"""Canonical error categories, provider-agnostic.
|
||||
|
||||
Upstream retry/circuit-breaker/scheduler logic should branch on these
|
||||
four values rather than inspecting provider-specific exception types.
|
||||
"""
|
||||
|
||||
AUTH = "auth"
|
||||
RATE_LIMIT = "rate_limit"
|
||||
SERVER = "server"
|
||||
CLIENT = "client"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ParsedError:
|
||||
"""Normalized error produced by a protocol's error parser.
|
||||
|
||||
Each protocol adapter (anthropic, openai, gemini, etc.) implements a
|
||||
parse_error(data: dict, status: int) -> ParsedError function that maps
|
||||
provider-specific error shapes into this canonical form.
|
||||
"""
|
||||
|
||||
category: ErrorCategory
|
||||
message: Optional[str] = None
|
||||
status_code: Optional[int] = None
|
||||
|
||||
|
||||
# ── Protocol-level parse_error functions ──
|
||||
|
||||
|
||||
def _extract_error_body(data: dict) -> dict:
|
||||
"""Normalize the error object from various provider shapes."""
|
||||
err = data.get("error", data)
|
||||
return err if isinstance(err, dict) else {}
|
||||
|
||||
|
||||
def default_parse_error(data: dict, status: int) -> ParsedError:
|
||||
"""Default HTTP-status-based error parser.
|
||||
|
||||
Used by OpenAI-compatible and Anthropic protocols where HTTP status
|
||||
codes alone are sufficient for categorization.
|
||||
"""
|
||||
err = _extract_error_body(data)
|
||||
message = err.get("message")
|
||||
|
||||
if status in (401, 403):
|
||||
return ParsedError(
|
||||
category=ErrorCategory.AUTH, message=message, status_code=status
|
||||
)
|
||||
if status == 429:
|
||||
return ParsedError(
|
||||
category=ErrorCategory.RATE_LIMIT, message=message, status_code=status
|
||||
)
|
||||
if status >= 500:
|
||||
return ParsedError(
|
||||
category=ErrorCategory.SERVER,
|
||||
message=message or "Server error",
|
||||
status_code=status,
|
||||
)
|
||||
return ParsedError(
|
||||
category=ErrorCategory.CLIENT,
|
||||
message=message or "Client error",
|
||||
status_code=status,
|
||||
)
|
||||
|
||||
|
||||
def google_parse_error(data: dict, status: int) -> ParsedError:
|
||||
"""Google Generative AI / Vertex AI error parser.
|
||||
|
||||
Google returns error status strings in the response body (e.g.
|
||||
'UNAUTHENTICATED', 'PERMISSION_DENIED', 'RESOURCE_EXHAUSTED',
|
||||
'UNAVAILABLE', 'DEADLINE_EXCEEDED') that override HTTP status
|
||||
for categorization.
|
||||
"""
|
||||
err = _extract_error_body(data)
|
||||
message = err.get("message")
|
||||
google_status = err.get("status", "").upper()
|
||||
|
||||
if status in (401, 403) or google_status in (
|
||||
"UNAUTHENTICATED",
|
||||
"PERMISSION_DENIED",
|
||||
):
|
||||
return ParsedError(
|
||||
category=ErrorCategory.AUTH, message=message, status_code=status
|
||||
)
|
||||
if status == 429 or google_status == "RESOURCE_EXHAUSTED":
|
||||
return ParsedError(
|
||||
category=ErrorCategory.RATE_LIMIT, message=message, status_code=status
|
||||
)
|
||||
if status >= 500 or google_status in (
|
||||
"UNAVAILABLE",
|
||||
"INTERNAL",
|
||||
"DEADLINE_EXCEEDED",
|
||||
):
|
||||
return ParsedError(
|
||||
category=ErrorCategory.SERVER,
|
||||
message=message or "Server error",
|
||||
status_code=status,
|
||||
)
|
||||
return ParsedError(
|
||||
category=ErrorCategory.CLIENT,
|
||||
message=message or "Client error",
|
||||
status_code=status,
|
||||
)
|
||||
|
||||
|
||||
# ── Integration with LiteLLM's existing ProviderError ──
|
||||
|
||||
|
||||
def categorize_exception(exc: Exception) -> Optional[ErrorCategory]:
|
||||
"""Extract canonical ErrorCategory from any LiteLLM exception.
|
||||
|
||||
Returns None if the exception cannot be categorized (caller should
|
||||
treat as CLIENT or re-raise).
|
||||
"""
|
||||
# If the exception already carries a category attribute, use it.
|
||||
category = getattr(exc, "error_category", None)
|
||||
if isinstance(category, ErrorCategory):
|
||||
return category
|
||||
|
||||
# Try status-code-based inference
|
||||
status = getattr(exc, "status_code", None)
|
||||
if status is not None:
|
||||
category_from_status = _categorize_by_status_code(status)
|
||||
if category_from_status is not None:
|
||||
return category_from_status
|
||||
|
||||
# Fall back to type-name heuristics
|
||||
return _categorize_by_exception_name(exc)
|
||||
|
||||
|
||||
def _categorize_by_status_code(status: any) -> Optional[ErrorCategory]:
|
||||
"""Categorize error by HTTP status code.
|
||||
|
||||
Handles both integer and string status codes.
|
||||
"""
|
||||
# Normalize to integer
|
||||
if not isinstance(status, int):
|
||||
try:
|
||||
status = int(status)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
# Auth errors
|
||||
if status in (401, 403):
|
||||
return ErrorCategory.AUTH
|
||||
|
||||
# Rate limiting
|
||||
if status == 429:
|
||||
return ErrorCategory.RATE_LIMIT
|
||||
|
||||
# Server errors (including 408 Request Timeout which should be retryable)
|
||||
if status == 408 or status >= 500:
|
||||
return ErrorCategory.SERVER
|
||||
|
||||
# Client errors (4xx except auth and rate limit)
|
||||
if 400 <= status < 500:
|
||||
return ErrorCategory.CLIENT
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _categorize_by_exception_name(exc: Exception) -> Optional[ErrorCategory]:
|
||||
"""Categorize error by exception class name patterns."""
|
||||
name = type(exc).__name__.lower()
|
||||
|
||||
if "auth" in name:
|
||||
return ErrorCategory.AUTH
|
||||
|
||||
if "rate" in name or "throttl" in name:
|
||||
return ErrorCategory.RATE_LIMIT
|
||||
|
||||
if "server" in name or "service" in name or "timeout" in name:
|
||||
return ErrorCategory.SERVER
|
||||
|
||||
return None
|
||||
226
tests/test_error_categories.py
Normal file
226
tests/test_error_categories.py
Normal file
|
|
@ -0,0 +1,226 @@
|
|||
"""Tests for litellm.error_categories — protocol-level error normalization."""
|
||||
|
||||
import pytest
|
||||
from litellm.error_categories import (
|
||||
ErrorCategory,
|
||||
ParsedError,
|
||||
categorize_exception,
|
||||
default_parse_error,
|
||||
google_parse_error,
|
||||
)
|
||||
|
||||
|
||||
class TestDefaultParseError:
|
||||
"""OpenAI-compatible & Anthropic error parsing (HTTP-status-based)."""
|
||||
|
||||
def test_auth_401(self):
|
||||
result = default_parse_error({}, 401)
|
||||
assert result.category == ErrorCategory.AUTH
|
||||
assert result.status_code == 401
|
||||
|
||||
def test_auth_403(self):
|
||||
result = default_parse_error({}, 403)
|
||||
assert result.category == ErrorCategory.AUTH
|
||||
|
||||
def test_rate_limit_429(self):
|
||||
result = default_parse_error({}, 429)
|
||||
assert result.category == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_server_500(self):
|
||||
result = default_parse_error({}, 500)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
assert result.message is not None
|
||||
|
||||
def test_server_502(self):
|
||||
result = default_parse_error({}, 502)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
def test_server_503(self):
|
||||
result = default_parse_error({}, 503)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
def test_client_400(self):
|
||||
result = default_parse_error({}, 400)
|
||||
assert result.category == ErrorCategory.CLIENT
|
||||
|
||||
def test_client_404(self):
|
||||
result = default_parse_error({}, 404)
|
||||
assert result.category == ErrorCategory.CLIENT
|
||||
|
||||
def test_extracts_message_from_body(self):
|
||||
result = default_parse_error({"error": {"message": "Bad request"}}, 400)
|
||||
assert result.message == "Bad request"
|
||||
|
||||
def test_no_crash_on_empty_body(self):
|
||||
result = default_parse_error({}, 500)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
def test_no_crash_on_invalid_body(self):
|
||||
"""Should not crash if error field is not a dict."""
|
||||
result = default_parse_error({"error": "string error"}, 500)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
|
||||
class TestGoogleParseError:
|
||||
"""Google Gemini / Vertex AI error parsing (status-string-aware)."""
|
||||
|
||||
def test_http_auth(self):
|
||||
result = google_parse_error({}, 401)
|
||||
assert result.category == ErrorCategory.AUTH
|
||||
|
||||
def test_body_unauthenticated(self):
|
||||
result = google_parse_error({"error": {"status": "UNAUTHENTICATED"}}, 200)
|
||||
assert result.category == ErrorCategory.AUTH
|
||||
|
||||
def test_rate_limit(self):
|
||||
result = google_parse_error({}, 429)
|
||||
assert result.category == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_body_resource_exhausted(self):
|
||||
result = google_parse_error({"error": {"status": "RESOURCE_EXHAUSTED"}}, 200)
|
||||
assert result.category == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_server_unavailable(self):
|
||||
result = google_parse_error({"error": {"status": "UNAVAILABLE"}}, 200)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
def test_server_internal(self):
|
||||
result = google_parse_error({"error": {"status": "INTERNAL"}}, 200)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
def test_client_fallback(self):
|
||||
result = google_parse_error({"error": {"message": "Invalid argument"}}, 400)
|
||||
assert result.category == ErrorCategory.CLIENT
|
||||
assert result.message == "Invalid argument"
|
||||
|
||||
def test_body_permission_denied(self):
|
||||
"""PERMISSION_DENIED should map to AUTH."""
|
||||
result = google_parse_error({"error": {"status": "PERMISSION_DENIED"}}, 200)
|
||||
assert result.category == ErrorCategory.AUTH
|
||||
|
||||
def test_body_deadline_exceeded(self):
|
||||
"""DEADLINE_EXCEEDED should map to SERVER (retryable)."""
|
||||
result = google_parse_error({"error": {"status": "DEADLINE_EXCEEDED"}}, 200)
|
||||
assert result.category == ErrorCategory.SERVER
|
||||
|
||||
|
||||
class TestCategorizeException:
|
||||
"""Integration: extract ErrorCategory from existing LiteLLM exceptions."""
|
||||
|
||||
def test_exception_with_category_attr(self):
|
||||
exc = Exception()
|
||||
exc.error_category = ErrorCategory.RATE_LIMIT # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_exception_with_status_code(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 429 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_exception_by_name_auth(self):
|
||||
class AuthenticationError(Exception):
|
||||
pass
|
||||
|
||||
assert categorize_exception(AuthenticationError()) == ErrorCategory.AUTH
|
||||
|
||||
def test_exception_by_name_rate(self):
|
||||
class RateLimitError(Exception):
|
||||
pass
|
||||
|
||||
assert categorize_exception(RateLimitError()) == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_exception_by_name_server(self):
|
||||
class ServiceUnavailableError(Exception):
|
||||
pass
|
||||
|
||||
assert categorize_exception(ServiceUnavailableError()) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_by_name_throttle(self):
|
||||
class ThrottlingError(Exception):
|
||||
pass
|
||||
|
||||
assert categorize_exception(ThrottlingError()) == ErrorCategory.RATE_LIMIT
|
||||
|
||||
def test_exception_by_name_timeout(self):
|
||||
class TimeoutError(Exception):
|
||||
pass
|
||||
|
||||
assert categorize_exception(TimeoutError()) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_with_status_code_401(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 401 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.AUTH
|
||||
|
||||
def test_exception_with_status_code_403(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 403 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.AUTH
|
||||
|
||||
def test_exception_with_status_code_500(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 500 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_with_status_code_503(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 503 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_with_status_code_400(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 400 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.CLIENT
|
||||
|
||||
def test_exception_with_status_code_404(self):
|
||||
exc = Exception()
|
||||
exc.status_code = 404 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.CLIENT
|
||||
|
||||
def test_exception_with_status_code_408_timeout(self):
|
||||
"""408 Request Timeout should be SERVER (retryable), not CLIENT."""
|
||||
exc = Exception()
|
||||
exc.status_code = 408 # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_with_string_status_code(self):
|
||||
"""String status_code should be converted to int."""
|
||||
exc = Exception()
|
||||
exc.status_code = "503" # type: ignore[attr-defined]
|
||||
assert categorize_exception(exc) == ErrorCategory.SERVER
|
||||
|
||||
def test_exception_with_invalid_status_code(self):
|
||||
"""Invalid status_code should fall through to name heuristics."""
|
||||
exc = Exception()
|
||||
exc.status_code = "invalid" # type: ignore[attr-defined]
|
||||
# Falls through to None since no name match
|
||||
assert categorize_exception(exc) is None
|
||||
|
||||
def test_unknown_returns_none(self):
|
||||
assert categorize_exception(ValueError("unexpected")) is None
|
||||
|
||||
|
||||
class TestErrorCategoryEnum:
|
||||
"""ErrorCategory is a string enum for easy serialization."""
|
||||
|
||||
def test_values(self):
|
||||
assert ErrorCategory.AUTH.value == "auth"
|
||||
assert ErrorCategory.RATE_LIMIT.value == "rate_limit"
|
||||
assert ErrorCategory.SERVER.value == "server"
|
||||
assert ErrorCategory.CLIENT.value == "client"
|
||||
|
||||
def test_is_str(self):
|
||||
assert isinstance(ErrorCategory.AUTH, str)
|
||||
|
||||
|
||||
class TestParsedError:
|
||||
"""ParsedError is an immutable value object."""
|
||||
|
||||
def test_frozen(self):
|
||||
err = ParsedError(category=ErrorCategory.AUTH, message="Unauthorized")
|
||||
with pytest.raises(Exception):
|
||||
err.category = ErrorCategory.CLIENT # type: ignore[misc]
|
||||
|
||||
def test_repr(self):
|
||||
err = ParsedError(category=ErrorCategory.SERVER, status_code=503)
|
||||
assert "server" in repr(err)
|
||||
Loading…
Add table
Reference in a new issue