mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: support OpenAI SDK 3 while retaining SDK 2 compatibility (#44927)
* fix: support OpenAI SDK 3.x while retaining 2.x compatibility * fix: preserve HTTPX compatibility across OpenAI SDK versions * refactor: trim OpenAI SDK compatibility patch * refactor: simplify SDK client defaults * test: cover SDK compatibility across provider routes * fix: use shared names for SDK transport helpers * fix: address SDK compatibility review feedback * test(openai): send requests through the SDK API factory clients * fix(exceptions): keep the provider response on PaymentRequiredError under httpx 2 * ci(base-sdk): count httpx2 as base-only when a base dependency brings it in OpenAI SDK 3 depends on httpx2, so the litellm-core install at highest resolution pulls it in through openai. The base-only check now resolves the declared base dependency closure of litellm and litellm-core and accepts extras-only modules owned by a distribution inside that closure * test: cover OpenAI SDK 3 client paths end to end and read SDK HTTP headers case-insensitively * test: read C02 spend rows once, wait out the worker healthcheck on respawn, and resolve module owners on Python 3.10 --------- Co-authored-by: Marcus Wood <marcuswood@openai.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
674208bb7d
commit
63d823d28a
32 changed files with 2851 additions and 153 deletions
71
.github/workflows/test-unit.yml
vendored
71
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
1200
tests/integration/compatibility/test_openai_sdk_client_paths.py
Normal file
1200
tests/integration/compatibility/test_openai_sdk_client_paths.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,424 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import signal
|
||||
import subprocess
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from itertools import product
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
from anthropic import Anthropic
|
||||
from anthropic import APIError as AnthropicApiError
|
||||
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import graceful_stop_seconds, owned_proxy_process, owned_upstream
|
||||
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
|
||||
from openai import APIError as OpenAiApiError
|
||||
from openai import OpenAI
|
||||
from pydantic import JsonValue
|
||||
|
||||
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, SseResponse
|
||||
|
||||
_Kind: TypeAlias = Literal["chat", "chat_stream", "messages", "responses", "completions", "moderations"]
|
||||
|
||||
_ALL_KINDS: Final[tuple[_Kind, ...]] = ("chat", "chat_stream", "messages", "responses", "completions", "moderations")
|
||||
_CHAT_KINDS: Final[tuple[_Kind, ...]] = ("chat", "chat_stream")
|
||||
_NO_CACHE: Final = {"cache": {"no-cache": True}}
|
||||
_SPEND_QUERY: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
_CHAT: Final = JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "chatcmpl-$UNIQUE_ID",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
_MODERATION: Final = JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "modr-$UNIQUE_ID",
|
||||
"model": "omni-moderation-latest",
|
||||
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
|
||||
},
|
||||
)
|
||||
_RESPONSES: Final = JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "resp_$UNIQUE_ID",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_$UNIQUE_ID",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "scripted", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
_PLAIN: Final = RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={"POST /chat/completions": _CHAT, "POST /moderations": _MODERATION, "POST /responses": _RESPONSES},
|
||||
)
|
||||
_STREAM: Final = SseResponse(
|
||||
content_type="text/event-stream",
|
||||
frames=(
|
||||
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
|
||||
'"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}',
|
||||
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
|
||||
'"choices":[{"index":0,"delta":{"content":"streamed"},"finish_reason":null}]}',
|
||||
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini",'
|
||||
'"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}',
|
||||
"data: [DONE]",
|
||||
),
|
||||
frame_delay_ms=150,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
kind: _Kind
|
||||
marker: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Completed:
|
||||
call: _Call
|
||||
response_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Failed:
|
||||
call: _Call
|
||||
error: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Clients:
|
||||
openai: OpenAI
|
||||
anthropic: Anthropic
|
||||
plain_model: str
|
||||
stream_model: str
|
||||
|
||||
|
||||
class _Observations:
|
||||
def __init__(self, url: str) -> None:
|
||||
self.url = url.rstrip("/")
|
||||
self.items: tuple[Mapping[str, JsonValue], ...] = ()
|
||||
|
||||
def read(self) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
with httpx.Client(timeout=10, trust_env=False) as client:
|
||||
payload: Final = JSON_OBJECT.validate_python(client.get(f"{self.url}/__observations").json())
|
||||
requests: Final = payload.get("requests")
|
||||
assert isinstance(requests, list)
|
||||
self.items = (*self.items, *(object_value(item) for item in requests if isinstance(item, dict)))
|
||||
return self.items
|
||||
|
||||
|
||||
def _register(
|
||||
scenario: Scenario, upstream_url: str, label: str, response: RoutedResponse | SseResponse
|
||||
) -> ScenarioHandle:
|
||||
handle: Final = register_scenario(f"chaos-{label}-{uuid.uuid4().hex}", response, control_url=upstream_url)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
return handle
|
||||
|
||||
|
||||
def _deployments(scenario: Scenario, upstream_url: str) -> tuple[str, str]:
|
||||
plain_handle: Final = _register(scenario, upstream_url, "plain", _PLAIN)
|
||||
stream_handle: Final = _register(scenario, upstream_url, "stream", _STREAM)
|
||||
return (
|
||||
scenario.model(api_base=plain_handle.api_base(), api_key=plain_handle.scenario_id),
|
||||
scenario.model(api_base=stream_handle.api_base(), api_key=stream_handle.scenario_id),
|
||||
)
|
||||
|
||||
|
||||
def _clients(gateway: Gateway, deployments: tuple[str, str]) -> _Clients:
|
||||
base: Final = str(gateway.client.base_url).rstrip("/")
|
||||
clients: Final = _Clients(
|
||||
openai=OpenAI(base_url=f"{base}/v1", api_key=gateway.key, max_retries=0),
|
||||
anthropic=Anthropic(base_url=base, api_key=gateway.key, max_retries=0),
|
||||
plain_model=deployments[0],
|
||||
stream_model=deployments[1],
|
||||
)
|
||||
_import_lazy_sdk_resources_on_this_thread(clients)
|
||||
return clients
|
||||
|
||||
|
||||
def _import_lazy_sdk_resources_on_this_thread(clients: _Clients) -> None:
|
||||
resources: Final = (
|
||||
clients.openai.chat.completions,
|
||||
clients.openai.completions,
|
||||
clients.openai.responses,
|
||||
clients.openai.moderations,
|
||||
clients.anthropic.messages,
|
||||
)
|
||||
assert all(resource is not None for resource in resources)
|
||||
|
||||
|
||||
def _invoke(clients: _Clients, call: _Call) -> str:
|
||||
match call.kind:
|
||||
case "chat":
|
||||
return clients.openai.chat.completions.create(
|
||||
model=clients.plain_model, messages=[{"role": "user", "content": call.marker}], extra_body=_NO_CACHE
|
||||
).id
|
||||
case "chat_stream":
|
||||
chunks: Final = tuple(
|
||||
clients.openai.chat.completions.create(
|
||||
model=clients.stream_model,
|
||||
messages=[{"role": "user", "content": call.marker}],
|
||||
stream=True,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
)
|
||||
assert chunks[-1].choices[0].finish_reason == "stop", chunks
|
||||
assert len({chunk.id for chunk in chunks}) == 1, chunks
|
||||
return chunks[0].id
|
||||
case "messages":
|
||||
return clients.anthropic.messages.create(
|
||||
model=clients.plain_model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": call.marker}],
|
||||
extra_body=_NO_CACHE,
|
||||
).id
|
||||
case "responses":
|
||||
return clients.openai.responses.create(
|
||||
model=clients.plain_model, input=call.marker, extra_body=_NO_CACHE
|
||||
).id
|
||||
case "completions":
|
||||
return clients.openai.completions.create(
|
||||
model=clients.plain_model, prompt=call.marker, extra_body=_NO_CACHE
|
||||
).id
|
||||
case "moderations":
|
||||
return clients.openai.moderations.create(
|
||||
model=clients.plain_model, input=call.marker, extra_body=_NO_CACHE
|
||||
).id
|
||||
|
||||
|
||||
def _attempt(clients: _Clients, call: _Call) -> _Completed | _Failed:
|
||||
try:
|
||||
return _Completed(call, _invoke(clients, call))
|
||||
except (OpenAiApiError, AnthropicApiError, httpx.HTTPError) as error:
|
||||
return _Failed(call, f"{type(error).__name__}: {str(error)[:200]}")
|
||||
|
||||
|
||||
def _every_worker_serves_every_kind(gateway: Gateway, deployments: tuple[str, str]) -> bool:
|
||||
calls: Final = _burst_calls(_ALL_KINDS, 2)
|
||||
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
|
||||
futures: Final = _submit(pool, _clients(gateway, deployments), calls)
|
||||
outcomes: Final = tuple(future.result(timeout=60) for future in futures)
|
||||
return all(isinstance(outcome, _Completed) for outcome in outcomes)
|
||||
|
||||
|
||||
def _burst_calls(kinds: tuple[_Kind, ...], per_kind: int) -> tuple[_Call, ...]:
|
||||
return tuple(_Call(kind, f"burst-{kind}-{uuid.uuid4().hex}") for kind, _ in product(kinds, range(per_kind)))
|
||||
|
||||
|
||||
def _submit(
|
||||
pool: ThreadPoolExecutor, clients: _Clients, calls: tuple[_Call, ...]
|
||||
) -> tuple[Future[_Completed | _Failed], ...]:
|
||||
return tuple(pool.submit(_attempt, clients, call) for call in calls)
|
||||
|
||||
|
||||
def _done_count(futures: tuple[Future[_Completed | _Failed], ...]) -> int:
|
||||
return sum(future.done() for future in futures)
|
||||
|
||||
|
||||
def _liveliness(gateway: Gateway) -> int:
|
||||
return gateway.client.get("/health/liveliness").status_code
|
||||
|
||||
|
||||
def _spend_rows(request_id: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(read_rows(_SPEND_QUERY, (request_id,)))
|
||||
|
||||
|
||||
def _marker_counts(items: tuple[Mapping[str, JsonValue], ...], markers: tuple[str, ...]) -> tuple[int, ...]:
|
||||
bodies: Final = tuple(json.dumps(item.get("body")) for item in items)
|
||||
return tuple(sum(marker in body for body in bodies) for marker in markers)
|
||||
|
||||
|
||||
def _landed_once(response_ids: tuple[str, ...]) -> tuple[tuple[Mapping[str, JsonValue], ...], ...]:
|
||||
rows: Final = tuple(
|
||||
eventually(partial(_spend_rows, response_id), lambda values: len(values) == 1, seconds=90)
|
||||
for response_id in response_ids
|
||||
)
|
||||
assert all(found[0]["request_id"] == response_id for found, response_id in zip(rows, response_ids)), rows
|
||||
return rows
|
||||
|
||||
|
||||
def _is_worker(child: psutil.Process) -> bool:
|
||||
try:
|
||||
return "spawn_main" in " ".join(child.cmdline())
|
||||
except psutil.Error:
|
||||
return False
|
||||
|
||||
|
||||
def _worker_alive(worker: psutil.Process) -> bool:
|
||||
try:
|
||||
return worker.is_running() and worker.status() != psutil.STATUS_ZOMBIE
|
||||
except psutil.NoSuchProcess:
|
||||
return False
|
||||
|
||||
|
||||
def _alive_workers(process: subprocess.Popen[bytes]) -> tuple[psutil.Process, ...]:
|
||||
children: Final = tuple(child for child in psutil.Process(process.pid).children() if _is_worker(child))
|
||||
return tuple(child for child in children if _worker_alive(child))
|
||||
|
||||
|
||||
def _process_tree(process: subprocess.Popen[bytes]) -> str:
|
||||
root: Final = psutil.Process(process.pid)
|
||||
return "\n".join(_process_line(member) for member in (root, *root.children(recursive=True)))
|
||||
|
||||
|
||||
def _process_line(member: psutil.Process) -> str:
|
||||
try:
|
||||
return f"{member.pid} {' '.join(member.cmdline())}"
|
||||
except psutil.Error:
|
||||
return f"{member.pid} <exited>"
|
||||
|
||||
|
||||
def _kind_counts(outcomes: tuple[_Completed | _Failed, ...]) -> str:
|
||||
kinds: Final = tuple(outcome.call.kind for outcome in outcomes)
|
||||
return str({kind: kinds.count(kind) for kind in _ALL_KINDS if kind in kinds})
|
||||
|
||||
|
||||
def test_c01_upstream_pause_mid_burst_keeps_liveliness_and_lands_every_id_once(
|
||||
gateway: Gateway,
|
||||
tmp_path: Path,
|
||||
record_property: pytest.RecordProperty,
|
||||
) -> None:
|
||||
with owned_upstream(tmp_path) as slot, gateway.scenario() as scenario:
|
||||
upstream: Final = slot.process
|
||||
assert upstream is not None
|
||||
deployments: Final = _deployments(scenario, slot.url)
|
||||
eventually(partial(_every_worker_serves_every_kind, gateway, deployments), lambda served: served, seconds=90)
|
||||
clients: Final = _clients(gateway, deployments)
|
||||
calls: Final = _burst_calls(_ALL_KINDS, 5)
|
||||
observations: Final = _Observations(slot.url)
|
||||
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
|
||||
futures: Final = _submit(pool, clients, calls)
|
||||
eventually(partial(_done_count, futures), lambda done: done >= 1, seconds=60)
|
||||
upstream.send_signal(signal.SIGSTOP)
|
||||
try:
|
||||
done_at_pause: Final = _done_count(futures)
|
||||
paused_liveliness: Final = tuple(_liveliness(gateway) for _ in range(3))
|
||||
finally:
|
||||
upstream.send_signal(signal.SIGCONT)
|
||||
done_at_resume: Final = _done_count(futures)
|
||||
outcomes: Final = tuple(future.result(timeout=180) for future in futures)
|
||||
record_property("c01_burst_size", len(calls))
|
||||
record_property("c01_done_at_pause", done_at_pause)
|
||||
record_property("c01_done_at_resume", done_at_resume)
|
||||
record_property("c01_paused_liveliness_statuses", str(paused_liveliness))
|
||||
assert paused_liveliness == (200, 200, 200), paused_liveliness
|
||||
assert done_at_resume < len(calls), done_at_resume
|
||||
failures: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Failed))
|
||||
assert not failures, failures
|
||||
completed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Completed))
|
||||
record_property("c01_completed_by_kind", _kind_counts(completed))
|
||||
response_ids: Final = tuple(outcome.response_id for outcome in completed)
|
||||
assert len(set(response_ids)) == len(calls), response_ids
|
||||
spend_rows: Final = _landed_once(response_ids)
|
||||
record_property("c01_spend_query", _SPEND_QUERY)
|
||||
record_property("c01_spend_row_counts", str(tuple(len(rows) for rows in spend_rows)))
|
||||
markers: Final = tuple(call.marker for call in calls)
|
||||
eventually(
|
||||
observations.read,
|
||||
lambda items: _marker_counts(items, markers) == (1,) * len(markers),
|
||||
seconds=60,
|
||||
)
|
||||
record_property("c01_upstream_marker_counts", str(_marker_counts(observations.items, markers)))
|
||||
|
||||
|
||||
@pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
|
||||
def test_c02_worker_sigkill_mid_burst_keeps_serving_and_respawns(
|
||||
gateway: Gateway,
|
||||
tmp_path: Path,
|
||||
record_property: pytest.RecordProperty,
|
||||
) -> None:
|
||||
with owned_upstream(tmp_path) as slot, gateway.scenario() as scenario:
|
||||
deployments: Final = _deployments(scenario, slot.url)
|
||||
with owned_proxy_process(gateway, tmp_path, {"INTEGRATION_UPSTREAM_URL": slot.url}, workers=2) as owned:
|
||||
eventually(
|
||||
partial(_every_worker_serves_every_kind, owned.gateway, deployments), lambda served: served, seconds=90
|
||||
)
|
||||
clients: Final = _clients(owned.gateway, deployments)
|
||||
workers: Final = eventually(
|
||||
partial(_alive_workers, owned.process), lambda found: len(found) == 2, seconds=60
|
||||
)
|
||||
record_property("c02_process_tree_before_kill", _process_tree(owned.process))
|
||||
victim: Final = workers[0]
|
||||
calls: Final = _burst_calls(_CHAT_KINDS, 15)
|
||||
observations: Final = _Observations(slot.url)
|
||||
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
|
||||
futures: Final = _submit(pool, clients, calls)
|
||||
eventually(partial(_done_count, futures), lambda done: done >= 1, seconds=60)
|
||||
victim.kill()
|
||||
eventually(partial(_worker_alive, victim), lambda alive: not alive, seconds=10)
|
||||
done_at_kill: Final = _done_count(futures)
|
||||
survivor: Final = eventually(
|
||||
lambda: _attempt(clients, _Call("chat", f"after-kill-{uuid.uuid4().hex}")),
|
||||
lambda outcome: isinstance(outcome, _Completed),
|
||||
seconds=30,
|
||||
)
|
||||
outcomes: Final = tuple(future.result(timeout=180) for future in futures)
|
||||
record_property("c02_burst_size", len(calls))
|
||||
record_property("c02_killed_worker_pid", victim.pid)
|
||||
record_property("c02_done_at_kill", done_at_kill)
|
||||
assert isinstance(survivor, _Completed), survivor
|
||||
completed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Completed))
|
||||
failed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, _Failed))
|
||||
record_property("c02_completed_by_kind", _kind_counts(completed))
|
||||
record_property("c02_inflight_failure_count", len(failed))
|
||||
record_property("c02_inflight_failure_errors", str(tuple(outcome.error for outcome in failed)))
|
||||
assert completed, outcomes
|
||||
respawned: Final = eventually(
|
||||
partial(_alive_workers, owned.process),
|
||||
lambda found: len(found) == 2 and victim.pid not in {worker.pid for worker in found},
|
||||
seconds=graceful_stop_seconds() + 60,
|
||||
)
|
||||
record_property("c02_process_tree_after_respawn", _process_tree(owned.process))
|
||||
record_property("c02_respawned_worker_pids", str(tuple(worker.pid for worker in respawned)))
|
||||
after_respawn: Final = eventually(
|
||||
lambda: _attempt(clients, _Call("chat", f"after-respawn-{uuid.uuid4().hex}")),
|
||||
lambda outcome: isinstance(outcome, _Completed),
|
||||
seconds=60,
|
||||
)
|
||||
assert isinstance(after_respawn, _Completed), after_respawn
|
||||
_landed_once((survivor.response_id, after_respawn.response_id))
|
||||
response_ids: Final = tuple(outcome.response_id for outcome in completed)
|
||||
assert len(set(response_ids)) == len(completed), response_ids
|
||||
spend_row_counts: Final = tuple((outcome, len(_spend_rows(outcome.response_id))) for outcome in completed)
|
||||
landed: Final = tuple(outcome for outcome, count in spend_row_counts if count == 1)
|
||||
lost: Final = tuple(outcome for outcome, count in spend_row_counts if count == 0)
|
||||
assert len(landed) + len(lost) == len(completed), spend_row_counts
|
||||
record_property("c02_spend_query", _SPEND_QUERY)
|
||||
record_property("c02_burst_ids_landed_once", len(landed))
|
||||
record_property("c02_burst_ids_lost", len(lost))
|
||||
record_property("c02_burst_ids_lost_by_kind", _kind_counts(lost))
|
||||
completed_markers: Final = tuple(outcome.call.marker for outcome in completed)
|
||||
eventually(
|
||||
observations.read,
|
||||
lambda items: _marker_counts(items, completed_markers) == (1,) * len(completed_markers),
|
||||
seconds=60,
|
||||
)
|
||||
failed_markers: Final = tuple(outcome.call.marker for outcome in failed)
|
||||
record_property(
|
||||
"c02_failed_calls_seen_upstream",
|
||||
sum(count == 1 for count in _marker_counts(observations.items, failed_markers)),
|
||||
)
|
||||
488
tests/integration/sdk/test_openai_sdk_client_sessions_wire.py
Normal file
488
tests/integration/sdk/test_openai_sdk_client_sessions_wire.py
Normal file
|
|
@ -0,0 +1,488 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import struct
|
||||
import threading
|
||||
import zlib
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping
|
||||
from queue import SimpleQueue
|
||||
from typing import Final, TypeVar
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.vertex import service_account_json
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.utils import ImageResponse, ModelResponse, TextCompletionResponse
|
||||
|
||||
_R: Final = TypeVar("_R")
|
||||
|
||||
_PROVIDER_KEY: Final = "sk-scripted-provider"
|
||||
_API_VERSION: Final = "2024-10-21"
|
||||
_GATEWAY_PATH: Final = "/gateway.ai.cloudflare.com/v1/scripted-account/scripted-gateway/azure-openai/scripted-resource"
|
||||
_DEPLOYMENT: Final = "gpt-4o-mini-gateway"
|
||||
_ANSWER: Final = "wire answer"
|
||||
_IMAGE_URL: Final = "https://images.example.invalid/variation.png"
|
||||
|
||||
|
||||
def _png_chunk(kind: bytes, data: bytes) -> bytes:
|
||||
return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", zlib.crc32(kind + data))
|
||||
|
||||
|
||||
_PNG: Final = (
|
||||
b"\x89PNG\r\n\x1a\n"
|
||||
+ _png_chunk(b"IHDR", struct.pack(">IIBBBBB", 1, 1, 8, 2, 0, 0, 0))
|
||||
+ _png_chunk(b"IDAT", zlib.compress(b"\x00\x00\x00\x00"))
|
||||
+ _png_chunk(b"IEND", b"")
|
||||
)
|
||||
|
||||
|
||||
def _json(body: Mapping[str, JsonValue]) -> Reply:
|
||||
return Reply(body=json.dumps(body).encode())
|
||||
|
||||
|
||||
def _chat_completion() -> Mapping[str, JsonValue]:
|
||||
return {
|
||||
"id": "chatcmpl-wire",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": _ANSWER}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
def _text_completion() -> Mapping[str, JsonValue]:
|
||||
return {
|
||||
"id": "cmpl-wire",
|
||||
"object": "text_completion",
|
||||
"created": 1,
|
||||
"model": "gpt-3.5-turbo-instruct",
|
||||
"choices": [{"index": 0, "text": _ANSWER, "finish_reason": "stop", "logprobs": None}],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
def _peer(expected_path: str, body: Mapping[str, JsonValue]) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert urlsplit(request.target).path == expected_path, request.target
|
||||
return _json(body)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _held_peer(gate: threading.Event) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert gate.wait(timeout=10), "the timeout cell never released its peer"
|
||||
return _json(_chat_completion())
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _only_request(wire: Wire) -> Request:
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, received
|
||||
return received[0]
|
||||
|
||||
|
||||
def _drain(seen: SimpleQueue[str]) -> tuple[str, ...]:
|
||||
return tuple(seen.get_nowait() for _ in range(seen.qsize()))
|
||||
|
||||
|
||||
def _content(response: ModelResponse) -> str | None:
|
||||
choice: Final = response.choices[0]
|
||||
return choice.message.content if isinstance(choice, litellm.Choices) else None
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sync_session(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[httpx.Client, SimpleQueue[str]]]:
|
||||
seen: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
def record(request: httpx.Request) -> None:
|
||||
seen.put(str(request.url))
|
||||
|
||||
with httpx.Client(event_hooks={"request": [record]}) as client:
|
||||
monkeypatch.setattr(litellm, "client_session", client)
|
||||
monkeypatch.setattr(litellm, "aclient_session", None)
|
||||
yield client, seen
|
||||
|
||||
|
||||
def _run_with_async_session(
|
||||
monkeypatch: pytest.MonkeyPatch, call: Callable[[], Awaitable[_R]]
|
||||
) -> tuple[_R, tuple[str, ...]]:
|
||||
seen: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
async def record(request: httpx.Request) -> None:
|
||||
seen.put(str(request.url))
|
||||
|
||||
async def run() -> _R:
|
||||
async with httpx.AsyncClient(event_hooks={"request": [record]}) as client:
|
||||
monkeypatch.setattr(litellm, "client_session", None)
|
||||
monkeypatch.setattr(litellm, "aclient_session", client)
|
||||
return await call()
|
||||
|
||||
result: Final = asyncio.run(run())
|
||||
return result, _drain(seen)
|
||||
|
||||
|
||||
def _deployment(timeout: httpx.Timeout | openai.Timeout, api_base: str, deployment_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": api_base,
|
||||
"api_key": _PROVIDER_KEY,
|
||||
"timeout": timeout,
|
||||
"max_retries": 0,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
|
||||
def test_f01_async_image_variation_with_only_a_sync_session_set_builds_its_own_async_client(
|
||||
sync_session: tuple[httpx.Client, SimpleQueue[str]],
|
||||
) -> None:
|
||||
with wire_server(_peer("/images/variations", {"created": 1, "data": [{"url": _IMAGE_URL}]})) as wire:
|
||||
response: Final = asyncio.run(
|
||||
litellm.aimage_variation(
|
||||
image=("probe.png", _PNG, "image/png"),
|
||||
model="dall-e-2",
|
||||
custom_llm_provider="openai",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
num_retries=0,
|
||||
)
|
||||
)
|
||||
assert isinstance(response, ImageResponse), type(response)
|
||||
assert response.data is not None and response.data[0].url == _IMAGE_URL, response
|
||||
request: Final = _only_request(wire)
|
||||
assert request.method == "POST", request.method
|
||||
assert _PNG in request.body
|
||||
assert request.headers.get("authorization") == f"Bearer {_PROVIDER_KEY}", request.headers
|
||||
assert _drain(sync_session[1]) == ()
|
||||
|
||||
|
||||
def test_f02_async_azure_cloudflare_gateway_call_with_only_a_sync_session_set_builds_its_own_async_client(
|
||||
sync_session: tuple[httpx.Client, SimpleQueue[str]],
|
||||
) -> None:
|
||||
with wire_server(_peer(f"{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions", _chat_completion())) as wire:
|
||||
response: Final = asyncio.run(
|
||||
litellm.acompletion(
|
||||
model=f"azure/{_DEPLOYMENT}",
|
||||
messages=[{"role": "user", "content": "via the gateway"}],
|
||||
api_base=f"{wire.url}{_GATEWAY_PATH}",
|
||||
api_key=_PROVIDER_KEY,
|
||||
api_version=_API_VERSION,
|
||||
num_retries=0,
|
||||
)
|
||||
)
|
||||
assert isinstance(response, ModelResponse), type(response)
|
||||
assert _content(response) == _ANSWER, response
|
||||
request: Final = _only_request(wire)
|
||||
assert urlsplit(request.target).query == f"api-version={_API_VERSION}", request.target
|
||||
assert request.headers.get("api-key") == _PROVIDER_KEY, request.headers
|
||||
assert _drain(sync_session[1]) == ()
|
||||
|
||||
|
||||
def test_f03_async_text_completion_uses_the_callers_async_session(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
with wire_server(_peer("/completions", _text_completion())) as wire:
|
||||
response, seen = _run_with_async_session(
|
||||
monkeypatch,
|
||||
lambda: litellm.atext_completion(
|
||||
model="gpt-3.5-turbo-instruct",
|
||||
prompt="complete this",
|
||||
custom_llm_provider="openai",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
num_retries=0,
|
||||
),
|
||||
)
|
||||
assert isinstance(response, TextCompletionResponse), type(response)
|
||||
assert response.choices[0].text == _ANSWER, response
|
||||
assert _only_request(wire).method == "POST"
|
||||
assert seen == (f"{wire.url}/completions",)
|
||||
|
||||
|
||||
def test_f04_sync_text_completion_uses_the_callers_sync_session(
|
||||
sync_session: tuple[httpx.Client, SimpleQueue[str]],
|
||||
) -> None:
|
||||
with wire_server(_peer("/completions", _text_completion())) as wire:
|
||||
response: Final = litellm.text_completion(
|
||||
model="gpt-3.5-turbo-instruct",
|
||||
prompt="complete this",
|
||||
custom_llm_provider="openai",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, TextCompletionResponse), type(response)
|
||||
assert response.choices[0].text == _ANSWER, response
|
||||
assert _only_request(wire).method == "POST"
|
||||
assert _drain(sync_session[1]) == (f"{wire.url}/completions",)
|
||||
|
||||
|
||||
def test_f05_sync_moderation_reaches_the_peer_and_returns_its_verdict() -> None:
|
||||
verdict: Final[Mapping[str, JsonValue]] = {
|
||||
"id": "modr-wire",
|
||||
"model": "omni-moderation-latest",
|
||||
"results": [{"flagged": True, "categories": {"harassment": True}, "category_scores": {"harassment": 0.91}}],
|
||||
}
|
||||
with wire_server(_peer("/moderations", verdict)) as wire:
|
||||
response: Final = litellm.moderation(
|
||||
input="moderate this",
|
||||
model="omni-moderation-latest",
|
||||
api_key=_PROVIDER_KEY,
|
||||
api_base=wire.url,
|
||||
)
|
||||
assert response.results[0].flagged is True, response
|
||||
request: Final = _only_request(wire)
|
||||
assert json.loads(request.body) == {"input": "moderate this", "model": "omni-moderation-latest"}, request.body
|
||||
|
||||
|
||||
@pytest.mark.parametrize("timeout_type", (httpx.Timeout, openai.Timeout), ids=("httpx", "openai"))
|
||||
def test_f06_sync_completion_accepts_a_structured_timeout(
|
||||
timeout_type: type[httpx.Timeout] | type[openai.Timeout],
|
||||
) -> None:
|
||||
with wire_server(_peer("/chat/completions", _chat_completion())) as wire:
|
||||
response: Final = litellm.completion(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "with a structured timeout"}],
|
||||
custom_llm_provider="openai",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
timeout=timeout_type(5.0, connect=2.0),
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, ModelResponse), type(response)
|
||||
assert _content(response) == _ANSWER, response
|
||||
assert _only_request(wire).method == "POST"
|
||||
|
||||
|
||||
def test_f07_async_azure_cloudflare_gateway_call_uses_the_callers_async_session(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
with wire_server(_peer(f"{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions", _chat_completion())) as wire:
|
||||
response, seen = _run_with_async_session(
|
||||
monkeypatch,
|
||||
lambda: litellm.acompletion(
|
||||
model=f"azure/{_DEPLOYMENT}",
|
||||
messages=[{"role": "user", "content": "via the gateway"}],
|
||||
api_base=f"{wire.url}{_GATEWAY_PATH}",
|
||||
api_key=_PROVIDER_KEY,
|
||||
api_version=_API_VERSION,
|
||||
num_retries=0,
|
||||
),
|
||||
)
|
||||
assert isinstance(response, ModelResponse), type(response)
|
||||
assert _content(response) == _ANSWER, response
|
||||
assert _only_request(wire).headers.get("api-key") == _PROVIDER_KEY
|
||||
assert seen == (f"{wire.url}{_GATEWAY_PATH}/{_DEPLOYMENT}/chat/completions?api-version={_API_VERSION}",)
|
||||
|
||||
|
||||
def test_f08_sync_image_variation_uses_the_callers_sync_session(
|
||||
sync_session: tuple[httpx.Client, SimpleQueue[str]],
|
||||
) -> None:
|
||||
with wire_server(_peer("/images/variations", {"created": 1, "data": [{"url": _IMAGE_URL}]})) as wire:
|
||||
response: Final = litellm.image_variation(
|
||||
image=("probe.png", _PNG, "image/png"),
|
||||
model="dall-e-2",
|
||||
custom_llm_provider="openai",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, ImageResponse), type(response)
|
||||
assert response.data is not None and response.data[0].url == _IMAGE_URL, response
|
||||
assert _PNG in _only_request(wire).body
|
||||
assert _drain(sync_session[1]) == (f"{wire.url}/images/variations",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("timeout_type", (httpx.Timeout, openai.Timeout), ids=("httpx", "openai"))
|
||||
def test_f09_router_deployment_keeps_a_structured_timeout_as_httpx_and_serves(
|
||||
timeout_type: type[httpx.Timeout] | type[openai.Timeout],
|
||||
) -> None:
|
||||
with wire_server(_peer("/chat/completions", _chat_completion())) as wire:
|
||||
router: Final = _deployment(timeout_type(5.0, connect=2.0), wire.url, "deployment-f09")
|
||||
response: Final = router.completion(
|
||||
model="gpt-4o-mini", messages=[{"role": "user", "content": "through the router"}]
|
||||
)
|
||||
assert isinstance(response, ModelResponse), type(response)
|
||||
assert _content(response) == _ANSWER, response
|
||||
assert _only_request(wire).method == "POST"
|
||||
deployment: Final = router.get_deployment("deployment-f09")
|
||||
assert deployment is not None
|
||||
stored: Final = deployment.litellm_params.timeout
|
||||
assert isinstance(stored, httpx.Timeout), type(stored)
|
||||
assert (stored.read, stored.connect) == (5.0, 2.0), stored
|
||||
|
||||
|
||||
def test_f10_router_deployment_structured_read_timeout_fires_once_as_a_408() -> None:
|
||||
gate: Final = threading.Event()
|
||||
with wire_server(_held_peer(gate)) as wire:
|
||||
router: Final = _deployment(httpx.Timeout(0.5, connect=2.0), wire.url, "deployment-f10")
|
||||
with pytest.raises(litellm.Timeout) as raised:
|
||||
router.completion(model="gpt-4o-mini", messages=[{"role": "user", "content": "held upstream"}])
|
||||
gate.set()
|
||||
assert raised.value.status_code == 408, raised.value
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
_SPEECH_AUDIO: Final = b"OggS" + bytes(range(60))
|
||||
_SPEECH_AUDIO_REPLY: Final = Reply(body=_SPEECH_AUDIO, content_type="audio/ogg")
|
||||
_ELEVENLABS_VOICE: Final = "21m00Tcm4TlvDq8ikWAM"
|
||||
_VERTEX_PROJECT: Final = "scripted-project"
|
||||
_SPEECH_ROUTES: Final[Mapping[str, tuple[Mapping[str, str], str, Reply]]] = {
|
||||
"openai": (
|
||||
{"model": "openai/gpt-4o-mini-tts", "voice": "alloy", "api_key": _PROVIDER_KEY},
|
||||
"/audio/speech",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
"azure": (
|
||||
{"model": "azure/tts-deployment", "voice": "alloy", "api_key": _PROVIDER_KEY, "api_version": _API_VERSION},
|
||||
"/openai/deployments/tts-deployment/audio/speech",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
"azure_ava": (
|
||||
{"model": "azure/speech/azure-tts", "voice": "alloy", "api_key": _PROVIDER_KEY},
|
||||
"/cognitiveservices/v1",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
"elevenlabs": (
|
||||
{"model": "elevenlabs/eleven_multilingual_v2", "voice": _ELEVENLABS_VOICE, "api_key": _PROVIDER_KEY},
|
||||
f"/v1/text-to-speech/{_ELEVENLABS_VOICE}",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
"edenai": (
|
||||
{"model": "edenai/openai", "voice": "alloy", "api_key": _PROVIDER_KEY},
|
||||
"/audio/speech",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
"minimax": (
|
||||
{"model": "minimax/speech-02-hd", "voice": "alloy", "api_key": _PROVIDER_KEY},
|
||||
"/v1/t2a_v2",
|
||||
_json({"data": {"audio": _SPEECH_AUDIO.hex()}, "base_resp": {"status_code": 0, "status_msg": "success"}}),
|
||||
),
|
||||
"mistral": (
|
||||
{"model": "mistral/voxtral-mini-tts-2603", "voice": "alloy", "api_key": _PROVIDER_KEY},
|
||||
"/v1/audio/speech",
|
||||
_json({"audio_data": base64.b64encode(_SPEECH_AUDIO).decode()}),
|
||||
),
|
||||
"aws_polly": (
|
||||
{
|
||||
"model": "aws_polly/neural",
|
||||
"voice": "Joanna",
|
||||
"aws_access_key_id": "AKIASCRIPTED",
|
||||
"aws_secret_access_key": "scripted-secret",
|
||||
"aws_region_name": "us-east-1",
|
||||
},
|
||||
"/v1/speech",
|
||||
_SPEECH_AUDIO_REPLY,
|
||||
),
|
||||
}
|
||||
_PAYMENT_REQUIRED_MODELS: Final[Mapping[str, type[litellm.BadRequestError]]] = {
|
||||
"openai/gpt-4o-mini": litellm.BadRequestError,
|
||||
"anthropic/claude-sonnet-4-5": litellm.PaymentRequiredError,
|
||||
}
|
||||
|
||||
|
||||
def _speech_peer(expected_path: str, reply: Reply) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert urlsplit(request.target).path == expected_path, request.target
|
||||
return reply
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _vertex_speech_peer(request: Request) -> Reply:
|
||||
if urlsplit(request.target).path == "/_oauth/token":
|
||||
return _json({"access_token": "scripted-vertex-token", "expires_in": 3600, "token_type": "Bearer"})
|
||||
return _json({"audioContent": base64.b64encode(_SPEECH_AUDIO).decode()})
|
||||
|
||||
|
||||
def _payment_required_peer(request: Request) -> Reply:
|
||||
return Reply(
|
||||
status=402,
|
||||
body=json.dumps({"error": {"message": "scripted 402", "type": "invalid_request_error"}}).encode(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", tuple(_SPEECH_ROUTES))
|
||||
def test_f11_sync_speech_accepts_an_sdk_timeout_on_every_provider_branch(provider: str) -> None:
|
||||
route, expected_path, reply = _SPEECH_ROUTES[provider]
|
||||
with wire_server(_speech_peer(expected_path, reply)) as wire:
|
||||
response: Final = litellm.speech(
|
||||
input="speak this",
|
||||
api_base=wire.url,
|
||||
timeout=openai.Timeout(5.0, connect=2.0),
|
||||
**route,
|
||||
)
|
||||
assert response.content == _SPEECH_AUDIO, response.content[:16]
|
||||
assert _only_request(wire).method == "POST"
|
||||
|
||||
|
||||
def test_f13_sync_vertex_speech_accepts_an_sdk_timeout() -> None:
|
||||
with wire_server(_vertex_speech_peer) as wire:
|
||||
response: Final = litellm.speech(
|
||||
model="vertex_ai/chirp",
|
||||
voice="alloy",
|
||||
input="speak this",
|
||||
api_base=wire.url,
|
||||
vertex_credentials=service_account_json(_VERTEX_PROJECT, wire.url),
|
||||
vertex_project=_VERTEX_PROJECT,
|
||||
vertex_location="us-central1",
|
||||
timeout=openai.Timeout(5.0, connect=2.0),
|
||||
)
|
||||
assert response.content == _SPEECH_AUDIO, response.content[:16]
|
||||
assert tuple((seen.method, urlsplit(seen.target).path) for seen in wire.drain()) == (
|
||||
("POST", "/_oauth/token"),
|
||||
("POST", "/"),
|
||||
)
|
||||
|
||||
|
||||
def _runwayml_rejecting_peer(request: Request) -> Reply:
|
||||
assert urlsplit(request.target).path == "/v1/text_to_speech", request.target
|
||||
return Reply(status=400, body=json.dumps({"error": "scripted 400"}).encode())
|
||||
|
||||
|
||||
def test_f14_sync_runwayml_speech_sends_its_task_with_an_sdk_timeout() -> None:
|
||||
with wire_server(_runwayml_rejecting_peer) as wire:
|
||||
with pytest.raises(BaseLLMException) as raised:
|
||||
litellm.speech(
|
||||
model="runwayml/eleven_multilingual_v2",
|
||||
voice="Maya",
|
||||
input="speak this",
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
timeout=openai.Timeout(5.0, connect=2.0),
|
||||
)
|
||||
assert raised.value.status_code == 400, raised.value
|
||||
assert "scripted 400" in raised.value.message, raised.value.message
|
||||
assert _only_request(wire).method == "POST"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", tuple(_PAYMENT_REQUIRED_MODELS))
|
||||
def test_f12_sync_completion_maps_an_upstream_402_with_its_message(model: str) -> None:
|
||||
with wire_server(_payment_required_peer) as wire:
|
||||
with pytest.raises(_PAYMENT_REQUIRED_MODELS[model]) as raised:
|
||||
litellm.completion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "out of credit"}],
|
||||
api_base=wire.url,
|
||||
api_key=_PROVIDER_KEY,
|
||||
num_retries=0,
|
||||
max_retries=0,
|
||||
)
|
||||
assert raised.value.status_code == 402, raised.value
|
||||
assert "scripted 402" in raised.value.message, raised.value.message
|
||||
assert len(wire.drain()) == 1
|
||||
66
tests/typing/openai_sdk_compat.py
Normal file
66
tests/typing/openai_sdk_compat.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai import DefaultAsyncHttpxClient, DefaultHttpxClient
|
||||
from openai import HttpxBinaryResponseContent as SDKBinaryResponse
|
||||
from openai import Timeout as SDKTimeout
|
||||
from openai._types import Response as SDKResponse
|
||||
from typing_extensions import assert_type
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import (
|
||||
APIConnectionError,
|
||||
APIResponseValidationError,
|
||||
BadGatewayError,
|
||||
InternalServerError,
|
||||
InvalidRequestError,
|
||||
RateLimitError,
|
||||
Timeout,
|
||||
)
|
||||
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
from litellm.types.router import GenericLiteLLMParams, LiteLLMParamsTypedDict
|
||||
|
||||
|
||||
def consume_binary_response(response: HttpxBinaryResponseContent) -> httpx.Response:
|
||||
assert_type(response.response, httpx.Response)
|
||||
assert_type(response.response.request, httpx.Request)
|
||||
return response.response
|
||||
|
||||
|
||||
def binary_response_types(legacy_response: httpx.Response, sdk_response: SDKResponse) -> None:
|
||||
legacy: Final = HttpxBinaryResponseContent(legacy_response)
|
||||
native: Final = HttpxBinaryResponseContent(sdk_response)
|
||||
|
||||
assert_type(legacy.response, httpx.Response)
|
||||
assert_type(native.response, SDKResponse)
|
||||
assert_type(legacy.read(), bytes)
|
||||
assert_type(native.read(), bytes)
|
||||
assert_type(consume_binary_response(legacy), httpx.Response)
|
||||
|
||||
sdk_wrapper: Final[SDKBinaryResponse] = legacy
|
||||
assert_type(sdk_wrapper.read(), bytes)
|
||||
|
||||
|
||||
def synthesized_response_types(
|
||||
error: RateLimitError | BadGatewayError | InternalServerError | APIResponseValidationError | InvalidRequestError,
|
||||
) -> None:
|
||||
assert_type(error.response, httpx.Response)
|
||||
assert_type(error.request, httpx.Request)
|
||||
|
||||
|
||||
def synthesized_request_types(error: APIConnectionError | Timeout) -> None:
|
||||
assert_type(error.request, httpx.Request)
|
||||
|
||||
|
||||
def configured_clients(
|
||||
sync_client: httpx.Client | DefaultHttpxClient,
|
||||
async_client: httpx.AsyncClient | DefaultAsyncHttpxClient,
|
||||
timeout: httpx.Timeout | SDKTimeout,
|
||||
) -> None:
|
||||
litellm.client_session = sync_client # test-quality-ok: [TQ005] Type-check-only fixture; never executed
|
||||
litellm.aclient_session = async_client # test-quality-ok: [TQ005] Type-check-only fixture; never executed
|
||||
assert_type(CompletionTimeout.normalize(timeout), httpx.Timeout)
|
||||
params: Final[LiteLLMParamsTypedDict] = {"timeout": 1.0}
|
||||
params["timeout"] = timeout
|
||||
GenericLiteLLMParams(timeout=timeout)
|
||||
8
tests/typing/pyrightconfig.json
Normal file
8
tests/typing/pyrightconfig.json
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
{
|
||||
"extends": "../../pyrightconfig.json",
|
||||
"include": ["openai_sdk_compat.py"],
|
||||
"exclude": [],
|
||||
"extraPaths": ["../.."],
|
||||
"pythonVersion": "3.10",
|
||||
"reportMissingImports": true
|
||||
}
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import logging
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai import Timeout as SDKTimeout
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.router import (
|
||||
|
|
@ -26,6 +29,13 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
def test_sdk_timeout_is_normalized_for_provider_clients() -> None:
|
||||
timeout: Final = SDKTimeout(connect=2.0, read=None, write=5.0, pool=7.0)
|
||||
params: Final = GenericLiteLLMParams(timeout=timeout)
|
||||
assert isinstance(params.timeout, httpx.Timeout)
|
||||
assert params.timeout.as_dict() == timeout.as_dict()
|
||||
|
||||
|
||||
def test_model_info_declares_mirrored_pricing_fields():
|
||||
"""The pricing keys Deployment mirrors onto model_info must be declared fields, not
|
||||
extras that only survive because ModelInfo sets extra="allow"."""
|
||||
|
|
|
|||
2
uv.lock
generated
2
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue