diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 43d7bdb61b1..f9215864d1e 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -315,7 +315,7 @@ class LLMCachingHandler: args = args or () final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False - cached_result: Any | None = None + cached_result: object | None = None kwargs = kwargs.copy() ######################################################### # Init cache timing metrics @@ -849,7 +849,7 @@ class LLMCachingHandler: 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) 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 isinstance(new_kwargs["input"], str): new_kwargs["input"] = [new_kwargs["input"]] diff --git a/litellm/harness/handlers/deepagents_handler.py b/litellm/harness/handlers/deepagents_handler.py index 9ffaf52ee59..e6ca7c161c7 100644 --- a/litellm/harness/handlers/deepagents_handler.py +++ b/litellm/harness/handlers/deepagents_handler.py @@ -11,7 +11,7 @@ from __future__ import annotations import asyncio import importlib -from collections.abc import AsyncIterator, Mapping +from collections.abc import AsyncIterator, Callable, Mapping from dataclasses import dataclass from types import MappingProxyType, ModuleType from typing import TYPE_CHECKING, Any @@ -59,10 +59,10 @@ _MODEL_NODE = "model" class DeepAgentsDeps: """The optional-dependency entrypoints this handler uses.""" - create_deep_agent: Any - chat_litellm: Any - checkpointer_cls: Any - command_cls: Any + create_deep_agent: Callable[..., CompiledStateGraph] + chat_litellm: Callable[..., BaseChatModel] + checkpointer_cls: Callable[[], BaseCheckpointSaver] + command_cls: Callable[..., Command] subagent_defaults: Mapping[str, object] convert_to_openai_messages: Any backend: ModuleType @@ -112,7 +112,7 @@ class DeepAgentsHandler(BaseHarnessHandler): def __init__(self, config: BaseHarnessConfig) -> None: super().__init__(config) self._deps: DeepAgentsDeps | None = None - self._agent: Any = None + self._agent: CompiledStateGraph | None = None self._thread_id: str | None = None self._skip_tools: frozenset[str] = frozenset() @@ -221,7 +221,7 @@ class DeepAgentsHandler(BaseHarnessHandler): if ctx.output is not None: 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: raise HarnessError("Deep Agents session is not started") return self._agent, self._deps diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 335e33e1f1a..1027a681d56 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -60,7 +60,7 @@ _GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__) _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. diff --git a/litellm/llms/cohere/chat/transformation.py b/litellm/llms/cohere/chat/transformation.py index a26cdc81695..fdf44d5333e 100644 --- a/litellm/llms/cohere/chat/transformation.py +++ b/litellm/llms/cohere/chat/transformation.py @@ -1,9 +1,11 @@ import json import time -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Sequence from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter, with_config +from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm from litellm.litellm_core_utils.prompt_templates.factory import cohere_messages_pt_v2 @@ -23,6 +25,35 @@ else: 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): def __init__( self, @@ -231,7 +262,7 @@ class CohereChatConfig(BaseConfig): json_mode: bool | None = None, ) -> ModelResponse: 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"] except Exception: raise CohereError(message=raw_response.text, status_code=raw_response.status_code) diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 7821756bc16..83867708c46 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -5,6 +5,8 @@ from datetime import datetime from typing import Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter, with_config +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_logger 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" +@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: def __init__(self) -> None: """Initialize the GitHub Copilot authenticator with configurable token paths.""" @@ -226,7 +246,7 @@ class Authenticator: return headers - def _get_device_code(self) -> dict[str, str]: + def _get_device_code(self) -> _DeviceCode: """ Get a device code for GitHub authentication. @@ -246,7 +266,7 @@ class Authenticator: json={"client_id": client_id, "scope": "read:user"}, ) 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"] if not all(field in resp_json for field in required_fields): @@ -256,7 +276,7 @@ class Authenticator: status_code=400, ) - return resp_json + return _DEVICE_CODE.validate_python(resp_json) except httpx.HTTPStatusError as e: verbose_logger.error("HTTP error getting device code: %s", e) raise GetDeviceCodeError( @@ -307,12 +327,13 @@ class Authenticator: }, ) 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!") - return resp_json["access_token"] - elif "error" in resp_json and resp_json.get("error") == "authorization_pending": + return poll_result["access_token"] + elif "error" in poll_result and poll_result.get("error") == "authorization_pending": verbose_logger.debug("Authorization pending (attempt %s/%s)", attempt + 1, max_attempts) else: verbose_logger.warning("Unexpected response: %s", resp_json) diff --git a/litellm/llms/manus/responses/transformation.py b/litellm/llms/manus/responses/transformation.py index 41ccee3589d..9785a50bbce 100644 --- a/litellm/llms/manus/responses/transformation.py +++ b/litellm/llms/manus/responses/transformation.py @@ -2,6 +2,7 @@ import uuid from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -29,6 +30,7 @@ else: LiteLLMLoggingObj = Any MANUS_API_BASE: Final = "https://api.manus.im" +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True)) class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): @@ -177,7 +179,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): original_response=raw_response.text, 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" 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, 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" if "createdAt" in raw_response_json and "created_at" not in raw_response_json: diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 2a17456ffa8..309589b367a 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -4,7 +4,8 @@ from collections.abc import AsyncIterator, Iterator from typing import TYPE_CHECKING, Any, Final 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 from litellm._logging import verbose_proxy_logger @@ -76,6 +77,21 @@ class _OllamaGenerateReasoning(LiteLLMBaseModel): 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): """ Reference: https://github.com/ollama/ollama/blob/main/docs/api.md#parameters @@ -159,7 +175,7 @@ class OllamaConfig(BaseConfig): system: str | None = None, template: str | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[dict[str, object]] = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @@ -288,7 +304,7 @@ class OllamaConfig(BaseConfig): api_key: str | None = None, json_mode: bool | None = None, ) -> ModelResponse: - response_json: Final = raw_response.json() + response_json: Final = _OLLAMA_GENERATE_RESPONSE.validate_python(raw_response.json()) ## RESPONSE OBJECT model_response.choices[0].finish_reason = "stop" if request_data.get("format", "") == "json": @@ -303,7 +319,7 @@ class OllamaConfig(BaseConfig): model_response.choices[0].finish_reason = "stop" else: 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 if ( @@ -356,7 +372,7 @@ class OllamaConfig(BaseConfig): len(tokenizer.encode(_prompt, disallowed_special=())), ) 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( model_response, diff --git a/litellm/llms/replicate/chat/handler.py b/litellm/llms/replicate/chat/handler.py index fc114104d32..764ea40db8d 100644 --- a/litellm/llms/replicate/chat/handler.py +++ b/litellm/llms/replicate/chat/handler.py @@ -4,6 +4,8 @@ import time from collections.abc import Callable from typing import Final +from pydantic import ConfigDict, TypeAdapter + import litellm from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -20,6 +22,7 @@ from ..common_utils import ReplicateError from .transformation import ReplicateConfig replicate_config: Final = ReplicateConfig() +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True)) # 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 response = await http_client.get(prediction_url, headers=headers) if response.status_code == 200: - response_data = response.json() + response_data = _JSON_OBJECT.validate_python(response.json()) status = response_data.get("status", "") # Check that "output" exists and is not None or empty 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 ) # 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) - 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", "starting", ]: @@ -280,7 +283,7 @@ async def async_completion( 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" 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", "starting", ]: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 1cf1d9e1b49..a291c8e0efb 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -3209,7 +3209,7 @@ class ModelResponseIterator: def _apply_stream_usage_metadata( self, processed_chunk: GenerateContentResponseBody, - model_response: Any, + model_response: "ModelResponseStream", grounding_metadata: list[dict], ) -> Usage | None: if "usageMetadata" not in processed_chunk: diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index b01487bffcf..6ea4317a174 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -2,7 +2,7 @@ import copy import functools import os import uuid -from collections.abc import AsyncGenerator, Iterator, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterable, Iterator, Mapping, Sequence from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ @@ -10,9 +10,11 @@ from typing import ( Final, Literal, Optional, + Protocol, ) import httpx +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.exceptions import GuardrailRaisedException @@ -80,6 +82,8 @@ _VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 +_SHIELD_BODY: Final = TypeAdapter(dict[str, object]) + def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, str]]: for call in calls: @@ -89,6 +93,15 @@ def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, st yield function, arguments +class _ContentDelta(Protocol): + content: object + + +class _FinishingDelta(Protocol): + content: object + tool_calls: object + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -215,7 +228,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): timeout=_DEFAULT_TIMEOUT_SECONDS, ) response.raise_for_status() - return response.json() + return _SHIELD_BODY.validate_python(response.json(), strict=True) except httpx.HTTPStatusError as exc: verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path) raise GuardrailRaisedException( @@ -329,8 +342,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): self, data: MutableRequest, user_api_key_dict: UserAPIKeyAuth, - response: Any, - ) -> Any: + response: object, + ) -> object: """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 @@ -345,8 +358,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): return response response = detached(response) # rebind-ok: everything below restores the copy. - if self._is_anthropic_message_response(response): - return await self._restore_anthropic_response(response, data) + anthropic_body: Final = as_object(response) + 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) if response_slots: @@ -448,7 +462,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, - response: Any, + response: AsyncIterable[object], request_data: MutableRequest, ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. @@ -545,7 +559,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async def _restore_content_window( self, - delta: Any, + delta: _ContentDelta, key: CarryKey, carries: CarryWindows, session_id: str, @@ -597,7 +611,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async def _flush_finished_choice( self, - delta: Any, + delta: _FinishingDelta, choice_index: int, carries: CarryWindows, session_id: str, diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 9ee53c05904..a4e50cc183a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -16,10 +16,11 @@ from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaita from contextlib import asynccontextmanager from dataclasses import dataclass 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 -from typing_extensions import NotRequired, ReadOnly +from pydantic import ConfigDict, TypeAdapter, with_config +from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm 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): entity_type: ReadOnly[str | None] +@with_config(ConfigDict(extra="allow", strict=True)) class _PresidioAnonymizeResponse(TypedDict): text: ReadOnly[str] items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]] +_PRESIDIO_ANONYMIZE_ADAPTER: Final[TypeAdapter[_PresidioAnonymizeResponse | None]] = TypeAdapter( + _PresidioAnonymizeResponse | None +) + + class _JsonResponse(Protocol): def json(self) -> Awaitable[object]: ... @@ -769,7 +777,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): raise Exception( 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( self, diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index b6d75b1f204..b7d83b407a4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -28,6 +28,8 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional 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.integrations.custom_guardrail import ( @@ -76,6 +78,12 @@ _DEFAULT_POLICIES: Final = [ "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): pass @@ -90,7 +98,7 @@ class XecGuardGuardrail(CustomGuardrail): policy_names: list[str] | None = None, block_on_error: bool | None = None, grounding_strictness: str | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: self.api_key = api_key or os.environ.get("XECGUARD_API_KEY") if not self.api_key: @@ -122,9 +130,12 @@ class XecGuardGuardrail(CustomGuardrail): 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 def get_config_model() -> type["GuardrailConfigModel"] | None: @@ -345,7 +356,7 @@ class XecGuardGuardrail(CustomGuardrail): path: str, payload: dict, suppress_errors: bool = False, - ) -> dict | None: + ) -> dict[str, object] | None: endpoint: Final = f"{self.api_base}{path}" verbose_proxy_logger.debug( "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, ) response.raise_for_status() - return response.json() + return _RESPONSE_BODY.validate_python(response.json(), strict=True) except Exception as exc: verbose_proxy_logger.error("XecGuard API error: %s", str(exc)) if suppress_errors: diff --git a/litellm/proxy/hooks/litellm_skills/main.py b/litellm/proxy/hooks/litellm_skills/main.py index 8848878d15b..6bd8417eb5e 100644 --- a/litellm/proxy/hooks/litellm_skills/main.py +++ b/litellm/proxy/hooks/litellm_skills/main.py @@ -29,6 +29,9 @@ import json from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Protocol +from pydantic import ConfigDict, TypeAdapter, with_config +from typing_extensions import ReadOnly, TypedDict + import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache @@ -81,6 +84,22 @@ class _ChatCompletion(Protocol): 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: """The first choice of an OpenAI shaped completion response.""" return response.choices[0] @@ -620,7 +639,7 @@ class SkillsInjectionHook(CustomLogger): # Collect generated 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: generated_files.append( { @@ -833,7 +852,7 @@ print('No executable skill module found') ) -> str: """Execute a litellm_code_execution tool call and return result string.""" 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", "") verbose_proxy_logger.debug("SkillsInjectionHook: Executing code (%s chars)", len(code)) @@ -849,7 +868,7 @@ print('No executable skill module found') # Collect generated files if exec_result.get("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: file_content = base64.b64decode(f["content_base64"]) generated_files.append( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ce23d464e68..e965a872ad9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -39,7 +39,6 @@ from typing import ( Optional, Protocol, TypeAlias, - TypedDict, Union, cast, get_args, @@ -50,9 +49,9 @@ from typing import ( import anyio import websockets 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 typing_extensions import NotRequired, ReadOnly, assert_never +from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never from litellm._uuid import uuid 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, ...]] +@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): keepalive_timeout: float ttl_dns_cache: int @@ -8471,9 +8498,10 @@ class ProxyConfig: if config_record is None or config_record.param_value is None: return - litellm_settings = config_record.param_value - if isinstance(litellm_settings, str): - litellm_settings = json.loads(litellm_settings) + raw_litellm_settings: Final = config_record.param_value + litellm_settings: Final = _WEBSEARCH_INTERCEPTION_LITELLM_SETTINGS.validate_python( + 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) @@ -8726,7 +8754,7 @@ class ProxyConfig: if config_record is None or config_record.param_value is None: 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") 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 try: - body: Final = await request.json() + body: Final = _LOGIN_REQUEST_BODY.validate_python(await request.json()) username: Final = str(body.get("username")) password: Final = str(body.get("password")) @@ -17199,7 +17227,7 @@ async def login_v3(request: Request): 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")) password: Final = str(body.get("password")) @@ -17271,7 +17299,7 @@ async def login_v3_exchange(request: Request): 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") if not code: raise ProxyException( diff --git a/litellm/router.py b/litellm/router.py index 9ac6eae9e3e..77cff6254ea 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9335,10 +9335,10 @@ class Router: if access_windows_error is not None: raise ValueError(access_windows_error) 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: 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) deployment = Deployment( **deployment_info, diff --git a/litellm/utils.py b/litellm/utils.py index bef8aaa72c2..72ce14912a9 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -45,7 +45,7 @@ from httpx import Proxy from httpx._utils import get_environment_proxies from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # OpenAI parser module is private 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.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) # 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 max_position_embeddings: Final = config_json.get("max_position_embeddings") 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 +@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): split_model: 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) # 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 max_position_embeddings: Final = config_json.get("max_position_embeddings")