refactor(types): replace Any with proven types in 16 files (#45029)

* refactor(types): replace Any with proven types in 16 files

* fix(types): import TypedDict from typing_extensions for pydantic on 3.10

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 04:31:51 -07:00 • committed by GitHub
parent 498a3e1a67
commit cb138ba92f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 226 additions and 65 deletions

View file

@ -315,7 +315,7 @@ class LLMCachingHandler:
args = args or () args = args or ()
final_embedding_cached_response: EmbeddingResponse | None = None final_embedding_cached_response: EmbeddingResponse | None = None
embedding_all_elements_cache_hit: bool = False embedding_all_elements_cache_hit: bool = False
cached_result: Any | None = None cached_result: object | None = None
kwargs = kwargs.copy() kwargs = kwargs.copy()
######################################################### #########################################################
# Init cache timing metrics # Init cache timing metrics
@ -849,7 +849,7 @@ class LLMCachingHandler:
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs) self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
cached_result: Any | None = None cached_result: object | None = None
if call_type == CallTypes.aembedding.value: if call_type == CallTypes.aembedding.value:
if isinstance(new_kwargs["input"], str): if isinstance(new_kwargs["input"], str):
new_kwargs["input"] = [new_kwargs["input"]] new_kwargs["input"] = [new_kwargs["input"]]

View file

@ -11,7 +11,7 @@ from __future__ import annotations
import asyncio import asyncio
import importlib import importlib
from collections.abc import AsyncIterator, Mapping from collections.abc import AsyncIterator, Callable, Mapping
from dataclasses import dataclass from dataclasses import dataclass
from types import MappingProxyType, ModuleType from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@ -59,10 +59,10 @@ _MODEL_NODE = "model"
class DeepAgentsDeps: class DeepAgentsDeps:
"""The optional-dependency entrypoints this handler uses.""" """The optional-dependency entrypoints this handler uses."""
create_deep_agent: Any create_deep_agent: Callable[..., CompiledStateGraph]
chat_litellm: Any chat_litellm: Callable[..., BaseChatModel]
checkpointer_cls: Any checkpointer_cls: Callable[[], BaseCheckpointSaver]
command_cls: Any command_cls: Callable[..., Command]
subagent_defaults: Mapping[str, object] subagent_defaults: Mapping[str, object]
convert_to_openai_messages: Any convert_to_openai_messages: Any
backend: ModuleType backend: ModuleType
@ -112,7 +112,7 @@ class DeepAgentsHandler(BaseHarnessHandler):
def __init__(self, config: BaseHarnessConfig) -> None: def __init__(self, config: BaseHarnessConfig) -> None:
super().__init__(config) super().__init__(config)
self._deps: DeepAgentsDeps | None = None self._deps: DeepAgentsDeps | None = None
self._agent: Any = None self._agent: CompiledStateGraph | None = None
self._thread_id: str | None = None self._thread_id: str | None = None
self._skip_tools: frozenset[str] = frozenset() self._skip_tools: frozenset[str] = frozenset()
@ -221,7 +221,7 @@ class DeepAgentsHandler(BaseHarnessHandler):
if ctx.output is not None: if ctx.output is not None:
ctx.output_json = structured_json(values.get("structured_response")) ctx.output_json = structured_json(values.get("structured_response"))
def _require_agent(self) -> tuple[Any, DeepAgentsDeps]: def _require_agent(self) -> tuple[CompiledStateGraph, DeepAgentsDeps]:
if self._agent is None or self._deps is None: if self._agent is None or self._deps is None:
raise HarnessError("Deep Agents session is not started") raise HarnessError("Deep Agents session is not started")
return self._agent, self._deps return self._agent, self._deps

View file

@ -60,7 +60,7 @@ _GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value}) _USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value})
def _next_sync_or_exhausted(it: Any) -> object: def _next_sync_or_exhausted(it: Iterator[object]) -> object:
""" """
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration. Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.

View file

@ -1,9 +1,11 @@
import json import json
import time import time
from collections.abc import AsyncIterator, Iterator from collections.abc import AsyncIterator, Iterator, Sequence
from typing import TYPE_CHECKING, Any, Final from typing import TYPE_CHECKING, Any, Final
import httpx import httpx
from pydantic import ConfigDict, TypeAdapter, with_config
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm import litellm
from litellm.litellm_core_utils.prompt_templates.factory import cohere_messages_pt_v2 from litellm.litellm_core_utils.prompt_templates.factory import cohere_messages_pt_v2
@ -23,6 +25,35 @@ else:
LiteLLMLoggingObj = Any LiteLLMLoggingObj = Any
@with_config(ConfigDict(extra="allow", strict=True))
class _CohereToolCall(TypedDict):
name: NotRequired[ReadOnly[object]]
generation_id: NotRequired[ReadOnly[object]]
parameters: NotRequired[ReadOnly[object]]
@with_config(ConfigDict(extra="allow", strict=True))
class _CohereBilledUnits(TypedDict):
input_tokens: NotRequired[ReadOnly[int | float]]
output_tokens: NotRequired[ReadOnly[int | float]]
@with_config(ConfigDict(extra="allow", strict=True))
class _CohereMeta(TypedDict):
billed_units: NotRequired[ReadOnly[_CohereBilledUnits]]
@with_config(ConfigDict(extra="allow", strict=True))
class _CohereChatResponse(TypedDict):
text: ReadOnly[str | None]
citations: NotRequired[ReadOnly[object]]
tool_calls: NotRequired[ReadOnly[Sequence[_CohereToolCall] | None]]
meta: NotRequired[ReadOnly[_CohereMeta]]
_COHERE_CHAT_RESPONSE: Final = TypeAdapter(_CohereChatResponse)
class CohereError(BaseLLMException): class CohereError(BaseLLMException):
def __init__( def __init__(
self, self,
@ -231,7 +262,7 @@ class CohereChatConfig(BaseConfig):
json_mode: bool | None = None, json_mode: bool | None = None,
) -> ModelResponse: ) -> ModelResponse:
try: try:
raw_response_json: Final = raw_response.json() raw_response_json: Final = _COHERE_CHAT_RESPONSE.validate_python(raw_response.json())
model_response.choices[0].message.content = raw_response_json["text"] model_response.choices[0].message.content = raw_response_json["text"]
except Exception: except Exception:
raise CohereError(message=raw_response.text, status_code=raw_response.status_code) raise CohereError(message=raw_response.text, status_code=raw_response.status_code)

View file

@ -5,6 +5,8 @@ from datetime import datetime
from typing import Any, Final from typing import Any, Final
import httpx import httpx
from pydantic import ConfigDict, TypeAdapter, with_config
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import can_block_current_thread from litellm.litellm_core_utils.asyncify import can_block_current_thread
@ -25,6 +27,24 @@ DEFAULT_GITHUB_ACCESS_TOKEN_URL: Final = "https://github.com/login/oauth/access_
DEFAULT_GITHUB_API_KEY_URL: Final = "https://api.github.com/copilot_internal/v2/token" DEFAULT_GITHUB_API_KEY_URL: Final = "https://api.github.com/copilot_internal/v2/token"
@with_config(ConfigDict(extra="allow", strict=True))
class _DeviceCode(TypedDict):
device_code: ReadOnly[str]
user_code: ReadOnly[object]
verification_uri: ReadOnly[object]
@with_config(ConfigDict(extra="allow", strict=True))
class _AccessTokenPoll(TypedDict):
access_token: ReadOnly[NotRequired[str]]
error: ReadOnly[NotRequired[object]]
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
_DEVICE_CODE: Final = TypeAdapter(_DeviceCode)
_ACCESS_TOKEN_POLL: Final = TypeAdapter(_AccessTokenPoll)
class Authenticator: class Authenticator:
def __init__(self) -> None: def __init__(self) -> None:
"""Initialize the GitHub Copilot authenticator with configurable token paths.""" """Initialize the GitHub Copilot authenticator with configurable token paths."""
@ -226,7 +246,7 @@ class Authenticator:
return headers return headers
def _get_device_code(self) -> dict[str, str]: def _get_device_code(self) -> _DeviceCode:
""" """
Get a device code for GitHub authentication. Get a device code for GitHub authentication.
@ -246,7 +266,7 @@ class Authenticator:
json={"client_id": client_id, "scope": "read:user"}, json={"client_id": client_id, "scope": "read:user"},
) )
resp.raise_for_status() resp.raise_for_status()
resp_json: Final = resp.json() resp_json: Final = _JSON_OBJECT.validate_python(resp.json())
required_fields: Final = ["device_code", "user_code", "verification_uri"] required_fields: Final = ["device_code", "user_code", "verification_uri"]
if not all(field in resp_json for field in required_fields): if not all(field in resp_json for field in required_fields):
@ -256,7 +276,7 @@ class Authenticator:
status_code=400, status_code=400,
) )
return resp_json return _DEVICE_CODE.validate_python(resp_json)
except httpx.HTTPStatusError as e: except httpx.HTTPStatusError as e:
verbose_logger.error("HTTP error getting device code: %s", e) verbose_logger.error("HTTP error getting device code: %s", e)
raise GetDeviceCodeError( raise GetDeviceCodeError(
@ -307,12 +327,13 @@ class Authenticator:
}, },
) )
resp.raise_for_status() resp.raise_for_status()
resp_json = resp.json() resp_json = _JSON_OBJECT.validate_python(resp.json())
poll_result = _ACCESS_TOKEN_POLL.validate_python(resp_json)
if "access_token" in resp_json: if "access_token" in poll_result:
verbose_logger.info("Authentication successful!") verbose_logger.info("Authentication successful!")
return resp_json["access_token"] return poll_result["access_token"]
elif "error" in resp_json and resp_json.get("error") == "authorization_pending": elif "error" in poll_result and poll_result.get("error") == "authorization_pending":
verbose_logger.debug("Authorization pending (attempt %s/%s)", attempt + 1, max_attempts) verbose_logger.debug("Authorization pending (attempt %s/%s)", attempt + 1, max_attempts)
else: else:
verbose_logger.warning("Unexpected response: %s", resp_json) verbose_logger.warning("Unexpected response: %s", resp_json)

View file

@ -2,6 +2,7 @@ import uuid
from typing import TYPE_CHECKING, Any, Final from typing import TYPE_CHECKING, Any, Final
import httpx import httpx
from pydantic import ConfigDict, TypeAdapter
import litellm import litellm
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
@ -29,6 +30,7 @@ else:
LiteLLMLoggingObj = Any LiteLLMLoggingObj = Any
MANUS_API_BASE: Final = "https://api.manus.im" MANUS_API_BASE: Final = "https://api.manus.im"
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
@ -177,7 +179,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
original_response=raw_response.text, original_response=raw_response.text,
additional_args={"complete_input_dict": {}}, additional_args={"complete_input_dict": {}},
) )
raw_response_json: Final = raw_response.json() raw_response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
# Manus uses camelCase "createdAt" instead of snake_case "created_at" # Manus uses camelCase "createdAt" instead of snake_case "created_at"
if "createdAt" in raw_response_json and "created_at" not in raw_response_json: if "createdAt" in raw_response_json and "created_at" not in raw_response_json:
@ -269,7 +271,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
original_response=raw_response.text, original_response=raw_response.text,
additional_args={"complete_input_dict": {}}, additional_args={"complete_input_dict": {}},
) )
raw_response_json: Final = raw_response.json() raw_response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
# Manus uses camelCase "createdAt" instead of snake_case "created_at" # Manus uses camelCase "createdAt" instead of snake_case "created_at"
if "createdAt" in raw_response_json and "created_at" not in raw_response_json: if "createdAt" in raw_response_json and "created_at" not in raw_response_json:

View file

@ -4,7 +4,8 @@ from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, Final from typing import TYPE_CHECKING, Any, Final
from httpx._models import Headers, Response from httpx._models import Headers, Response
from pydantic import ConfigDict, ValidationError from pydantic import ConfigDict, TypeAdapter, ValidationError, with_config
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm import litellm
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
@ -76,6 +77,21 @@ class _OllamaGenerateReasoning(LiteLLMBaseModel):
return parse_content_for_reasoning(self.response) return parse_content_for_reasoning(self.response)
@with_config(ConfigDict(extra="allow", strict=True))
class _OllamaGenerateMessage(TypedDict):
content: ReadOnly[NotRequired[str]]
@with_config(ConfigDict(extra="allow", strict=True))
class _OllamaGenerateResponse(TypedDict):
message: ReadOnly[NotRequired[_OllamaGenerateMessage]]
prompt_eval_count: ReadOnly[NotRequired[int]]
eval_count: ReadOnly[NotRequired[int]]
_OLLAMA_GENERATE_RESPONSE: Final = TypeAdapter(_OllamaGenerateResponse)
class OllamaConfig(BaseConfig): class OllamaConfig(BaseConfig):
""" """
Reference: https://github.com/ollama/ollama/blob/main/docs/api.md#parameters Reference: https://github.com/ollama/ollama/blob/main/docs/api.md#parameters
@ -159,7 +175,7 @@ class OllamaConfig(BaseConfig):
system: str | None = None, system: str | None = None,
template: str | None = None, template: str | None = None,
) -> None: ) -> None:
locals_: Final = locals().copy() locals_: Final[dict[str, object]] = locals().copy()
for key, value in locals_.items(): for key, value in locals_.items():
if key != "self" and value is not None: if key != "self" and value is not None:
setattr(self.__class__, key, value) setattr(self.__class__, key, value)
@ -288,7 +304,7 @@ class OllamaConfig(BaseConfig):
api_key: str | None = None, api_key: str | None = None,
json_mode: bool | None = None, json_mode: bool | None = None,
) -> ModelResponse: ) -> ModelResponse:
response_json: Final = raw_response.json() response_json: Final = _OLLAMA_GENERATE_RESPONSE.validate_python(raw_response.json())
## RESPONSE OBJECT ## RESPONSE OBJECT
model_response.choices[0].finish_reason = "stop" model_response.choices[0].finish_reason = "stop"
if request_data.get("format", "") == "json": if request_data.get("format", "") == "json":
@ -303,7 +319,7 @@ class OllamaConfig(BaseConfig):
model_response.choices[0].finish_reason = "stop" model_response.choices[0].finish_reason = "stop"
else: else:
try: try:
response_content: Final = json.loads(response_text) response_content: Final[object] = json.loads(response_text)
# Check if this is a function call format with name/arguments structure # Check if this is a function call format with name/arguments structure
if ( if (
@ -356,7 +372,7 @@ class OllamaConfig(BaseConfig):
len(tokenizer.encode(_prompt, disallowed_special=())), len(tokenizer.encode(_prompt, disallowed_special=())),
) )
completion_tokens: Final = response_json.get( completion_tokens: Final = response_json.get(
"eval_count", len(response_json.get("message", dict()).get("content", "")) "eval_count", len((response_json.get("message") or {}).get("content", ""))
) )
setattr( setattr(
model_response, model_response,

View file

@ -4,6 +4,8 @@ import time
from collections.abc import Callable from collections.abc import Callable
from typing import Final from typing import Final
from pydantic import ConfigDict, TypeAdapter
import litellm import litellm
from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -20,6 +22,7 @@ from ..common_utils import ReplicateError
from .transformation import ReplicateConfig from .transformation import ReplicateConfig
replicate_config: Final = ReplicateConfig() replicate_config: Final = ReplicateConfig()
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
# Function to handle prediction response (streaming) # Function to handle prediction response (streaming)
@ -81,7 +84,7 @@ async def async_handle_prediction_response_streaming(
await asyncio.sleep(REPLICATE_POLLING_DELAY_SECONDS) # prevent being rate limited by replicate await asyncio.sleep(REPLICATE_POLLING_DELAY_SECONDS) # prevent being rate limited by replicate
response = await http_client.get(prediction_url, headers=headers) response = await http_client.get(prediction_url, headers=headers)
if response.status_code == 200: if response.status_code == 200:
response_data = response.json() response_data = _JSON_OBJECT.validate_python(response.json())
status = response_data.get("status", "") status = response_data.get("status", "")
# Check that "output" exists and is not None or empty # Check that "output" exists and is not None or empty
output_present = "output" in response_data and response_data["output"] is not None output_present = "output" in response_data and response_data["output"] is not None
@ -211,7 +214,7 @@ def completion(
litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry
) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing" ) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing"
response = httpx_client.get(url=prediction_url, headers=headers) response = httpx_client.get(url=prediction_url, headers=headers)
if response.status_code == 200 and response.json().get("status") in [ if response.status_code == 200 and _JSON_OBJECT.validate_python(response.json()).get("status") in [
"processing", "processing",
"starting", "starting",
]: ]:
@ -280,7 +283,7 @@ async def async_completion(
litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry
) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing" ) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing"
response = await async_handler.get(url=prediction_url, headers=headers) response = await async_handler.get(url=prediction_url, headers=headers)
if response.status_code == 200 and response.json().get("status") in [ if response.status_code == 200 and _JSON_OBJECT.validate_python(response.json()).get("status") in [
"processing", "processing",
"starting", "starting",
]: ]:

View file

@ -3209,7 +3209,7 @@ class ModelResponseIterator:
def _apply_stream_usage_metadata( def _apply_stream_usage_metadata(
self, self,
processed_chunk: GenerateContentResponseBody, processed_chunk: GenerateContentResponseBody,
model_response: Any, model_response: "ModelResponseStream",
grounding_metadata: list[dict], grounding_metadata: list[dict],
) -> Usage | None: ) -> Usage | None:
if "usageMetadata" not in processed_chunk: if "usageMetadata" not in processed_chunk:

View file

@ -2,7 +2,7 @@ import copy
import functools import functools
import os import os
import uuid import uuid
from collections.abc import AsyncGenerator, Iterator, Mapping, Sequence from collections.abc import AsyncGenerator, AsyncIterable, Iterator, Mapping, Sequence
from typing import ( from typing import (
TYPE_CHECKING, TYPE_CHECKING,
Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__
@ -10,9 +10,11 @@ from typing import (
Final, Final,
Literal, Literal,
Optional, Optional,
Protocol,
) )
import httpx import httpx
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException from litellm.exceptions import GuardrailRaisedException
@ -80,6 +82,8 @@ _VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}"
_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 _DEFAULT_TIMEOUT_SECONDS: Final = 10.0
_SHIELD_BODY: Final = TypeAdapter(dict[str, object])
def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, str]]: def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, str]]:
for call in calls: for call in calls:
@ -89,6 +93,15 @@ def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, st
yield function, arguments yield function, arguments
class _ContentDelta(Protocol):
content: object
class _FinishingDelta(Protocol):
content: object
tool_calls: object
class LLMShieldProxyGuardrail(CustomGuardrail): class LLMShieldProxyGuardrail(CustomGuardrail):
"""Redacts PII before it leaves the proxy and restores it in the response. """Redacts PII before it leaves the proxy and restores it in the response.
@ -215,7 +228,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
timeout=_DEFAULT_TIMEOUT_SECONDS, timeout=_DEFAULT_TIMEOUT_SECONDS,
) )
response.raise_for_status() response.raise_for_status()
return response.json() return _SHIELD_BODY.validate_python(response.json(), strict=True)
except httpx.HTTPStatusError as exc: except httpx.HTTPStatusError as exc:
verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path) verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path)
raise GuardrailRaisedException( raise GuardrailRaisedException(
@ -329,8 +342,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
self, self,
data: MutableRequest, data: MutableRequest,
user_api_key_dict: UserAPIKeyAuth, user_api_key_dict: UserAPIKeyAuth,
response: Any, response: object,
) -> Any: ) -> object:
"""Restores the original values in a copy of a non-streaming response. """Restores the original values in a copy of a non-streaming response.
The copy is what keeps plaintext out of the response cache. LiteLLM caches the The copy is what keeps plaintext out of the response cache. LiteLLM caches the
@ -345,8 +358,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
return response return response
response = detached(response) # rebind-ok: everything below restores the copy. response = detached(response) # rebind-ok: everything below restores the copy.
if self._is_anthropic_message_response(response): anthropic_body: Final = as_object(response)
return await self._restore_anthropic_response(response, data) if anthropic_body is not None and self._is_anthropic_message_response(anthropic_body):
return await self._restore_anthropic_response(anthropic_body, data)
response_slots: Final = self._responses_api_slots(response) response_slots: Final = self._responses_api_slots(response)
if response_slots: if response_slots:
@ -448,7 +462,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
async def async_post_call_streaming_iterator_hook( async def async_post_call_streaming_iterator_hook(
self, self,
user_api_key_dict: UserAPIKeyAuth, user_api_key_dict: UserAPIKeyAuth,
response: Any, response: AsyncIterable[object],
request_data: MutableRequest, request_data: MutableRequest,
) -> AsyncGenerator[Any, None]: ) -> AsyncGenerator[Any, None]:
"""Restores original values incrementally, without buffering the stream. """Restores original values incrementally, without buffering the stream.
@ -545,7 +559,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
async def _restore_content_window( async def _restore_content_window(
self, self,
delta: Any, delta: _ContentDelta,
key: CarryKey, key: CarryKey,
carries: CarryWindows, carries: CarryWindows,
session_id: str, session_id: str,
@ -597,7 +611,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
async def _flush_finished_choice( async def _flush_finished_choice(
self, self,
delta: Any, delta: _FinishingDelta,
choice_index: int, choice_index: int,
carries: CarryWindows, carries: CarryWindows,
session_id: str, session_id: str,

View file

@ -16,10 +16,11 @@ from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaita
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
import aiohttp import aiohttp
from typing_extensions import NotRequired, ReadOnly from pydantic import ConfigDict, TypeAdapter, with_config
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm import litellm
from litellm import get_secret from litellm import get_secret
@ -68,15 +69,22 @@ from litellm.utils import (
) )
@with_config(ConfigDict(extra="allow", strict=True))
class _PresidioAnonymizeItem(TypedDict, total=False): class _PresidioAnonymizeItem(TypedDict, total=False):
entity_type: ReadOnly[str | None] entity_type: ReadOnly[str | None]
@with_config(ConfigDict(extra="allow", strict=True))
class _PresidioAnonymizeResponse(TypedDict): class _PresidioAnonymizeResponse(TypedDict):
text: ReadOnly[str] text: ReadOnly[str]
items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]] items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
_PRESIDIO_ANONYMIZE_ADAPTER: Final[TypeAdapter[_PresidioAnonymizeResponse | None]] = TypeAdapter(
_PresidioAnonymizeResponse | None
)
class _JsonResponse(Protocol): class _JsonResponse(Protocol):
def json(self) -> Awaitable[object]: ... def json(self) -> Awaitable[object]: ...
@ -769,7 +777,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
raise Exception( raise Exception(
f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'" f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'"
) )
return await response.json() return _PRESIDIO_ANONYMIZE_ADAPTER.validate_python(await response.json())
def _finalize_presidio_anonymize_simple( def _finalize_presidio_anonymize_simple(
self, self,

View file

@ -28,6 +28,8 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from fastapi.exceptions import HTTPException from fastapi.exceptions import HTTPException
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict, Unpack
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import ( from litellm.integrations.custom_guardrail import (
@ -76,6 +78,12 @@ _DEFAULT_POLICIES: Final = [
"Default_Policy_GeneralPromptAttackProtection", "Default_Policy_GeneralPromptAttackProtection",
] ]
_RESPONSE_BODY: Final = TypeAdapter(dict[str, object])
class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
supported_event_hooks: ReadOnly[list[GuardrailEventHooks]]
class XecGuardMissingCredentials(Exception): class XecGuardMissingCredentials(Exception):
pass pass
@ -90,7 +98,7 @@ class XecGuardGuardrail(CustomGuardrail):
policy_names: list[str] | None = None, policy_names: list[str] | None = None,
block_on_error: bool | None = None, block_on_error: bool | None = None,
grounding_strictness: str | None = None, grounding_strictness: str | None = None,
**kwargs: Any, **kwargs: Unpack[_CustomGuardrailOptions],
) -> None: ) -> None:
self.api_key = api_key or os.environ.get("XECGUARD_API_KEY") self.api_key = api_key or os.environ.get("XECGUARD_API_KEY")
if not self.api_key: if not self.api_key:
@ -122,9 +130,12 @@ class XecGuardGuardrail(CustomGuardrail):
llm_provider=httpxSpecialProvider.GuardrailCallback, llm_provider=httpxSpecialProvider.GuardrailCallback,
) )
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) forwarded: Final[_CustomGuardrailOptions] = {
"supported_event_hooks": list(self.get_supported_event_hooks()),
**kwargs,
}
super().__init__(**kwargs) super().__init__(**forwarded)
@staticmethod @staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None: def get_config_model() -> type["GuardrailConfigModel"] | None:
@ -345,7 +356,7 @@ class XecGuardGuardrail(CustomGuardrail):
path: str, path: str,
payload: dict, payload: dict,
suppress_errors: bool = False, suppress_errors: bool = False,
) -> dict | None: ) -> dict[str, object] | None:
endpoint: Final = f"{self.api_base}{path}" endpoint: Final = f"{self.api_base}{path}"
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
"XecGuard: POST %s payload_keys=%s", "XecGuard: POST %s payload_keys=%s",
@ -363,7 +374,7 @@ class XecGuardGuardrail(CustomGuardrail):
timeout=self.timeout if self.timeout is not None else 10.0, timeout=self.timeout if self.timeout is not None else 10.0,
) )
response.raise_for_status() response.raise_for_status()
return response.json() return _RESPONSE_BODY.validate_python(response.json(), strict=True)
except Exception as exc: except Exception as exc:
verbose_proxy_logger.error("XecGuard API error: %s", str(exc)) verbose_proxy_logger.error("XecGuard API error: %s", str(exc))
if suppress_errors: if suppress_errors:

View file

@ -29,6 +29,9 @@ import json
from collections.abc import Mapping, Sequence from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Protocol from typing import TYPE_CHECKING, Any, Final, Protocol
from pydantic import ConfigDict, TypeAdapter, with_config
from typing_extensions import ReadOnly, TypedDict
import litellm import litellm
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache from litellm.caching.caching import DualCache
@ -81,6 +84,22 @@ class _ChatCompletion(Protocol):
def choices(self) -> Sequence[_ChatChoice]: ... def choices(self) -> Sequence[_ChatChoice]: ...
@with_config(ConfigDict(extra="allow", strict=True))
class _CodeExecutionToolArgs(TypedDict, total=False):
code: ReadOnly[str]
@with_config(ConfigDict(extra="allow", strict=True))
class _GeneratedFileEntry(TypedDict):
name: ReadOnly[str]
mime_type: ReadOnly[str]
content_base64: ReadOnly[str]
_TOOL_CALL_ARGS_ADAPTER: Final = TypeAdapter(_CodeExecutionToolArgs)
_EXEC_RESULT_FILES_ADAPTER: Final = TypeAdapter(Sequence[_GeneratedFileEntry])
def _first_choice(response: _ChatCompletion) -> _ChatChoice: def _first_choice(response: _ChatCompletion) -> _ChatChoice:
"""The first choice of an OpenAI shaped completion response.""" """The first choice of an OpenAI shaped completion response."""
return response.choices[0] return response.choices[0]
@ -620,7 +639,7 @@ class SkillsInjectionHook(CustomLogger):
# Collect generated files # Collect generated files
if exec_result.get("files"): if exec_result.get("files"):
files: Final[Sequence[Mapping[str, str]]] = exec_result["files"] files: Final = _EXEC_RESULT_FILES_ADAPTER.validate_python(exec_result["files"])
for f in files: for f in files:
generated_files.append( generated_files.append(
{ {
@ -833,7 +852,7 @@ print('No executable skill module found')
) -> str: ) -> str:
"""Execute a litellm_code_execution tool call and return result string.""" """Execute a litellm_code_execution tool call and return result string."""
try: try:
args: Final[Mapping[str, str]] = json.loads(tool_call.function.arguments) args: Final = _TOOL_CALL_ARGS_ADAPTER.validate_python(json.loads(tool_call.function.arguments))
code: Final[str] = args.get("code", "") code: Final[str] = args.get("code", "")
verbose_proxy_logger.debug("SkillsInjectionHook: Executing code (%s chars)", len(code)) verbose_proxy_logger.debug("SkillsInjectionHook: Executing code (%s chars)", len(code))
@ -849,7 +868,7 @@ print('No executable skill module found')
# Collect generated files # Collect generated files
if exec_result.get("files"): if exec_result.get("files"):
tool_result += "\n\nGenerated files:" tool_result += "\n\nGenerated files:"
files: Final[Sequence[Mapping[str, str]]] = exec_result["files"] files: Final = _EXEC_RESULT_FILES_ADAPTER.validate_python(exec_result["files"])
for f in files: for f in files:
file_content = base64.b64decode(f["content_base64"]) file_content = base64.b64decode(f["content_base64"])
generated_files.append( generated_files.append(

View file

@ -39,7 +39,6 @@ from typing import (
Optional, Optional,
Protocol, Protocol,
TypeAlias, TypeAlias,
TypedDict,
Union, Union,
cast, cast,
get_args, get_args,
@ -50,9 +49,9 @@ from typing import (
import anyio import anyio
import websockets import websockets
import websockets.exceptions import websockets.exceptions
from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError from pydantic import BaseModel, ConfigDict, Json, JsonValue, TypeAdapter, ValidationError, with_config
from pydantic.fields import FieldInfo, PydanticUndefined from pydantic.fields import FieldInfo, PydanticUndefined
from typing_extensions import NotRequired, ReadOnly, assert_never from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
from litellm._uuid import uuid from litellm._uuid import uuid
from litellm.constants import ( from litellm.constants import (
@ -1191,6 +1190,34 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N
_AiohttpAddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind, int, str, tuple[object, ...]] _AiohttpAddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind, int, str, tuple[object, ...]]
@with_config(ConfigDict(extra="allow", strict=True))
class _LoginRequestBody(TypedDict, total=False):
username: ReadOnly[object]
password: ReadOnly[object]
@with_config(ConfigDict(extra="allow", strict=True))
class _LoginExchangeRequestBody(TypedDict, total=False):
code: ReadOnly[object]
@with_config(ConfigDict(extra="allow", strict=True))
class _AnthropicBetaHeadersReloadConfig(TypedDict, total=False):
interval_hours: ReadOnly[int | float | None]
force_reload: ReadOnly[object]
@with_config(ConfigDict(extra="allow", strict=True))
class _WebsearchInterceptionLitellmSettings(TypedDict, total=False):
websearch_interception_params: ReadOnly[object]
_LOGIN_REQUEST_BODY: Final = TypeAdapter(_LoginRequestBody)
_LOGIN_EXCHANGE_REQUEST_BODY: Final = TypeAdapter(_LoginExchangeRequestBody)
_ANTHROPIC_BETA_HEADERS_RELOAD_CONFIG: Final = TypeAdapter(_AnthropicBetaHeadersReloadConfig)
_WEBSEARCH_INTERCEPTION_LITELLM_SETTINGS: Final = TypeAdapter(_WebsearchInterceptionLitellmSettings)
class _AiohttpConnectorKwargs(TypedDict, total=False): class _AiohttpConnectorKwargs(TypedDict, total=False):
keepalive_timeout: float keepalive_timeout: float
ttl_dns_cache: int ttl_dns_cache: int
@ -8471,9 +8498,10 @@ class ProxyConfig:
if config_record is None or config_record.param_value is None: if config_record is None or config_record.param_value is None:
return return
litellm_settings = config_record.param_value raw_litellm_settings: Final = config_record.param_value
if isinstance(litellm_settings, str): litellm_settings: Final = _WEBSEARCH_INTERCEPTION_LITELLM_SETTINGS.validate_python(
litellm_settings = json.loads(litellm_settings) json.loads(raw_litellm_settings) if isinstance(raw_litellm_settings, str) else raw_litellm_settings
)
websearch_config: Final = litellm_settings.get("websearch_interception_params", None) websearch_config: Final = litellm_settings.get("websearch_interception_params", None)
@ -8726,7 +8754,7 @@ class ProxyConfig:
if config_record is None or config_record.param_value is None: if config_record is None or config_record.param_value is None:
return # No configuration found, skip reload return # No configuration found, skip reload
config: Final = config_record.param_value config: Final = _ANTHROPIC_BETA_HEADERS_RELOAD_CONFIG.validate_python(config_record.param_value)
interval_hours: Final = config.get("interval_hours") interval_hours: Final = config.get("interval_hours")
force_reload: Final = config.get("force_reload", False) force_reload: Final = config.get("force_reload", False)
@ -17127,7 +17155,7 @@ async def login_v2(request: Request):
from litellm.proxy.utils import get_custom_url from litellm.proxy.utils import get_custom_url
try: try:
body: Final = await request.json() body: Final = _LOGIN_REQUEST_BODY.validate_python(await request.json())
username: Final = str(body.get("username")) username: Final = str(body.get("username"))
password: Final = str(body.get("password")) password: Final = str(body.get("password"))
@ -17199,7 +17227,7 @@ async def login_v3(request: Request):
code=status.HTTP_404_NOT_FOUND, code=status.HTTP_404_NOT_FOUND,
) )
body: Final = await request.json() body: Final = _LOGIN_REQUEST_BODY.validate_python(await request.json())
username: Final = str(body.get("username")) username: Final = str(body.get("username"))
password: Final = str(body.get("password")) password: Final = str(body.get("password"))
@ -17271,7 +17299,7 @@ async def login_v3_exchange(request: Request):
code=status.HTTP_404_NOT_FOUND, code=status.HTTP_404_NOT_FOUND,
) )
body: Final = await request.json() body: Final = _LOGIN_EXCHANGE_REQUEST_BODY.validate_python(await request.json())
code: Final = body.get("code") code: Final = body.get("code")
if not code: if not code:
raise ProxyException( raise ProxyException(

View file

@ -9335,10 +9335,10 @@ class Router:
if access_windows_error is not None: if access_windows_error is not None:
raise ValueError(access_windows_error) raise ValueError(access_windows_error)
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
merged_params: Final[Mapping[str, Any]] = ( merged_params: Final[Mapping[str, object]] = (
_litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing}) _litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing})
) )
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**merged_params) litellm_params: Final[LiteLLM_Params] = LiteLLM_Params.model_validate(dict(**merged_params))
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params) warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
deployment = Deployment( deployment = Deployment(
**deployment_info, **deployment_info,

View file

@ -45,7 +45,7 @@ from httpx import Proxy
from httpx._utils import get_environment_proxies from httpx._utils import get_environment_proxies
from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # OpenAI parser module is private from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # OpenAI parser module is private
from openai.types.chat.completion_create_params import ResponseFormat from openai.types.chat.completion_create_params import ResponseFormat
from pydantic import BaseModel, JsonValue from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, with_config
import litellm import litellm
import litellm.litellm_core_utils import litellm.litellm_core_utils
@ -5506,7 +5506,7 @@ def get_max_tokens(model: str) -> int | None:
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx) response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
# Parse the JSON response # Parse the JSON response
config_json: Final[Mapping[str, int]] = response.json() config_json: Final = _HUGGINGFACE_MODEL_CONFIG.validate_python(response.json())
# Extract and return the max_position_embeddings # Extract and return the max_position_embeddings
max_position_embeddings: Final = config_json.get("max_position_embeddings") max_position_embeddings: Final = config_json.get("max_position_embeddings")
if max_position_embeddings is not None: if max_position_embeddings is not None:
@ -5776,6 +5776,14 @@ def _check_provider_match(model_info: dict, custom_llm_provider: str | None) ->
from typing_extensions import ReadOnly, TypedDict from typing_extensions import ReadOnly, TypedDict
@with_config(ConfigDict(extra="allow", strict=True, hide_input_in_errors=True))
class _HuggingFaceModelConfig(TypedDict, total=False):
max_position_embeddings: ReadOnly[int | None]
_HUGGINGFACE_MODEL_CONFIG: Final = TypeAdapter(_HuggingFaceModelConfig)
class PotentialModelNamesAndCustomLLMProvider(TypedDict): class PotentialModelNamesAndCustomLLMProvider(TypedDict):
split_model: str split_model: str
combined_model_name: str combined_model_name: str
@ -5923,7 +5931,7 @@ def _get_max_position_embeddings(model_name: str) -> int | None:
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx) response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
# Parse the JSON response # Parse the JSON response
config_json: Final[Mapping[str, int]] = response.json() config_json: Final = _HUGGINGFACE_MODEL_CONFIG.validate_python(response.json())
# Extract and return the max_position_embeddings # Extract and return the max_position_embeddings
max_position_embeddings: Final = config_json.get("max_position_embeddings") max_position_embeddings: Final = config_json.get("max_position_embeddings")