fix: support OpenAI SDK 3 while retaining SDK 2 compatibility (#44927)

* fix: support OpenAI SDK 3.x while retaining 2.x compatibility

* fix: preserve HTTPX compatibility across OpenAI SDK versions

* refactor: trim OpenAI SDK compatibility patch

* refactor: simplify SDK client defaults

* test: cover SDK compatibility across provider routes

* fix: use shared names for SDK transport helpers

* fix: address SDK compatibility review feedback

* test(openai): send requests through the SDK API factory clients

* fix(exceptions): keep the provider response on PaymentRequiredError under httpx 2

* ci(base-sdk): count httpx2 as base-only when a base dependency brings it in

OpenAI SDK 3 depends on httpx2, so the litellm-core install at highest resolution pulls it in through openai. The base-only check now resolves the declared base dependency closure of litellm and litellm-core and accepts extras-only modules owned by a distribution inside that closure

* test: cover OpenAI SDK 3 client paths end to end and read SDK HTTP headers case-insensitively

* test: read C02 spend rows once, wait out the worker healthcheck on respawn, and resolve module owners on Python 3.10

---------

Co-authored-by: Marcus Wood <marcuswood@openai.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 19:45:50 -07:00 • committed by GitHub
parent 674208bb7d
commit 63d823d28a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 2851 additions and 153 deletions

View file

@ -1268,6 +1268,75 @@ jobs:
fi
echo "$chart: all $declared declared test suites ran"
done
openai-sdk-compat:
name: OpenAI ${{ matrix.openai-version }} compatibility
permissions:
contents: read
pull-requests: read
runs-on: ubuntu-latest
timeout-minutes: 20
strategy:
fail-fast: false
matrix:
openai-version: ["2.20.0", "3.0.0", "3.25.0"]
env:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
PYTHONPATH: ${{ github.workspace }}
UV_PYTHON: "3.10"
OPENAI_VERSION: ${{ matrix.openai-version }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: ${{ env.UV_PYTHON }}
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group dev --group proxy-dev --extra proxy --no-install-workspace
uv pip install "openai==${OPENAI_VERSION:?}" --exclude-newer-package openai=2026-10-07T00:00:00Z
uv run --no-sync python -c 'import os, openai; assert openai.__version__ == os.environ["OPENAI_VERSION"], openai.__version__'
- name: Run OpenAI SDK compatibility tests
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync pytest \
tests/unit/llms/openai/test_openai.py \
tests/unit/llms/openai/test_openai_common_utils.py \
tests/unit/llms/openai/test_openai_file_content_streaming.py \
tests/unit/test_openai_embedding_encoding_format_default.py \
tests/unit/llms/inception/test_inception_completion_transformation.py::test_inception_fim_targets_fim_endpoint \
tests/unit/llms/inception/test_inception_completion_transformation.py::test_inception_fim_async \
tests/unit/llms/azure/test_azure_exception_mapping.py \
tests/unit/test_completion_timeout_resolution.py \
tests/unit/types/test_router.py \
tests/unit/test_exception_header_preservation.py \
tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py::test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far \
--tb=short -q
- name: Run workload identity tests with SDK 3
if: steps.changes.outputs.decision != 'skip' && startsWith(matrix.openai-version, '3.')
run: uv run --no-sync pytest tests/unit/llms/openai/test_openai_workload_identity.py --tb=short -q
- name: Check consumer types with the installed SDK
if: steps.changes.outputs.decision != 'skip'
run: uv run --no-sync basedpyright --pythonpath .venv/bin/python --project tests/typing/pyrightconfig.json
coverage:
name: coverage
needs: unit
@ -1316,7 +1385,7 @@ jobs:
unit-passed:
name: unit passed
permissions: {}
needs: [rust-bridge, assert-shard-coverage, unit, ui-unit, docs, helm]
needs: [rust-bridge, assert-shard-coverage, unit, ui-unit, docs, helm, openai-sdk-compat]
if: always()
runs-on: ubuntu-latest
timeout-minutes: 2

View file

@ -100,6 +100,7 @@ from litellm.constants import (
DEFAULT_ALLOWED_FAILS,
)
import httpx
from openai import DefaultAsyncHttpxClient, DefaultHttpxClient
# register_async_client_cleanup is lazy-loaded and called on first access
@ -420,8 +421,8 @@ error_logs: Dict = {}
add_function_to_prompt: bool = (
False # if function calling not supported by api, append function call details to system prompt
)
client_session: Optional[httpx.Client] = None
aclient_session: Optional[httpx.AsyncClient] = None
client_session: Optional[Union[httpx.Client, DefaultHttpxClient]] = None
aclient_session: Optional[Union[httpx.AsyncClient, DefaultAsyncHttpxClient]] = None
model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks'
model_cost_map_url: str = os.getenv(
"LITELLM_MODEL_COST_MAP_URL",

View file

@ -241,14 +241,7 @@ class BadRequestError(openai.BadRequestError):
self.litellm_debug_info = litellm_debug_info
self.max_retries = max_retries
self.num_retries = num_retries
# Use response if it's a valid httpx.Response with a request, otherwise use minimal error response
# Note: We check _request (not .request property) to avoid RuntimeError when _request is None
if (
response is not None
and isinstance(response, httpx.Response)
and hasattr(response, "_request")
and getattr(response, "_request", None) is not None
):
if response is not None and getattr(response, "_request", None) is not None:
self.response = response
else:
self.response = _get_minimal_error_response()
@ -360,6 +353,8 @@ class UnprocessableEntityError(openai.UnprocessableEntityError):
class Timeout(openai.APITimeoutError):
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -458,6 +453,9 @@ class RateLimitError(openai.RateLimitError):
:class:`RateLimitErrorCategory` for the available values.
"""
response: httpx.Response # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX response
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -515,7 +513,9 @@ class RateLimitError(openai.RateLimitError):
),
)
super().__init__(
self.message, response=self.response, body=body
self.message,
response=self.response, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
body=body,
) # Call the base class constructor with the parameters it needs
self.code = "429"
self.type = "throttling_error"
@ -594,12 +594,7 @@ class PaymentRequiredError(BadRequestError):
response: httpx.Response | None = None,
litellm_debug_info: str | None = None,
) -> None:
response_is_valid: Final = (
response is not None
and isinstance(response, httpx.Response)
and hasattr(response, "_request")
and getattr(response, "_request", None) is not None
)
response_is_valid: Final = response is not None and getattr(response, "_request", None) is not None
response_for_parent: Final = (
response
if response_is_valid
@ -770,6 +765,9 @@ class ServiceUnavailableError(openai.APIStatusError):
class BadGatewayError(openai.APIStatusError):
response: httpx.Response # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX response
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -797,7 +795,9 @@ class BadGatewayError(openai.APIStatusError):
),
)
super().__init__(
self.message, response=self.response, body=None
self.message,
response=self.response, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
body=None,
) # Call the base class constructor with the parameters it needs
def __str__(self):
@ -818,6 +818,9 @@ class BadGatewayError(openai.APIStatusError):
class InternalServerError(openai.InternalServerError):
response: httpx.Response # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX response
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -846,7 +849,9 @@ class InternalServerError(openai.InternalServerError):
),
)
super().__init__(
self.message, response=self.response, body=body
self.message,
response=self.response, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
body=body,
) # Call the base class constructor with the parameters it needs
self.type = "internal_server_error"
@ -911,6 +916,8 @@ class APIError(openai.APIError):
# raised if an invalid request (not get, delete, put, post) is made
class APIConnectionError(openai.APIConnectionError):
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -929,7 +936,10 @@ class APIConnectionError(openai.APIConnectionError):
self.request = httpx.Request(method="POST", url="https://api.openai.com/v1")
self.max_retries = max_retries
self.num_retries = num_retries
super().__init__(message=self.message, request=self.request)
super().__init__(
message=self.message,
request=self.request, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
)
def __str__(self):
_message = self.message
@ -950,6 +960,9 @@ class APIConnectionError(openai.APIConnectionError):
# raised if an invalid request (not get, delete, put, post) is made
class APIResponseValidationError(openai.APIResponseValidationError):
response: httpx.Response # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX response
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(
self,
message,
@ -967,7 +980,11 @@ class APIResponseValidationError(openai.APIResponseValidationError):
self.litellm_debug_info = litellm_debug_info
self.max_retries = max_retries
self.num_retries = num_retries
super().__init__(response=response, body=None, message=message)
super().__init__(
response=response, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
body=None,
message=message,
)
def __str__(self):
_message = self.message
@ -1087,6 +1104,9 @@ class BudgetExceededError(Exception):
## DEPRECATED ##
class InvalidRequestError(openai.BadRequestError):
response: httpx.Response # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX response
request: httpx.Request # pyright: ignore[reportIncompatibleVariableOverride] # LiteLLM constructs an HTTPX request
def __init__(self, message, model, llm_provider):
self.status_code = 400
self.message = message
@ -1097,7 +1117,9 @@ class InvalidRequestError(openai.BadRequestError):
request=httpx.Request(method="GET", url="https://litellm.ai"), # mock request object
)
super().__init__(
message=self.message, response=self.response, body=None
message=self.message,
response=self.response, # pyright: ignore[reportArgumentType] # SDK accepts HTTPX at runtime
body=None,
) # Call the base class constructor with the parameters it needs

View file

@ -6,6 +6,7 @@ from collections.abc import Callable
from typing import Final
import httpx
from openai import Timeout as SDKTimeout
from litellm.constants import COMPLETION_HTTP_FALLBACK_SECONDS
@ -13,6 +14,12 @@ from litellm.constants import COMPLETION_HTTP_FALLBACK_SECONDS
class CompletionTimeout:
"""Resolves HTTP timeout for ``completion()`` from model vs global settings."""
@staticmethod
def normalize(timeout: httpx.Timeout | SDKTimeout) -> httpx.Timeout:
if isinstance(timeout, httpx.Timeout): # pyright: ignore[reportUnnecessaryIsInstance] # SDK 2 aliases both classes.
return timeout
return httpx.Timeout(connect=timeout.connect, read=timeout.read, write=timeout.write, pool=timeout.pool)
@staticmethod
def _fallback_when_no_explicit_timeout(
global_timeout: float | str | None,
@ -31,7 +38,7 @@ class CompletionTimeout:
@staticmethod
def resolve(
model_timeout: float | str | httpx.Timeout | None,
model_timeout: float | str | httpx.Timeout | SDKTimeout | None,
kwargs: dict,
custom_llm_provider: str,
*,
@ -49,7 +56,7 @@ class CompletionTimeout:
Coerce :class:`httpx.Timeout` when the provider does not support it.
"""
resolved: float | str | httpx.Timeout
resolved: float | str | httpx.Timeout | SDKTimeout
if model_timeout is not None:
resolved = model_timeout
elif kwargs.get("timeout") is not None:
@ -59,12 +66,12 @@ class CompletionTimeout:
else:
resolved = CompletionTimeout._fallback_when_no_explicit_timeout(global_timeout)
if isinstance(resolved, httpx.Timeout) and not supports_httpx_timeout(custom_llm_provider):
read_timeout: Final = resolved.read
resolved = (
if isinstance(resolved, (httpx.Timeout, SDKTimeout)):
normalized: Final = CompletionTimeout.normalize(resolved)
if supports_httpx_timeout(custom_llm_provider):
return normalized
read_timeout: Final = normalized.read
return (
float(read_timeout) if read_timeout is not None else COMPLETION_HTTP_FALLBACK_SECONDS
) # default 10 min timeout
elif not isinstance(resolved, httpx.Timeout):
resolved = float(resolved)
return resolved
return float(resolved)

View file

@ -16,7 +16,7 @@ from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.constants import DEFAULT_MAX_RETRIES
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.common_utils import BaseOpenAILLM
from litellm.llms.openai.common_utils import BaseOpenAILLM, OpenAIAsyncHTTPClient, OpenAIHTTPClient
from litellm.secret_managers.get_azure_ad_token_provider import (
get_azure_ad_token_provider,
)
@ -733,7 +733,11 @@ class BaseAzureLLM(BaseOpenAILLM):
azure_client_params: Final[_AzureGatewayClientParams] = {
"api_version": api_version,
"base_url": f"{api_base}",
"http_client": litellm.client_session,
"http_client": (
litellm.aclient_session or OpenAIAsyncHTTPClient()
if acompletion
else litellm.client_session or OpenAIHTTPClient()
),
"max_retries": max_retries,
"timeout": timeout,
}

View file

@ -2,6 +2,7 @@
Common helpers / utils across al OpenAI endpoints
"""
import asyncio
import hashlib
import inspect
import json
@ -10,12 +11,13 @@ import ssl
import time
import uuid
from collections.abc import AsyncIterator, Iterator, Mapping
from contextlib import suppress
from typing import TYPE_CHECKING, Final, Literal, NamedTuple, Optional
from urllib.parse import urlsplit
import httpx
import openai
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, DefaultAsyncHttpxClient, DefaultHttpxClient, OpenAI
from openai.types.chat import ChatCompletion, ChatCompletionChunk, ChatCompletionMessage
from openai.types.chat.chat_completion import Choice
from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice
@ -26,6 +28,7 @@ if TYPE_CHECKING:
from aiohttp import ClientSession
import litellm
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.token_counter import token_counter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
@ -48,6 +51,42 @@ _AZURE_OPENAI_INIT_PARAMS: Final[tuple[str, ...]] = _get_client_init_params(Azur
_OPENAI_API_HOST: Final[str] = "api.openai.com"
_OPENAI_HTTPX_DEFAULT_TIMEOUT: Final = CompletionTimeout.normalize(openai.DEFAULT_TIMEOUT)
_OPENAI_HTTPX_CONNECTION_LIMITS: Final = httpx.Limits(
max_connections=openai.DEFAULT_CONNECTION_LIMITS.max_connections,
max_keepalive_connections=openai.DEFAULT_CONNECTION_LIMITS.max_keepalive_connections,
keepalive_expiry=openai.DEFAULT_CONNECTION_LIMITS.keepalive_expiry,
)
class OpenAIHTTPClient(httpx.Client):
def __init__(self) -> None:
super().__init__(
timeout=_OPENAI_HTTPX_DEFAULT_TIMEOUT,
limits=_OPENAI_HTTPX_CONNECTION_LIMITS,
follow_redirects=True,
)
def __del__(self) -> None:
with suppress(Exception):
if not self.is_closed:
self.close()
class OpenAIAsyncHTTPClient(httpx.AsyncClient):
def __init__(self) -> None:
super().__init__(
timeout=_OPENAI_HTTPX_DEFAULT_TIMEOUT,
limits=_OPENAI_HTTPX_CONNECTION_LIMITS,
follow_redirects=True,
)
def __del__(self) -> None:
with suppress(Exception):
if not self.is_closed:
asyncio.get_running_loop().create_task(self.aclose())
def is_openai_backed_api_base(api_base: str) -> bool:
hostname: Final = urlsplit(api_base).hostname
return hostname is not None and (hostname == _OPENAI_API_HOST or hostname.endswith(f".{_OPENAI_API_HOST}"))
@ -221,7 +260,9 @@ class BaseOpenAILLM:
return _cached_client
@staticmethod
def owns_wrapped_http_client(http_client: httpx.Client | httpx.AsyncClient | None) -> bool:
def owns_wrapped_http_client(
http_client: httpx.Client | httpx.AsyncClient | DefaultHttpxClient | DefaultAsyncHttpxClient | None,
) -> bool:
"""Whether litellm may close an SDK client built around ``http_client``.
``_get_async_http_client`` / ``_get_sync_http_client`` hand back
@ -304,7 +345,7 @@ class BaseOpenAILLM:
@staticmethod
def _get_async_http_client(
shared_session: Optional["ClientSession"] = None,
) -> httpx.AsyncClient | None:
) -> httpx.AsyncClient | DefaultAsyncHttpxClient | None:
if litellm.aclient_session is not None:
return litellm.aclient_session
@ -337,7 +378,7 @@ class BaseOpenAILLM:
return cls._get_async_http_client(shared_session)
@staticmethod
def _get_sync_http_client() -> httpx.Client | None:
def _get_sync_http_client() -> httpx.Client | DefaultHttpxClient | None:
if litellm.client_session is not None:
return litellm.client_session

View file

@ -12,7 +12,7 @@ from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUser
from litellm.types.utils import LlmProviders, ModelResponse, TextCompletionResponse
from litellm.utils import ProviderConfigManager
from ..common_utils import BaseOpenAILLM, OpenAIError
from ..common_utils import BaseOpenAILLM, OpenAIAsyncHTTPClient, OpenAIError, OpenAIHTTPClient
from .transformation import OpenAITextCompletionConfig
@ -129,7 +129,7 @@ class OpenAITextCompletion(BaseLLM):
openai_client = OpenAI(
api_key=api_key,
base_url=api_base,
http_client=litellm.client_session,
http_client=litellm.client_session or OpenAIHTTPClient(),
timeout=timeout,
max_retries=max_retries,
organization=organization,
@ -233,7 +233,7 @@ class OpenAITextCompletion(BaseLLM):
openai_client = OpenAI(
api_key=api_key,
base_url=api_base,
http_client=litellm.client_session,
http_client=litellm.client_session or OpenAIHTTPClient(),
timeout=timeout,
max_retries=max_retries,
organization=organization,
@ -290,7 +290,7 @@ class OpenAITextCompletion(BaseLLM):
openai_client = AsyncOpenAI(
api_key=api_key,
base_url=api_base,
http_client=litellm.aclient_session,
http_client=litellm.aclient_session or OpenAIAsyncHTTPClient(),
timeout=timeout,
max_retries=max_retries,
organization=organization,

View file

@ -6,6 +6,7 @@ from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from openai.types.fine_tuning import FineTuningJob
from litellm._logging import verbose_logger
from litellm.llms.openai.common_utils import OpenAIAsyncHTTPClient, OpenAIHTTPClient
from litellm.types.utils import LiteLLMFineTuningJob
_AZURE_STATUS_MAP: Final[Mapping[object, str]] = {
@ -87,9 +88,9 @@ class OpenAIFineTuningAPI:
elif v is not None:
data[k] = v
if _is_async is True:
openai_client = AsyncOpenAI(**data)
openai_client = AsyncOpenAI(**data, http_client=OpenAIAsyncHTTPClient())
else:
openai_client = OpenAI(**data)
openai_client = OpenAI(**data, http_client=OpenAIHTTPClient())
else:
openai_client = client

View file

@ -14,7 +14,7 @@ from litellm.utils import ProviderConfigManager
from ...base_llm.image_variations.transformation import BaseImageVariationConfig
from ...custom_httpx.llm_http_handler import LiteLLMLoggingObj
from ..common_utils import OpenAIError
from ..common_utils import OpenAIAsyncHTTPClient, OpenAIError, OpenAIHTTPClient
class OpenAIImageVariationsHandler:
@ -23,22 +23,14 @@ class OpenAIImageVariationsHandler:
client: OpenAI | None,
init_client_params: dict,
):
if client is None:
openai_client = OpenAI(
**init_client_params,
)
else:
openai_client = client
return openai_client
if client is not None:
return client
return OpenAI(**init_client_params, http_client=litellm.client_session or OpenAIHTTPClient())
def get_async_client(self, client: AsyncOpenAI | None, init_client_params: dict) -> AsyncOpenAI:
if client is None:
openai_client = AsyncOpenAI(
**init_client_params,
)
else:
openai_client = client
return openai_client
if client is not None:
return client
return AsyncOpenAI(**init_client_params, http_client=litellm.aclient_session or OpenAIAsyncHTTPClient())
async def async_image_variations(
self,
@ -62,7 +54,6 @@ class OpenAIImageVariationsHandler:
init_client_params: Final = {
"api_key": api_key,
"base_url": api_base,
"http_client": litellm.client_session,
"timeout": timeout,
"max_retries": max_retries,
"organization": organization,
@ -179,7 +170,6 @@ class OpenAIImageVariationsHandler:
init_client_params: Final = {
"api_key": api_key,
"base_url": api_base,
"http_client": litellm.client_session,
"timeout": timeout,
"max_retries": max_retries,
"organization": organization,

View file

@ -53,7 +53,9 @@ from .chat.gpt_transformation import OpenAIGPTConfig, OpenAIUnknownModelConfig
from .chat.o_series_transformation import OpenAIOSeriesConfig
from .common_utils import (
BaseOpenAILLM,
OpenAIAsyncHTTPClient,
OpenAIError,
OpenAIHTTPClient,
build_output_token_limit_response,
drop_params_from_unprocessable_entity_error,
is_openai_backed_api_base,
@ -1749,9 +1751,9 @@ class OpenAIFilesAPI(BaseLLM):
elif v is not None:
data[k] = v
if _is_async is True:
openai_client = AsyncOpenAI(**data)
openai_client = AsyncOpenAI(**data, http_client=OpenAIAsyncHTTPClient())
else:
openai_client = OpenAI(**data)
openai_client = OpenAI(**data, http_client=OpenAIHTTPClient())
else:
openai_client = client
@ -2107,9 +2109,9 @@ class OpenAIBatchesAPI(BaseLLM):
elif v is not None:
data[k] = v
if _is_async is True:
openai_client = AsyncOpenAI(**data)
openai_client = AsyncOpenAI(**data, http_client=OpenAIAsyncHTTPClient())
else:
openai_client = OpenAI(**data)
openai_client = OpenAI(**data, http_client=OpenAIHTTPClient())
else:
openai_client = client
@ -2317,7 +2319,7 @@ class OpenAIAssistantsAPI(BaseLLM):
data["base_url"] = v
elif v is not None:
data[k] = v
openai_client = OpenAI(**data)
openai_client = OpenAI(**data, http_client=OpenAIHTTPClient())
else:
openai_client = client
@ -2342,7 +2344,7 @@ class OpenAIAssistantsAPI(BaseLLM):
data["base_url"] = v
elif v is not None:
data[k] = v
openai_client = AsyncOpenAI(**data)
openai_client = AsyncOpenAI(**data, http_client=OpenAIAsyncHTTPClient())
else:
openai_client = client

View file

@ -432,7 +432,7 @@ async def acompletion(
messages: list = [],
functions: list | None = None,
function_call: str | None = None,
timeout: float | None = None,
timeout: float | httpx.Timeout | openai.Timeout | None = None,
temperature: float | None = None,
top_p: float | None = None,
n: int | None = None,
@ -5225,7 +5225,7 @@ def completion(
model: str,
# Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create
messages: list = [],
timeout: float | str | httpx.Timeout | None = None,
timeout: float | str | httpx.Timeout | openai.Timeout | None = None,
temperature: float | None = None,
top_p: float | None = None,
n: int | None = None,
@ -7880,6 +7880,8 @@ def adapter_completion(*, adapter_id: str, **kwargs) -> BaseModel | AdapterCompl
def moderation(input: str, model: str | None = None, api_key: str | None = None, **kwargs) -> OpenAIModerationResponse:
from litellm.llms.openai.common_utils import OpenAIHTTPClient
# only supports open ai for now
api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
@ -7888,10 +7890,7 @@ def moderation(input: str, model: str | None = None, api_key: str | None = None,
openai_client = kwargs.get("client", None)
if openai_client is None:
if api_base is not None:
openai_client = openai.OpenAI(api_key=api_key, base_url=api_base)
else:
openai_client = openai.OpenAI(api_key=api_key)
openai_client = openai.OpenAI(api_key=api_key, base_url=api_base, http_client=OpenAIHTTPClient())
if model is not None:
response = openai_client.moderations.create(input=input, model=model)
@ -8364,7 +8363,7 @@ def speech(
project: str | None = None,
max_retries: int | None = None,
metadata: dict | None = None,
timeout: float | httpx.Timeout | None = None,
timeout: float | httpx.Timeout | openai.Timeout | None = None,
response_format: str | None = None,
speed: int | None = None,
instructions: str | None = None,
@ -8394,9 +8393,12 @@ def speech(
if instructions is not None:
optional_params["instructions"] = instructions
if timeout is None:
timeout = litellm.request_timeout
timeout_or_default: Final = litellm.request_timeout if timeout is None else timeout
http_timeout: Final = (
CompletionTimeout.normalize(timeout_or_default)
if isinstance(timeout_or_default, openai.Timeout)
else timeout_or_default
)
if max_retries is None:
max_retries = litellm.num_retries or openai.DEFAULT_MAX_RETRIES
litellm_params_dict: Final = get_litellm_params(metadata=metadata, api_key=api_key or dynamic_api_key, **kwargs)
@ -8445,7 +8447,7 @@ def speech(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
client=client,
_is_async=aspeech or False,
@ -8501,7 +8503,7 @@ def speech(
organization=organization,
project=project,
max_retries=max_retries,
timeout=timeout,
timeout=http_timeout,
logging_obj=logging_obj,
client=client, # pass AsyncOpenAI, OpenAI client
aspeech=aspeech,
@ -8532,7 +8534,7 @@ def speech(
optional_params=optional_params,
litellm_params_dict=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
base_llm_http_handler=base_llm_http_handler,
aspeech=aspeech or False,
@ -8580,7 +8582,7 @@ def speech(
azure_ad_token_provider=azure_ad_token_provider,
organization=organization,
max_retries=max_retries,
timeout=timeout,
timeout=http_timeout,
logging_obj=logging_obj,
client=client, # pass AsyncOpenAI, OpenAI client
aspeech=aspeech,
@ -8625,7 +8627,7 @@ def speech(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
client=client,
_is_async=aspeech or False,
@ -8682,7 +8684,7 @@ def speech(
optional_params=optional_params,
litellm_params_dict=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=headers,
base_llm_http_handler=base_llm_http_handler,
aspeech=aspeech or False,
@ -8728,7 +8730,7 @@ def speech(
optional_params=optional_params,
litellm_params_dict=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
base_llm_http_handler=base_llm_http_handler,
aspeech=aspeech or False,
@ -8769,7 +8771,7 @@ def speech(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
client=client,
_is_async=aspeech or False,
@ -8797,7 +8799,7 @@ def speech(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
client=client,
_is_async=aspeech or False,
@ -8821,7 +8823,7 @@ def speech(
optional_params=optional_params,
litellm_params_dict=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
timeout=http_timeout,
extra_headers=extra_headers,
base_llm_http_handler=base_llm_http_handler,
aspeech=aspeech or False,

View file

@ -12919,7 +12919,7 @@ async def audio_speech(
requested_format: Final = data.get("response_format")
upstream_content_type: Final = (
response.response.headers.get("content-type") if isinstance(response, HttpxBinaryResponseContent) else None
response.content_type if isinstance(response, HttpxBinaryResponseContent) else None
)
media_type: Final = resolve_speech_media_type(
upstream_content_type=upstream_content_type,

View file

@ -2,13 +2,14 @@ import builtins
from collections.abc import Iterable, Mapping
from enum import Enum
from os import PathLike
from typing import IO, Any, Final, Literal, Optional, TypeAlias, Union
from typing import IO, Any, Final, Generic, Literal, Optional, TypeAlias, Union
import httpx
from openai import Omit
from openai._legacy_response import (
HttpxBinaryResponseContent as _HttpxBinaryResponseContent,
)
from openai._types import Response as SDKResponse
from openai.lib.streaming._assistants import (
AssistantEventHandler,
AssistantStreamManager,
@ -79,6 +80,7 @@ from typing_extensions import (
ReadOnly,
Required,
TypedDict,
TypeVar,
override,
)
@ -118,8 +120,12 @@ class BinaryResponseSummary(TypedDict):
num_bytes: ReadOnly[int]
class HttpxBinaryResponseContent(_HttpxBinaryResponseContent):
_ResponseT = TypeVar("_ResponseT", bound=httpx.Response | SDKResponse, default=httpx.Response)
class HttpxBinaryResponseContent(_HttpxBinaryResponseContent, Generic[_ResponseT]):
_hidden_params: dict
response: _ResponseT # pyright: ignore[reportIncompatibleVariableOverride] # SDK accepts both backends at runtime
@property
def hidden_params(self) -> dict[str, builtins.object]: # mutable-ok: API requires mutation
@ -129,22 +135,28 @@ class HttpxBinaryResponseContent(_HttpxBinaryResponseContent):
def hidden_params(self, hidden_params: dict[str, builtins.object]) -> None: # mutable-ok: API requires mutation
self._hidden_params = hidden_params
def __init__(self, response: httpx.Response) -> None:
super().__init__(response)
def __init__(self, response: _ResponseT) -> None:
super().__init__(response) # pyright: ignore[reportArgumentType] # SDK accepts both backends at runtime
self._hidden_params = {}
@property
def content_type(self) -> str | None:
headers: Final[Mapping[str, str]] = self.response.headers
return headers.get("content-type")
def logging_summary(self) -> BinaryResponseSummary:
return {
"object": "binary",
"content_type": self.response.headers.get("content-type"),
"content_type": self.content_type,
"num_bytes": self._num_bytes(),
}
def _num_bytes(self) -> int:
try:
return len(self.response.content)
except httpx.ResponseNotRead:
content: Final = self.response.content
except RuntimeError:
return self.response.num_bytes_downloaded
return len(content)
def set_response_cost(self, response_cost: float | None) -> None:
if response_cost is None:

View file

@ -21,11 +21,13 @@ from typing import (
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
import httpx
from openai import Timeout as SDKTimeout
from pydantic import ConfigDict, Field, JsonValue, field_validator, model_validator
from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_checkable
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.litellm_core_utils.provider_affinity import validate_provider_affinity_header_name
from litellm.types.llms.base import LiteLLMBaseModel
@ -481,7 +483,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
rpm: int | None = None
itpm: int | None = None
otpm: int | None = None
timeout: float | str | httpx.Timeout | None = None # if str, pass in as os.environ/
timeout: float | str | httpx.Timeout | SDKTimeout | None = None # if str, pass in as os.environ/
stream_timeout: float | str | None = None # timeout when making stream=True calls, if str, pass in as os.environ/
max_retries: int | None = None
drop_params: bool | str | None = None
@ -564,6 +566,13 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
valkey_text_field: str | None = None
valkey_embedding_field: str | None = None
@field_validator("timeout")
@classmethod
def normalize_timeout(
cls, value: float | str | httpx.Timeout | SDKTimeout | None
) -> float | str | httpx.Timeout | None:
return CompletionTimeout.normalize(value) if isinstance(value, (httpx.Timeout, SDKTimeout)) else value
@field_validator("provider_affinity_header")
@classmethod
def validate_provider_affinity_header(cls, value: str | None) -> str | None:
@ -677,7 +686,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
api_key: str | None
api_base: str | None
api_version: str | None
timeout: float | str | httpx.Timeout | None
timeout: float | str | httpx.Timeout | SDKTimeout | None # writable-ok: preserve timeout updates
stream_timeout: float | str | None
max_retries: int | None
organization: list | str | None # for openai orgs

View file

@ -13,7 +13,7 @@ dependencies = [
"fastuuid>=0.14.0,<1.0",
"filelock>=3.16.1,<4.0",
"httpx[http2]>=0.28.0,<1.0",
"openai>=2.20.0,<3.0.0",
"openai>=2.20.0,<4.0.0",
"python-dateutil>=2.8.2,<3.0",
"python-dotenv>=1.0.0,<2.0",
"pyyaml>=6.0.3,<7.0",

View file

@ -17,7 +17,7 @@ dependencies = [
"fastuuid>=0.14.0,<1.0",
"filelock>=3.16.1,<4.0",
"httpx[http2]>=0.28.0,<1.0",
"openai>=2.20.0,<3.0.0",
"openai>=2.20.0,<4.0.0",
"python-dotenv>=1.0.0,<2.0",
"pyyaml>=6.0.3,<7.0",
"packaging>=24.0",

View file

@ -8,11 +8,18 @@ the very class of undeclared-dependency bug this guards against.
import argparse
import importlib.util
import re
import sys
import traceback
from collections.abc import Callable
from importlib.metadata import PackageNotFoundError, packages_distributions, requires
from itertools import chain
from typing import Final
EXTRAS_ONLY_MODULES = ("fastapi", "uvicorn", "keyring", "mcp", "mcp_types", "httpx2", "httpcore2")
SDK_DISTRIBUTIONS: Final = frozenset({"litellm", "litellm-core"})
REQUIREMENT_NAME: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]*")
EXTRA_MARKER: Final = re.compile(r";.*\bextra\s*==")
def _require(condition: bool, message: str) -> None:
@ -20,13 +27,43 @@ def _require(condition: bool, message: str) -> None:
raise AssertionError(message)
def _canonical(name: str) -> str:
return re.sub(r"[-_.]+", "-", name).lower()
def _base_requirements(distribution: str) -> frozenset[str]:
try:
declared: Final = requires(distribution) or ()
except PackageNotFoundError:
return frozenset()
return frozenset(
_canonical(match.group(0))
for requirement in declared
if not EXTRA_MARKER.search(requirement) and (match := REQUIREMENT_NAME.match(requirement))
)
def _base_closure(pending: frozenset[str], reached: frozenset[str] = frozenset()) -> frozenset[str]:
if not pending:
return reached
expanded: Final = reached | pending
return _base_closure(frozenset(chain.from_iterable(_base_requirements(name) for name in pending)) - expanded, expanded)
def check_environment_is_base_only() -> str:
present = tuple(name for name in EXTRAS_ONLY_MODULES if importlib.util.find_spec(name) is not None)
from_base: Final = _base_closure(SDK_DISTRIBUTIONS)
owners: Final = packages_distributions()
installed: Final = tuple(name for name in EXTRAS_ONLY_MODULES if importlib.util.find_spec(name) is not None)
inherited: Final = tuple(
name for name in installed if from_base & {_canonical(owner) for owner in owners.get(name) or (name,)}
)
present: Final = tuple(name for name in installed if name not in inherited)
_require(
not present,
f"{', '.join(present)} installed, so this environment is not base-only and the run proves nothing",
)
return f"no extras-only packages present ({', '.join(EXTRAS_ONLY_MODULES)})"
from_dependencies: Final = f"; {', '.join(inherited)} come from base dependencies" if inherited else ""
return f"no extras-only packages present ({', '.join(EXTRAS_ONLY_MODULES)}){from_dependencies}"
def check_import() -> str:

View file

@ -37,7 +37,8 @@ class RecordingUpstream(BaseHTTPRequestHandler):
pass
def do_POST(self) -> None:
REQUESTS.put((self.path, dict(self.headers), self.rfile.read(int(self.headers["Content-Length"]))))
headers: Final = {name.lower(): value for name, value in self.headers.items()}
REQUESTS.put((self.path, headers, self.rfile.read(int(self.headers["Content-Length"]))))
status, body = RESPONSES.get(timeout=10)
ARRIVED.set()
if status == 0:
@ -78,7 +79,7 @@ def check_http(base: str) -> None:
assert response.usage.total_tokens == 5
path, headers, body = REQUESTS.get(timeout=10)
assert path == "/v1/chat/completions"
assert headers["Authorization"] == "Bearer test-key"
assert headers["authorization"] == "Bearer test-key"
assert json.loads(body)["messages"] == MESSAGES
enqueue_stream()
assert "".join(part.choices[0].delta.content or "" for part in litellm.completion(**arguments, stream=True)) == "pong"
@ -156,7 +157,7 @@ def check_http(base: str) -> None:
assert result.usage.total_tokens == 5
path, headers, body = REQUESTS.get(timeout=10)
assert path.endswith("/converse")
assert {key.lower(): value for key, value in headers.items()}["authorization"].startswith(
assert headers["authorization"].startswith(
"AWS4-HMAC-SHA256" if signed else "Bearer bearer-key"
)
assert json.loads(body)["messages"][0]["content"][0]["text"] == "ping"

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,424 @@
from __future__ import annotations
import json
import signal
import subprocess
import uuid
from collections.abc import Mapping
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass
from functools import partial
from itertools import product
from pathlib import Path
from typing import Final, Literal, TypeAlias
import httpx
import psutil
import pytest
from anthropic import Anthropic
from anthropic import APIError as AnthropicApiError
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value
from integration._support.database import read_rows
from integration._support.process import graceful_stop_seconds, owned_proxy_process, owned_upstream
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
from openai import APIError as OpenAiApiError
from openai import OpenAI
from pydantic import JsonValue
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, SseResponse
_Kind: TypeAlias = Literal["chat", "chat_stream", "messages", "responses", "completions", "moderations"]
_ALL_KINDS: Final[tuple[_Kind, ...]] = ("chat", "chat_stream", "messages", "responses", "completions", "moderations")
_CHAT_KINDS: Final[tuple[_Kind, ...]] = ("chat", "chat_stream")
_NO_CACHE: Final = {"cache": {"no-cache": True}}
_SPEND_QUERY: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
_CHAT: Final = JsonResponse(
content_type="application/json",
body={
"id": "chatcmpl-$UNIQUE_ID",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
)
_MODERATION: Final = JsonResponse(
content_type="application/json",
body={
"id": "modr-$UNIQUE_ID",
"model": "omni-moderation-latest",
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
},
)
_RESPONSES: Final = JsonResponse(
content_type="application/json",
body={
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_$UNIQUE_ID",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "scripted", "annotations": []}],
}
],
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
},
)
_PLAIN: Final = RoutedResponse(
content_type="application/x-routed",
routes={"POST /chat/completions": _CHAT, "POST /moderations": _MODERATION, "POST /responses": _RESPONSES},
)
_STREAM: Final = SseResponse(
content_type="text/event-stream",
frames=(
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
'"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}',
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
'"choices":[{"index":0,"delta":{"content":"streamed"},"finish_reason":null}]}',
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
'"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}',
"data: [DONE]",
),
frame_delay_ms=150,
)
@dataclass(frozen=True, slots=True)
class _Call:
kind: _Kind
marker: str
@dataclass(frozen=True, slots=True)
class _Completed:
call: _Call
response_id: str
@dataclass(frozen=True, slots=True)
class _Failed:
call: _Call
error: str
@dataclass(frozen=True, slots=True)
class _Clients:
openai: OpenAI
anthropic: Anthropic
plain_model: str
stream_model: str
class _Observations:
def __init__(self, url: str) -> None:
self.url = url.rstrip("/")
self.items: tuple[Mapping[str, JsonValue], ...] = ()
def read(self) -> tuple[Mapping[str, JsonValue], ...]:
with httpx.Client(timeout=10, trust_env=False) as client:
payload: Final = JSON_OBJECT.validate_python(client.get(f"{self.url}/__observations").json())
requests: Final = payload.get("requests")
assert isinstance(requests, list)
self.items = (*self.items, *(object_value(item) for item in requests if isinstance(item, dict)))
return self.items
def _register(
scenario: Scenario, upstream_url: str, label: str, response: RoutedResponse | SseResponse
) -> ScenarioHandle:
handle: Final = register_scenario(f"chaos-{label}-{uuid.uuid4().hex}", response, control_url=upstream_url)
scenario.cleanups.callback(delete_scenario, handle)
return handle
def _deployments(scenario: Scenario, upstream_url: str) -> tuple[str, str]:
plain_handle: Final = _register(scenario, upstream_url, "plain", _PLAIN)
stream_handle: Final = _register(scenario, upstream_url, "stream", _STREAM)
return (
scenario.model(api_base=plain_handle.api_base(), api_key=plain_handle.scenario_id),
scenario.model(api_base=stream_handle.api_base(), api_key=stream_handle.scenario_id),
)
def _clients(gateway: Gateway, deployments: tuple[str, str]) -> _Clients:
base: Final = str(gateway.client.base_url).rstrip("/")
clients: Final = _Clients(
openai=OpenAI(base_url=f"{base}/v1", api_key=gateway.key, max_retries=0),
anthropic=Anthropic(base_url=base, api_key=gateway.key, max_retries=0),
plain_model=deployments[0],
stream_model=deployments[1],
)
_import_lazy_sdk_resources_on_this_thread(clients)
return clients
def _import_lazy_sdk_resources_on_this_thread(clients: _Clients) -> None:
resources: Final = (
clients.openai.chat.completions,
clients.openai.completions,
clients.openai.responses,
clients.openai.moderations,
clients.anthropic.messages,
)
assert all(resource is not None for resource in resources)
def _invoke(clients: _Clients, call: _Call) -> str:
match call.kind:
case "chat":
return clients.openai.chat.completions.create(
model=clients.plain_model, messages=[{"role": "user", "content": call.marker}], extra_body=_NO_CACHE
).id
case "chat_stream":
chunks: Final = tuple(
clients.openai.chat.completions.create(
model=clients.stream_model,
messages=[{"role": "user", "content": call.marker}],
stream=True,
extra_body=_NO_CACHE,
)
)
assert chunks[-1].choices[0].finish_reason == "stop", chunks
assert len({chunk.id for chunk in chunks}) == 1, chunks
return chunks[0].id
case "messages":
return clients.anthropic.messages.create(
model=clients.plain_model,
max_tokens=64,
messages=[{"role": "user", "content": call.marker}],
extra_body=_NO_CACHE,
).id
case "responses":
return clients.openai.responses.create(
model=clients.plain_model, input=call.marker, extra_body=_NO_CACHE
).id
case "completions":
return clients.openai.completions.create(
model=clients.plain_model, prompt=call.marker, extra_body=_NO_CACHE
).id
case "moderations":
return clients.openai.moderations.create(
model=clients.plain_model, input=call.marker, extra_body=_NO_CACHE
).id
def _attempt(clients: _Clients, call: _Call) -> _Completed | _Failed:
try:
return _Completed(call, _invoke(clients, call))
except (OpenAiApiError, AnthropicApiError, httpx.HTTPError) as error:
return _Failed(call, f"{type(error).__name__}: {str(error)[:200]}")
def _every_worker_serves_every_kind(gateway: Gateway, deployments: tuple[str, str]) -> bool:
calls: Final = _burst_calls(_ALL_KINDS, 2)
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
futures: Final = _submit(pool, _clients(gateway, deployments), calls)
outcomes: Final = tuple(future.result(timeout=60) for future in futures)
return all(isinstance(outcome, _Completed) for outcome in outcomes)
def _burst_calls(kinds: tuple[_Kind, ...], per_kind: int) -> tuple[_Call, ...]:
return tuple(_Call(kind, f"burst-{kind}-{uuid.uuid4().hex}") for kind, _ in product(kinds, range(per_kind)))
def _submit(
pool: ThreadPoolExecutor, clients: _Clients, calls: tuple[_Call, ...]
) -> tuple[Future[_Completed | _Failed], ...]:
return tuple(pool.submit(_attempt, clients, call) for call in calls)
def _done_count(futures: tuple[Future[_Completed | _Failed], ...]) -> int:
return sum(future.done() for future in futures)
def _liveliness(gateway: Gateway) -> int:
return gateway.client.get("/health/liveliness").status_code
def _spend_rows(request_id: str) -> tuple[Mapping[str, JsonValue], ...]:
return tuple(read_rows(_SPEND_QUERY, (request_id,)))
def _marker_counts(items: tuple[Mapping[str, JsonValue], ...], markers: tuple[str, ...]) -> tuple[int, ...]:
bodies: Final = tuple(json.dumps(item.get("body")) for item in items)
return tuple(sum(marker in body for body in bodies) for marker in markers)
def _landed_once(response_ids: tuple[str, ...]) -> tuple[tuple[Mapping[str, JsonValue], ...], ...]:
rows: Final = tuple(
eventually(partial(_spend_rows, response_id), lambda values: len(values) == 1, seconds=90)
for response_id in response_ids
)
assert all(found[0]["request_id"] == response_id for found, response_id in zip(rows, response_ids)), rows
return rows
def _is_worker(child: psutil.Process) -> bool:
try:
return "spawn_main" in " ".join(child.cmdline())
except psutil.Error:
return False
def _worker_alive(worker: psutil.Process) -> bool:
try:
return worker.is_running() and worker.status() != psutil.STATUS_ZOMBIE
except psutil.NoSuchProcess:
return False
def _alive_workers(process: subprocess.Popen[bytes]) -> tuple[psutil.Process, ...]:
children: Final = tuple(child for child in psutil.Process(process.pid).children() if _is_worker(child))
return tuple(child for child in children if _worker_alive(child))
def _process_tree(process: subprocess.Popen[bytes]) -> str:
root: Final = psutil.Process(process.pid)
return "\n".join(_process_line(member) for member in (root, *root.children(recursive=True)))
def _process_line(member: psutil.Process) -> str:
try:
return f"{member.pid} {' '.join(member.cmdline())}"
except psutil.Error:
return f"{member.pid} <exited>"
def _kind_counts(outcomes: tuple[_Completed | _Failed, ...]) -> str:
kinds: Final = tuple(outcome.call.kind for outcome in outcomes)
return str({kind: kinds.count(kind) for kind in _ALL_KINDS if kind in kinds})
def test_c01_upstream_pause_mid_burst_keeps_liveliness_and_lands_every_id_once(
gateway: Gateway,
tmp_path: Path,
record_property: pytest.RecordProperty,
) -> None:
with owned_upstream(tmp_path) as slot, gateway.scenario() as scenario:
upstream: Final = slot.process
assert upstream is not None
deployments: Final = _deployments(scenario, slot.url)
eventually(partial(_every_worker_serves_every_kind, gateway, deployments), lambda served: served, seconds=90)
clients: Final = _clients(gateway, deployments)
calls: Final = _burst_calls(_ALL_KINDS, 5)
observations: Final = _Observations(slot.url)
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
futures: Final = _submit(pool, clients, calls)
eventually(partial(_done_count, futures), lambda done: done >= 1, seconds=60)
upstream.send_signal(signal.SIGSTOP)
try:
done_at_pause: Final = _done_count(futures)
paused_liveliness: Final = tuple(_liveliness(gateway) for _ in range(3))
finally:
upstream.send_signal(signal.SIGCONT)
done_at_resume: Final = _done_count(futures)
outcomes: Final = tuple(future.result(timeout=180) for future in futures)
record_property("c01_burst_size", len(calls))
record_property("c01_done_at_pause", done_at_pause)
record_property("c01_done_at_resume", done_at_resume)
record_property("c01_paused_liveliness_statuses", str(paused_liveliness))
assert paused_liveliness == (200, 200, 200), paused_liveliness
assert done_at_resume < len(calls), done_at_resume
failures: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Failed))
assert not failures, failures
completed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Completed))
record_property("c01_completed_by_kind", _kind_counts(completed))
response_ids: Final = tuple(outcome.response_id for outcome in completed)
assert len(set(response_ids)) == len(calls), response_ids
spend_rows: Final = _landed_once(response_ids)
record_property("c01_spend_query", _SPEND_QUERY)
record_property("c01_spend_row_counts", str(tuple(len(rows) for rows in spend_rows)))
markers: Final = tuple(call.marker for call in calls)
eventually(
observations.read,
lambda items: _marker_counts(items, markers) == (1,) * len(markers),
seconds=60,
)
record_property("c01_upstream_marker_counts", str(_marker_counts(observations.items, markers)))
@pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
def test_c02_worker_sigkill_mid_burst_keeps_serving_and_respawns(
gateway: Gateway,
tmp_path: Path,
record_property: pytest.RecordProperty,
) -> None:
with owned_upstream(tmp_path) as slot, gateway.scenario() as scenario:
deployments: Final = _deployments(scenario, slot.url)
with owned_proxy_process(gateway, tmp_path, {"INTEGRATION_UPSTREAM_URL": slot.url}, workers=2) as owned:
eventually(
partial(_every_worker_serves_every_kind, owned.gateway, deployments), lambda served: served, seconds=90
)
clients: Final = _clients(owned.gateway, deployments)
workers: Final = eventually(
partial(_alive_workers, owned.process), lambda found: len(found) == 2, seconds=60
)
record_property("c02_process_tree_before_kill", _process_tree(owned.process))
victim: Final = workers[0]
calls: Final = _burst_calls(_CHAT_KINDS, 15)
observations: Final = _Observations(slot.url)
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
futures: Final = _submit(pool, clients, calls)
eventually(partial(_done_count, futures), lambda done: done >= 1, seconds=60)
victim.kill()
eventually(partial(_worker_alive, victim), lambda alive: not alive, seconds=10)
done_at_kill: Final = _done_count(futures)
survivor: Final = eventually(
lambda: _attempt(clients, _Call("chat", f"after-kill-{uuid.uuid4().hex}")),
lambda outcome: isinstance(outcome, _Completed),
seconds=30,
)
outcomes: Final = tuple(future.result(timeout=180) for future in futures)
record_property("c02_burst_size", len(calls))
record_property("c02_killed_worker_pid", victim.pid)
record_property("c02_done_at_kill", done_at_kill)
assert isinstance(survivor, _Completed), survivor
completed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Completed))
failed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Failed))
record_property("c02_completed_by_kind", _kind_counts(completed))
record_property("c02_inflight_failure_count", len(failed))
record_property("c02_inflight_failure_errors", str(tuple(outcome.error for outcome in failed)))
assert completed, outcomes
respawned: Final = eventually(
partial(_alive_workers, owned.process),
lambda found: len(found) == 2 and victim.pid not in {worker.pid for worker in found},
seconds=graceful_stop_seconds() + 60,
)
record_property("c02_process_tree_after_respawn", _process_tree(owned.process))
record_property("c02_respawned_worker_pids", str(tuple(worker.pid for worker in respawned)))
after_respawn: Final = eventually(
lambda: _attempt(clients, _Call("chat", f"after-respawn-{uuid.uuid4().hex}")),
lambda outcome: isinstance(outcome, _Completed),
seconds=60,
)
assert isinstance(after_respawn, _Completed), after_respawn
_landed_once((survivor.response_id, after_respawn.response_id))
response_ids: Final = tuple(outcome.response_id for outcome in completed)
assert len(set(response_ids)) == len(completed), response_ids
spend_row_counts: Final = tuple((outcome, len(_spend_rows(outcome.response_id))) for outcome in completed)
landed: Final = tuple(outcome for outcome, count in spend_row_counts if count == 1)
lost: Final = tuple(outcome for outcome, count in spend_row_counts if count == 0)
assert len(landed) + len(lost) == len(completed), spend_row_counts
record_property("c02_spend_query", _SPEND_QUERY)
record_property("c02_burst_ids_landed_once", len(landed))
record_property("c02_burst_ids_lost", len(lost))
record_property("c02_burst_ids_lost_by_kind", _kind_counts(lost))
completed_markers: Final = tuple(outcome.call.marker for outcome in completed)
eventually(
observations.read,
lambda items: _marker_counts(items, completed_markers) == (1,) * len(completed_markers),
seconds=60,
)
failed_markers: Final = tuple(outcome.call.marker for outcome in failed)
record_property(
"c02_failed_calls_seen_upstream",
sum(count == 1 for count in _marker_counts(observations.items, failed_markers)),
)

View file

@ -0,0 +1,488 @@
import asyncio
import base64
import json
import struct
import threading
import zlib
from collections.abc import Awaitable, Callable, Iterator, Mapping
from queue import SimpleQueue
from typing import Final, TypeVar
from urllib.parse import urlsplit
import httpx
import openai
import pytest
from integration._support.vertex import service_account_json
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
import litellm
from litellm import Router
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.utils import ImageResponse, ModelResponse, TextCompletionResponse
_R: Final = TypeVar("_R")
_PROVIDER_KEY: Final = "sk-scripted-provider"
_API_VERSION: Final = "2024-10-21"
_GATEWAY_PATH: Final = "/gateway.ai.cloudflare.com/v1/scripted-account/scripted-gateway/azure-openai/scripted-resource"
_DEPLOYMENT: Final = "gpt-4o-mini-gateway"
_ANSWER: Final = "wire answer"
_IMAGE_URL: Final = "https://images.example.invalid/variation.png"
def _png_chunk(kind: bytes, data: bytes) -> bytes:
return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", zlib.crc32(kind + data))
_PNG: Final = (
b"\x89PNG\r\n\x1a\n"
+ _png_chunk(b"IHDR", struct.pack(">IIBBBBB", 1, 1, 8, 2, 0, 0, 0))
+ _png_chunk(b"IDAT", zlib.compress(b"\x00\x00\x00\x00"))
+ _png_chunk(b"IEND", b"")
)
def _json(body: Mapping[str, JsonValue]) -> Reply:
return Reply(body=json.dumps(body).encode())
def _chat_completion() -> Mapping[str, JsonValue]:
return {
"id": "chatcmpl-wire",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": _ANSWER}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
}
def _text_completion() -> Mapping[str, JsonValue]:
return {
"id": "cmpl-wire",
"object": "text_completion",
"created": 1,
"model": "gpt-3.5-turbo-instruct",
"choices": [{"index": 0, "text": _ANSWER, "finish_reason": "stop", "logprobs": None}],
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
}
def _peer(expected_path: str, body: Mapping[str, JsonValue]) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert urlsplit(request.target).path == expected_path, request.target
return _json(body)
return respond
def _held_peer(gate: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert gate.wait(timeout=10), "the timeout cell never released its peer"
return _json(_chat_completion())
return respond
def _only_request(wire: Wire) -> Request:
received: Final = wire.drain()
assert len(received) == 1, received
return received[0]
def _drain(seen: SimpleQueue[str]) -> tuple[str, ...]:
return tuple(seen.get_nowait() for _ in range(seen.qsize()))
def _content(response: ModelResponse) -> str | None:
choice: Final = response.choices[0]
return choice.message.content if isinstance(choice, litellm.Choices) else None
@pytest.fixture
def sync_session(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[httpx.Client, SimpleQueue[str]]]:
seen: Final[SimpleQueue[str]] = SimpleQueue()
def record(request: httpx.Request) -> None:
seen.put(str(request.url))
with httpx.Client(event_hooks={"request": [record]}) as client:
monkeypatch.setattr(litellm, "client_session", client)
monkeypatch.setattr(litellm, "aclient_session", None)
yield client, seen
def _run_with_async_session(
monkeypatch: pytest.MonkeyPatch, call: Callable[[], Awaitable[_R]]
) -> tuple[_R, tuple[str, ...]]:
seen: Final[SimpleQueue[str]] = SimpleQueue()
async def record(request: httpx.Request) -> None:
seen.put(str(request.url))
async def run() -> _R:
async with httpx.AsyncClient(event_hooks={"request": [record]}) as client:
monkeypatch.setattr(litellm, "client_session", None)
monkeypatch.setattr(litellm, "aclient_session", client)
return await call()
result: Final = asyncio.run(run())
return result, _drain(seen)
def _deployment(timeout: httpx.Timeout | openai.Timeout, api_base: str, deployment_id: str) -> Router:
return Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": api_base,
"api_key": _PROVIDER_KEY,
"timeout": timeout,
"max_retries": 0,
},
"model_info": {"id": deployment_id},
}
],
num_retries=0,
)
def test_f01_async_image_variation_with_only_a_sync_session_set_builds_its_own_async_client(
sync_session: tuple[httpx.Client, SimpleQueue[str]],
) -> None:
with wire_server(_peer("/images/variations", {"created": 1, "data": [{"url": _IMAGE_URL}]})) as wire:
response: Final = asyncio.run(
litellm.aimage_variation(
image=("probe.png", _PNG, "image/png"),
model="dall-e-2",
custom_llm_provider="openai",
api_base=wire.url,
api_key=_PROVIDER_KEY,
num_retries=0,
)
)
assert isinstance(response, ImageResponse), type(response)
assert response.data is not None and response.data[0].url == _IMAGE_URL, response
request: Final = _only_request(wire)
assert request.method == "POST", request.method
assert _PNG in request.body
assert request.headers.get("authorization") == f"Bearer {_PROVIDER_KEY}", request.headers
assert _drain(sync_session[1]) == ()
def test_f02_async_azure_cloudflare_gateway_call_with_only_a_sync_session_set_builds_its_own_async_client(
sync_session: tuple[httpx.Client, SimpleQueue[str]],
) -> None:
with wire_server(_peer(f"{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions", _chat_completion())) as wire:
response: Final = asyncio.run(
litellm.acompletion(
model=f"azure/{_DEPLOYMENT}",
messages=[{"role": "user", "content": "via the gateway"}],
api_base=f"{wire.url}{_GATEWAY_PATH}",
api_key=_PROVIDER_KEY,
api_version=_API_VERSION,
num_retries=0,
)
)
assert isinstance(response, ModelResponse), type(response)
assert _content(response) == _ANSWER, response
request: Final = _only_request(wire)
assert urlsplit(request.target).query == f"api-version={_API_VERSION}", request.target
assert request.headers.get("api-key") == _PROVIDER_KEY, request.headers
assert _drain(sync_session[1]) == ()
def test_f03_async_text_completion_uses_the_callers_async_session(monkeypatch: pytest.MonkeyPatch) -> None:
with wire_server(_peer("/completions", _text_completion())) as wire:
response, seen = _run_with_async_session(
monkeypatch,
lambda: litellm.atext_completion(
model="gpt-3.5-turbo-instruct",
prompt="complete this",
custom_llm_provider="openai",
api_base=wire.url,
api_key=_PROVIDER_KEY,
num_retries=0,
),
)
assert isinstance(response, TextCompletionResponse), type(response)
assert response.choices[0].text == _ANSWER, response
assert _only_request(wire).method == "POST"
assert seen == (f"{wire.url}/completions",)
def test_f04_sync_text_completion_uses_the_callers_sync_session(
sync_session: tuple[httpx.Client, SimpleQueue[str]],
) -> None:
with wire_server(_peer("/completions", _text_completion())) as wire:
response: Final = litellm.text_completion(
model="gpt-3.5-turbo-instruct",
prompt="complete this",
custom_llm_provider="openai",
api_base=wire.url,
api_key=_PROVIDER_KEY,
num_retries=0,
)
assert isinstance(response, TextCompletionResponse), type(response)
assert response.choices[0].text == _ANSWER, response
assert _only_request(wire).method == "POST"
assert _drain(sync_session[1]) == (f"{wire.url}/completions",)
def test_f05_sync_moderation_reaches_the_peer_and_returns_its_verdict() -> None:
verdict: Final[Mapping[str, JsonValue]] = {
"id": "modr-wire",
"model": "omni-moderation-latest",
"results": [{"flagged": True, "categories": {"harassment": True}, "category_scores": {"harassment": 0.91}}],
}
with wire_server(_peer("/moderations", verdict)) as wire:
response: Final = litellm.moderation(
input="moderate this",
model="omni-moderation-latest",
api_key=_PROVIDER_KEY,
api_base=wire.url,
)
assert response.results[0].flagged is True, response
request: Final = _only_request(wire)
assert json.loads(request.body) == {"input": "moderate this", "model": "omni-moderation-latest"}, request.body
@pytest.mark.parametrize("timeout_type", (httpx.Timeout, openai.Timeout), ids=("httpx", "openai"))
def test_f06_sync_completion_accepts_a_structured_timeout(
timeout_type: type[httpx.Timeout] | type[openai.Timeout],
) -> None:
with wire_server(_peer("/chat/completions", _chat_completion())) as wire:
response: Final = litellm.completion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "with a structured timeout"}],
custom_llm_provider="openai",
api_base=wire.url,
api_key=_PROVIDER_KEY,
timeout=timeout_type(5.0, connect=2.0),
num_retries=0,
)
assert isinstance(response, ModelResponse), type(response)
assert _content(response) == _ANSWER, response
assert _only_request(wire).method == "POST"
def test_f07_async_azure_cloudflare_gateway_call_uses_the_callers_async_session(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with wire_server(_peer(f"{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions", _chat_completion())) as wire:
response, seen = _run_with_async_session(
monkeypatch,
lambda: litellm.acompletion(
model=f"azure/{_DEPLOYMENT}",
messages=[{"role": "user", "content": "via the gateway"}],
api_base=f"{wire.url}{_GATEWAY_PATH}",
api_key=_PROVIDER_KEY,
api_version=_API_VERSION,
num_retries=0,
),
)
assert isinstance(response, ModelResponse), type(response)
assert _content(response) == _ANSWER, response
assert _only_request(wire).headers.get("api-key") == _PROVIDER_KEY
assert seen == (f"{wire.url}{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions?api-version={_API_VERSION}",)
def test_f08_sync_image_variation_uses_the_callers_sync_session(
sync_session: tuple[httpx.Client, SimpleQueue[str]],
) -> None:
with wire_server(_peer("/images/variations", {"created": 1, "data": [{"url": _IMAGE_URL}]})) as wire:
response: Final = litellm.image_variation(
image=("probe.png", _PNG, "image/png"),
model="dall-e-2",
custom_llm_provider="openai",
api_base=wire.url,
api_key=_PROVIDER_KEY,
num_retries=0,
)
assert isinstance(response, ImageResponse), type(response)
assert response.data is not None and response.data[0].url == _IMAGE_URL, response
assert _PNG in _only_request(wire).body
assert _drain(sync_session[1]) == (f"{wire.url}/images/variations",)
@pytest.mark.parametrize("timeout_type", (httpx.Timeout, openai.Timeout), ids=("httpx", "openai"))
def test_f09_router_deployment_keeps_a_structured_timeout_as_httpx_and_serves(
timeout_type: type[httpx.Timeout] | type[openai.Timeout],
) -> None:
with wire_server(_peer("/chat/completions", _chat_completion())) as wire:
router: Final = _deployment(timeout_type(5.0, connect=2.0), wire.url, "deployment-f09")
response: Final = router.completion(
model="gpt-4o-mini", messages=[{"role": "user", "content": "through the router"}]
)
assert isinstance(response, ModelResponse), type(response)
assert _content(response) == _ANSWER, response
assert _only_request(wire).method == "POST"
deployment: Final = router.get_deployment("deployment-f09")
assert deployment is not None
stored: Final = deployment.litellm_params.timeout
assert isinstance(stored, httpx.Timeout), type(stored)
assert (stored.read, stored.connect) == (5.0, 2.0), stored
def test_f10_router_deployment_structured_read_timeout_fires_once_as_a_408() -> None:
gate: Final = threading.Event()
with wire_server(_held_peer(gate)) as wire:
router: Final = _deployment(httpx.Timeout(0.5, connect=2.0), wire.url, "deployment-f10")
with pytest.raises(litellm.Timeout) as raised:
router.completion(model="gpt-4o-mini", messages=[{"role": "user", "content": "held upstream"}])
gate.set()
assert raised.value.status_code == 408, raised.value
assert len(wire.drain()) == 1
_SPEECH_AUDIO: Final = b"OggS" + bytes(range(60))
_SPEECH_AUDIO_REPLY: Final = Reply(body=_SPEECH_AUDIO, content_type="audio/ogg")
_ELEVENLABS_VOICE: Final = "21m00Tcm4TlvDq8ikWAM"
_VERTEX_PROJECT: Final = "scripted-project"
_SPEECH_ROUTES: Final[Mapping[str, tuple[Mapping[str, str], str, Reply]]] = {
"openai": (
{"model": "openai/gpt-4o-mini-tts", "voice": "alloy", "api_key": _PROVIDER_KEY},
"/audio/speech",
_SPEECH_AUDIO_REPLY,
),
"azure": (
{"model": "azure/tts-deployment", "voice": "alloy", "api_key": _PROVIDER_KEY, "api_version": _API_VERSION},
"/openai/deployments/tts-deployment/audio/speech",
_SPEECH_AUDIO_REPLY,
),
"azure_ava": (
{"model": "azure/speech/azure-tts", "voice": "alloy", "api_key": _PROVIDER_KEY},
"/cognitiveservices/v1",
_SPEECH_AUDIO_REPLY,
),
"elevenlabs": (
{"model": "elevenlabs/eleven_multilingual_v2", "voice": _ELEVENLABS_VOICE, "api_key": _PROVIDER_KEY},
f"/v1/text-to-speech/{_ELEVENLABS_VOICE}",
_SPEECH_AUDIO_REPLY,
),
"edenai": (
{"model": "edenai/openai", "voice": "alloy", "api_key": _PROVIDER_KEY},
"/audio/speech",
_SPEECH_AUDIO_REPLY,
),
"minimax": (
{"model": "minimax/speech-02-hd", "voice": "alloy", "api_key": _PROVIDER_KEY},
"/v1/t2a_v2",
_json({"data": {"audio": _SPEECH_AUDIO.hex()}, "base_resp": {"status_code": 0, "status_msg": "success"}}),
),
"mistral": (
{"model": "mistral/voxtral-mini-tts-2603", "voice": "alloy", "api_key": _PROVIDER_KEY},
"/v1/audio/speech",
_json({"audio_data": base64.b64encode(_SPEECH_AUDIO).decode()}),
),
"aws_polly": (
{
"model": "aws_polly/neural",
"voice": "Joanna",
"aws_access_key_id": "AKIASCRIPTED",
"aws_secret_access_key": "scripted-secret",
"aws_region_name": "us-east-1",
},
"/v1/speech",
_SPEECH_AUDIO_REPLY,
),
}
_PAYMENT_REQUIRED_MODELS: Final[Mapping[str, type[litellm.BadRequestError]]] = {
"openai/gpt-4o-mini": litellm.BadRequestError,
"anthropic/claude-sonnet-4-5": litellm.PaymentRequiredError,
}
def _speech_peer(expected_path: str, reply: Reply) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert urlsplit(request.target).path == expected_path, request.target
return reply
return respond
def _vertex_speech_peer(request: Request) -> Reply:
if urlsplit(request.target).path == "/_oauth/token":
return _json({"access_token": "scripted-vertex-token", "expires_in": 3600, "token_type": "Bearer"})
return _json({"audioContent": base64.b64encode(_SPEECH_AUDIO).decode()})
def _payment_required_peer(request: Request) -> Reply:
return Reply(
status=402,
body=json.dumps({"error": {"message": "scripted 402", "type": "invalid_request_error"}}).encode(),
)
@pytest.mark.parametrize("provider", tuple(_SPEECH_ROUTES))
def test_f11_sync_speech_accepts_an_sdk_timeout_on_every_provider_branch(provider: str) -> None:
route, expected_path, reply = _SPEECH_ROUTES[provider]
with wire_server(_speech_peer(expected_path, reply)) as wire:
response: Final = litellm.speech(
input="speak this",
api_base=wire.url,
timeout=openai.Timeout(5.0, connect=2.0),
**route,
)
assert response.content == _SPEECH_AUDIO, response.content[:16]
assert _only_request(wire).method == "POST"
def test_f13_sync_vertex_speech_accepts_an_sdk_timeout() -> None:
with wire_server(_vertex_speech_peer) as wire:
response: Final = litellm.speech(
model="vertex_ai/chirp",
voice="alloy",
input="speak this",
api_base=wire.url,
vertex_credentials=service_account_json(_VERTEX_PROJECT, wire.url),
vertex_project=_VERTEX_PROJECT,
vertex_location="us-central1",
timeout=openai.Timeout(5.0, connect=2.0),
)
assert response.content == _SPEECH_AUDIO, response.content[:16]
assert tuple((seen.method, urlsplit(seen.target).path) for seen in wire.drain()) == (
("POST", "/_oauth/token"),
("POST", "/"),
)
def _runwayml_rejecting_peer(request: Request) -> Reply:
assert urlsplit(request.target).path == "/v1/text_to_speech", request.target
return Reply(status=400, body=json.dumps({"error": "scripted 400"}).encode())
def test_f14_sync_runwayml_speech_sends_its_task_with_an_sdk_timeout() -> None:
with wire_server(_runwayml_rejecting_peer) as wire:
with pytest.raises(BaseLLMException) as raised:
litellm.speech(
model="runwayml/eleven_multilingual_v2",
voice="Maya",
input="speak this",
api_base=wire.url,
api_key=_PROVIDER_KEY,
timeout=openai.Timeout(5.0, connect=2.0),
)
assert raised.value.status_code == 400, raised.value
assert "scripted 400" in raised.value.message, raised.value.message
assert _only_request(wire).method == "POST"
@pytest.mark.parametrize("model", tuple(_PAYMENT_REQUIRED_MODELS))
def test_f12_sync_completion_maps_an_upstream_402_with_its_message(model: str) -> None:
with wire_server(_payment_required_peer) as wire:
with pytest.raises(_PAYMENT_REQUIRED_MODELS[model]) as raised:
litellm.completion(
model=model,
messages=[{"role": "user", "content": "out of credit"}],
api_base=wire.url,
api_key=_PROVIDER_KEY,
num_retries=0,
max_retries=0,
)
assert raised.value.status_code == 402, raised.value
assert "scripted 402" in raised.value.message, raised.value.message
assert len(wire.drain()) == 1

View file

@ -0,0 +1,66 @@
from typing import Final
import httpx
from openai import DefaultAsyncHttpxClient, DefaultHttpxClient
from openai import HttpxBinaryResponseContent as SDKBinaryResponse
from openai import Timeout as SDKTimeout
from openai._types import Response as SDKResponse
from typing_extensions import assert_type
import litellm
from litellm.exceptions import (
APIConnectionError,
APIResponseValidationError,
BadGatewayError,
InternalServerError,
InvalidRequestError,
RateLimitError,
Timeout,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.router import GenericLiteLLMParams, LiteLLMParamsTypedDict
def consume_binary_response(response: HttpxBinaryResponseContent) -> httpx.Response:
assert_type(response.response, httpx.Response)
assert_type(response.response.request, httpx.Request)
return response.response
def binary_response_types(legacy_response: httpx.Response, sdk_response: SDKResponse) -> None:
legacy: Final = HttpxBinaryResponseContent(legacy_response)
native: Final = HttpxBinaryResponseContent(sdk_response)
assert_type(legacy.response, httpx.Response)
assert_type(native.response, SDKResponse)
assert_type(legacy.read(), bytes)
assert_type(native.read(), bytes)
assert_type(consume_binary_response(legacy), httpx.Response)
sdk_wrapper: Final[SDKBinaryResponse] = legacy
assert_type(sdk_wrapper.read(), bytes)
def synthesized_response_types(
error: RateLimitError | BadGatewayError | InternalServerError | APIResponseValidationError | InvalidRequestError,
) -> None:
assert_type(error.response, httpx.Response)
assert_type(error.request, httpx.Request)
def synthesized_request_types(error: APIConnectionError | Timeout) -> None:
assert_type(error.request, httpx.Request)
def configured_clients(
sync_client: httpx.Client | DefaultHttpxClient,
async_client: httpx.AsyncClient | DefaultAsyncHttpxClient,
timeout: httpx.Timeout | SDKTimeout,
) -> None:
litellm.client_session = sync_client # test-quality-ok: [TQ005] Type-check-only fixture; never executed
litellm.aclient_session = async_client # test-quality-ok: [TQ005] Type-check-only fixture; never executed
assert_type(CompletionTimeout.normalize(timeout), httpx.Timeout)
params: Final[LiteLLMParamsTypedDict] = {"timeout": 1.0}
params["timeout"] = timeout
GenericLiteLLMParams(timeout=timeout)

View file

@ -0,0 +1,8 @@
{
"extends": "../../pyrightconfig.json",
"include": ["openai_sdk_compat.py"],
"exclude": [],
"extraPaths": ["../.."],
"pythonVersion": "3.10",
"reportMissingImports": true
}

View file

@ -1204,8 +1204,10 @@ def test_speech_response_without_a_byte_count_produces_no_output() -> None:
assert data.choices_out == ()
def test_speech_binary_response_is_logged_as_its_summary_not_dropped() -> None:
import httpx
@pytest.mark.parametrize("http_module", ("httpx", "httpx2"))
def test_speech_binary_response_is_logged_as_its_summary_not_dropped(http_module: str) -> None:
httpx: Final = pytest.importorskip(http_module)
from openai import HttpxBinaryResponseContent as SDKBinaryResponse
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import _extract_response_obj_and_hidden_params
@ -1216,13 +1218,17 @@ def test_speech_binary_response_is_logged_as_its_summary_not_dropped() -> None:
set_provider_response_headers_in_hidden_params(speech, raw.headers)
response_obj, hidden_params = _extract_response_obj_and_hidden_params(speech, None)
assert isinstance(speech, SDKBinaryResponse)
assert speech.response is raw
assert speech.read() == raw.content
assert response_obj == {"object": "binary", "content_type": "audio/mpeg", "num_bytes": 1234}
assert hidden_params is not None
assert hidden_params["headers"]["content-type"] == "audio/mpeg"
def test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far() -> None:
import httpx
@pytest.mark.parametrize("http_module", ("httpx", "httpx2"))
def test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far(http_module: str) -> None:
httpx: Final = pytest.importorskip(http_module)
from litellm.types.llms.openai import HttpxBinaryResponseContent

View file

@ -1,11 +1,16 @@
import asyncio, importlib, os
import asyncio
import importlib
import os
import json
from typing import Final
from unittest.mock import AsyncMock, Mock, patch
from collections.abc import Mapping
from types import ModuleType
from typing import Final, Protocol
from unittest.mock import Mock, patch
import httpx
import pytest
from openai import AsyncOpenAI, OpenAI
from openai._types import Response as SDKResponse
import litellm
from litellm.llms.openai.openai import (
@ -77,19 +82,48 @@ def test_get_stream_options_passes_caller_stream_options_through_on_any_host(api
}
@pytest.fixture(params=("httpx", "sdk"))
def http_backend(request: pytest.FixtureRequest) -> ModuleType:
return importlib.import_module("httpx" if request.param == "httpx" else SDKResponse.__module__)
@pytest.fixture(params=("openai", "perplexity", "cerebras", "nvidia_nim"))
def sdk_provider(request: pytest.FixtureRequest) -> str:
return str(request.param)
class _Request(Protocol):
@property
def content(self) -> bytes: ...
@property
def extensions(self) -> Mapping[str, object]: ...
@property
def url(self) -> object: ...
@property
def headers(self) -> Mapping[str, str]: ...
@pytest.mark.asyncio
async def test_acompletion_returns_json_reply_over_injected_transport():
async def test_acompletion_returns_json_reply_over_injected_transport(
http_backend: ModuleType, sdk_provider: str, monkeypatch: pytest.MonkeyPatch
):
outbound: Final = asyncio.Queue()
def respond(request: httpx.Request) -> httpx.Response:
def respond(request: _Request) -> httpx.Response | SDKResponse:
assert str(request.url) == f"https://{sdk_provider}.example/v1/chat/completions"
assert request.headers["authorization"] == f"Bearer {sdk_provider}-key"
assert request.extensions["timeout"] == {"connect": 7, "read": 7, "write": 7, "pool": 7}
outbound.put_nowait(json.loads(request.content))
return httpx.Response(
return http_backend.Response(
200,
json={
"id": "chatcmpl-smoke",
"object": "chat.completion",
"created": 0,
"model": "gpt-5.6",
"model": "sdk-compat",
"choices": [
{
"index": 0,
@ -101,31 +135,71 @@ async def test_acompletion_returns_json_reply_over_injected_transport():
},
)
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client)
async with http_backend.AsyncClient(transport=http_backend.MockTransport(respond)) as http_client:
monkeypatch.setattr(litellm, "aclient_session", http_client)
response: Final = await asyncio.wait_for(
litellm.acompletion(
model="openai/gpt-5.6",
api_key="transport-only",
client=client,
model=f"{sdk_provider}/sdk-compat",
api_key=f"{sdk_provider}-key",
api_base=f"https://{sdk_provider}.example/v1",
messages=[{"role": "user", "content": "smoke-json-request"}],
timeout=7,
num_retries=0,
max_retries=0,
),
timeout=10,
)
request: Final = await asyncio.wait_for(outbound.get(), timeout=10)
assert request["model"] == "gpt-5.6"
assert request["model"] == "sdk-compat"
assert request["messages"] == [{"role": "user", "content": "smoke-json-request"}]
assert not request.get("stream")
assert outbound.empty()
assert response.choices[0].message.content == "smoke-json-reply"
assert response.choices[0].finish_reason == "stop"
assert response.usage.total_tokens == 15
assert not http_client.is_closed
@pytest.mark.asyncio
async def test_acompletion_streams_text_deltas_over_injected_transport():
@pytest.mark.parametrize(
("status", "expected_error"),
(
(400, litellm.BadRequestError),
(429, litellm.RateLimitError),
(503, litellm.ServiceUnavailableError),
(None, litellm.Timeout),
),
)
async def test_acompletion_preserves_public_errors_for_both_http_clients(
http_backend: ModuleType,
sdk_provider: str,
status: int | None,
expected_error: type[Exception],
monkeypatch: pytest.MonkeyPatch,
) -> None:
def respond(request: _Request) -> httpx.Response | SDKResponse:
if status is None:
raise http_backend.ReadTimeout("upstream timeout", request=request)
return http_backend.Response(status, json={"error": {"message": "upstream failure", "type": "api_error"}})
async with http_backend.AsyncClient(transport=http_backend.MockTransport(respond)) as http_client:
monkeypatch.setattr(litellm, "aclient_session", http_client)
with pytest.raises(expected_error, match="timed out" if status is None else "upstream failure"):
await litellm.acompletion(
model=f"{sdk_provider}/sdk-compat",
api_key=f"{sdk_provider}-key",
api_base=f"https://{sdk_provider}.example/v1",
messages=[{"role": "user", "content": "request"}],
num_retries=0,
max_retries=0,
)
assert not http_client.is_closed
@pytest.mark.asyncio
async def test_acompletion_streams_text_deltas_over_injected_transport(
http_backend: ModuleType, sdk_provider: str, monkeypatch: pytest.MonkeyPatch
):
outbound: Final = asyncio.Queue()
def chunk(delta: dict, finish: str | None) -> bytes:
@ -133,18 +207,18 @@ async def test_acompletion_streams_text_deltas_over_injected_transport():
"id": "chatcmpl-smoke",
"object": "chat.completion.chunk",
"created": 0,
"model": "gpt-5.6",
"model": "sdk-compat",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish}],
}
return f"data: {json.dumps(body)}\n\n".encode()
def respond(request: httpx.Request) -> httpx.Response:
def respond(request: _Request) -> httpx.Response | SDKResponse:
outbound.put_nowait(json.loads(request.content))
usage: Final = {
"id": "chatcmpl-smoke",
"object": "chat.completion.chunk",
"created": 0,
"model": "gpt-5.6",
"model": "sdk-compat",
"choices": [],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
@ -156,14 +230,14 @@ async def test_acompletion_streams_text_deltas_over_injected_transport():
b"data: [DONE]\n\n",
)
)
return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=content)
return http_backend.Response(200, headers={"content-type": "text/event-stream"}, content=content)
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client)
async with http_backend.AsyncClient(transport=http_backend.MockTransport(respond)) as http_client:
monkeypatch.setattr(litellm, "aclient_session", http_client)
stream: Final = await litellm.acompletion(
model="openai/gpt-5.6",
api_key="transport-only",
client=client,
model=f"{sdk_provider}/sdk-compat",
api_key=f"{sdk_provider}-key",
api_base=f"https://{sdk_provider}.example/v1",
messages=[{"role": "user", "content": "smoke-stream-request"}],
stream=True,
num_retries=0,
@ -190,7 +264,7 @@ async def test_acompletion_streams_text_deltas_over_injected_transport():
@pytest.mark.asyncio
async def test_acompletion_streams_tool_call_arguments_over_injected_transport():
async def test_acompletion_streams_tool_call_arguments_over_injected_transport(http_backend: ModuleType):
outbound: Final = asyncio.Queue()
tools: Final = [
{
@ -217,7 +291,8 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport()
}
return f"data: {json.dumps(body)}\n\n".encode()
def respond(request: httpx.Request) -> httpx.Response:
def respond(request: _Request) -> httpx.Response | SDKResponse:
assert request.extensions["timeout"] == {"connect": 1, "read": None, "write": 2, "pool": 3}
outbound.put_nowait(json.loads(request.content))
content: Final = b"".join(
(
@ -239,9 +314,9 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport()
b"data: [DONE]\n\n",
)
)
return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=content)
return http_backend.Response(200, headers={"content-type": "text/event-stream"}, content=content)
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
async with http_backend.AsyncClient(transport=http_backend.MockTransport(respond)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client)
messages: Final = [{"role": "user", "content": "weather in Paris"}]
stream: Final = await litellm.acompletion(
@ -251,6 +326,7 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport()
messages=messages,
tools=tools,
stream=True,
timeout=http_backend.Timeout(connect=1, read=None, write=2, pool=3),
num_retries=0,
max_retries=0,
)

View file

@ -1,9 +1,10 @@
from unittest.mock import MagicMock, call, patch
from typing import Final
from unittest.mock import MagicMock, patch
import httpx
import openai
import pytest
import respx
import litellm
from litellm.litellm_core_utils.token_counter import token_counter
@ -413,3 +414,124 @@ def test_is_openai_backed_api_base_decides_by_hostname_only(api_base, expected):
assert is_openai_backed_api_base(api_base) is expected
def _sdk_api_client(
api: str,
is_async: bool,
timeout: float | httpx.Timeout | openai.Timeout | None,
client: openai.OpenAI | openai.AsyncOpenAI | None = None,
) -> openai.OpenAI | openai.AsyncOpenAI | None:
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.fine_tuning.handler import OpenAIFineTuningAPI
from litellm.llms.openai.image_variations.handler import OpenAIImageVariationsHandler
from litellm.llms.openai.openai import OpenAIAssistantsAPI, OpenAIBatchesAPI, OpenAIFilesAPI
kwargs: Final = {
"api_key": "transport-only",
"api_base": "https://sdk-default.example/v1",
"timeout": timeout,
"max_retries": 0,
"organization": None,
"client": client,
}
if api == "assistants":
assistant_factory: Final = OpenAIAssistantsAPI()
return (
assistant_factory.async_get_openai_client(**kwargs)
if is_async
else assistant_factory.get_openai_client(**kwargs)
)
if api == "image_variations":
variation_factory: Final = OpenAIImageVariationsHandler()
params: Final = {
"api_key": "transport-only",
"base_url": kwargs["api_base"],
"timeout": timeout,
}
return (
variation_factory.get_async_client(client=client, init_client_params=params)
if is_async
else variation_factory.get_sync_client(client=client, init_client_params=params)
)
if api == "azure_gateway":
return BaseAzureLLM()._init_azure_client_for_cloudflare_ai_gateway(
api_base="https://sdk-default.example",
model="deployment",
api_version="2024-02-01",
max_retries=0,
timeout=timeout,
litellm_params={},
api_key="transport-only",
azure_ad_token=None,
azure_ad_token_provider=None,
acompletion=is_async,
client=client,
)
factory: Final = {"files": OpenAIFilesAPI, "batches": OpenAIBatchesAPI, "fine_tuning": OpenAIFineTuningAPI}[api]()
return factory.get_openai_client(**kwargs, _is_async=is_async)
async def _list_model_ids(sdk_client: openai.OpenAI | openai.AsyncOpenAI) -> list[str]:
page: Final = (
await sdk_client.models.list() if isinstance(sdk_client, openai.AsyncOpenAI) else sdk_client.models.list()
)
return [model.id for model in page.data]
@pytest.mark.parametrize("api", ["files", "batches", "assistants", "fine_tuning", "image_variations", "azure_gateway"])
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.asyncio
async def test_sdk_api_factories_keep_httpx_transport_and_request_timeouts(api: str, is_async: bool) -> None:
timeout: Final = openai.Timeout(connect=1, read=7, write=2, pool=3)
with respx.mock(assert_all_called=True) as mock:
mock.get(host="sdk-default.example", path__regex=r".*/models$").mock(
return_value=httpx.Response(307, headers={"location": "https://sdk-default.example/v1/moved-models"})
)
moved: Final = mock.get(host="sdk-default.example", path="/v1/moved-models").mock(
return_value=httpx.Response(
200,
json={"object": "list", "data": [{"id": "sdk-compat", "object": "model", "created": 0, "owned_by": "t"}]},
)
)
sdk_client: Final = _sdk_api_client(api, is_async, timeout)
assert sdk_client is not None
try:
assert _sdk_api_client(api, is_async, 19, client=sdk_client) is sdk_client
model_ids: Final = await _list_model_ids(sdk_client)
finally:
if is_async:
await sdk_client.close()
else:
sdk_client.close()
assert model_ids == ["sdk-compat"]
assert moved.calls.last.request.extensions["timeout"] == {"connect": 1.0, "read": 7.0, "write": 2.0, "pool": 3.0}
@pytest.mark.parametrize("backend", ["httpx", "sdk_default"])
@pytest.mark.asyncio
async def test_azure_gateway_and_image_variations_use_the_callers_async_session(
backend: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from importlib import import_module
from io import BytesIO
from openai._types import Response as SDKResponse
from litellm.images.main import aimage_variation
http_module: Final = httpx if backend == "httpx" else import_module(SDKResponse.__module__)
transport: Final = http_module.MockTransport(
lambda request: http_module.Response(
200, json={"created": 0, "data": [{"url": "https://example.com/image.png"}]}
)
)
async with http_module.AsyncClient(transport=transport) as session:
monkeypatch.setattr(litellm, "aclient_session", session)
monkeypatch.setattr(litellm, "client_session", None)
azure_client: Final = _sdk_api_client("azure_gateway", True, 7)
assert azure_client is not None
assert azure_client._client is session
response: Final = await aimage_variation(
image=BytesIO(b"image-bytes"), api_key="transport-only", api_base="https://sdk-default.example/v1"
)
assert response.data[0].url == "https://example.com/image.png"
assert not session.is_closed

View file

@ -1,5 +1,6 @@
import json
import sys
from functools import partial
from pathlib import Path
from typing import Final
@ -36,6 +37,24 @@ CHAT_COMPLETION_BODY: Final = {
}
@pytest.fixture(autouse=True)
def mock_sdk_token_exchange_transport(monkeypatch: pytest.MonkeyPatch) -> None:
native_client_factory: Final = getattr(sys.modules.get("openai.auth._workload"), "DefaultHttpx2Client", None)
if native_client_factory is not None:
import httpx2
def handle_exchange(request: httpx2.Request) -> httpx2.Response:
response: Final = respx.mock.handler(
httpx.Request(request.method, str(request.url), headers=dict(request.headers), content=request.content)
)
return httpx2.Response(response.status_code, headers=dict(response.headers), content=response.content)
monkeypatch.setattr(
"openai.auth._workload.DefaultHttpx2Client",
partial(native_client_factory, transport=httpx2.MockTransport(handle_exchange), trust_env=False),
)
@pytest.fixture
def wif_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> OpenAIWorkloadIdentityConfig:
token_file: Final = tmp_path / "subject_token.jwt"
@ -602,7 +621,6 @@ class TestDiscoverModels:
assert models_route.calls.last.request.headers["Authorization"] == "Bearer None"
assert not exchange_route.called
@respx.mock
def test_empty_static_key_never_borrows_the_env_key(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "sk-env-key-that-must-stay-home")

View file

@ -2,8 +2,11 @@
import os
import sys
from typing import Final, Literal
import httpx
import pytest
from openai import Timeout as SDKTimeout
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../..")))
@ -142,3 +145,35 @@ def test_httpx_timeout_preserved_for_openai():
)
assert out is t
assert isinstance(out, httpx.Timeout)
@pytest.mark.parametrize("source", ("model", "timeout", "request_timeout"))
def test_sdk_timeout_preserves_components_for_httpx_providers(
source: Literal["model", "timeout", "request_timeout"],
) -> None:
timeout: Final = SDKTimeout(connect=2.0, read=None, write=5.0, pool=7.0)
resolved: Final = CompletionTimeout.resolve(
timeout if source == "model" else None,
{} if source == "model" else {source: timeout},
"openai",
global_timeout=None,
supports_httpx_timeout=supports_httpx_timeout,
)
assert isinstance(resolved, httpx.Timeout)
assert resolved.as_dict() == timeout.as_dict()
@pytest.mark.parametrize("read_timeout, expected", ((23.0, 23.0), (None, 600.0)))
def test_sdk_timeout_coerces_read_timeout_for_providers_without_httpx_support(
read_timeout: float | None, expected: float
) -> None:
assert (
CompletionTimeout.resolve(
SDKTimeout(connect=2.0, read=read_timeout, write=5.0, pool=7.0),
{},
"azure_ai",
global_timeout=None,
supports_httpx_timeout=supports_httpx_timeout,
)
== expected
)

View file

@ -8,6 +8,8 @@ This is important for debugging and observability - headers like x-request-id,
x-ms-region, rate limit headers, etc. should be available even when errors occur.
"""
from typing import Final
import httpx
import pytest
@ -17,6 +19,7 @@ from litellm.exceptions import (
ContextWindowExceededError,
ImageFetchError,
MidStreamFallbackError,
PaymentRequiredError,
RateLimitError,
ServiceUnavailableError,
)
@ -25,10 +28,11 @@ from litellm.exceptions import (
class TestExceptionHeaderPreservation:
"""Test that exception classes preserve headers from provider responses."""
@pytest.fixture
def mock_response_with_headers(self) -> httpx.Response:
@pytest.fixture(params=("httpx", "httpx2"))
def mock_response_with_headers(self, request: pytest.FixtureRequest) -> httpx.Response:
"""Create a mock response with typical provider headers."""
return httpx.Response(
transport: Final = pytest.importorskip(request.param)
return transport.Response(
status_code=400,
headers={
"x-request-id": "req-abc123",
@ -36,7 +40,7 @@ class TestExceptionHeaderPreservation:
"x-ratelimit-remaining-requests": "99",
"x-ratelimit-remaining-tokens": "9999",
},
request=httpx.Request("POST", "https://api.openai.com/v1/chat/completions"),
request=transport.Request("POST", "https://api.openai.com/v1/chat/completions"),
)
def test_bad_request_error_preserves_headers(
@ -50,11 +54,30 @@ class TestExceptionHeaderPreservation:
response=mock_response_with_headers,
)
assert error.response is not None
assert error.response is mock_response_with_headers
assert error.request is mock_response_with_headers.request
assert error.request_id == mock_response_with_headers.headers["x-request-id"]
assert error.response.headers.get("x-request-id") == "req-abc123"
assert error.response.headers.get("x-ms-region") == "eastus"
assert error.response.headers.get("x-ratelimit-remaining-requests") == "99"
def test_payment_required_error_preserves_headers(
self, mock_response_with_headers: httpx.Response
):
"""PaymentRequiredError should keep the provider response like its BadRequestError parent."""
error = PaymentRequiredError(
message="Insufficient credit",
model="gpt-4",
llm_provider="openrouter",
response=mock_response_with_headers,
)
assert error.status_code == 402
assert error.response is mock_response_with_headers
assert error.request is mock_response_with_headers.request
assert error.request_id == mock_response_with_headers.headers["x-request-id"]
assert error.response.headers.get("x-ms-region") == "eastus"
def test_content_policy_violation_error_preserves_headers(
self, mock_response_with_headers: httpx.Response
):

View file

@ -7,9 +7,19 @@ import respx
import litellm
_SDK_EMBEDDING_MODELS: Final = (
"openai/sdk-compat",
"mistral/sdk-compat",
"fireworks_ai/accounts/fireworks/models/sdk-compat",
"together_ai/sdk-compat",
"nvidia_nim/sdk-compat",
)
def _mock_openai_embedding_route(respx_mock: respx.MockRouter) -> respx.Route:
return respx_mock.post("https://api.openai.com/v1/embeddings").mock(
def _mock_openai_embedding_route(
respx_mock: respx.MockRouter, api_base: str = "https://api.openai.com/v1"
) -> respx.Route:
return respx_mock.post(f"{api_base}/embeddings").mock(
return_value=httpx.Response(
200,
json={
@ -27,13 +37,21 @@ def clear_default_encoding_format_env(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
def test_embedding_openai_omits_encoding_format_when_client_omits_it(respx_mock: respx.MockRouter) -> None:
mock_route: Final = _mock_openai_embedding_route(respx_mock)
@pytest.mark.parametrize("model", _SDK_EMBEDDING_MODELS)
def test_embedding_sdk_providers_omit_encoding_format_when_client_omits_it(
respx_mock: respx.MockRouter, model: str
) -> None:
provider, upstream_model = model.split("/", 1)
api_base: Final = f"https://{provider.replace('_', '-')}.example/v1"
mock_route: Final = _mock_openai_embedding_route(respx_mock, api_base)
response: Final = litellm.embedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
response: Final = litellm.embedding(model=model, input=["hello"], api_key="sk-test", api_base=api_base)
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert "encoding_format" not in request_body
assert request_body["model"] == upstream_model
assert request_body["input"] == ["hello"]
assert mock_route.calls.last.request.headers["authorization"] == "Bearer sk-test"
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
@ -89,16 +107,22 @@ def test_embedding_openai_env_none_omits_encoding_format(
@pytest.mark.asyncio
async def test_aembedding_openai_omits_encoding_format_when_client_omits_it(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
@pytest.mark.parametrize("model", _SDK_EMBEDDING_MODELS)
async def test_aembedding_sdk_providers_omit_encoding_format_when_client_omits_it(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, model: str
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
mock_route: Final = _mock_openai_embedding_route(respx_mock)
provider, upstream_model = model.split("/", 1)
api_base: Final = f"https://{provider.replace('_', '-')}.example/v1"
mock_route: Final = _mock_openai_embedding_route(respx_mock, api_base)
response: Final = await litellm.aembedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
response: Final = await litellm.aembedding(model=model, input=["hello"], api_key="sk-test", api_base=api_base)
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert "encoding_format" not in request_body
assert request_body["model"] == upstream_model
assert request_body["input"] == ["hello"]
assert mock_route.calls.last.request.headers["authorization"] == "Bearer sk-test"
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]

View file

@ -1,6 +1,9 @@
import logging
from typing import Final
import httpx
import pytest
from openai import Timeout as SDKTimeout
from pydantic import ValidationError
from litellm.types.router import (
@ -26,6 +29,13 @@ from litellm.types.utils import (
)
def test_sdk_timeout_is_normalized_for_provider_clients() -> None:
timeout: Final = SDKTimeout(connect=2.0, read=None, write=5.0, pool=7.0)
params: Final = GenericLiteLLMParams(timeout=timeout)
assert isinstance(params.timeout, httpx.Timeout)
assert params.timeout.as_dict() == timeout.as_dict()
def test_model_info_declares_mirrored_pricing_fields():
"""The pricing keys Deployment mirrors onto model_info must be declared fields, not
extras that only survive because ModelInfo sets extra="allow"."""

2
uv.lock generated
View file

@ -4814,7 +4814,7 @@ requires-dist = [
{ name = "numpy", marker = "extra == 'stt-nvidia-riva'", specifier = ">=1.26.0" },
{ name = "numpydoc", marker = "extra == 'utils'", specifier = ">=1.8.0,<2.0" },
{ name = "nvidia-riva-client", marker = "extra == 'stt-nvidia-riva'", specifier = ">=2.15.0" },
{ name = "openai", specifier = ">=2.20.0,<3.0.0" },
{ name = "openai", specifier = ">=2.20.0,<4.0.0" },
{ name = "opentelemetry-api", marker = "extra == 'proxy-runtime'", specifier = "==1.33.1" },
{ name = "opentelemetry-exporter-otlp", marker = "extra == 'proxy-runtime'", specifier = "==1.33.1" },
{ name = "opentelemetry-instrumentation-fastapi", marker = "extra == 'proxy-runtime'", specifier = "==0.54b1" },