mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: validate Codex route inputs for type gates
This commit is contained in:
parent
8b036f636c
commit
90dabda182
10 changed files with 74 additions and 46 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue