From 63d823d28aff0fef8b0630e59f52c4d3adcdf7d2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 9 Oct 2026 19:45:50 -0700 Subject: [PATCH] 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 Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .github/workflows/test-unit.yml | 71 +- litellm/__init__.py | 5 +- litellm/exceptions.py | 62 +- .../litellm_core_utils/completion_timeout.py | 25 +- litellm/llms/azure/common_utils.py | 8 +- litellm/llms/openai/common_utils.py | 49 +- litellm/llms/openai/completion/handler.py | 8 +- litellm/llms/openai/fine_tuning/handler.py | 5 +- .../llms/openai/image_variations/handler.py | 24 +- litellm/llms/openai/openai.py | 14 +- litellm/main.py | 42 +- litellm/proxy/proxy_server.py | 2 +- litellm/types/llms/openai.py | 26 +- litellm/types/router.py | 13 +- packaging/litellm-core/pyproject.toml | 2 +- pyproject.toml | 2 +- .../base_sdk_tests/check_base_sdk_install.py | 41 +- tests/base_sdk_tests/check_sdk_http.py | 7 +- .../test_openai_sdk_client_paths.py | 1200 +++++++++++++++++ .../test_openai_sdk_client_paths_chaos.py | 424 ++++++ .../test_openai_sdk_client_sessions_wire.py | 488 +++++++ tests/typing/openai_sdk_compat.py | 66 + tests/typing/pyrightconfig.json | 8 + .../otel/test_otel_v2_sources_of_truth.py | 14 +- tests/unit/llms/openai/test_openai.py | 130 +- .../llms/openai/test_openai_common_utils.py | 126 +- .../openai/test_openai_workload_identity.py | 20 +- .../test_completion_timeout_resolution.py | 35 + .../test_exception_header_preservation.py | 33 +- ...penai_embedding_encoding_format_default.py | 42 +- tests/unit/types/test_router.py | 10 + uv.lock | 2 +- 32 files changed, 2851 insertions(+), 153 deletions(-) create mode 100644 tests/integration/compatibility/test_openai_sdk_client_paths.py create mode 100644 tests/integration/compatibility/test_openai_sdk_client_paths_chaos.py create mode 100644 tests/integration/sdk/test_openai_sdk_client_sessions_wire.py create mode 100644 tests/typing/openai_sdk_compat.py create mode 100644 tests/typing/pyrightconfig.json diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 1cf79049a66..2b9849793e9 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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 diff --git a/litellm/__init__.py b/litellm/__init__.py index ac57a9d0fee..4196912488f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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", diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 439eadce13d..ac461a2472e 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -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 diff --git a/litellm/litellm_core_utils/completion_timeout.py b/litellm/litellm_core_utils/completion_timeout.py index 163a4a6b9d6..e6f18257138 100644 --- a/litellm/litellm_core_utils/completion_timeout.py +++ b/litellm/litellm_core_utils/completion_timeout.py @@ -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) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index c640405c269..3564d7a1364 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -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, } diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index ffc5a5a71e0..839ca04e082 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -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 diff --git a/litellm/llms/openai/completion/handler.py b/litellm/llms/openai/completion/handler.py index 1bc752b6714..174aefd29e5 100644 --- a/litellm/llms/openai/completion/handler.py +++ b/litellm/llms/openai/completion/handler.py @@ -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, diff --git a/litellm/llms/openai/fine_tuning/handler.py b/litellm/llms/openai/fine_tuning/handler.py index c54263f4c35..851113bbd37 100644 --- a/litellm/llms/openai/fine_tuning/handler.py +++ b/litellm/llms/openai/fine_tuning/handler.py @@ -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 diff --git a/litellm/llms/openai/image_variations/handler.py b/litellm/llms/openai/image_variations/handler.py index bc02d274f24..664d16c48bf 100644 --- a/litellm/llms/openai/image_variations/handler.py +++ b/litellm/llms/openai/image_variations/handler.py @@ -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, diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index a782bf35ac8..3a12326d433 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index b80819354f0..4744a502194 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 53f49f78c9f..a7198fd1a95 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index cd77d6088f6..26f65fee4fe 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -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: diff --git a/litellm/types/router.py b/litellm/types/router.py index 53df3ce6773..cbe332dab12 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/packaging/litellm-core/pyproject.toml b/packaging/litellm-core/pyproject.toml index 9653e6f983a..cf87cc6b936 100644 --- a/packaging/litellm-core/pyproject.toml +++ b/packaging/litellm-core/pyproject.toml @@ -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", diff --git a/pyproject.toml b/pyproject.toml index 924e3a939b2..133fc735df1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/tests/base_sdk_tests/check_base_sdk_install.py b/tests/base_sdk_tests/check_base_sdk_install.py index e56d48f5e6e..3abea2a9707 100644 --- a/tests/base_sdk_tests/check_base_sdk_install.py +++ b/tests/base_sdk_tests/check_base_sdk_install.py @@ -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: diff --git a/tests/base_sdk_tests/check_sdk_http.py b/tests/base_sdk_tests/check_sdk_http.py index d773e252127..c7fb6a542ff 100644 --- a/tests/base_sdk_tests/check_sdk_http.py +++ b/tests/base_sdk_tests/check_sdk_http.py @@ -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" diff --git a/tests/integration/compatibility/test_openai_sdk_client_paths.py b/tests/integration/compatibility/test_openai_sdk_client_paths.py new file mode 100644 index 00000000000..f8d7449038a --- /dev/null +++ b/tests/integration/compatibility/test_openai_sdk_client_paths.py @@ -0,0 +1,1200 @@ +from __future__ import annotations + +import asyncio +import base64 +import binascii +import json +import os +import threading +import uuid +from collections.abc import Awaitable, Callable, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, TypeVar + +import anthropic +import httpx +import openai +import pytest +import yaml +from anthropic.types import RawContentBlockDeltaEvent, RawMessageStartEvent, RawMessageStreamEvent, TextDelta +from integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.openai_wire import chat_reply, openai_error, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy +from integration._support.responses_vendor import response_identities +from integration._support.wire import Reply, Request, Wire, wire_server +from openai.types import Completion, ModerationCreateResponse +from openai.types.chat import ChatCompletionChunk +from openai.types.responses import ResponseStreamEvent +from pydantic import JsonValue + +from litellm.constants import PROXY_CONFIG_RELOAD_INTERVAL_SECONDS + +PROXY_WORKERS: Final = int(os.environ.get("INTEGRATION_PROXY_WORKERS", "1")) +WORKER_SYNC_SECONDS: Final = 0.0 if PROXY_WORKERS == 1 else PROXY_CONFIG_RELOAD_INTERVAL_SECONDS + 5.0 +ROOT: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) +TEXT: Final = "openai sdk client path" +PROVIDER_KEY: Final = "integration-provider-key" +PROVIDER_AUTHORIZATION: Final = f"Bearer {PROVIDER_KEY}" +UPSTREAM_MODEL: Final = "gpt-4o-mini" +DEPLOYMENT_MODEL: Final = f"openai/{UPSTREAM_MODEL}" +INSTRUCT_MODEL: Final = "gpt-3.5-turbo-instruct" +INSTRUCT_DEPLOYMENT_MODEL: Final = f"openai/{INSTRUCT_MODEL}" +NO_CACHE: Final[Mapping[str, JsonValue]] = MappingProxyType({"cache": {"no-cache": True}}) +USAGE: Final[Mapping[str, JsonValue]] = MappingProxyType( + {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8} +) +FIVE_KB: Final = "x" * 5120 +REJECTED_CHAT_TIMEOUTS: Final[Mapping[str, JsonValue]] = MappingProxyType( + {"list": [1, 2], "dict": {"read": 5}, "5 KB string": FIVE_KB} +) +ACCEPTED_CHAT_TIMEOUTS: Final[Mapping[str, JsonValue]] = MappingProxyType({"empty string": "", "int": 30}) +NO_EXTRA: Final[Mapping[str, JsonValue]] = MappingProxyType({}) +REJECTED_DEPLOYMENT_TIMEOUTS: Final[Mapping[str, JsonValue]] = MappingProxyType({"list": [1, 2], "dict": {"read": 5}}) +ACCEPTED_DEPLOYMENT_TIMEOUTS: Final[Mapping[str, JsonValue]] = MappingProxyType( + {"empty string": "", "5 KB string": FIVE_KB} +) +MODERATION_CATEGORIES: Final = ("hate", "harassment", "self-harm", "sexual", "violence") +LISTED_JOB: Final = "ftjob-listed" +ASSISTANT: Final = "asst_integration" +CLOUDFLARE_GATEWAY: Final = "/gateway.ai.cloudflare.com/v1/account/gateway/azure-openai/resource" +AZURE_API_VERSION: Final = "2024-10-21" +STREAM_HEAD: Final = b"h" * 65536 +STREAM_TAIL: Final = b'{"custom_id": "tail", "response": {"status_code": 200}}\n' +SPEECH_AUDIO: Final = b"OggS" + bytes(range(60)) +SPEECH_UPSTREAM_MODEL: Final = "gpt-4o-mini-tts" +PAYMENT_REQUIRED_DEPLOYMENTS: Final[Mapping[str, str]] = MappingProxyType( + {"openai": "openai/gpt-4o-mini", "anthropic": "anthropic/claude-sonnet-4-5"} +) +_R: Final = TypeVar("_R") + + +def _identity(prefix: str) -> str: + return f"{prefix}-{uuid.uuid4().hex}" + + +def _path(request: Request) -> str: + return request.target.split("?", 1)[0] + + +def _body(request: Request) -> Mapping[str, JsonValue]: + return JSON_OBJECT.validate_json(request.body) + + +def _streams(request: Request) -> bool: + return _body(request).get("stream") is True + + +def _upstream_model(request: Request) -> str: + return string_value(_body(request).get("model") or UPSTREAM_MODEL) + + +def _json( + payload: Mapping[str, JsonValue], *, status: int = 200, headers: Mapping[str, str] = MappingProxyType({}) +) -> Reply: + return Reply(status=status, body=json.dumps(payload).encode(), headers=headers) + + +def _peer(route: Callable[[Request], Reply | None]) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.method == "GET" and _path(request).endswith("/models"): + return _json({"object": "list", "data": []}) + reply: Final = route(request) + if reply is None: + return _json({"error": {"message": f"unrouted {request.method} {request.target}"}}, status=404) + return reply + + return respond + + +def _is_model_discovery(request: Request) -> bool: + return request.method == "GET" and _path(request).endswith("/models") + + +def _served(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if not _is_model_discovery(request)) + + +def _routes(served: tuple[Request, ...]) -> tuple[tuple[str, str], ...]: + return tuple((request.method, _path(request)) for request in served) + + +def _is_chat(request: Request) -> bool: + return request.method == "POST" and _path(request).endswith("/chat/completions") + + +def _chat_peer(identity: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if not _is_chat(request): + return None + return chat_reply(identity, _upstream_model(request), TEXT, stream=_streams(request)) + + return _peer(route) + + +def _fresh_chat_peer(served: SimpleQueue[str]) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if not _is_chat(request): + return None + identity: Final = _identity("chatcmpl") + served.put(identity) + return chat_reply(identity, _upstream_model(request), TEXT, stream=False) + + return _peer(route) + + +def _held_chat_peer(gate: threading.Event, identity: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if not _is_chat(request): + return None + assert gate.wait(timeout=10), "the timeout cell never released its peer" + return chat_reply(identity, _upstream_model(request), TEXT, stream=False) + + return _peer(route) + + +def _status_peer(reply: Reply) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + return reply if request.method == "POST" else None + + return _peer(route) + + +def _completion_frame(identity: str, model: str, text: str, finish_reason: str | None) -> bytes: + frame: Final = { + "id": identity, + "object": "text_completion", + "created": 1, + "model": model, + "choices": [{"text": text, "index": 0, "finish_reason": finish_reason, "logprobs": None}], + } + return b"data: " + json.dumps(frame).encode() + b"\n\n" + + +def _completion_reply(identity: str, model: str, *, stream: bool) -> Reply: + if stream: + return Reply( + content_type="text/event-stream", + chunks=( + _completion_frame(identity, model, TEXT, None), + _completion_frame(identity, model, "", "stop") + b"data: [DONE]\n\n", + ), + ) + return _json( + { + "id": identity, + "object": "text_completion", + "created": 1, + "model": model, + "choices": [{"text": TEXT, "index": 0, "finish_reason": "stop", "logprobs": None}], + "usage": dict(USAGE), + } + ) + + +def _completion_peer(identity: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if request.method != "POST" or _path(request) != "/v1/completions": + return None + return _completion_reply(identity, _upstream_model(request), stream=_streams(request)) + + return _peer(route) + + +def _responses_peer(identity: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if request.method != "POST" or _path(request) != "/v1/responses": + return None + return responses_reply(identity, _upstream_model(request), TEXT, stream=_streams(request)) + + return _peer(route) + + +def _moderation(identity: str, model: str) -> Mapping[str, JsonValue]: + return { + "id": identity, + "model": model, + "results": [ + { + "flagged": False, + "categories": {category: False for category in MODERATION_CATEGORIES}, + "category_scores": {category: 0.0 for category in MODERATION_CATEGORIES}, + "category_applied_input_types": {category: ["text"] for category in MODERATION_CATEGORIES}, + } + ], + } + + +def _moderation_peer(identity: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if request.method != "POST" or _path(request) != "/v1/moderations": + return None + return _json(_moderation(identity, _upstream_model(request))) + + return _peer(route) + + +def _batch_line(custom_id: str) -> bytes: + line: Final = { + "custom_id": custom_id, + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": UPSTREAM_MODEL, "messages": [_user_message()]}, + } + return json.dumps(line).encode() + b"\n" + + +def _file_object(file_id: str, size: int) -> Mapping[str, JsonValue]: + return { + "id": file_id, + "object": "file", + "purpose": "batch", + "bytes": size, + "created_at": 1, + "filename": "batch.jsonl", + "status": "processed", + } + + +def _files_peer(file_id: str, content: bytes, chunks: tuple[bytes, ...] | None = None) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + path: Final = _path(request) + if request.method == "POST" and path == "/v1/files": + return _json(_file_object(file_id, len(content))) + if request.method == "GET" and path == f"/v1/files/{file_id}": + return _json(_file_object(file_id, len(content))) + if request.method == "GET" and path == f"/v1/files/{file_id}/content": + return Reply(content_type="application/octet-stream", body=content, chunks=chunks) + return None + + return _peer(route) + + +def _batch_object(batch_id: str, input_file_id: str, status: str) -> Mapping[str, JsonValue]: + return { + "id": batch_id, + "object": "batch", + "endpoint": "/v1/chat/completions", + "errors": None, + "input_file_id": input_file_id, + "completion_window": "24h", + "status": status, + "output_file_id": None, + "error_file_id": None, + "created_at": 1, + "in_progress_at": None, + "expires_at": None, + "finalizing_at": None, + "completed_at": None, + "failed_at": None, + "expired_at": None, + "cancelling_at": None, + "cancelled_at": None, + "request_counts": {"total": 1, "completed": 0, "failed": 0}, + "metadata": None, + } + + +def _batches_peer(file_id: str, batch_id: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + path: Final = _path(request) + if request.method == "POST" and path == "/v1/files": + return _json(_file_object(file_id, 1)) + if request.method == "POST" and path == "/v1/batches": + return _json(_batch_object(batch_id, string_value(_body(request)["input_file_id"]), "validating")) + if request.method == "GET" and path == f"/v1/batches/{batch_id}": + return _json(_batch_object(batch_id, file_id, "in_progress")) + if request.method == "GET" and path == "/v1/batches": + return _json( + { + "object": "list", + "data": [_batch_object(batch_id, file_id, "in_progress")], + "first_id": batch_id, + "last_id": batch_id, + "has_more": False, + } + ) + if request.method == "POST" and path == f"/v1/batches/{batch_id}/cancel": + return _json(_batch_object(batch_id, file_id, "cancelling")) + return None + + return _peer(route) + + +def _fine_tuning_job(job_id: str, status: str) -> Mapping[str, JsonValue]: + return { + "id": job_id, + "object": "fine_tuning.job", + "model": UPSTREAM_MODEL, + "created_at": 1, + "fine_tuned_model": None, + "finished_at": None, + "hyperparameters": {"n_epochs": 1, "batch_size": 1, "learning_rate_multiplier": 1.0}, + "organization_id": "org-integration", + "result_files": [], + "status": status, + "trained_tokens": None, + "training_file": f"file-{job_id.removeprefix('ftjob-')}", + "validation_file": None, + "seed": 1, + "error": None, + } + + +def _assistant(identity: str) -> Mapping[str, JsonValue]: + return { + "id": identity, + "object": "assistant", + "created_at": 1, + "name": "integration", + "description": None, + "model": UPSTREAM_MODEL, + "instructions": None, + "tools": [], + "metadata": {}, + "temperature": 1.0, + "top_p": 1.0, + "response_format": "auto", + "tool_resources": None, + } + + +def _configured_peer() -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + path: Final = _path(request) + if request.method == "POST" and path == "/v1/fine_tuning/jobs": + training_file: Final = string_value(_body(request)["training_file"]) + return _json(_fine_tuning_job(f"ftjob-{training_file.removeprefix('file-')}", "queued")) + if request.method == "GET" and path == "/v1/fine_tuning/jobs": + return _json({"object": "list", "data": [_fine_tuning_job(LISTED_JOB, "running")], "has_more": False}) + if request.method == "POST" and path.startswith("/v1/fine_tuning/jobs/") and path.endswith("/cancel"): + return _json( + _fine_tuning_job(path.removeprefix("/v1/fine_tuning/jobs/").removesuffix("/cancel"), "cancelled") + ) + if request.method == "GET" and path.startswith("/v1/fine_tuning/jobs/"): + return _json(_fine_tuning_job(path.removeprefix("/v1/fine_tuning/jobs/"), "running")) + if request.method == "GET" and path == "/v1/assistants": + return _json( + { + "object": "list", + "data": [_assistant(ASSISTANT)], + "first_id": ASSISTANT, + "last_id": ASSISTANT, + "has_more": False, + } + ) + return None + + return _peer(route) + + +def _config_with_provider_settings(directory: Path, api_base: str) -> Path: + shared: Final = object_value(yaml.safe_load((ROOT / "tests/integration/proxy_config.yaml").read_text())) + path: Final = directory / "proxy_config_with_provider_settings.yaml" + path.write_text( + yaml.safe_dump( + { + **shared, + "finetune_settings": [{"custom_llm_provider": "openai", "api_base": api_base, "api_key": PROVIDER_KEY}], + "assistant_settings": { + "custom_llm_provider": "openai", + "litellm_params": {"api_base": api_base, "api_key": PROVIDER_KEY}, + }, + } + ) + ) + return path + + +@dataclass(frozen=True, slots=True) +class ConfiguredProxy: + gateway: Gateway + wire: Wire + + +@pytest.fixture(scope="module") +def configured_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ConfiguredProxy]: + with gateway_from_environment() as rig_gateway, wire_server(_configured_peer()) as wire: + directory: Final = tmp_path_factory.mktemp("openai_sdk_client_paths") + config: Final = _config_with_provider_settings(directory, f"{wire.url}/v1") + with owned_proxy(rig_gateway, directory, {}, config=config, workers=PROXY_WORKERS) as owned: + yield ConfiguredProxy(owned, wire) + + +def _v1(gateway: Gateway) -> str: + return f"{str(gateway.client.base_url).rstrip('/')}/v1" + + +def _origin(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +@contextmanager +def _openai(gateway: Gateway) -> Iterator[openai.OpenAI]: + with ( + httpx.Client(timeout=15, trust_env=False) as transport, + openai.OpenAI(api_key=gateway.key, base_url=_v1(gateway), max_retries=0, http_client=transport) as client, + ): + yield client + + +def _async_openai(gateway: Gateway, call: Callable[[openai.AsyncOpenAI], Awaitable[_R]]) -> _R: + async def run() -> _R: + async with ( + httpx.AsyncClient(timeout=15, trust_env=False) as transport, + openai.AsyncOpenAI( + api_key=gateway.key, base_url=_v1(gateway), max_retries=0, http_client=transport + ) as client, + ): + return await call(client) + + return asyncio.run(run()) + + +@contextmanager +def _anthropic(gateway: Gateway) -> Iterator[anthropic.Anthropic]: + with ( + httpx.Client(timeout=15, trust_env=False) as transport, + anthropic.Anthropic( + api_key=gateway.key, base_url=_origin(gateway), max_retries=0, http_client=transport + ) as client, + ): + yield client + + +def _async_anthropic(gateway: Gateway, call: Callable[[anthropic.AsyncAnthropic], Awaitable[_R]]) -> _R: + async def run() -> _R: + async with ( + httpx.AsyncClient(timeout=15, trust_env=False) as transport, + anthropic.AsyncAnthropic( + api_key=gateway.key, base_url=_origin(gateway), max_retries=0, http_client=transport + ) as client, + ): + return await call(client) + + return asyncio.run(run()) + + +def _user_message() -> Mapping[str, JsonValue]: + return {"role": "user", "content": TEXT} + + +def _chat_body(model: str, extra: Mapping[str, JsonValue] = NO_EXTRA) -> Mapping[str, JsonValue]: + return {"model": model, "messages": [_user_message()], **NO_CACHE, **extra} + + +def _duplicated_timeout_body(model: str) -> bytes: + fields: Final = json.dumps(_chat_body(model)).removesuffix("}") + return f'{fields}, "timeout": 30, "timeout": 30}}'.encode() + + +def _chat(gateway: Gateway, model: str, extra: Mapping[str, JsonValue] = NO_EXTRA) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", _chat_body(model, extra)) + + +def _response_id(response: httpx.Response) -> str: + assert response.status_code == 200, response.text + return string_value(JSON_OBJECT.validate_json(response.content)["id"]) + + +def _error_message(response: httpx.Response) -> str: + body: Final = JSON_OBJECT.validate_json(response.content) + if "error" in body: + return string_value(object_value(body["error"])["message"]) + detail: Final = body["detail"] + if isinstance(detail, str): + return detail + assert isinstance(detail, list) and detail, response.text + return string_value(object_value(detail[0])["msg"]) + + +def _assert_rejected(response: httpx.Response, shape: str) -> None: + assert response.status_code >= 400, (shape, response.status_code, response.text) + assert _error_message(response), (shape, response.text) + + +def _assert_served(response: httpx.Response, identity: str, shape: str) -> None: + assert response.status_code == 200, (shape, response.status_code, response.text) + assert _response_id(response) == identity, (shape, response.text) + + +def _moderation_through_the_sdk(gateway: Gateway, model: str) -> ModerationCreateResponse | openai.APIStatusError: + with _openai(gateway) as client: + try: + return client.moderations.create(model=model, input=TEXT) + except openai.APIStatusError as error: + return error + + +def _moderation_after_worker_sync(gateway: Gateway, model: str) -> ModerationCreateResponse | openai.APIStatusError: + return eventually( + lambda: _moderation_through_the_sdk(gateway, model), + lambda outcome: not isinstance(outcome, openai.InternalServerError), + seconds=WORKER_SYNC_SECONDS + 10, + ) + + +def _speech_peer(content_type: str) -> Callable[[Request], Reply]: + def route(request: Request) -> Reply | None: + if request.method != "POST" or not _path(request).endswith("/audio/speech"): + return None + return Reply(body=SPEECH_AUDIO, content_type=content_type) + + return _peer(route) + + +def _speech_through_the_sdk(gateway: Gateway, model: str) -> httpx.Response | openai.APIStatusError: + with _openai(gateway) as client: + try: + return client.audio.speech.create(model=model, voice="alloy", input=TEXT, response_format="mp3").response + except openai.APIStatusError as error: + return error + + +def _speech_after_worker_sync(gateway: Gateway, model: str) -> httpx.Response | openai.APIStatusError: + return eventually( + lambda: _speech_through_the_sdk(gateway, model), + lambda outcome: not isinstance(outcome, openai.InternalServerError), + seconds=WORKER_SYNC_SECONDS + 10, + ) + + +def _attempt(call: Callable[[], _R]) -> _R | openai.APIStatusError: + try: + return call() + except openai.APIStatusError as error: + return error + + +def _unknown_to_the_serving_worker(outcome: object) -> bool: + return isinstance(outcome, openai.BadRequestError) and "Invalid model name" in outcome.response.text + + +def _outcome_after_worker_sync(call: Callable[[], _R]) -> _R | openai.APIStatusError: + return eventually( + lambda: _attempt(call), + lambda outcome: not _unknown_to_the_serving_worker(outcome), + seconds=WORKER_SYNC_SECONDS + 10, + ) + + +def _after_worker_sync(call: Callable[[], _R]) -> _R: + outcome: Final = _outcome_after_worker_sync(call) + assert not isinstance(outcome, openai.APIStatusError), outcome.response.text + return outcome + + +def _async_after_worker_sync(gateway: Gateway, call: Callable[[openai.AsyncOpenAI], Awaitable[_R]]) -> _R: + return _after_worker_sync(lambda: _async_openai(gateway, call)) + + +def _streamed_file_content(gateway: Gateway, file_id: str, model: str) -> tuple[int, tuple[bytes, ...]]: + with gateway.client.stream( + "GET", + f"/v1/files/{file_id}/content", + params={"model": model}, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + return response.status_code, tuple(response.iter_bytes()) + + +def _upstream_id(managed: str, prefix: str) -> str: + encoded: Final = managed.removeprefix(prefix) + try: + decoded: Final = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode() + except (binascii.Error, UnicodeDecodeError): + return managed + return decoded.removeprefix("litellm:").split(";", 1)[0] if decoded.startswith("litellm:") else managed + + +def _spend_row(identity: str) -> Mapping[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, call_type, model FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (identity,) + ), + lambda rows: len(rows) == 1, + seconds=70, + ) + return rows[0] + + +def _spend_row_keyed_by_any(identities: frozenset[str]) -> Mapping[str, JsonValue]: + keys: Final = tuple(sorted(identities)) + placeholders: Final = ", ".join("%s" for _ in keys) + rows: Final = eventually( + lambda: read_rows( + f'SELECT request_id, call_type, model FROM "LiteLLM_SpendLogs" WHERE request_id IN ({placeholders})', keys + ), + lambda rows: len(rows) == 1, + seconds=70, + ) + return rows[0] + + +def _spend_request_ids(identities: tuple[str, str, str]) -> frozenset[str]: + rows: Final = eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id IN (%s, %s, %s)', identities), + lambda rows: len(rows) == 3, + seconds=70, + ) + return frozenset(string_value(row["request_id"]) for row in rows) + + +def _anthropic_text(event: RawMessageStreamEvent) -> str: + if isinstance(event, RawContentBlockDeltaEvent) and isinstance(event.delta, TextDelta): + return event.delta.text + return "" + + +def _chunk_text(chunk: ChatCompletionChunk) -> str: + return chunk.choices[0].delta.content or "" if chunk.choices else "" + + +def _chunk_finished(chunk: ChatCompletionChunk) -> bool: + return bool(chunk.choices) and chunk.choices[0].finish_reason == "stop" + + +def _new_model(gateway: Gateway, api_base: str, timeout: JsonValue) -> httpx.Response: + return gateway.request( + "POST", + "/model/new", + { + "model_name": f"integration-{uuid.uuid4().hex}", + "litellm_params": { + "model": f"openai/{UPSTREAM_MODEL}", + "api_key": PROVIDER_KEY, + "api_base": api_base, + "timeout": timeout, + }, + "model_info": {}, + }, + ) + + +def _register_accepted_model(scenario: Scenario, response: httpx.Response, shape: str) -> None: + assert response.status_code == 200, (shape, response.status_code, response.text) + identity: Final = string_value(object_value(JSON_OBJECT.validate_json(response.content)["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, identity) + + +def test_r01_chat_completions_non_stream_through_the_sync_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + with _openai(gateway) as client: + response: Final = client.chat.completions.create( + model=model, messages=[{"role": "user", "content": TEXT}], extra_body=dict(NO_CACHE) + ) + assert response.id == identity, response + assert response.choices[0].message.content == TEXT, response + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/chat/completions"),), served + assert served[0].headers["authorization"] == PROVIDER_AUTHORIZATION, served[0].headers + upstream_body: Final = _body(served[0]) + assert upstream_body["model"] == UPSTREAM_MODEL, upstream_body + assert upstream_body["messages"] == [_user_message()], upstream_body + assert _spend_row(identity)["model"] == DEPLOYMENT_MODEL + + +def test_r02_chat_completions_stream_through_the_async_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + + async def stream(client: openai.AsyncOpenAI) -> tuple[ChatCompletionChunk, ...]: + chunks: Final = await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": TEXT}], stream=True, extra_body=dict(NO_CACHE) + ) + return tuple([chunk async for chunk in chunks]) + + received: Final = _async_openai(gateway, stream) + assert {chunk.id for chunk in received} == {identity}, received + assert "".join(_chunk_text(chunk) for chunk in received) == TEXT, received + assert any(_chunk_finished(chunk) for chunk in received), received + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/chat/completions"),), served + assert _body(served[0])["stream"] is True, served[0].body + assert _spend_row(identity)["model"] == DEPLOYMENT_MODEL + + +def test_r03_chat_completions_through_raw_httpx_with_a_body_timeout(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + assert _response_id(_chat(gateway, model, {"timeout": 30})) == identity + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/chat/completions"),), served + assert "timeout" not in _body(served[0]), served[0].body + assert _spend_row(identity)["model"] == DEPLOYMENT_MODEL + + +def test_r04_messages_non_stream_through_the_sync_anthropic_sdk(gateway: Gateway) -> None: + identity: Final = _identity("resp") + with gateway.scenario() as scenario, wire_server(_responses_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + with _anthropic(gateway) as client: + message: Final = client.messages.create( + model=model, max_tokens=64, messages=[{"role": "user", "content": TEXT}] + ) + assert identity in response_identities(message.id), message.id + first_block: Final = message.content[0] + assert first_block.type == "text" and first_block.text == TEXT, message + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/responses"),), served + assert TEXT in served[0].body.decode(), served[0].body + assert _spend_row(message.id)["call_type"] == "anthropic_messages" + + +def test_r05_messages_stream_through_the_async_anthropic_sdk(gateway: Gateway) -> None: + identity: Final = _identity("resp") + with gateway.scenario() as scenario, wire_server(_responses_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + + async def stream(client: anthropic.AsyncAnthropic) -> tuple[RawMessageStreamEvent, ...]: + events: Final = await client.messages.create( + model=model, max_tokens=64, messages=[{"role": "user", "content": TEXT}], stream=True + ) + return tuple([event async for event in events]) + + received: Final = _async_anthropic(gateway, stream) + first: Final = received[0] + assert isinstance(first, RawMessageStartEvent), received + assert received[-1].type == "message_stop", received + assert "".join(_anthropic_text(event) for event in received) == TEXT, received + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/responses"),), served + assert _body(served[0])["stream"] is True, served[0].body + assert _spend_row(first.message.id)["call_type"] == "anthropic_messages" + + +def test_r06_responses_non_stream_through_the_sync_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("resp") + with gateway.scenario() as scenario, wire_server(_responses_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + with _openai(gateway) as client: + response: Final = client.responses.create(model=model, input=TEXT, extra_body=dict(NO_CACHE)) + assert identity in response_identities(response.id), response.id + assert response.output_text == TEXT, response + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/responses"),), served + upstream_body: Final = _body(served[0]) + assert upstream_body["model"] == UPSTREAM_MODEL, upstream_body + assert upstream_body["input"] == TEXT, upstream_body + assert _spend_row(response.id)["model"] == DEPLOYMENT_MODEL + + +def test_r07_responses_stream_through_the_async_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("resp") + with gateway.scenario() as scenario, wire_server(_responses_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + + async def stream(client: openai.AsyncOpenAI) -> tuple[ResponseStreamEvent, ...]: + events: Final = await client.responses.create( + model=model, input=TEXT, stream=True, extra_body=dict(NO_CACHE) + ) + return tuple([event async for event in events]) + + received: Final = _async_openai(gateway, stream) + assert received[0].type == "response.created", received + last: Final = received[-1] + assert last.type == "response.completed", received + assert identity in response_identities(last.response.id), last + assert last.response.output_text == TEXT, last + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/responses"),), served + assert _body(served[0])["stream"] is True, served[0].body + assert _spend_row_keyed_by_any(response_identities(last.response.id))["model"] == DEPLOYMENT_MODEL + + +def test_r08_completions_non_stream_through_the_sync_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("cmpl") + with gateway.scenario() as scenario, wire_server(_completion_peer(identity)) as wire: + model: Final = scenario.model(model=INSTRUCT_DEPLOYMENT_MODEL, api_base=f"{wire.url}/v1") + with _openai(gateway) as client: + completion: Final = client.completions.create(model=model, prompt=TEXT, extra_body=dict(NO_CACHE)) + assert completion.id == identity, completion + assert completion.choices[0].text == TEXT, completion + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/completions"),), served + upstream_body: Final = _body(served[0]) + assert upstream_body["model"] == INSTRUCT_MODEL, upstream_body + assert upstream_body["prompt"] == TEXT, upstream_body + assert _spend_row(identity)["model"] == INSTRUCT_DEPLOYMENT_MODEL + + +def test_r09_completions_stream_through_the_async_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("cmpl") + with gateway.scenario() as scenario, wire_server(_completion_peer(identity)) as wire: + model: Final = scenario.model(model=INSTRUCT_DEPLOYMENT_MODEL, api_base=f"{wire.url}/v1") + + async def stream(client: openai.AsyncOpenAI) -> tuple[Completion, ...]: + chunks: Final = await client.completions.create( + model=model, prompt=TEXT, stream=True, extra_body=dict(NO_CACHE) + ) + return tuple([chunk async for chunk in chunks]) + + received: Final = _async_openai(gateway, stream) + assert {chunk.id for chunk in received} == {identity}, received + assert "".join(chunk.choices[0].text or "" for chunk in received if chunk.choices) == TEXT, received + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/completions"),), served + assert _body(served[0])["stream"] is True, served[0].body + assert _spend_row(identity)["model"] == INSTRUCT_DEPLOYMENT_MODEL + + +def test_r10_moderations_through_the_sync_openai_sdk(gateway: Gateway) -> None: + identity: Final = _identity("modr") + with gateway.scenario() as scenario, wire_server(_moderation_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + moderation: Final = _moderation_after_worker_sync(gateway, model) + assert isinstance(moderation, ModerationCreateResponse), moderation + assert moderation.id == identity, moderation + assert [result.flagged for result in moderation.results] == [False], moderation + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/moderations"),), served + assert served[0].headers["authorization"] == PROVIDER_AUTHORIZATION, served[0].headers + assert _body(served[0])["input"] == TEXT, served[0].body + assert _spend_row(moderation.id)["call_type"] == "amoderation" + + +def test_r11_files_create_retrieve_and_content_through_the_sync_openai_sdk(gateway: Gateway) -> None: + file_id: Final = _identity("file") + content: Final = _batch_line(file_id) + with gateway.scenario() as scenario, wire_server(_files_peer(file_id, content)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + with _openai(gateway) as client: + created: Final = _after_worker_sync( + lambda: client.files.create(file=("batch.jsonl", content), purpose="batch", extra_body={"model": model}) + ) + retrieved: Final = _after_worker_sync( + lambda: client.files.retrieve(created.id, extra_query={"model": model}) + ) + downloaded: Final = _after_worker_sync( + lambda: client.files.content(created.id, extra_query={"model": model}) + ).content + assert _upstream_id(created.id, "file-") == file_id, created + assert _upstream_id(retrieved.id, "file-") == file_id, retrieved + assert retrieved.bytes == len(content), retrieved + assert downloaded == content, downloaded + served: Final = _served(wire) + assert _routes(served) == ( + ("POST", "/v1/files"), + ("GET", f"/v1/files/{file_id}"), + ("GET", f"/v1/files/{file_id}/content"), + ), served + assert content in served[0].body, served[0].body + assert {request.headers["authorization"] for request in served} == {PROVIDER_AUTHORIZATION}, served + + +def test_r12_file_content_streams_through_raw_httpx(gateway: Gateway) -> None: + file_id: Final = _identity("file") + content: Final = STREAM_HEAD + STREAM_TAIL + with ( + gateway.scenario() as scenario, + wire_server(_files_peer(file_id, content, chunks=(STREAM_HEAD, STREAM_TAIL))) as wire, + ): + model: Final = scenario.model(api_base=f"{wire.url}/v1") + status, received = eventually( + lambda: _streamed_file_content(gateway, file_id, model), + lambda outcome: not (outcome[0] == 400 and b"Invalid model name" in b"".join(outcome[1])), + seconds=WORKER_SYNC_SECONDS + 10, + ) + assert status == 200, b"".join(received).decode() + assert b"".join(received) == content, (len(received), sum(len(chunk) for chunk in received)) + served: Final = _served(wire) + assert _routes(served) == (("GET", f"/v1/files/{file_id}/content"),), served + assert served[0].headers["authorization"] == PROVIDER_AUTHORIZATION, served[0].headers + + +def test_r13_batches_create_retrieve_list_and_cancel_through_the_async_openai_sdk(gateway: Gateway) -> None: + file_id: Final = _identity("file") + batch_id: Final = _identity("batch") + with gateway.scenario() as scenario, wire_server(_batches_peer(file_id, batch_id)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + uploaded: Final = _async_after_worker_sync( + gateway, + lambda client: client.files.create( + file=("batch.jsonl", _batch_line("r1")), purpose="batch", extra_body={"model": model} + ), + ) + created: Final = _async_after_worker_sync( + gateway, + lambda client: client.batches.create( + input_file_id=uploaded.id, + endpoint="/v1/chat/completions", + completion_window="24h", + extra_body={"model": model}, + ), + ) + retrieved: Final = _async_after_worker_sync(gateway, lambda client: client.batches.retrieve(created.id)) + listed: Final = _async_after_worker_sync( + gateway, lambda client: client.batches.list(extra_query={"model": model}) + ) + cancelled: Final = _async_after_worker_sync(gateway, lambda client: client.batches.cancel(created.id)) + listed_ids: Final = tuple(batch.id for batch in listed.data) + assert _upstream_id(uploaded.id, "file-") == file_id, uploaded + assert _upstream_id(created.id, "batch_") == batch_id, created + assert _upstream_id(retrieved.id, "batch_") == batch_id, retrieved + assert batch_id in [_upstream_id(listed_id, "batch_") for listed_id in listed_ids], listed_ids + assert _upstream_id(cancelled.id, "batch_") == batch_id, cancelled + assert cancelled.status == "cancelling", cancelled + served: Final = _served(wire) + assert _routes(served) == ( + ("POST", "/v1/files"), + ("POST", "/v1/batches"), + ("GET", f"/v1/batches/{batch_id}"), + ("POST", f"/v1/batches/{batch_id}/cancel"), + ), served + assert _body(served[1])["input_file_id"] == file_id, served[1].body + assert {request.headers["authorization"] for request in served} == {PROVIDER_AUTHORIZATION}, served + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 120) +def test_r14_fine_tuning_jobs_create_list_retrieve_and_cancel_through_the_sync_openai_sdk( + configured_proxy: ConfiguredProxy, +) -> None: + suffix: Final = uuid.uuid4().hex + job_id: Final = f"ftjob-{suffix}" + with _openai(configured_proxy.gateway) as client: + created: Final = client.fine_tuning.jobs.create( + model=UPSTREAM_MODEL, training_file=f"file-{suffix}", extra_body={"custom_llm_provider": "openai"} + ) + listed: Final = client.fine_tuning.jobs.list(extra_query={"custom_llm_provider": "openai"}) + retrieved: Final = client.fine_tuning.jobs.retrieve(created.id, extra_query={"custom_llm_provider": "openai"}) + cancelled: Final = client.fine_tuning.jobs.cancel(created.id, extra_body={"custom_llm_provider": "openai"}) + assert created.id == job_id, created + assert created.status == "queued", created + assert [job.id for job in listed.data] == [LISTED_JOB], listed + assert retrieved.id == job_id, retrieved + assert retrieved.training_file == f"file-{suffix}", retrieved + assert cancelled.id == job_id and cancelled.status == "cancelled", cancelled + served: Final = _served(configured_proxy.wire) + assert _routes(served) == ( + ("POST", "/v1/fine_tuning/jobs"), + ("GET", "/v1/fine_tuning/jobs"), + ("GET", f"/v1/fine_tuning/jobs/{job_id}"), + ("POST", f"/v1/fine_tuning/jobs/{job_id}/cancel"), + ), served + assert _body(served[0])["training_file"] == f"file-{suffix}", served[0].body + assert {request.headers["authorization"] for request in served} == {PROVIDER_AUTHORIZATION}, served + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 120) +def test_r15_assistants_list_through_raw_httpx(configured_proxy: ConfiguredProxy) -> None: + response: Final = configured_proxy.gateway.request("GET", "/v1/assistants") + assert response.status_code == 200, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + assert body["object"] == "list", response.text + listed: Final = body["data"] + assert isinstance(listed, list), response.text + assert [string_value(object_value(entry)["id"]) for entry in listed] == [ASSISTANT], response.text + served: Final = _served(configured_proxy.wire) + assert _routes(served) == (("GET", "/v1/assistants"),), served + assert served[0].headers["authorization"] == PROVIDER_AUTHORIZATION, served[0].headers + + +def test_r17_chat_completions_on_an_azure_cloudflare_gateway_deployment_through_the_sync_openai_sdk( + gateway: Gateway, +) -> None: + identity: Final = _identity("chatcmpl") + deployment: Final = f"gpt-4o-mini-{uuid.uuid4().hex[:8]}" + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model( + model=f"azure/{deployment}", api_base=f"{wire.url}{CLOUDFLARE_GATEWAY}", api_version=AZURE_API_VERSION + ) + with _openai(gateway) as client: + response: Final = client.chat.completions.create( + model=model, messages=[{"role": "user", "content": TEXT}], extra_body=dict(NO_CACHE) + ) + streamed: Final = tuple( + client.chat.completions.create( + model=model, messages=[{"role": "user", "content": TEXT}], stream=True, extra_body=dict(NO_CACHE) + ) + ) + assert response.id == identity, response + assert response.choices[0].message.content == TEXT, response + assert {chunk.id for chunk in streamed} == {identity}, streamed + assert "".join(_chunk_text(chunk) for chunk in streamed) == TEXT, streamed + served: Final = _served(wire) + assert _routes(served) == (("POST", f"{CLOUDFLARE_GATEWAY}/{deployment}/chat/completions"),) * 2, served + assert {request.target.split("?", 1)[1] for request in served} == {f"api-version={AZURE_API_VERSION}"}, served + assert {request.headers["api-key"] for request in served} == {PROVIDER_KEY}, served + assert [_streams(request) for request in served] == [False, True], served + assert _spend_row(identity)["model"] == f"azure/{deployment}" + + +def test_r18_audio_speech_answers_with_the_upstream_audio_type_through_the_sync_openai_sdk(gateway: Gateway) -> None: + with gateway.scenario() as scenario, wire_server(_speech_peer("audio/ogg")) as wire: + model: Final = scenario.model(model=f"openai/{SPEECH_UPSTREAM_MODEL}", api_base=f"{wire.url}/v1") + speech: Final = _speech_after_worker_sync(gateway, model) + assert not isinstance(speech, openai.APIStatusError), speech + assert speech.status_code == 200, speech.text + assert speech.headers["content-type"] == "audio/ogg", dict(speech.headers) + assert speech.content == SPEECH_AUDIO, speech.content[:16] + served: Final = _served(wire) + assert _routes(served) == (("POST", "/v1/audio/speech"),), served + assert served[0].headers["authorization"] == PROVIDER_AUTHORIZATION, served[0].headers + upstream_body: Final = _body(served[0]) + assert (upstream_body["model"], upstream_body["input"], upstream_body["response_format"]) == ( + SPEECH_UPSTREAM_MODEL, + TEXT, + "mp3", + ), upstream_body + assert _spend_row(speech.headers["x-litellm-call-id"])["call_type"] == "aspeech" + + +def test_s01_chat_completions_upstream_400_reaches_the_sync_openai_sdk_as_a_bad_request(gateway: Gateway) -> None: + with gateway.scenario() as scenario, wire_server(_status_peer(openai_error(400))) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1", max_retries=0) + with _openai(gateway) as client, pytest.raises(openai.BadRequestError) as raised: + client.chat.completions.create(model=model, messages=[_user_message()], extra_body=dict(NO_CACHE)) + assert raised.value.status_code == 400, raised.value + assert "scripted 400" in raised.value.response.text, raised.value.response.text + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) + + +def test_s02_chat_completions_upstream_429_keeps_retry_after_through_raw_httpx(gateway: Gateway) -> None: + limited: Final = Reply(status=429, body=openai_error(429).body, headers={"retry-after": "7"}) + with gateway.scenario() as scenario, wire_server(_status_peer(limited)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1", max_retries=0) + response: Final = _chat(gateway, model) + assert response.status_code == 429, response.text + assert response.headers.get("llm_provider-retry-after") == "7", dict(response.headers) + assert "scripted 429" in _error_message(response), response.text + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) + + +def test_s03_chat_completions_upstream_500_reaches_the_sync_openai_sdk_as_a_server_error(gateway: Gateway) -> None: + with gateway.scenario() as scenario, wire_server(_status_peer(openai_error(500))) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1", max_retries=0) + with _openai(gateway) as client, pytest.raises(openai.InternalServerError) as raised: + client.chat.completions.create(model=model, messages=[_user_message()], extra_body=dict(NO_CACHE)) + assert raised.value.status_code == 500, raised.value + assert "scripted 500" in raised.value.response.text, raised.value.response.text + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) + + +def test_s04_chat_completions_body_timeout_shapes_answer_without_breaking_the_proxy(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + for shape, timeout in REJECTED_CHAT_TIMEOUTS.items(): + _assert_rejected(_chat(gateway, model, {"timeout": timeout}), shape) + for shape, timeout in ACCEPTED_CHAT_TIMEOUTS.items(): + _assert_served(_chat(gateway, model, {"timeout": timeout}), identity, shape) + duplicated: Final = gateway.client.post( + "/v1/chat/completions", + content=_duplicated_timeout_body(model), + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + _assert_served(duplicated, identity, "duplicated key") + _assert_served(_chat(gateway, model), identity, "follow-up") + assert len(_served(wire)) == len(ACCEPTED_CHAT_TIMEOUTS) + 2 + + +def test_s05_model_new_timeout_shapes_are_rejected_or_stored_without_breaking_the_proxy(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + for shape, timeout in REJECTED_DEPLOYMENT_TIMEOUTS.items(): + _assert_rejected(_new_model(gateway, f"{wire.url}/v1", timeout), shape) + for shape, timeout in ACCEPTED_DEPLOYMENT_TIMEOUTS.items(): + _register_accepted_model(scenario, _new_model(gateway, f"{wire.url}/v1", timeout), shape) + _assert_served(_chat(gateway, model), identity, "follow-up") + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) + + +def test_s06_upstream_errors_on_completions_moderations_and_files_reach_the_sync_openai_sdk( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + with wire_server(_status_peer(openai_error(401))) as unauthorized: + model: Final = scenario.model( + model=INSTRUCT_DEPLOYMENT_MODEL, api_base=f"{unauthorized.url}/v1", max_retries=0 + ) + with _openai(gateway) as client, pytest.raises(openai.AuthenticationError) as completion_error: + client.completions.create(model=model, prompt=TEXT, extra_body=dict(NO_CACHE)) + assert completion_error.value.status_code == 401, completion_error.value + assert "scripted 401" in completion_error.value.response.text, completion_error.value.response.text + assert _routes(_served(unauthorized)) == (("POST", "/v1/completions"),) + with wire_server(_status_peer(openai_error(400))) as rejecting: + rejected_model: Final = scenario.model(api_base=f"{rejecting.url}/v1", max_retries=0) + moderation_error: Final = _moderation_after_worker_sync(gateway, rejected_model) + assert isinstance(moderation_error, openai.BadRequestError), moderation_error + assert "scripted 400" in moderation_error.response.text, moderation_error.response.text + with _openai(gateway) as client: + file_error: Final = _outcome_after_worker_sync( + lambda: client.files.create( + file=("batch.jsonl", _batch_line("rejected")), + purpose="batch", + extra_body={"model": rejected_model}, + ) + ) + assert isinstance(file_error, openai.BadRequestError), file_error + assert "scripted 400" in file_error.response.text, file_error.response.text + assert _routes(_served(rejecting)) == (("POST", "/v1/moderations"), ("POST", "/v1/files")) + + +@pytest.mark.parametrize("provider", tuple(PAYMENT_REQUIRED_DEPLOYMENTS)) +def test_s07_chat_completions_upstream_402_reaches_the_sync_openai_sdk_as_payment_required( + gateway: Gateway, provider: str +) -> None: + with gateway.scenario() as scenario, wire_server(_status_peer(openai_error(402))) as wire: + model: Final = scenario.model(model=PAYMENT_REQUIRED_DEPLOYMENTS[provider], api_base=wire.url, max_retries=0) + with _openai(gateway) as client, pytest.raises(openai.APIStatusError) as raised: + client.chat.completions.create(model=model, messages=[_user_message()], extra_body=dict(NO_CACHE)) + assert raised.value.status_code == 402, raised.value + assert "scripted 402" in raised.value.response.text, raised.value.response.text + assert len(_served(wire)) == 1 + + +def test_e01_deployment_timeout_answers_408_and_a_body_timeout_on_the_same_deployment_recovers( + gateway: Gateway, +) -> None: + identity: Final = _identity("chatcmpl") + gate: Final = threading.Event() + with gateway.scenario() as scenario, wire_server(_held_chat_peer(gate, identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1", timeout=0.5, max_retries=0) + timed_out: Final = _chat(gateway, model) + gate.set() + assert timed_out.status_code == 408, timed_out.text + assert "timeout" in _error_message(timed_out).lower(), timed_out.text + recovered: Final = _chat(gateway, model, {"timeout": 5}) + assert _response_id(recovered) == identity, recovered.text + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) * 2 + + +def test_e02_three_identical_no_cache_requests_through_raw_httpx_land_three_spend_rows(gateway: Gateway) -> None: + served_ids: Final[SimpleQueue[str]] = SimpleQueue() + with gateway.scenario() as scenario, wire_server(_fresh_chat_peer(served_ids)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + received: Final = ( + _response_id(_chat(gateway, model)), + _response_id(_chat(gateway, model)), + _response_id(_chat(gateway, model)), + ) + assert len(set(received)) == 3, received + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) * 3 + assert {served_ids.get_nowait() for _ in range(3)} == set(received), received + assert _spend_request_ids(received) == frozenset(received) + + +def test_e03_deployment_timeout_as_a_string_still_serves_through_raw_httpx(gateway: Gateway) -> None: + identity: Final = _identity("chatcmpl") + with gateway.scenario() as scenario, wire_server(_chat_peer(identity)) as wire: + model: Final = scenario.model(api_base=f"{wire.url}/v1", timeout="30") + assert _response_id(_chat(gateway, model)) == identity + assert _routes(_served(wire)) == (("POST", "/v1/chat/completions"),) + assert _spend_row(identity)["model"] == DEPLOYMENT_MODEL diff --git a/tests/integration/compatibility/test_openai_sdk_client_paths_chaos.py b/tests/integration/compatibility/test_openai_sdk_client_paths_chaos.py new file mode 100644 index 00000000000..0bcc9fc4340 --- /dev/null +++ b/tests/integration/compatibility/test_openai_sdk_client_paths_chaos.py @@ -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} " + + +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)), + ) diff --git a/tests/integration/sdk/test_openai_sdk_client_sessions_wire.py b/tests/integration/sdk/test_openai_sdk_client_sessions_wire.py new file mode 100644 index 00000000000..e8ec1a37082 --- /dev/null +++ b/tests/integration/sdk/test_openai_sdk_client_sessions_wire.py @@ -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 diff --git a/tests/typing/openai_sdk_compat.py b/tests/typing/openai_sdk_compat.py new file mode 100644 index 00000000000..03f44db6a9f --- /dev/null +++ b/tests/typing/openai_sdk_compat.py @@ -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) diff --git a/tests/typing/pyrightconfig.json b/tests/typing/pyrightconfig.json new file mode 100644 index 00000000000..0910ca64f5b --- /dev/null +++ b/tests/typing/pyrightconfig.json @@ -0,0 +1,8 @@ +{ + "extends": "../../pyrightconfig.json", + "include": ["openai_sdk_compat.py"], + "exclude": [], + "extraPaths": ["../.."], + "pythonVersion": "3.10", + "reportMissingImports": true +} diff --git a/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py index 57f4557c6f7..ea0416cecbd 100644 --- a/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py @@ -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 diff --git a/tests/unit/llms/openai/test_openai.py b/tests/unit/llms/openai/test_openai.py index 3774545488d..053ec7e141e 100644 --- a/tests/unit/llms/openai/test_openai.py +++ b/tests/unit/llms/openai/test_openai.py @@ -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, ) diff --git a/tests/unit/llms/openai/test_openai_common_utils.py b/tests/unit/llms/openai/test_openai_common_utils.py index fb605c6da7a..3c7dff34296 100644 --- a/tests/unit/llms/openai/test_openai_common_utils.py +++ b/tests/unit/llms/openai/test_openai_common_utils.py @@ -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 diff --git a/tests/unit/llms/openai/test_openai_workload_identity.py b/tests/unit/llms/openai/test_openai_workload_identity.py index 7d78a2c5238..d0cf52d0e28 100644 --- a/tests/unit/llms/openai/test_openai_workload_identity.py +++ b/tests/unit/llms/openai/test_openai_workload_identity.py @@ -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") diff --git a/tests/unit/test_completion_timeout_resolution.py b/tests/unit/test_completion_timeout_resolution.py index 7eb79e90e60..77166d6e947 100644 --- a/tests/unit/test_completion_timeout_resolution.py +++ b/tests/unit/test_completion_timeout_resolution.py @@ -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 + ) diff --git a/tests/unit/test_exception_header_preservation.py b/tests/unit/test_exception_header_preservation.py index dd142d9d40b..7acacf05226 100644 --- a/tests/unit/test_exception_header_preservation.py +++ b/tests/unit/test_exception_header_preservation.py @@ -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 ): diff --git a/tests/unit/test_openai_embedding_encoding_format_default.py b/tests/unit/test_openai_embedding_encoding_format_default.py index 7a42eaf0f0a..9788eb30b1b 100644 --- a/tests/unit/test_openai_embedding_encoding_format_default.py +++ b/tests/unit/test_openai_embedding_encoding_format_default.py @@ -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] diff --git a/tests/unit/types/test_router.py b/tests/unit/types/test_router.py index 89238e04a5a..6188bb966ee 100644 --- a/tests/unit/types/test_router.py +++ b/tests/unit/types/test_router.py @@ -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".""" diff --git a/uv.lock b/uv.lock index 96fa27289fa..471460a32b5 100644 --- a/uv.lock +++ b/uv.lock @@ -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" },