fix: validate Codex route inputs for type gates

This commit is contained in:
jibanez-staticduo 2026-10-01 10:55:16 +02:00
parent 8b036f636c
commit 90dabda182
No known key found for this signature in database
10 changed files with 74 additions and 46 deletions

View file

@ -9,6 +9,7 @@ if TYPE_CHECKING:
from litellm.images.utils import ImageEditRequestUtils
import httpx
from pydantic import TypeAdapter
import litellm
@ -398,7 +399,9 @@ def image_generation(
model=model,
prompt=prompt,
image_generation_provider_config=image_generation_config,
extra_headers=extra_headers,
extra_headers=TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
extra_headers
),
image_generation_optional_request_params=optional_params,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,

View file

@ -9,6 +9,7 @@ from types import MappingProxyType
from typing import Final, Literal, TypedDict, cast
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
@ -1867,10 +1868,13 @@ def calculate_image_response_cost_from_usage(
if cached_details is None:
return prompt_cost + completion_cost
catalog_model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
cached_text: Final = _get_token_detail_value(cached_details, "text_tokens") or 0
cached_image: Final = _get_token_detail_value(cached_details, "image_tokens") or 0
input_text_tokens: Final = _get_token_detail_value(input_tokens_details, "text_tokens") or 0
input_image_tokens: Final = _get_token_detail_value(input_tokens_details, "image_tokens") or 0
details_adapter: Final = TypeAdapter[object](object)
cached_token_details: Final = details_adapter.validate_python(cached_details)
input_token_details: Final = details_adapter.validate_python(input_tokens_details)
cached_text: Final = _get_token_detail_value(cached_token_details, "text_tokens") or 0
cached_image: Final = _get_token_detail_value(cached_token_details, "image_tokens") or 0
input_text_tokens: Final = _get_token_detail_value(input_token_details, "text_tokens") or 0
input_image_tokens: Final = _get_token_detail_value(input_token_details, "image_tokens") or 0
if not (0 <= cached_text <= input_text_tokens and 0 <= cached_image <= input_image_tokens):
raise ValueError("Image cached token counts exceed their input modality counts")
text_rate: Final = catalog_model_info.get("input_cost_per_token") or 0.0

View file

@ -5,6 +5,7 @@ from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence
from contextvars import ContextVar
from dataclasses import dataclass
from enum import Enum, auto
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
from pydantic import TypeAdapter
@ -284,27 +285,31 @@ class RealTimeStreaming:
nested: Final = message_obj["event"]
if nested.get("type") in ("response.completed", "response.incomplete", "response.failed"):
response: Final = nested.get("response")
# Retain billing evidence even when response content is excluded from logging.
stored: Final = (
message_obj
if self._should_store_message(message_obj)
else {
"type": "response.event",
"event": {
"type": nested["type"],
"response": {
**{
key: value
for key, value in response.items()
if key in ("id", "created_at", "model", "usage", "service_tier")
},
"output": [],
}
if isinstance(response, Mapping)
else None,
},
}
response_mapping: Final = (
TypeAdapter(Mapping[str, object]).validate_python(response)
if isinstance(response, Mapping)
else None
)
# Retain billing evidence even when response content is excluded from logging.
filtered: Final[OpenAILiveResponseEvent] = {
"type": "response.event",
"event": {
"type": nested["type"],
"response": {
**MappingProxyType(
{
key: value
for key, value in response_mapping.items()
if key in ("id", "created_at", "model", "usage", "service_tier")
}
),
"output": TypeAdapter(list[object]).validate_python(()),
}
if response_mapping is not None
else None,
},
}
stored: Final = message_obj if self._should_store_message(message_obj) else filtered
self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(stored))
return
if not self._should_store_message(message_obj):

View file

@ -3,7 +3,7 @@ from typing import Final
from urllib.parse import urlsplit
import httpx
from pydantic import BaseModel, Field
from pydantic import BaseModel, Field, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig
@ -69,7 +69,8 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire
if not routing_data:
raise ValueError("Direct call signaling requires a ChatGPT deployment")
routing: Final = ChatGPTCallRouting.model_validate(routing_data)
call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1]
location: Final = TypeAdapter(str).validate_python(response.headers.get("location", ""))
call_id: Final[str] = urlsplit(location).path.rstrip("/").rsplit("/", 1)[-1]
return CodexRealtimeCall(
call_id=call_id,
model=routing.model,

View file

@ -1921,7 +1921,7 @@ def _extract_model_candidates_from_request(
if uses_body_target_model_sources or not body_model:
_append_model_candidates(candidates, request_data.get("target_model_names"))
if uses_session_model:
_append_model_candidates(candidates, session_model)
_append_model_candidates(candidates, TypeAdapter[object](object).validate_python(session_model))
if uses_completion_model_sources and isinstance(request_data.get("completion"), dict):
_append_model_candidates(candidates, request_data["completion"].get("model"))

View file

@ -110,7 +110,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
local: Final = self.internal_usage_cache.dual_cache.in_memory_cache
remote: Final = self.internal_usage_cache.dual_cache.redis_cache
raw: Final[object] = local.get_cache(key)
current: Final = TypeAdapter(Mapping[str, int] | None).validate_python(raw)
current: Final = TypeAdapter[Mapping[str, int] | None](Mapping[str, int] | None).validate_python(raw)
updated: Final = (
{ # mutable-ok: shared cache counter dict
**current,

View file

@ -4253,7 +4253,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
effective_descriptors: Final = (
tuple(_without_parallel_limit(descriptor) for descriptor in descriptors)
if call_type == "_arealtime" and is_realtime_call_attachment(data.get("websocket"))
if call_type == "_arealtime"
and is_realtime_call_attachment(TypeAdapter[object](object).validate_python(data.get("websocket")))
else descriptors
)

View file

@ -112,7 +112,9 @@ async def _start_codex_supervisor(
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(
TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers"))
TypeAdapter[Mapping[str, object] | None](Mapping[str, object] | None).validate_python(
processed.get("extra_headers")
)
),
**configured_realtime_headers(call.extra_headers),
}
@ -290,11 +292,11 @@ async def process_codex_request(
proxy_logging_obj=server.proxy_logging_obj,
proxy_config=server.proxy_config,
llm_router=server.llm_router,
user_model=server.user_model,
user_temperature=server.user_temperature,
user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model),
user_temperature=TypeAdapter[float | None](float | None).validate_python(server.user_temperature),
user_request_timeout=server.user_request_timeout,
user_max_tokens=server.user_max_tokens,
user_api_base=server.user_api_base,
user_api_base=TypeAdapter[str | None](str | None).validate_python(server.user_api_base),
model=model,
route_type=route_type,
**(
@ -358,7 +360,9 @@ async def _create_codex_realtime_call(request: Request) -> Response:
try:
await can_key_call_resolved_model(
model=model,
llm_model_list=server.llm_model_list,
llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python(
server.llm_model_list
),
valid_token=auth,
llm_router=server.llm_router,
)
@ -388,7 +392,7 @@ async def _create_codex_realtime_call(request: Request) -> Response:
data=processed,
route_type="arealtime_calls",
llm_router=server.llm_router,
user_model=server.user_model,
user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model),
)
try:
response: Final = await result
@ -462,7 +466,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
call: Final = decode_call(token, f"Bearer {api_key}")
await can_key_call_resolved_model(
model=call.alias,
llm_model_list=server.llm_model_list,
llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python(
server.llm_model_list
),
valid_token=auth,
llm_router=server.llm_router,
)
@ -530,9 +536,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(
TypeAdapter(Mapping[str, object] | None).validate_python(
processed.get("extra_headers")
)
TypeAdapter[Mapping[str, object] | None](
Mapping[str, object] | None
).validate_python(processed.get("extra_headers"))
),
**configured_realtime_headers(call.extra_headers),
}

View file

@ -6,6 +6,8 @@ from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from pydantic import TypeAdapter
import litellm
from litellm.constants import (
AZURE_OPENAI_AUDIO_PROVIDERS,
@ -292,10 +294,13 @@ async def arealtime_calls(
)
if session is not None:
session = _with_resolved_session_model(session, model_name)
supplied_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
kwargs.get("extra_headers")
)
call_headers: Final = (
provider_config.get_realtime_calls_extra_headers(kwargs.get("extra_headers"))
provider_config.get_realtime_calls_extra_headers(supplied_headers)
if provider_config is not None
else kwargs.get("extra_headers")
else supplied_headers
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
@ -379,8 +384,10 @@ async def _arealtime(
For PROXY use only.
"""
headers = cast(dict | None, kwargs.get("headers"))
extra_headers: Final = cast(dict | None, kwargs.get("extra_headers"))
headers = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(kwargs.get("headers"))
extra_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
kwargs.get("extra_headers")
)
if headers is None:
headers = {}
if extra_headers is not None:
@ -428,6 +435,7 @@ async def _arealtime(
else None
)
if provider_handler is not None:
user_api_key_dict: Final = TypeAdapter[object](object).validate_python(kwargs.get("user_api_key_dict"))
await provider_handler.async_realtime(
model=model,
websocket=websocket,
@ -436,7 +444,7 @@ async def _arealtime(
api_key=api_key,
timeout=timeout,
query_params=query_params,
user_api_key_dict=kwargs.get("user_api_key_dict"),
user_api_key_dict=user_api_key_dict,
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif provider_config is not None:

View file

@ -93,7 +93,7 @@ from litellm.types.responses.main import (
from .base import CachedTokensDetails
FileContent = IO[bytes] | bytes | PathLike
FileContent = IO[bytes] | bytes | PathLike[str]
FileTypes = (
# file (or bytes)