From d4a791d650c15e49c3a51b301ea017bbfc5af28c Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 6 Oct 2026 15:31:54 +0000 Subject: [PATCH] refactor(types): replace Any with proven types in 157 files (#44798) * refactor(types): replace Any with proven types in 299 files Clears 871 basedpyright Any errors (reportAny 6,703 to 6,240, reportExplicitAny 1,651 to 1,243) without adding a cast, an ignore or a suppression, and without touching any budget file Most edits are annotation-only: a parameter, return or local goes from Any to object, Mapping[str, object] or the concrete type the value always held. Fourteen files validate untyped JSON once where it enters, through a module-level pydantic TypeAdapter or model_validate, and then use real types No HTTP status, error type or response shape changes. Mistral speech and fal.ai Bria image generation now report a pydantic ValidationError instead of an AttributeError when the provider answers 2xx with a body that is not a JSON object * refactor(types): make the config locals fix effective and trim no-op edits The Mapping[str, object] annotation on `locals().copy()` removed no error, because the checker narrows the variable back to the dict[str, Any] the call returns. 21 provider config constructors now build the same copy with dict(locals()), which the checker infers as object values under that annotation, so each file loses one reportAny. The same annotation is reverted in 31 other config files where it stayed a no-op, together with the tests that were added only to cover those lines, and the one MCP server manager line that no CI coverage shard executes is reverted too. The pull request drops from 345 to 297 changed files. * refactor(types): accept only int in the proxy state setter get_proxy_state_variable is annotated to return int, but set_proxy_state_variable still took Any, so the checker could not hold callers to the type the getter promises. The setter now takes int, which is what its only caller already passes. * refactor(types): index the proxy state key so the getter returns int * refactor(types): keep public annotations and provider error text unchanged Restore every public return, public method parameter, public attribute and exported alias to its annotation on main so code that type-checks against the package keeps type-checking, and take the Mistral speech and fal.ai Bria changes back out so no provider error message differs from main * refactor(types): leave the Vertex RAG chunking read as it is on main Take the chunking format validation back out of the Vertex RAG ingestion path. It needs the vertexai SDK, a storage bucket and a RAG corpus to execute, so nothing here could run it end to end, and it cleared only two errors * test(integration): pin the validated provider boundaries on a live proxy * test(integration): give the held burst a client that outlasts the gate The fault cell holds a burst at the upstream for up to 60 seconds while it kills a worker, but sent the burst through the shared 15 second client, so a slow box could time the survivors out before the gate opened. The burst now goes through its own client whose timeout is twice the gate, and the gate length is one named constant. * refactor(types): keep the license reply handling and experimental MCP signatures as they were The license check validated the whole reply as a mapping, which changed the error text logged for a reply that is not an object. It now validates only the verify value, so every reply is handled and logged exactly as before while the value is still typed. Three files under the experimental MCP server changed annotations on public functions and methods (three returns and three parameters). They go back to their previous content so no public signature in the diff is narrowed. * test(integration): answer the proxy's model-list call in the OpenAI stand-ins Every 300 seconds each proxy worker asks an OpenAI deployment for GET /v1/models. Four new cells own an OpenAI stand-in that accepted only the call under test, so a refresh landing inside a cell failed it. The stand-ins now answer that call through the suite's own helper and the cells count only the provider calls they drive. --- litellm/assistants/main.py | 8 +- litellm/caching/valkey_semantic_cache.py | 2 +- .../compression/scoring/embedding_scorer.py | 4 +- litellm/containers/endpoint_factory.py | 5 +- litellm/google_genai/streaming_iterator.py | 2 +- litellm/harness/endpoint.py | 4 +- .../SlackAlerting/slack_alerting.py | 2 +- .../clickhouse/clickhouse_batch_logger.py | 4 +- .../integrations/cloudzero/cz_stream_api.py | 4 +- litellm/integrations/cloudzero/database.py | 4 +- .../focus/destinations/factory.py | 2 +- .../generic_prompt_manager.py | 2 +- litellm/integrations/gitlab/__init__.py | 4 +- .../integrations/langfuse/langfuse_handler.py | 2 +- .../integrations/otel/plumbing/providers.py | 4 +- litellm/integrations/otel/presets/agentops.py | 4 +- .../prometheus_helpers/__init__.py | 4 +- .../transformation.py | 6 +- .../cloud_storage_security.py | 2 +- litellm/litellm_core_utils/core_helpers.py | 2 +- .../litellm_core_utils/get_model_cost_map.py | 7 +- .../initialize_dynamic_callback_params.py | 2 +- litellm/litellm_core_utils/logging_utils.py | 2 +- litellm/litellm_core_utils/redact_messages.py | 2 +- litellm/litellm_core_utils/safe_json_loads.py | 2 +- .../sensitive_data_masker.py | 2 +- .../aiml/image_generation/cost_calculator.py | 4 +- litellm/llms/anthropic/chat/transformation.py | 2 +- litellm/llms/anthropic/common_utils.py | 6 +- .../anthropic/completion/transformation.py | 4 +- .../pass_through/adapters/transformation.py | 2 +- .../pass_through/messages/transformation.py | 2 +- .../responses_adapters/handler.py | 4 +- .../responses_adapters/transformation.py | 6 +- litellm/llms/azure/chat/gpt_transformation.py | 2 +- litellm/llms/azure/common_utils.py | 4 +- litellm/llms/azure/completion/handler.py | 6 +- litellm/llms/azure/exception_mapping.py | 4 +- .../llms/azure_ai/agents/transformation.py | 2 +- litellm/llms/azure_ai/anthropic/handler.py | 2 +- .../image_generation/cost_calculator.py | 4 +- .../base_llm/guardrail_translation/utils.py | 4 +- .../llms/base_llm/responses/transformation.py | 2 +- litellm/llms/bedrock/base_aws_llm.py | 2 +- .../bedrock/chat/converse_transformation.py | 2 +- .../amazon_ai21_transformation.py | 3 +- .../amazon_cohere_transformation.py | 3 +- .../amazon_llama_transformation.py | 3 +- .../amazon_mistral_transformation.py | 3 +- .../amazon_qwen3_transformation.py | 3 +- .../amazon_titan_transformation.py | 3 +- .../anthropic_claude2_transformation.py | 3 +- litellm/llms/bedrock/common_utils.py | 14 +- .../embed/amazon_titan_g1_transformation.py | 3 +- .../embed/amazon_titan_v2_transformation.py | 3 +- .../anthropic_claude3_transformation.py | 8 +- litellm/llms/codex/harness/transformation.py | 2 +- litellm/llms/cometapi/chat/transformation.py | 2 +- .../image_generation/cost_calculator.py | 4 +- .../llms/custom_httpx/container_handler.py | 6 +- litellm/llms/dashscope/chat/transformation.py | 6 +- litellm/llms/deepinfra/chat/transformation.py | 10 +- litellm/llms/deepseek/chat/transformation.py | 6 +- .../llms/deepseek/messages/transformation.py | 2 +- .../chat/transformation.py | 6 +- .../bytedance_transformation.py | 2 +- .../flux_pro_v11_transformation.py | 2 +- .../flux_schnell_transformation.py | 2 +- .../ideogram_v3_transformation.py | 2 +- .../llms/fireworks_ai/chat/transformation.py | 2 +- litellm/llms/gemini/chat/transformation.py | 3 +- litellm/llms/gemini/common_utils.py | 6 +- litellm/llms/gemini/count_tokens/handler.py | 2 +- .../llms/gemini/image_edit/cost_calculator.py | 4 +- .../image_generation/cost_calculator.py | 4 +- .../llms/gemini/image_usage_transformation.py | 6 +- .../github_copilot/chat/transformation.py | 9 +- litellm/llms/groq/chat/transformation.py | 8 +- litellm/llms/heroku/chat/transformation.py | 6 +- litellm/llms/langflow/chat/transformation.py | 2 +- litellm/llms/langgraph/chat/transformation.py | 8 +- .../litellm_proxy/skills/prompt_injection.py | 2 +- litellm/llms/mistral/chat/transformation.py | 10 +- .../llms/modelscope/chat/transformation.py | 6 +- litellm/llms/moonshot/chat/transformation.py | 6 +- .../audio_transcription/handler.py | 2 +- litellm/llms/nvidia_riva/common_utils.py | 6 +- litellm/llms/oci/chat/generic.py | 4 +- litellm/llms/oci/chat/transformation.py | 4 +- litellm/llms/ollama/completion/handler.py | 4 +- .../llms/openai/chat/gpt_transformation.py | 2 +- .../chat/guardrail_translation/handler.py | 2 +- .../openai/chat/o_series_transformation.py | 6 +- litellm/llms/openai/completion/handler.py | 8 +- .../llms/openai/completion/transformation.py | 3 +- .../llms/openai/responses/transformation.py | 2 +- litellm/llms/openai_like/dynamic_config.py | 6 +- .../llms/opencode/harness/transformation.py | 2 +- litellm/llms/recraft/cost_calculator.py | 4 +- litellm/llms/runwayml/cost_calculator.py | 4 +- litellm/llms/sambanova/chat.py | 10 +- litellm/llms/sap/credentials.py | 4 +- .../vertex_ai/agent_engine/transformation.py | 4 +- .../llms/vertex_ai/files/transformation.py | 2 +- .../llms/vertex_ai/gemini/transformation.py | 2 +- .../vertex_ai/image_edit/cost_calculator.py | 4 +- .../vertex_gemma_models/transformation.py | 2 +- litellm/llms/xai/chat/transformation.py | 5 +- .../_experimental/mcp_server/oauth_utils.py | 6 +- .../proxy/_experimental/mcp_server/utils.py | 2 +- litellm/proxy/_lazy_features.py | 2 +- .../proxy/agent_endpoints/a2a_endpoints.py | 4 +- litellm/proxy/auth/handle_jwt.py | 8 +- litellm/proxy/auth/litellm_license.py | 6 +- litellm/proxy/auth/model_checks.py | 4 +- litellm/proxy/auth/network.py | 6 +- litellm/proxy/auth/trusted_proxy_utils.py | 2 +- .../guardrail_hooks/aporia_ai/aporia_ai.py | 4 +- .../guardrails/guardrail_hooks/azure/base.py | 6 +- .../guardrail_hooks/bedrock_guardrails.py | 2 +- .../cisco_ai_defense/cisco_ai_defense.py | 4 +- .../litellm_content_filter/content_filter.py | 3 +- .../guardrail_benchmarks/test_eval.py | 4 +- .../llm_shield_proxy/llm_shield_proxy.py | 10 +- .../mcp_end_user_permission.py | 4 +- .../guardrails/guardrail_hooks/presidio.py | 4 +- .../guardrail_hooks/singulr/singulr.py | 6 +- .../guardrail_hooks/straiker/straiker.py | 4 +- .../guardrails/guardrail_initializers.py | 4 +- litellm/proxy/guardrails/init_guardrails.py | 4 +- litellm/proxy/hooks/batch_rate_limiter.py | 2 +- .../hooks/parallel_request_limiter_v3.py | 4 +- .../proxy/hooks/proxy_track_cost_callback.py | 4 +- litellm/proxy/litellm_pre_call_utils.py | 4 +- .../callback_logs_endpoints.py | 4 +- .../per_request_root_path_middleware.py | 2 +- .../middleware/prometheus_auth_middleware.py | 8 +- ...end_medical_passthrough_logging_handler.py | 2 +- .../gemini_passthrough_logging_handler.py | 6 +- .../transcribe_passthrough_logging_handler.py | 2 +- .../vertex_passthrough_logging_handler.py | 5 +- .../pass_through_endpoints/success_handler.py | 4 +- .../upstream_usage_headers.py | 4 +- litellm/proxy/proxy_cli.py | 2 +- litellm/proxy/proxy_server.py | 2 +- litellm/proxy/rag_endpoints/endpoints.py | 2 +- .../search_tool_management.py | 2 +- litellm/proxy/types_utils/utils.py | 2 +- .../streaming_iterator.py | 2 +- .../transformation.py | 2 +- litellm/responses/streaming_iterator.py | 2 +- litellm/responses/utils.py | 12 +- .../router_strategy/adaptive_router/hooks.py | 2 +- litellm/router_strategy/budget_limiter.py | 4 +- litellm/search/main.py | 4 +- litellm/videos/utils.py | 4 +- .../test_bundled_cost_map_price_wire.py | 59 ++ ...test_bedrock_anthropic_beta_header_wire.py | 154 ++++++ .../test_container_file_routes_wire.py | 91 ++++ ...test_gemini_image_config_and_usage_wire.py | 244 +++++++++ .../test_interactions_openai_bridge_wire.py | 93 ++++ .../test_responses_reasoning_item_wire.py | 126 +++++ ...validated_reply_boundaries_under_faults.py | 503 ++++++++++++++++++ .../test_xai_server_tool_usage_wire.py | 124 +++++ tests/unit/integrations/gitlab/test_init.py | 56 ++ .../test_litellm_responses_bridge.py | 64 +++ tests/unit/llms/azure/completion/__init__.py | 0 .../llms/azure/completion/test_handler.py | 52 ++ tests/unit/llms/azure_ai/agents/__init__.py | 0 .../azure_ai/agents/test_transformation.py | 67 +++ .../claude/test_azure_anthropic_handler.py | 26 + .../guardrail_translation/__init__.py | 0 .../guardrail_translation/test_utils.py | 34 ++ .../llms/bedrock/test_bedrock_common_utils.py | 18 + .../custom_httpx/test_container_handler.py | 24 + .../gemini/test_image_usage_transformation.py | 47 ++ .../skills/test_prompt_injection.py | 83 +++ .../completion/test_completion_handler.py | 59 ++ .../test_openai_responses_transformation.py | 26 + .../agent_engine/test_transformation.py | 45 ++ .../llms/xai/test_xai_chat_transformation.py | 52 ++ tests/unit/proxy/auth/test_handle_jwt.py | 27 + tests/unit/proxy/auth/test_litellm_license.py | 31 ++ .../guardrail_hooks/aporia_ai/__init__.py | 0 .../aporia_ai/test_aporia_ai.py | 87 +++ .../content_filter/test_content_filter.py | 55 ++ .../guardrails/test_guardrail_initializers.py | 84 +++ .../proxy/guardrails/test_init_guardrails.py | 45 +- ...test_gemini_passthrough_logging_handler.py | 33 ++ ...test_vertex_passthrough_logging_handler.py | 34 ++ .../test_success_handler.py | 78 +++ 191 files changed, 2852 insertions(+), 290 deletions(-) create mode 100644 tests/integration/pricing/test_bundled_cost_map_price_wire.py create mode 100644 tests/integration/providers/test_bedrock_anthropic_beta_header_wire.py create mode 100644 tests/integration/providers/test_container_file_routes_wire.py create mode 100644 tests/integration/providers/test_gemini_image_config_and_usage_wire.py create mode 100644 tests/integration/providers/test_interactions_openai_bridge_wire.py create mode 100644 tests/integration/providers/test_responses_reasoning_item_wire.py create mode 100644 tests/integration/providers/test_validated_reply_boundaries_under_faults.py create mode 100644 tests/integration/providers/test_xai_server_tool_usage_wire.py create mode 100644 tests/unit/integrations/gitlab/test_init.py create mode 100644 tests/unit/llms/azure/completion/__init__.py create mode 100644 tests/unit/llms/azure/completion/test_handler.py create mode 100644 tests/unit/llms/azure_ai/agents/__init__.py create mode 100644 tests/unit/llms/azure_ai/agents/test_transformation.py create mode 100644 tests/unit/llms/base_llm/guardrail_translation/__init__.py create mode 100644 tests/unit/llms/base_llm/guardrail_translation/test_utils.py create mode 100644 tests/unit/llms/gemini/test_image_usage_transformation.py create mode 100644 tests/unit/llms/litellm_proxy/skills/test_prompt_injection.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/test_aporia_ai.py create mode 100644 tests/unit/proxy/guardrails/test_guardrail_initializers.py create mode 100644 tests/unit/proxy/pass_through_endpoints/test_success_handler.py diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index c14c4aec093..ba542bebaf8 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -239,7 +239,7 @@ def create_assistants( temperature: float | None = None, top_p: float | None = None, response_format: str | dict[str, str] | None = None, - client: Any | None = None, + client: object | None = None, api_key: str | None = None, api_base: str | None = None, api_version: str | None = None, @@ -410,7 +410,7 @@ async def adelete_assistant( def delete_assistant( custom_llm_provider: Literal["openai", "azure"], assistant_id: str, - client: Any | None = None, + client: object | None = None, api_key: str | None = None, api_base: str | None = None, api_version: str | None = None, @@ -1181,7 +1181,7 @@ async def arun_thread( model: str | None = None, stream: bool | None = None, tools: Iterable[AssistantToolParam] | None = None, - client: Any | None = None, + client: object | None = None, **kwargs, ) -> Run: loop: Final = asyncio.get_event_loop() @@ -1246,7 +1246,7 @@ def run_thread( model: str | None = None, stream: bool | None = None, tools: Iterable[AssistantToolParam] | None = None, - client: Any | None = None, + client: object | None = None, event_handler: AssistantEventHandler | None = None, # for stream=True calls **kwargs, ) -> Run: diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index 58b76d98d6d..5a3f9a4184c 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -239,7 +239,7 @@ class ValkeySemanticCache(RedisSemanticCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity @staticmethod - def _embedding_metadata(kwargs: dict[str, Any]) -> dict[str, Any] | None: + def _embedding_metadata(kwargs: dict[str, Any]) -> dict[str, object] | None: """The request metadata forwarded to the embedding call.""" return kwargs.get("metadata") diff --git a/litellm/compression/scoring/embedding_scorer.py b/litellm/compression/scoring/embedding_scorer.py index 7e645ba3f9c..7dc070e5ca4 100644 --- a/litellm/compression/scoring/embedding_scorer.py +++ b/litellm/compression/scoring/embedding_scorer.py @@ -6,7 +6,7 @@ Computes cosine similarity between the query embedding and each message embeddin import math from collections.abc import Mapping -from typing import Any, Final +from typing import Final from litellm.caching.dual_cache import DualCache @@ -75,7 +75,7 @@ def embedding_score_messages( # Filter out empty texts — replace with a placeholder to maintain indexing processed_texts: Final = [t if t.strip() else "empty" for t in texts] - kwargs: dict[str, Any] = { + kwargs: dict[str, object] = { "model": model, "input": processed_texts, "caching": cache is not None, diff --git a/litellm/containers/endpoint_factory.py b/litellm/containers/endpoint_factory.py index 25fc223cde0..a969a9cb877 100644 --- a/litellm/containers/endpoint_factory.py +++ b/litellm/containers/endpoint_factory.py @@ -13,6 +13,8 @@ from functools import partial from pathlib import Path from typing import Final, Literal +from pydantic import ConfigDict, TypeAdapter + import litellm from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT from litellm.containers.utils import decode_managed_container_id_for_request @@ -33,13 +35,14 @@ RESPONSE_TYPES: Final[dict[str, type]] = { "ContainerFileObject": ContainerFileObject, "DeleteContainerFileResponse": DeleteContainerFileResponse, } +_ENDPOINTS_DOCUMENT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" config_path: Final = Path(__file__).parent / "endpoints.json" with open(config_path) as f: - return json.load(f) + return _ENDPOINTS_DOCUMENT.validate_python(json.load(f)) def create_sync_endpoint_function(endpoint_config: dict) -> Callable: diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index a49e43e7bdc..27883c938db 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -78,7 +78,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: self.endpoint_type: Final = ( EndpointType.GEMINI if custom_llm_provider == litellm.LlmProviders.GEMINI.value else EndpointType.VERTEX_AI ) - self._hidden_params: dict[str, Any] = hidden_params or {} + self._hidden_params: dict[str, object] = hidden_params or {} async def _handle_async_streaming_logging( self, diff --git a/litellm/harness/endpoint.py b/litellm/harness/endpoint.py index 192c97533ea..1f500fafba2 100644 --- a/litellm/harness/endpoint.py +++ b/litellm/harness/endpoint.py @@ -628,8 +628,8 @@ class ModelEndpoint: def _sdk_kwargs( self, body: Mapping[str, object] - ) -> dict[str, Any]: # mutable-ok: SDK call kwargs, mutated by _invoke_sdk then splatted - kwargs: dict[str, Any] = {**body} # mutable-ok: SDK call kwargs built from the JSON body, then overridden + ) -> dict[str, object]: # mutable-ok: SDK call kwargs, mutated by _invoke_sdk then splatted + kwargs: dict[str, object] = {**body} # mutable-ok: SDK call kwargs built from the JSON body, then overridden if self.model: kwargs["model"] = self.model if self.api_key: diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 34ba857ecb6..cac8c7cb337 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -1807,7 +1807,7 @@ Model Info: async def _run_scheduled_daily_report( self, - llm_router: Any | None = None, + llm_router: object | None = None, pod_lock_manager: "PodLockManager | None" = None, ): """ diff --git a/litellm/integrations/clickhouse/clickhouse_batch_logger.py b/litellm/integrations/clickhouse/clickhouse_batch_logger.py index 84354ecc65c..421ed55189f 100644 --- a/litellm/integrations/clickhouse/clickhouse_batch_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_batch_logger.py @@ -11,7 +11,7 @@ gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as import asyncio from collections.abc import Mapping, Sequence from contextlib import suppress -from typing import Any, ClassVar, Final +from typing import ClassVar, Final from litellm._logging import verbose_logger from litellm.constants import ( @@ -89,7 +89,7 @@ class ClickHouseBatchLogger(CustomBatchLogger): async def async_send_batch(self) -> None: await self.flush_queue() - async def _insert(self, batch: list[dict[str, Any]]) -> bool: + async def _insert(self, batch: list[dict[str, object]]) -> bool: try: await self.storage.insert_rows(self.table, batch) self.rows_written += len(batch) diff --git a/litellm/integrations/cloudzero/cz_stream_api.py b/litellm/integrations/cloudzero/cz_stream_api.py index 2213c5fe275..81d08e6a4f4 100644 --- a/litellm/integrations/cloudzero/cz_stream_api.py +++ b/litellm/integrations/cloudzero/cz_stream_api.py @@ -68,7 +68,7 @@ class CloudZeroStreamer: def _group_by_date(self, data: pl.DataFrame) -> dict[str, pl.DataFrame]: """Group data by date, converting to UTC and validating dates.""" - daily_batches: Final[dict[str, list[dict[str, Any]]]] = {} + daily_batches: Final[dict[str, list[dict[str, object]]]] = {} # Ensure we have the required columns if "time/usage_start" not in data.columns: @@ -209,7 +209,7 @@ class CloudZeroStreamer: return payload - def _convert_cbf_to_api_format(self, row: dict[str, Any]) -> dict[str, Any] | None: + def _convert_cbf_to_api_format(self, row: dict[str, object]) -> dict[str, Any] | None: """Convert CBF row to CloudZero API format - keeping CBF field names as CloudZero expects them.""" try: # CloudZero expects CBF format field names directly, not converted names diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 2fb10ad8a96..e825aee4e15 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -19,7 +19,7 @@ """Database connection and data extraction for LiteLLM.""" from datetime import datetime -from typing import Any, Final +from typing import Final import polars as pl @@ -80,7 +80,7 @@ class LiteLLMDatabase: ORDER BY dus.date DESC, dus.created_at DESC """ - params: Final[list[Any]] = [ + params: Final[list[object]] = [ start_time_utc, end_time_utc, ] diff --git a/litellm/integrations/focus/destinations/factory.py b/litellm/integrations/focus/destinations/factory.py index a46b1ea2dcb..0ecf04326b3 100644 --- a/litellm/integrations/focus/destinations/factory.py +++ b/litellm/integrations/focus/destinations/factory.py @@ -39,7 +39,7 @@ class FocusDestinationFactory: def _resolve_config( *, provider: str, - overrides: dict[str, Any], + overrides: dict[str, object], ) -> dict[str, Any]: if provider == "s3": resolved = { diff --git a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py index de01b2bb02c..d78d1e7d223 100644 --- a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py +++ b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py @@ -93,7 +93,7 @@ class GenericPromptManager(CustomPromptManagement): headers["Authorization"] = f"Bearer {self.api_key}" return headers - def _fetch_prompt_from_api(self, prompt_id: str | None, prompt_spec: PromptSpec | None) -> dict[str, Any]: + def _fetch_prompt_from_api(self, prompt_id: str | None, prompt_spec: PromptSpec | None) -> dict[str, object]: """ Fetch a prompt from the API. diff --git a/litellm/integrations/gitlab/__init__.py b/litellm/integrations/gitlab/__init__.py index cba69d2df83..64acda610a1 100644 --- a/litellm/integrations/gitlab/__init__.py +++ b/litellm/integrations/gitlab/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final if TYPE_CHECKING: from litellm.integrations.custom_prompt_management import CustomPromptManagement @@ -66,7 +66,7 @@ def _gitlab_prompt_initializer( # You can store arbitrary integration-specific config on PromptLiteLLMParams. # If your dataclass doesn't have these attributes, add them or put inside # `litellm_params.extra` and pull them from there. - gitlab_config: Final[dict[str, Any]] = getattr(litellm_params, "gitlab_config", None) or {} + gitlab_config: Final[dict[str, object]] = getattr(litellm_params, "gitlab_config", None) or {} git_ref: Final[str | None] = getattr(litellm_params, "git_ref", None) if not gitlab_config: diff --git a/litellm/integrations/langfuse/langfuse_handler.py b/litellm/integrations/langfuse/langfuse_handler.py index c74866c7a9e..b09236e0874 100644 --- a/litellm/integrations/langfuse/langfuse_handler.py +++ b/litellm/integrations/langfuse/langfuse_handler.py @@ -81,7 +81,7 @@ class LangFuseHandler: return globalLangfuseLogger credentials_dict: dict[ - str, Any + str, object ] = {} # the global langfuse logger uses Environment Variables, there are no dynamic credentials globalLangfuseLogger = in_memory_dynamic_logger_cache.get_cache( credentials=credentials_dict, diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index 13bfa5a20ea..37a8f1951d9 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -7,7 +7,7 @@ import time from collections import OrderedDict from collections.abc import Callable, Iterable, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Final, Literal from opentelemetry import _logs, baggage, metrics, trace from opentelemetry._logs import Logger, LoggerProvider, NoOpLoggerProvider @@ -1072,7 +1072,7 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": ) tls: Final = resolve_otlp_http_tls("METRICS") - exporter: Any = HTTPMetricExporter( + exporter: object = HTTPMetricExporter( endpoint=_otlp_metrics_endpoint(config.endpoint), headers=parse_headers(config.headers), certificate_file=tls.certificate_file, diff --git a/litellm/integrations/otel/presets/agentops.py b/litellm/integrations/otel/presets/agentops.py index 58123656caa..3befb73b0cc 100644 --- a/litellm/integrations/otel/presets/agentops.py +++ b/litellm/integrations/otel/presets/agentops.py @@ -10,7 +10,7 @@ worker thread, off any event loop — and caches it for the process lifetime. """ from collections.abc import Sequence -from typing import Any, Final +from typing import Final import httpx from opentelemetry.sdk.trace import ReadableSpan @@ -117,7 +117,7 @@ def _build_agentops_exporter(spec: ExporterSpec) -> SpanExporter: return _LazyAuthAgentOpsExporter(endpoint=spec.endpoint, api_key=options.get("api_key")) -def _fetch_agentops_jwt(api_key: str) -> dict[str, Any]: +def _fetch_agentops_jwt(api_key: str) -> dict[str, object]: # Own a short-lived client rather than ``_get_httpx_client()``: that returns # a process-wide cached ``HTTPHandler`` whose connection pool is shared by # every caller, so closing it here would break concurrent/subsequent diff --git a/litellm/integrations/prometheus_helpers/__init__.py b/litellm/integrations/prometheus_helpers/__init__.py index 1f3d224237f..2e6a350e3c2 100644 --- a/litellm/integrations/prometheus_helpers/__init__.py +++ b/litellm/integrations/prometheus_helpers/__init__.py @@ -6,7 +6,7 @@ Helpers for the Prometheus integration (extracted to keep ``prometheus.py`` smal from __future__ import annotations -from typing import Any, Final, cast +from typing import Final, cast from litellm.types.integrations.prometheus import ( UserAPIKeyLabelValues, @@ -66,7 +66,7 @@ class PrometheusLabelFactoryContext: for k, v in get_custom_labels_from_tags(enum_values.tags).items(): self._tag_labels[k] = _sanitize_prometheus_label_value(v) # Use a dedicated sentinel so `None` can be cached as a computed result. - self._resolved_end_user: Any = self._END_USER_NOT_COMPUTED + self._resolved_end_user: object = self._END_USER_NOT_COMPUTED def get_resolved_end_user(self) -> str | None: if self._resolved_end_user is self._END_USER_NOT_COMPUTED: diff --git a/litellm/interactions/litellm_responses_transformation/transformation.py b/litellm/interactions/litellm_responses_transformation/transformation.py index 39ccc26c38c..ade33831cdc 100644 --- a/litellm/interactions/litellm_responses_transformation/transformation.py +++ b/litellm/interactions/litellm_responses_transformation/transformation.py @@ -8,7 +8,7 @@ This module handles transforming between: from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import Any, Final, cast +from typing import Final, cast from pydantic import BaseModel @@ -251,7 +251,7 @@ class LiteLLMResponsesInteractionsConfig: # Build interactions response — populate both `outputs` (legacy schema) and # `steps` (new schema) so callers work regardless of which schema they expect. - interactions_response_dict: Final[dict[str, Any]] = { + interactions_response_dict: Final[dict[str, object]] = { "id": getattr(responses_response, "id", ""), "object": "interaction", "status": interactions_status, @@ -274,4 +274,4 @@ class LiteLLMResponsesInteractionsConfig: # Add updated (same as created for now) interactions_response_dict["updated"] = created - return InteractionsAPIResponse(**interactions_response_dict) + return InteractionsAPIResponse.model_validate(interactions_response_dict) diff --git a/litellm/litellm_core_utils/cloud_storage_security.py b/litellm/litellm_core_utils/cloud_storage_security.py index 58457c628c7..c913e4de02e 100644 --- a/litellm/litellm_core_utils/cloud_storage_security.py +++ b/litellm/litellm_core_utils/cloud_storage_security.py @@ -116,7 +116,7 @@ def encode_s3_object_key_for_url(object_key: str) -> str: def should_allow_legacy_cloud_file_ids( - litellm_params: Mapping[str, Any] | None = None, + litellm_params: Mapping[str, object] | None = None, ) -> bool: value = None if isinstance(litellm_params, Mapping): diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 39e95fbf687..4fa78640e24 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -725,7 +725,7 @@ def redact_nested_match_and_regex_keys( if payload is None or isinstance(payload, str): return payload try: - redacted: Final[dict | list[Any] | str | None] = copy.deepcopy(payload) + redacted: Final[dict | list[object] | str | None] = copy.deepcopy(payload) except Exception: return payload diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index 159590da0f4..c8f78e9a814 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -26,7 +26,7 @@ from types import MappingProxyType from typing import Final, Protocol import httpx -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly, TypedDict from litellm import verbose_logger @@ -40,6 +40,7 @@ from litellm.litellm_core_utils.fallback_generalizations import ( FALLBACK_GENERALIZATIONS_KEY: Final = "fallback_generalizations" _CATALOG_ADAPTER: Final = TypeAdapter(dict[str, dict[str, object]]) +_BUNDLED_CATALOG_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _CLI_ENTRYPOINT_NAMES: Final = frozenset({"lite", "litellm-proxy"}) @@ -83,7 +84,7 @@ class GetModelCostMap: @staticmethod def load_local_model_cost_map_with_revision() -> "ModelCostMapReloaded": body: Final = GetModelCostMap.read_local_model_cost_map_bytes() - content: Final = json.loads(body) + content: Final = _BUNDLED_CATALOG_ADAPTER.validate_python(json.loads(body)) return ModelCostMapReloaded(model_cost_map=content, revision=git_blob_id(body)) @staticmethod @@ -230,7 +231,7 @@ class _FetchAttemptRetryable: def _parse_retry_after_seconds(response: httpx.Response) -> float | None: - header: Final = response.headers.get("Retry-After") + header: Final[str | None] = response.headers.get("Retry-After") if header is None: return None try: diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 3100ca6fba1..9ac78e3b917 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -148,7 +148,7 @@ _trusted_overlay_callback_params: Final = frozenset( ) -def get_trusted_callback_params(kwargs: Mapping[str, Any] | None) -> tuple[tuple[str, str], ...]: +def get_trusted_callback_params(kwargs: Mapping[str, object] | None) -> tuple[tuple[str, str], ...]: """ Read callback params the proxy itself stamped from admin-configured team/key callback settings. diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index a2d4d91a6db..6aaa1ef692b 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -50,7 +50,7 @@ _DATA_URI_RE: Final = re.compile(r"data:([^;]+);base64,([A-Za-z0-9+/=]+)") _MAX_TRUNCATION_DEPTH: Final = 20 -def _base64_data_uri_replacer(match: re.Match) -> str: +def _base64_data_uri_replacer(match: re.Match[str]) -> str: """Replace a single base64 data-URI match with a size placeholder if too long.""" mime_type: Final = match.group(1) payload: Final = match.group(2) diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 8c77ef32cfb..85ed0a40687 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -377,7 +377,7 @@ def should_redact_message_logging(model_call_details: dict) -> bool: return litellm.turn_off_message_logging is True -def redact_message_input_output_from_logging(model_call_details: dict, result, input: Any | None = None) -> Any: +def redact_message_input_output_from_logging(model_call_details: dict, result, input: object | None = None) -> Any: """ Removes messages, prompts, input, response from logging. This modifies the data in-place only redacts when litellm.turn_off_message_logging == True diff --git a/litellm/litellm_core_utils/safe_json_loads.py b/litellm/litellm_core_utils/safe_json_loads.py index bb973064134..a0ef6f6fdd9 100644 --- a/litellm/litellm_core_utils/safe_json_loads.py +++ b/litellm/litellm_core_utils/safe_json_loads.py @@ -6,7 +6,7 @@ import json from typing import Any -def safe_json_loads(data: str, default: Any = None) -> Any: +def safe_json_loads(data: str, default: object = None) -> Any: """ Safely parse a JSON string. If parsing fails, return the default value (None by default). """ diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index b83ecc6929b..747ad9bf5c9 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -124,7 +124,7 @@ class SensitiveDataMasker: if depth >= max_depth: return data - masked_data: Final[dict[str, Any]] = {} + masked_data: Final[dict[str, object]] = {} for k, v in data.items(): try: key_is_sensitive = self.is_sensitive_key(k, excluded_keys) diff --git a/litellm/llms/aiml/image_generation/cost_calculator.py b/litellm/llms/aiml/image_generation/cost_calculator.py index abf4216807a..ff49052dc25 100644 --- a/litellm/llms/aiml/image_generation/cost_calculator.py +++ b/litellm/llms/aiml/image_generation/cost_calculator.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import resolve_image_model_info @@ -7,7 +7,7 @@ from litellm.types.utils import ImageResponse, ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index fd67ebd5293..5aa14b0442a 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -316,7 +316,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): metadata: dict | None = None, system: str | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 049f474032c..3b46fa9b9df 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -1645,7 +1645,7 @@ def strip_thinking_blocks_from_anthropic_messages_request_dict( def strip_empty_content_blocks_from_anthropic_messages( - messages: list[Any], + messages: Sequence[object], ) -> list[Any]: """ Return a new message list with empty or whitespace-only ``{"type": "text"}`` @@ -1760,7 +1760,7 @@ def _sanitize_tool_use_id_content_block(block: object) -> object: return block -def sanitize_tool_use_ids_in_anthropic_messages(messages: list[Any]) -> list[Any]: +def sanitize_tool_use_ids_in_anthropic_messages(messages: Sequence[object]) -> list[Any]: """ Return a new message list with ``tool_use`` / ``server_tool_use`` ``id`` and ``tool_result`` ``tool_use_id`` values rewritten to satisfy Anthropic's @@ -1934,7 +1934,7 @@ def _flatten_web_search_results_in_message(message: object) -> object: def flatten_unencrypted_web_search_results_in_anthropic_messages( - messages: list[Any], + messages: Sequence[object], ) -> list[Any]: """ Return a new message list with replayed ``web_search_tool_result`` blocks that diff --git a/litellm/llms/anthropic/completion/transformation.py b/litellm/llms/anthropic/completion/transformation.py index 46ab27ab0c7..8edd4354f20 100644 --- a/litellm/llms/anthropic/completion/transformation.py +++ b/litellm/llms/anthropic/completion/transformation.py @@ -6,7 +6,7 @@ Litellm provider slug: `anthropic_text/` import json import time -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Mapping from typing import TYPE_CHECKING, Final import httpx @@ -73,7 +73,7 @@ class AnthropicTextConfig(BaseConfig): top_k: int | None = None, metadata: dict | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/anthropic/pass_through/adapters/transformation.py b/litellm/llms/anthropic/pass_through/adapters/transformation.py index 022bc6337b5..88b99eb733e 100644 --- a/litellm/llms/anthropic/pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/adapters/transformation.py @@ -962,7 +962,7 @@ class LiteLLMAnthropicMessagesAdapter: ) elif isinstance(system_content, list): # Convert Anthropic system content blocks to OpenAI format - openai_system_content: Final[list[dict[str, Any]]] = [] + openai_system_content: Final[list[dict[str, object]]] = [] model_name: Final = anthropic_message_request.get("model", "") for block in system_content: if isinstance(block, dict) and block.get("type") == "text": diff --git a/litellm/llms/anthropic/pass_through/messages/transformation.py b/litellm/llms/anthropic/pass_through/messages/transformation.py index 1cf5d2fcc6e..11698477a33 100644 --- a/litellm/llms/anthropic/pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/pass_through/messages/transformation.py @@ -361,7 +361,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): self, headers: dict, # mutable-ok: out-param optional_params: dict, # mutable-ok: out-param - messages: list[Any], # mutable-ok: mirrors the validate_anthropic_messages_environment contract + messages: list[object], # mutable-ok: mirrors the validate_anthropic_messages_environment contract ) -> dict: # mutable-ok: out-param if "anthropic-version" not in headers: headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION diff --git a/litellm/llms/anthropic/pass_through/responses_adapters/handler.py b/litellm/llms/anthropic/pass_through/responses_adapters/handler.py index 627a027b72e..166ad916ea2 100644 --- a/litellm/llms/anthropic/pass_through/responses_adapters/handler.py +++ b/litellm/llms/anthropic/pass_through/responses_adapters/handler.py @@ -5,7 +5,7 @@ Used when the target model is an OpenAI or Azure model. """ from collections.abc import AsyncIterator, Coroutine, Mapping -from typing import Any, Final, TypeAlias +from typing import Final, TypeAlias import litellm from litellm.types.llms.anthropic import ( @@ -63,7 +63,7 @@ def _build_responses_kwargs( top_p: float | None = None, output_format: AnthropicOutputSchema | None = None, extra_kwargs: Mapping[str, object] | None = None, -) -> dict[str, Any]: +) -> dict[str, object]: """ Build the kwargs dict to pass directly to litellm.responses() / litellm.aresponses(). """ diff --git a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py index a3ebbbcb830..83bb146360e 100644 --- a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py @@ -257,7 +257,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: ) elif isinstance(content, list): user_parts: list[Mapping[str, object]] = [] - tool_image_parts: list[dict[str, Any]] = [] # mutable-ok: json content parts + tool_image_parts: list[dict[str, object]] = [] # mutable-ok: json content parts for block in content: if not isinstance(block, dict): continue @@ -367,7 +367,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: for _, group in groupby(enumerate(blocks), key=self._assistant_block_group_key) for item in self._assistant_group_to_input_items(tuple(block for _, block in group)) ) - asst_parts: list[dict[str, Any]] = [ # mutable-ok: API message payload + asst_parts: list[dict[str, object]] = [ # mutable-ok: API message payload {"type": "output_text", "text": block.get("text", "")} for block in blocks if block.get("type") == "text" @@ -533,7 +533,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: }, ) - responses_kwargs: Final[dict[str, Any]] = { + responses_kwargs: Final[dict[str, object]] = { "model": model, "input": input_items, } diff --git a/litellm/llms/azure/chat/gpt_transformation.py b/litellm/llms/azure/chat/gpt_transformation.py index 831e2c56f46..03fe27437a5 100644 --- a/litellm/llms/azure/chat/gpt_transformation.py +++ b/litellm/llms/azure/chat/gpt_transformation.py @@ -93,7 +93,7 @@ class AzureOpenAIConfig(BaseConfig): temperature: int | None = None, top_p: int | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 7436a0e1b00..5d1a66873a3 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -5,7 +5,7 @@ import os from collections.abc import Callable, Mapping from functools import lru_cache from types import MappingProxyType -from typing import Any, Final, Literal, NamedTuple, cast +from typing import Final, Literal, NamedTuple, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -549,7 +549,7 @@ class BaseAzureLLM(BaseOpenAILLM): # on every request (via `_refresh_api_key`), so passing # `azure_ad_token_provider` directly preserves Azure AD token refresh # behavior that the regular AzureOpenAI client provides. - v1_api_key: str | Callable[[], Any] | None = ( + v1_api_key: str | Callable[[], object] | None = ( azure_client_params.get("api_key") or azure_client_params.get("azure_ad_token_provider") or azure_client_params.get("azure_ad_token") diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 23eef51e7ee..476f6536ffa 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -180,7 +180,7 @@ class AzureTextCompletion(BaseAzureLLM): except Exception as e: status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise AzureOpenAIError(status_code=status_code, message=str(e), headers=error_headers) @@ -241,7 +241,7 @@ class AzureTextCompletion(BaseAzureLLM): except Exception as e: status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise AzureOpenAIError(status_code=status_code, message=str(e), headers=error_headers) @@ -352,7 +352,7 @@ class AzureTextCompletion(BaseAzureLLM): except Exception as e: status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise AzureOpenAIError(status_code=status_code, message=str(e), headers=error_headers) diff --git a/litellm/llms/azure/exception_mapping.py b/litellm/llms/azure/exception_mapping.py index b199fcd03d5..dcc3601cd8e 100644 --- a/litellm/llms/azure/exception_mapping.py +++ b/litellm/llms/azure/exception_mapping.py @@ -27,7 +27,7 @@ class AzureOpenAIExceptionMapping: # Keep the OpenAI-style body fields populated so downstream (proxy + SDK) # can surface `type` / `code` correctly. - openai_style_body: Final[dict[str, Any]] = { + openai_style_body: Final[dict[str, object]] = { "message": provider_message, "type": provider_type or "invalid_request_error", "code": provider_code or "content_policy_violation", @@ -54,7 +54,7 @@ class AzureOpenAIExceptionMapping: @staticmethod def _extract_azure_error( original_exception: Exception, - ) -> tuple[dict[str, Any], dict | None]: + ) -> tuple[dict[str, object], dict | None]: """Extract Azure OpenAI error payload and inner error details. Azure error formats can vary by endpoint/version. Common shapes: diff --git a/litellm/llms/azure_ai/agents/transformation.py b/litellm/llms/azure_ai/agents/transformation.py index baba3149963..0b6ea3ba717 100644 --- a/litellm/llms/azure_ai/agents/transformation.py +++ b/litellm/llms/azure_ai/agents/transformation.py @@ -216,7 +216,7 @@ class AzureAIAgentsConfig(BaseConfig): converted_messages.append({"role": role, "content": content}) - payload: Final[dict[str, Any]] = { + payload: Final[dict[str, object]] = { "agent_id": agent_id, "messages": converted_messages, "api_version": self._get_api_version(optional_params), diff --git a/litellm/llms/azure_ai/anthropic/handler.py b/litellm/llms/azure_ai/anthropic/handler.py index 24ee76b31d0..a132df62b08 100644 --- a/litellm/llms/azure_ai/anthropic/handler.py +++ b/litellm/llms/azure_ai/anthropic/handler.py @@ -196,7 +196,7 @@ class AzureAnthropicChatCompletion(AnthropicChatCompletion): status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) error_text = getattr(e, "text", str(e)) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) if error_response and hasattr(error_response, "text"): diff --git a/litellm/llms/azure_ai/image_generation/cost_calculator.py b/litellm/llms/azure_ai/image_generation/cost_calculator.py index 08b732197f3..629c83b7054 100644 --- a/litellm/llms/azure_ai/image_generation/cost_calculator.py +++ b/litellm/llms/azure_ai/image_generation/cost_calculator.py @@ -1,5 +1,5 @@ from collections.abc import Mapping -from typing import Any, Final +from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import ( @@ -23,7 +23,7 @@ def _input_cost_per_pixel(resolved: ModelInfo) -> float: def cost_calculator( model: str, - image_response: Any, + image_response: object, size: str | None = None, n: int | None = None, optional_params: Mapping[str, object] | None = None, diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index f5631128f7d..39ff2fa0528 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -192,7 +192,7 @@ def blocked_responses_stream_usage(original_response: object) -> ResponseAPIUsag def effective_skip_system_message_for_guardrail(guardrail_to_apply: object) -> bool: - per: Final = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None) + per: Final[object] = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None) if per is not None: return bool(per) import litellm @@ -201,7 +201,7 @@ def effective_skip_system_message_for_guardrail(guardrail_to_apply: object) -> b def effective_skip_tool_message_for_guardrail(guardrail_to_apply: object) -> bool: - per: Final = getattr(guardrail_to_apply, "skip_tool_message_in_guardrail", None) + per: Final[object] = getattr(guardrail_to_apply, "skip_tool_message_in_guardrail", None) if per is not None: return bool(per) import litellm diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 1b1ea75572e..a03cd0df913 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -357,7 +357,7 @@ class BaseResponsesAPIConfig(ABC): """ if not isinstance(input, list): return input - out: Final[list[Any]] = [] + out: Final[list[object]] = [] for item in input: if isinstance(item, dict) and item.get("type") == "custom_tool_call": out.append({k: v for k, v in item.items() if k != "namespace"}) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 626b0d78589..c79b7b6e07d 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -283,7 +283,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): self, credential_args: Mapping[str, str | bool | tuple[AwsSessionTag, ...] | None], credential_fetcher: Callable[[], tuple[Credentials, int | None]], - ) -> Any: + ) -> Credentials: """ Read-through IAM cache on the process-wide ``DualCache``. diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 5600e25750f..886d6d1aea1 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -164,7 +164,7 @@ class AmazonConverseConfig(BaseConfig): topP: int | None = None, topK: int | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py index 9b3827974c5..3005bc82299 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py @@ -1,4 +1,5 @@ import types +from collections.abc import Mapping from typing import Final from litellm.llms.base_llm.chat.transformation import BaseConfig @@ -46,7 +47,7 @@ class AmazonAI21Config(AmazonInvokeConfig, BaseConfig): presencePenalty: dict | None = None, countPenalty: dict | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py index 0dfe590c2de..00c2eaf0f8f 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py @@ -1,4 +1,5 @@ import types +from collections.abc import Mapping from typing import Final from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( @@ -27,7 +28,7 @@ class AmazonCohereConfig(AmazonInvokeConfig, CohereChatConfig): temperature: float | None = None, return_likelihood: str | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py index 578e8ce88ac..d2eb6eef9ef 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py @@ -1,4 +1,5 @@ import types +from collections.abc import Mapping from typing import Final from litellm.llms.base_llm.chat.transformation import BaseConfig @@ -28,7 +29,7 @@ class AmazonLlamaConfig(AmazonInvokeConfig, BaseConfig): temperature: float | None = None, topP: int | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py index e17c302a00d..44511eeaa3e 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py @@ -1,4 +1,5 @@ import types +from collections.abc import Mapping from typing import TYPE_CHECKING, Final from litellm.llms.base_llm.chat.transformation import BaseConfig @@ -37,7 +38,7 @@ class AmazonMistralConfig(AmazonInvokeConfig, BaseConfig): top_k: float | None = None, stop: list[str] | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py index 2e19dbf77af..04248a785dd 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py @@ -6,6 +6,7 @@ Inherits from `AmazonInvokeConfig` Qwen3 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Final import httpx @@ -43,7 +44,7 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): top_k: int | None = None, stop: list[str] | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py index 50934c73d34..a18e97c63a6 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py @@ -1,5 +1,6 @@ import re import types +from collections.abc import Mapping from typing import Final import litellm @@ -33,7 +34,7 @@ class AmazonTitanConfig(AmazonInvokeConfig, BaseConfig): temperature: float | None = None, topP: int | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py index 1346f7b3dd1..b29ddf01d1e 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py @@ -1,4 +1,5 @@ import types +from collections.abc import Mapping from typing import Final import litellm @@ -36,7 +37,7 @@ class AmazonAnthropicConfig(AmazonInvokeConfig): top_p: int | None = None, anthropic_version: str | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 333bce0b967..238b1ccaa30 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -22,7 +22,7 @@ if TYPE_CHECKING: from litellm.types.llms.bedrock import BedrockCreateBatchRequest import httpx -from pydantic import TypeAdapter, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError import litellm from litellm import verbose_logger @@ -1691,6 +1691,9 @@ def get_bedrock_chat_config(model: str): return litellm.AmazonInvokeConfig() +_BOTOCORE_SERVICE_DESCRIPTION: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + + def _load_bedrock_response_stream_shape(): """ Load the ResponseStream shape from botocore's bundled bedrock-runtime schema. @@ -1703,7 +1706,9 @@ def _load_bedrock_response_stream_shape(): from botocore.model import ServiceModel loader: Final = Loader() - service_dict: Final = loader.load_service_model("bedrock-runtime", "service-2") + service_dict: Final = _BOTOCORE_SERVICE_DESCRIPTION.validate_python( + loader.load_service_model("bedrock-runtime", "service-2") + ) return ServiceModel(service_dict).shape_for("ResponseStream") except Exception as e: verbose_logger.warning( @@ -1856,9 +1861,12 @@ class BedrockEventStreamDecoderBase: return chunk.decode() +_JSON_VALUE: Final = TypeAdapter(object) + + def _decoded_json_value(raw: str) -> object: """Decode a JSON document into an opaque value for isinstance narrowing.""" - return json.loads(raw) + return _JSON_VALUE.validate_python(json.loads(raw)) def get_anthropic_beta_from_headers(headers: dict) -> list[str]: diff --git a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py index ee02754b6c6..143e0b2f623 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py @@ -10,6 +10,7 @@ Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-tit """ import types +from collections.abc import Mapping from typing import Final from litellm.types.llms.bedrock import ( @@ -27,7 +28,7 @@ class AmazonTitanG1Config: def __init__( self, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py index 8d7a19671b1..3f70bbc3bc3 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py @@ -10,6 +10,7 @@ Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-tit """ import types +from collections.abc import Mapping from typing import Final from litellm.types.llms.bedrock import ( @@ -31,7 +32,7 @@ class AmazonTitanV2Config: dimensions: int | None = None def __init__(self, normalize: bool | None = None, dimensions: int | None = None) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 4baa77f0ed7..02f02d97e63 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -855,8 +855,8 @@ class AmazonAnthropicClaudeMessagesConfig( @staticmethod def _merge_message_start_cache_into_delta_usage( - delta_usage: dict[str, Any], - start_usage: dict[str, Any] | None, + delta_usage: dict[str, object], + start_usage: Mapping[str, object] | None, ) -> None: """ Copy cache breakdown from message_start onto message_delta usage when @@ -885,7 +885,7 @@ class AmazonAnthropicClaudeMessagesConfig( """ _CACHE_FIELDS: Final = ("cache_creation_input_tokens", "cache_read_input_tokens") pending_delta: dict[str, Any] | None = None - start_usage_snapshot: dict[str, Any] | None = None + start_usage_snapshot: Mapping[str, object] | None = None async for chunk in completion_stream: if not isinstance(chunk, dict): @@ -898,7 +898,7 @@ class AmazonAnthropicClaudeMessagesConfig( chunk_type = chunk.get("type") if chunk_type == "message_start": - msg: dict[str, Any] = cast(dict[str, Any], chunk.get("message") or {}) + msg: dict[str, object] = cast(dict[str, Any], chunk.get("message") or {}) u = msg.get("usage") if isinstance(u, dict): start_usage_snapshot = dict(u) diff --git a/litellm/llms/codex/harness/transformation.py b/litellm/llms/codex/harness/transformation.py index 86a5f8274e6..4512e39189e 100644 --- a/litellm/llms/codex/harness/transformation.py +++ b/litellm/llms/codex/harness/transformation.py @@ -77,7 +77,7 @@ class CodexStreamState: started: set[str] = field(default_factory=set) # mutable-ok: parser records announced tool items -def _tool_input(item: Mapping[str, Any]) -> tuple[str, str, Mapping[str, Any], bool]: +def _tool_input(item: Mapping[str, Any]) -> tuple[str, str, Mapping[str, object], bool]: """(normalized name, native name, input, builtin) for a tool-like item.""" item_type = item.get("type") if item_type == "command_execution": diff --git a/litellm/llms/cometapi/chat/transformation.py b/litellm/llms/cometapi/chat/transformation.py index 6335b7be8e0..26050700efc 100644 --- a/litellm/llms/cometapi/chat/transformation.py +++ b/litellm/llms/cometapi/chat/transformation.py @@ -40,7 +40,7 @@ class CometAPIConfig(OpenAIGPTConfig): mapped_openai_params: Final = super().map_openai_params(non_default_params, optional_params, model, drop_params) # CometAPI-specific parameters (if any) - extra_body: Final[dict[str, Any]] = {} + extra_body: Final[dict[str, object]] = {} # TODO: Add CometAPI-specific parameter handling here # Example: # custom_param = non_default_params.pop("custom_param", None) diff --git a/litellm/llms/cometapi/image_generation/cost_calculator.py b/litellm/llms/cometapi/image_generation/cost_calculator.py index 8f767cf311e..d9248eac598 100644 --- a/litellm/llms/cometapi/image_generation/cost_calculator.py +++ b/litellm/llms/cometapi/image_generation/cost_calculator.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import resolve_image_model_info @@ -7,7 +7,7 @@ from litellm.types.utils import ImageResponse, ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/custom_httpx/container_handler.py b/litellm/llms/custom_httpx/container_handler.py index 01ad8755915..6c95816846d 100644 --- a/litellm/llms/custom_httpx/container_handler.py +++ b/litellm/llms/custom_httpx/container_handler.py @@ -11,6 +11,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import TypeAdapter from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm @@ -53,6 +54,9 @@ class EndpointsConfig(TypedDict): endpoints: ReadOnly[Sequence[EndpointConfig]] +_ENDPOINTS_CONFIG: Final = TypeAdapter(EndpointsConfig) + + class ContainerErrorDetail(TypedDict, total=False): """The ``error`` object of a container API error body.""" @@ -81,7 +85,7 @@ def _load_endpoints_config() -> EndpointsConfig: """Load the endpoints configuration from JSON file.""" config_path: Final = Path(__file__).parent.parent.parent / "containers" / "endpoints.json" with open(config_path) as f: - return json.load(f) + return _ENDPOINTS_CONFIG.validate_python(json.load(f)) def _get_endpoint_config(endpoint_name: str) -> EndpointConfig | None: diff --git a/litellm/llms/dashscope/chat/transformation.py b/litellm/llms/dashscope/chat/transformation.py index bcd2a5d5320..5ef482a6c2e 100644 --- a/litellm/llms/dashscope/chat/transformation.py +++ b/litellm/llms/dashscope/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to DashScope's `/v1/chat/complet """ from collections.abc import Coroutine -from typing import Any, Final, Literal, overload +from typing import Final, Literal, overload from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam @@ -33,7 +33,7 @@ class DashScopeChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -45,7 +45,7 @@ class DashScopeChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: if is_async: return super()._transform_messages(messages=messages, model=model, is_async=True) else: diff --git a/litellm/llms/deepinfra/chat/transformation.py b/litellm/llms/deepinfra/chat/transformation.py index 311c17d49f6..e746ce0a5c9 100644 --- a/litellm/llms/deepinfra/chat/transformation.py +++ b/litellm/llms/deepinfra/chat/transformation.py @@ -1,6 +1,6 @@ import json -from collections.abc import Coroutine -from typing import Any, Final, Literal, cast, overload +from collections.abc import Coroutine, Mapping +from typing import Final, Literal, cast, overload import litellm from litellm.constants import MIN_NON_ZERO_TEMPERATURE @@ -50,7 +50,7 @@ class DeepInfraConfig(OpenAIGPTConfig): tools: list | None = None, tool_choice: str | dict | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @@ -154,7 +154,7 @@ class DeepInfraConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -166,7 +166,7 @@ class DeepInfraConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Transform messages for DeepInfra compatibility. Handles both sync and async transformations. diff --git a/litellm/llms/deepseek/chat/transformation.py b/litellm/llms/deepseek/chat/transformation.py index 4e428a23392..f8813739ae4 100644 --- a/litellm/llms/deepseek/chat/transformation.py +++ b/litellm/llms/deepseek/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to DeepSeek's `/v1/chat/completi """ from collections.abc import Coroutine, Mapping, Sequence -from typing import Any, Final, Literal, cast, overload +from typing import Final, Literal, cast, overload import litellm from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -104,7 +104,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -116,7 +116,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ DeepSeek vision models accept image_url content blocks in user messages (https://api-docs.deepseek.com/guides/vision), so those diff --git a/litellm/llms/deepseek/messages/transformation.py b/litellm/llms/deepseek/messages/transformation.py index 85b9ac66b5f..ad58070df4a 100644 --- a/litellm/llms/deepseek/messages/transformation.py +++ b/litellm/llms/deepseek/messages/transformation.py @@ -94,7 +94,7 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): return f"{base_url}/v1/messages" @staticmethod - def _sanitize_tools_for_deepseek(tools: Any) -> Any: + def _sanitize_tools_for_deepseek(tools: Any) -> object: if not isinstance(tools, list): return tools diff --git a/litellm/llms/docker_model_runner/chat/transformation.py b/litellm/llms/docker_model_runner/chat/transformation.py index b1e2c9638c5..303710a7108 100644 --- a/litellm/llms/docker_model_runner/chat/transformation.py +++ b/litellm/llms/docker_model_runner/chat/transformation.py @@ -5,7 +5,7 @@ Docker Model Runner API Reference: https://docs.docker.com/ai/model-runner/api-r """ from collections.abc import Coroutine -from typing import Any, Final, Literal, overload +from typing import Final, Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -27,7 +27,7 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -39,7 +39,7 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Docker Model Runner is OpenAI-compatible, so we use standard message transformation. """ diff --git a/litellm/llms/fal_ai/image_generation/bytedance_transformation.py b/litellm/llms/fal_ai/image_generation/bytedance_transformation.py index 8e7630f2cbe..44cf816c1aa 100644 --- a/litellm/llms/fal_ai/image_generation/bytedance_transformation.py +++ b/litellm/llms/fal_ai/image_generation/bytedance_transformation.py @@ -60,7 +60,7 @@ class FalAIBytedanceBaseConfig(FalAIFluxProV11UltraConfig): return optional_params - def _map_image_size(self, size: Any) -> Any: + def _map_image_size(self, size: Any) -> object: if isinstance(size, dict): return size diff --git a/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py b/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py index 5119169467c..48b5d392747 100644 --- a/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py @@ -68,7 +68,7 @@ class FalAIFluxProV11Config(FalAIFluxProV11UltraConfig): return optional_params - def _map_image_size(self, size: Any) -> Any: + def _map_image_size(self, size: Any) -> object: if isinstance(size, dict): return size if not isinstance(size, str): diff --git a/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py b/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py index fe458153d9c..ed26cc4bc12 100644 --- a/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py @@ -65,7 +65,7 @@ class FalAIFluxSchnellConfig(FalAIFluxProV11UltraConfig): return optional_params - def _map_image_size(self, size: Any) -> Any: + def _map_image_size(self, size: Any) -> object: if isinstance(size, dict): return size diff --git a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py index 970addbd5a2..2ca423fafff 100644 --- a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py @@ -92,7 +92,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): return optional_params - def _map_image_size(self, size: Any) -> Any: + def _map_image_size(self, size: Any) -> object: if isinstance(size, dict): width = size.get("width") height = size.get("height") diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 29a989a5cb0..e8f17fe0ba9 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -189,7 +189,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): top_p=top_p, response_format=response_format, ) - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index cae0ba49c3f..f833837f97c 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Final, cast import litellm @@ -68,7 +69,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): candidate_count: int | None = None, stop_sequences: list | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index 78e6e6aaf82..c44db8cac26 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -6,6 +6,7 @@ from collections.abc import Mapping, Sequence from typing import Any, Final import httpx +from pydantic import TypeAdapter import litellm from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -128,9 +129,12 @@ def is_gemini_image_model(model: str) -> bool: return "gemini" in base_model +_JSON_VALUE: Final = TypeAdapter(object) + + def _parse_image_config_string(raw_image_config: str, model: str) -> object: try: - return json.loads(raw_image_config) + return _JSON_VALUE.validate_python(json.loads(raw_image_config)) except json.JSONDecodeError as exc: raise litellm.UnsupportedParamsError( model=model, diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index c2f0ef473ae..67e84f6da24 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -13,7 +13,7 @@ else: class GoogleAIStudioTokenCounter: - def _clean_contents_for_gemini_api(self, contents: Any) -> Any: + def _clean_contents_for_gemini_api(self, contents: Any) -> object: """ Clean up contents to remove unsupported fields for the Gemini API. diff --git a/litellm/llms/gemini/image_edit/cost_calculator.py b/litellm/llms/gemini/image_edit/cost_calculator.py index 321fbbaeb37..615c4da9cea 100644 --- a/litellm/llms/gemini/image_edit/cost_calculator.py +++ b/litellm/llms/gemini/image_edit/cost_calculator.py @@ -2,8 +2,6 @@ Gemini Image Edit Cost Calculator """ -from typing import Any - from litellm.llms.gemini.image_generation.cost_calculator import ( cost_calculator as image_generation_cost_calculator, ) @@ -12,7 +10,7 @@ from litellm.types.utils import ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/gemini/image_generation/cost_calculator.py b/litellm/llms/gemini/image_generation/cost_calculator.py index e3232f8fb7b..4522fee8c46 100644 --- a/litellm/llms/gemini/image_generation/cost_calculator.py +++ b/litellm/llms/gemini/image_generation/cost_calculator.py @@ -2,7 +2,7 @@ Google AI Image Generation Cost Calculator """ -from typing import Any, Final +from typing import Final from litellm.litellm_core_utils.llm_cost_calc.utils import ( calculate_image_response_cost_from_usage, @@ -14,7 +14,7 @@ from litellm.types.utils import ImageResponse, ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/gemini/image_usage_transformation.py b/litellm/llms/gemini/image_usage_transformation.py index 9bce72513c7..5f08d8403af 100644 --- a/litellm/llms/gemini/image_usage_transformation.py +++ b/litellm/llms/gemini/image_usage_transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Final from litellm.types.utils import ImageUsage, ImageUsageInputTokensDetails @@ -51,7 +51,7 @@ def transform_gemini_image_usage(usage_metadata: dict) -> ImageUsage: if output_tokens > known_output_tokens: output_tokens_details.text_tokens += output_tokens - known_output_tokens - usage_payload: Final[dict[str, Any]] = { + usage_payload: Final[dict[str, object]] = { "input_tokens": usage_metadata.get("promptTokenCount", 0), "input_tokens_details": input_tokens_details, "output_tokens": output_tokens, @@ -62,4 +62,4 @@ def transform_gemini_image_usage(usage_metadata: dict) -> ImageUsage: "completion_tokens_details": output_tokens_details.model_dump(), "output_tokens_details": output_tokens_details.model_dump(), } - return ImageUsage(**usage_payload) + return ImageUsage.model_validate(usage_payload) diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 169b9a037a5..568cd4365cc 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -1,6 +1,7 @@ import json import os -from typing import TYPE_CHECKING, Any, Final +from collections.abc import Sequence +from typing import TYPE_CHECKING, Final import httpx @@ -177,8 +178,8 @@ class GithubCopilotConfig(OpenAIConfig): @staticmethod def _parse_anthropic_native_content( - content_blocks: list[Any], - ) -> tuple[str, list[ChatCompletionToolCallChunk], list[Any] | None]: + content_blocks: list[object], + ) -> tuple[str, list[ChatCompletionToolCallChunk], Sequence[object] | None]: """ Parse Anthropic-native content blocks into OpenAI-compatible fields. @@ -225,7 +226,7 @@ class GithubCopilotConfig(OpenAIConfig): content = "" tool_calls: list[ChatCompletionToolCallChunk] = [] - thinking_blocks: list[Any] | None = None + thinking_blocks: Sequence[object] | None = None raw_content: Final = response_json.get("content") if isinstance(raw_content, list): content, tool_calls, thinking_blocks = cls._parse_anthropic_native_content(raw_content) diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index cb3f1ac2d0c..e3a60acebe2 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to Groq's `/v1/chat/completions` """ -from collections.abc import AsyncIterator, Coroutine, Iterator +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload import httpx @@ -72,7 +72,7 @@ class GroqChatConfig(OpenAILikeChatConfig): tools: list | None = None, tool_choice: str | dict | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @@ -125,7 +125,7 @@ class GroqChatConfig(OpenAILikeChatConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -137,7 +137,7 @@ class GroqChatConfig(OpenAILikeChatConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: for idx, message in enumerate(messages): """ 1. Don't pass 'null' function_call assistant message to groq - https://github.com/BerriAI/litellm/issues/5839 diff --git a/litellm/llms/heroku/chat/transformation.py b/litellm/llms/heroku/chat/transformation.py index fd0c29b080b..23b458304ab 100644 --- a/litellm/llms/heroku/chat/transformation.py +++ b/litellm/llms/heroku/chat/transformation.py @@ -6,7 +6,7 @@ this is OpenAI compatible - no translation needed / occurs import os from collections.abc import Coroutine -from typing import Any, Literal, overload +from typing import Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -24,7 +24,7 @@ class HerokuChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -36,7 +36,7 @@ class HerokuChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Heroku does not support content in list format. See: https://devcenter.heroku.com/articles/heroku-inference-api-v1-chat-completions#content-object diff --git a/litellm/llms/langflow/chat/transformation.py b/litellm/llms/langflow/chat/transformation.py index a53887b36af..c5bb293951a 100644 --- a/litellm/llms/langflow/chat/transformation.py +++ b/litellm/llms/langflow/chat/transformation.py @@ -159,7 +159,7 @@ class LangFlowConfig(BaseConfig): input_value: Final = self._get_last_user_message(messages) - payload: Final[dict[str, Any]] = { + payload: Final[dict[str, object]] = { "input_value": input_value, "input_type": optional_params.get("input_type", "chat"), "output_type": optional_params.get("output_type", "chat"), diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index 07c90a73a9d..36c761e308e 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -135,7 +135,7 @@ class LangGraphConfig(BaseConfig): return parts[1] return model - def _convert_messages_to_langgraph_format(self, messages: list[AllMessageValues]) -> list[dict[str, Any]]: + def _convert_messages_to_langgraph_format(self, messages: list[AllMessageValues]) -> list[dict[str, object]]: """ Convert OpenAI-format messages to LangGraph format. @@ -144,7 +144,7 @@ class LangGraphConfig(BaseConfig): Preserves per-message ``metadata`` when present (e.g. A2A ``skillId``). """ - langgraph_messages: Final[list[dict[str, Any]]] = [] + langgraph_messages: Final[list[dict[str, object]]] = [] for msg in messages: role = msg.get("role", "user") content = msg.get("content", "") @@ -167,7 +167,7 @@ class LangGraphConfig(BaseConfig): if not isinstance(content, str): content = str(content) - langgraph_message: dict[str, Any] = { + langgraph_message: dict[str, object] = { "role": langgraph_role, "content": content, } @@ -202,7 +202,7 @@ class LangGraphConfig(BaseConfig): assistant_id: Final = self._get_assistant_id(model, optional_params) langgraph_messages: Final = self._convert_messages_to_langgraph_format(messages) - payload: Final[dict[str, Any]] = { + payload: Final[dict[str, object]] = { "assistant_id": assistant_id, "input": {"messages": langgraph_messages}, } diff --git a/litellm/llms/litellm_proxy/skills/prompt_injection.py b/litellm/llms/litellm_proxy/skills/prompt_injection.py index 4fdaa468ca1..2ba18f130e6 100644 --- a/litellm/llms/litellm_proxy/skills/prompt_injection.py +++ b/litellm/llms/litellm_proxy/skills/prompt_injection.py @@ -289,7 +289,7 @@ class SkillPromptInjectionHandler: if len(description) > max_desc_length: description = description[: max_desc_length - 3] + "..." - input_schema: dict[str, Any] = { + input_schema: dict[str, object] = { "type": "object", "properties": {}, "required": [], diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 33b567e9710..b78797c35a3 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -6,8 +6,8 @@ Why separate file? Make it easy to see how transformation works Docs - https://docs.mistral.ai/api/ """ -from collections.abc import AsyncIterator, Coroutine, Iterator -from typing import TYPE_CHECKING, Any, Final, Literal, cast, get_type_hints, overload +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping +from typing import TYPE_CHECKING, Final, Literal, cast, get_type_hints, overload import httpx @@ -99,7 +99,7 @@ class MistralConfig(OpenAIGPTConfig): response_format: dict | None = None, stop: str | list | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @@ -229,7 +229,7 @@ class MistralConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload @@ -244,7 +244,7 @@ class MistralConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ - handles scenario where content is list and not string - content list is just text, and no images diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py index d345b8efc56..27c3bff5346 100644 --- a/litellm/llms/modelscope/chat/transformation.py +++ b/litellm/llms/modelscope/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/comple """ from collections.abc import Coroutine -from typing import Any, Final, Literal, cast, overload +from typing import Final, Literal, cast, overload from typing_extensions import override @@ -27,7 +27,7 @@ class ModelScopeChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -39,7 +39,7 @@ class ModelScopeChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Flatten text-only content lists to strings for ModelScope. diff --git a/litellm/llms/moonshot/chat/transformation.py b/litellm/llms/moonshot/chat/transformation.py index 7b0fcd24770..8e428d93b2d 100644 --- a/litellm/llms/moonshot/chat/transformation.py +++ b/litellm/llms/moonshot/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Moonshot AI's `/v1/chat/compl """ from collections.abc import Coroutine, Mapping -from typing import Any, Final, Literal, cast, overload +from typing import Final, Literal, cast, overload import litellm from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -29,7 +29,7 @@ class MoonshotChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -41,7 +41,7 @@ class MoonshotChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Moonshot text-only models don't support content in list format. Multimodal models (kimi-k2.5, kimi-latest, etc.) accept the diff --git a/litellm/llms/nvidia_riva/audio_transcription/handler.py b/litellm/llms/nvidia_riva/audio_transcription/handler.py index bea77a6761c..651a2215052 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/handler.py +++ b/litellm/llms/nvidia_riva/audio_transcription/handler.py @@ -250,7 +250,7 @@ class NvidiaRivaAudioTranscription: message="NvidiaRivaAudioTranscriptionConfig produced an unexpected request payload type.", ) - recognition_config_dict: Final[dict[str, Any]] = request_payload["recognition_config"] + recognition_config_dict: Final[dict[str, object]] = request_payload["recognition_config"] # The wire format is fixed by our resampler; override anything stale # the caller passed in so the gRPC config matches the bytes we send. recognition_config_dict["sample_rate_hertz"] = RIVA_TARGET_SAMPLE_RATE_HZ diff --git a/litellm/llms/nvidia_riva/common_utils.py b/litellm/llms/nvidia_riva/common_utils.py index a9737822b17..b654e725dc8 100644 --- a/litellm/llms/nvidia_riva/common_utils.py +++ b/litellm/llms/nvidia_riva/common_utils.py @@ -2,7 +2,7 @@ Common utilities and exceptions for the NVIDIA Riva STT provider """ -from typing import Any, Final +from typing import Final from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -41,7 +41,7 @@ _GRPC_STATUS_CODE_TO_HTTP: Final[dict] = { } -def _extract_grpc_status_name(error: Any) -> str | None: +def _extract_grpc_status_name(error: object) -> str | None: """ Best-effort extraction of a gRPC StatusCode name from an arbitrary error. @@ -60,7 +60,7 @@ def _extract_grpc_status_name(error: Any) -> str | None: return None -def _extract_grpc_details(error: Any) -> str | None: +def _extract_grpc_details(error: object) -> str | None: """Best-effort extraction of a human-readable detail string from a gRPC error.""" details_fn: Final = getattr(error, "details", None) if callable(details_fn): diff --git a/litellm/llms/oci/chat/generic.py b/litellm/llms/oci/chat/generic.py index 8ff8ef9abc4..8db5803deb5 100644 --- a/litellm/llms/oci/chat/generic.py +++ b/litellm/llms/oci/chat/generic.py @@ -8,7 +8,7 @@ parsing, and streaming chunk parsing for models served with import datetime import hashlib -from typing import Any, Final +from typing import Final import httpx from pydantic import ValidationError @@ -404,7 +404,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: # same minimal ``{"id", "type", "function": {"name", "arguments"}}`` # shape keeps downstream stream-mergers behaving identically across # GENERIC and Cohere chunks. - tool_calls: list[dict[str, Any]] | None = None + tool_calls: list[dict[str, object]] | None = None if typed_chunk.message and typed_chunk.message.toolCalls: tool_calls = [ { diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 01b39ba3c40..e3c4b8d93f2 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -197,7 +197,7 @@ def _normalize_response_format(selected_params: dict, vendor: OCIVendors) -> Non if vendor == OCIVendors.COHERE: # OCI Cohere has no JSON_SCHEMA type; a schema rides on JSON_OBJECT. - payload: Final[dict[str, Any]] = {"type": "JSON_OBJECT"} + payload: Final[dict[str, object]] = {"type": "JSON_OBJECT"} if json_schema is not None and json_schema.get("schema") is not None: payload["schema"] = json_schema["schema"] selected_params["responseFormat"] = payload @@ -212,7 +212,7 @@ def _normalize_response_format(selected_params: dict, vendor: OCIVendors) -> Non # OCI's ResponseJsonSchema accepts only name/description/schema/isStrict. # OpenAI sends `strict` instead of `isStrict`; forwarding it (or any # other extra key) makes OCI reject the whole request with HTTP 400. - oci_schema: Final[dict[str, Any]] = {"name": json_schema.get("name") or "response"} + oci_schema: Final[dict[str, object]] = {"name": json_schema.get("name") or "response"} if json_schema.get("description") is not None: oci_schema["description"] = json_schema["description"] if json_schema.get("schema") is not None: diff --git a/litellm/llms/ollama/completion/handler.py b/litellm/llms/ollama/completion/handler.py index 449952217b5..ad2f15cdd7c 100644 --- a/litellm/llms/ollama/completion/handler.py +++ b/litellm/llms/ollama/completion/handler.py @@ -87,7 +87,7 @@ async def ollama_aembeddings( prompts: list[str], model_response: EmbeddingResponse, optional_params: dict, - logging_obj: Any, + logging_obj: object, encoding: TokenEncoder | None, ): if not api_base.endswith("/api/embed"): @@ -114,7 +114,7 @@ def ollama_embeddings( prompts: list[str], optional_params: dict, model_response: EmbeddingResponse, - logging_obj: Any, + logging_obj: object, encoding: TokenEncoder | None = None, ): if not api_base.endswith("/api/embed"): diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index e7a7cda0b67..2dfb397f427 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -135,7 +135,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): top_p: int | None = None, response_format: dict | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index aee439862db..249d6008f87 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -235,7 +235,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def _not_run_reason( self, - messages: Sequence[dict[str, Any]], # mutable-ok: raw request messages consumed by _extract_inputs + messages: Sequence[Mapping[str, object]], ) -> str | None: """Why nothing was scanned, or None when the only unscoped content is images, which this handler never scans.""" texts: Final[list[str]] = [] # mutable-ok: filled by _extract_inputs diff --git a/litellm/llms/openai/chat/o_series_transformation.py b/litellm/llms/openai/chat/o_series_transformation.py index 34122188855..c7552d47985 100644 --- a/litellm/llms/openai/chat/o_series_transformation.py +++ b/litellm/llms/openai/chat/o_series_transformation.py @@ -12,7 +12,7 @@ Translations handled by LiteLLM: """ from collections.abc import Coroutine -from typing import Any, Final, Literal, cast, overload +from typing import Final, Literal, cast, overload import litellm from litellm import verbose_logger @@ -126,7 +126,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -138,7 +138,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Handles limitations of O-1 model family. - modalities: image => drop param (if user opts in to dropping param) diff --git a/litellm/llms/openai/completion/handler.py b/litellm/llms/openai/completion/handler.py index c7b59509eb0..f9267873ea2 100644 --- a/litellm/llms/openai/completion/handler.py +++ b/litellm/llms/openai/completion/handler.py @@ -157,7 +157,7 @@ class OpenAITextCompletion(BaseLLM): status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) error_text: Final = getattr(e, "text", str(e)) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers) @@ -210,7 +210,7 @@ class OpenAITextCompletion(BaseLLM): status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) error_text: Final = getattr(e, "text", str(e)) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers) @@ -248,7 +248,7 @@ class OpenAITextCompletion(BaseLLM): status_code = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) + error_response: object = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers) @@ -315,7 +315,7 @@ class OpenAITextCompletion(BaseLLM): status_code: Final = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) error_text: Final = getattr(e, "text", str(e)) - error_response: Final = getattr(e, "response", None) + error_response: Final[object] = getattr(e, "response", None) if error_headers is None and error_response: error_headers = getattr(error_response, "headers", None) raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers) diff --git a/litellm/llms/openai/completion/transformation.py b/litellm/llms/openai/completion/transformation.py index 383a67fd913..19ba79720a5 100644 --- a/litellm/llms/openai/completion/transformation.py +++ b/litellm/llms/openai/completion/transformation.py @@ -2,6 +2,7 @@ Support for gpt model family """ +from collections.abc import Mapping from typing import Final from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig @@ -69,7 +70,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): temperature: float | None = None, top_p: float | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 434df9fae6a..53f57456126 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -549,7 +549,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): ) # Create ResponseReasoningItem object from the item data - reasoning_item: Final = ResponseReasoningItem(**item_data) + reasoning_item: Final = ResponseReasoningItem.model_validate(item_data) # Convert back to dict with exclude_none=True to exclude None fields dict_reasoning_item: Final = reasoning_item.model_dump(exclude_none=True) diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index 07d3d078180..b8517a4257c 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -3,7 +3,7 @@ Dynamic configuration class generator for JSON-based providers. """ from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Final, Literal, overload +from typing import TYPE_CHECKING, Final, Literal, overload from litellm._logging import verbose_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -30,7 +30,7 @@ def create_config_class(provider: SimpleProviderConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -42,7 +42,7 @@ def create_config_class(provider: SimpleProviderConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """Transform messages based on special_handling config""" # Handle content list to string conversion if configured diff --git a/litellm/llms/opencode/harness/transformation.py b/litellm/llms/opencode/harness/transformation.py index 158af0fae0a..787bc99a2d7 100644 --- a/litellm/llms/opencode/harness/transformation.py +++ b/litellm/llms/opencode/harness/transformation.py @@ -252,7 +252,7 @@ def build_opencode_config( extra_skills: Final = (skills_path,) if skills_path else () skill_paths: Final = [*(user_skills.get("paths") or ()), *extra_skills] # mutable-ok: opencode config JSON options: Final = {"baseURL": base_url, "apiKey": "{file:" + token_path + "}"} # mutable-ok: opencode config JSON - models: Final[dict[str, Any]] = {model: {}} # mutable-ok: opencode config JSON + models: Final[dict[str, Mapping[str, object]]] = {model: {}} # mutable-ok: opencode config JSON provider: Final = { # mutable-ok: opencode config JSON "npm": OPENCODE_PROVIDER_NPM, "name": "LiteLLM", diff --git a/litellm/llms/recraft/cost_calculator.py b/litellm/llms/recraft/cost_calculator.py index 221d4a35913..052e09a30c4 100644 --- a/litellm/llms/recraft/cost_calculator.py +++ b/litellm/llms/recraft/cost_calculator.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import resolve_image_model_info @@ -7,7 +7,7 @@ from litellm.types.utils import ImageResponse, ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/runwayml/cost_calculator.py b/litellm/llms/runwayml/cost_calculator.py index 07f7d564ac2..729ca516213 100644 --- a/litellm/llms/runwayml/cost_calculator.py +++ b/litellm/llms/runwayml/cost_calculator.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import resolve_image_model_info @@ -7,7 +7,7 @@ from litellm.types.utils import ImageResponse, ModelInfo def cost_calculator( model: str, - image_response: Any, + image_response: object, model_info: ModelInfo | None = None, ) -> float: """ diff --git a/litellm/llms/sambanova/chat.py b/litellm/llms/sambanova/chat.py index 0e7dadf7062..0350950f342 100644 --- a/litellm/llms/sambanova/chat.py +++ b/litellm/llms/sambanova/chat.py @@ -4,8 +4,8 @@ Sambanova Chat Completions API this is OpenAI compatible - no translation needed / occurs """ -from collections.abc import Coroutine -from typing import Any, Final, Literal, overload +from collections.abc import Coroutine, Mapping +from typing import Final, Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -45,7 +45,7 @@ class SambanovaConfig(OpenAIGPTConfig): tool_choice: str | None = None, tools: list | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[Mapping[str, object]] = dict(locals()) for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @@ -101,7 +101,7 @@ class SambanovaConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -113,7 +113,7 @@ class SambanovaConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Transform messages to handle content list conversion. diff --git a/litellm/llms/sap/credentials.py b/litellm/llms/sap/credentials.py index a2a93b6114a..b5ff34b81d7 100644 --- a/litellm/llms/sap/credentials.py +++ b/litellm/llms/sap/credentials.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import os import tempfile -from collections.abc import Callable, Sequence +from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path @@ -76,7 +76,7 @@ def _load_vcap() -> dict[str, Any]: return _load_json_env(VCAP_SERVICES_ENV_VAR) or {} -def _get_vcap_service(label: str) -> dict[str, Any] | None: +def _get_vcap_service(label: str) -> Mapping[str, object] | None: for services in _load_vcap().values(): for svc in services: if svc.get("label") == label: diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index fd1ff24fd52..81b2c3d5109 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -11,7 +11,7 @@ API Reference: import json from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Final, Optional, Union import httpx @@ -439,7 +439,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): from litellm.utils import CustomStreamWrapper if client is None or not isinstance(client, AsyncHTTPHandler): - client = get_async_httpx_client(llm_provider=cast(Any, "vertex_ai"), params={}) + client = get_async_httpx_client(llm_provider="vertex_ai", params={}) # Avoid logging sensitive api_base directly verbose_logger.debug("Making async streaming request to Vertex AI endpoint.") diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 2b0694697a4..a7e0bef37df 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -770,7 +770,7 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]: def _iter_openai_jsonl_entries( openai_file_content: FileTypes, -) -> Iterator[dict[str, Any]]: +) -> Iterator[dict[str, object]]: for line in _iter_openai_jsonl_lines(openai_file_content): yield json.loads(line) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 2521a462fc0..d641d3da454 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -65,7 +65,7 @@ from ..common_utils import ( # Typed as Any to avoid introducing a module-load-time cyclic import to # vertex_llm_base. The instance is lazily constructed by _get_vertex_base() # the first time GCS metadata needs to be fetched. -_GCS_METADATA_VERTEX_BASE: Any | None = None +_GCS_METADATA_VERTEX_BASE: object | None = None # Shared sync client for GCS JSON API metadata reads so proxy/SSL settings # from litellm's HTTP stack apply (see Greptile review on PR #27278). _GCS_METADATA_HTTP_HANDLER: HTTPHandler | None = None diff --git a/litellm/llms/vertex_ai/image_edit/cost_calculator.py b/litellm/llms/vertex_ai/image_edit/cost_calculator.py index 7e67ba8f317..7d1c2c1fc51 100644 --- a/litellm/llms/vertex_ai/image_edit/cost_calculator.py +++ b/litellm/llms/vertex_ai/image_edit/cost_calculator.py @@ -2,7 +2,7 @@ Vertex AI Image Edit Cost Calculator """ -from typing import Any, Final +from typing import Final import litellm from litellm.types.utils import ImageResponse @@ -10,7 +10,7 @@ from litellm.types.utils import ImageResponse def cost_calculator( model: str, - image_response: Any, + image_response: object, ) -> float: """ Vertex AI image edit cost calculator. diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index abd47173608..08d95430041 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -143,7 +143,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): def _unwrap_predictions_response( self, response_json: dict[str, Any], - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Unwrap the Vertex Gemma predictions format to OpenAI format. diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 49c8ee2ac55..86073eb672d 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -3,6 +3,7 @@ from types import MappingProxyType from typing import Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -30,6 +31,8 @@ from ...openai.chat.gpt_transformation import ( OpenAIGPTConfig, ) +_RESPONSE_BODY: Final = TypeAdapter(dict[object, object], config=ConfigDict(hide_input_in_errors=True)) + def _usage_restated_from_xai_ticks(usage: Usage | None) -> Usage | None: reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None)) @@ -289,7 +292,7 @@ class XAIChatConfig(OpenAIGPTConfig): # Handle X.AI web search usage tracking try: - raw_response_json: Final = raw_response.json() + raw_response_json: Final = _RESPONSE_BODY.validate_python(raw_response.json()) self._enhance_usage_with_xai_web_search_fields(response, raw_response_json) except Exception as e: verbose_logger.debug("Error extracting X.AI web search usage: %s", e) diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 3b5eda6cff2..ccbf69d0f54 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -3,7 +3,7 @@ import os from ipaddress import ip_address -from typing import TYPE_CHECKING, Any, Final, NoReturn +from typing import TYPE_CHECKING, Final, NoReturn from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit from fastapi import HTTPException, Request @@ -61,7 +61,7 @@ def _oauth_invalid_request( error_description: str, *, hint: str | None = None, - **extra: Any, + **extra: object, ) -> NoReturn: """Raise ``invalid_request`` (RFC 6749) with a debuggable description. @@ -69,7 +69,7 @@ def _oauth_invalid_request( ``invalid_request``; ``error_description`` and ``hint`` explain what failed and how to fix it (e.g. reverse-proxy / PROXY_BASE_URL issues). """ - detail: Final[dict[str, Any]] = { + detail: Final[dict[str, object]] = { "error": "invalid_request", "error_description": error_description, } diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 464e3043458..23a6e193d5d 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -686,7 +686,7 @@ def parse_admin_env_vars( Unknown / malformed entries are skipped silently. """ global_values: Final[dict[str, str]] = {} - user_specs: Final[list[dict[str, Any]]] = [] + user_specs: Final[list[dict[str, object]]] = [] if not env_vars: return global_values, user_specs for raw in env_vars: diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 4a7ca257d9e..f018029058f 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -353,7 +353,7 @@ class LazyFeatureMiddleware: async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: # Short-circuit once every feature has loaded. if scope["type"] in ("http", "websocket") and len(self._loaded) < len(self._features): - path = scope.get("path", "") + path: str = scope.get("path", "") # Strip the request's root_path so prefix matching works under a # server root path. Without this, requests like /api/v1/policies/... # never match the registered prefixes (/policies/...) and lazy diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index dca48babc81..40476230879 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -674,7 +674,7 @@ async def invoke_agent_a2a( ) body: dict[str, Any] = {} - request_data: dict[str, Any] = body + request_data: dict[str, object] = body try: body = await request.json() request_data = body @@ -892,7 +892,7 @@ async def invoke_agent_a2a( logging_obj._enqueue_deferred_logging = None _enqueue_fn() - response_dict: Final[dict[str, Any]] = ( + response_dict: Final[dict[str, object]] = ( response.model_dump(mode="json", exclude_none=True) if hasattr(response, "model_dump") else response diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index cfa9d10b74c..d600fb3486b 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -16,7 +16,7 @@ import re import time from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence from dataclasses import dataclass -from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast +from typing import Final, Literal, NoReturn, Protocol, TypeVar, cast import httpx import jwt @@ -428,7 +428,7 @@ class JWTHandler: default-team behavior should still go through ``get_team_id``. """ team_ids: Final[list[str]] = list(self.get_team_ids_from_jwt(token)) - singular: Any = None + singular: object = None if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_ID_CLAIM): singular = token.get(self.LITELLM_TEAM_ID_CLAIM) elif self.litellm_jwtauth.team_id_jwt_field is not None: @@ -1151,7 +1151,7 @@ class JWTHandler: # its jwks_uri, matching JWTIssuerConfig.jwks_url's documented fallback. return f"{issuer_config.issuer.rstrip('/')}/.well-known/openid-configuration" - def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> Any: + def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> object: """Resolve a mapped claim from ``token``. Returns ``None`` when the field is absent or empty so that mapped claims @@ -1480,7 +1480,7 @@ class JWTAuthManager: litellm_proxy_roles=jwt_handler.litellm_jwtauth, ) if not is_allowed: - allowed_routes: Final[list[Any]] = jwt_handler.litellm_jwtauth.admin_allowed_routes + allowed_routes: Final = jwt_handler.litellm_jwtauth.admin_allowed_routes actual_routes: Final = get_actual_routes(allowed_routes=allowed_routes) raise Exception(f"Admin not allowed to access this route. Route={route}, Allowed Routes={actual_routes}") diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index 64608567f92..a404b6926bd 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -7,6 +7,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Final import httpx +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.constants import NON_LLM_CONNECTION_TIMEOUT @@ -19,6 +20,7 @@ if TYPE_CHECKING: AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router" LICENSE_ALL_FEATURES: Final = "*" AUTO_ROUTER_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit." +_LICENSE_VERDICT: Final = TypeAdapter(object) class LicenseCheck: @@ -79,9 +81,7 @@ class LicenseCheck: if response is None: raise Exception("No response from license server") - response_json: Final = response.json() - - premium: Final = response_json["verify"] + premium: Final = _LICENSE_VERDICT.validate_python(response.json()["verify"]) assert isinstance(premium, bool) diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index de2ca4762f1..f48f5dd7162 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -333,11 +333,11 @@ def expand_wildcard_deployments_for_model_info( on top of that: a wildcard deployment like model_name="*" / litellm_params.model="openai/*" becomes one entry per known openai model, matching /v1/models behaviour. """ - expanded: Final[list[dict[str, Any]]] = [] + expanded: Final[list[dict[str, object]]] = [] for deployment in deployments: model_name = str(deployment.get("model_name") or "") raw_params = deployment.get("litellm_params") - litellm_params_dict: dict[str, Any] = raw_params if isinstance(raw_params, dict) else {} + litellm_params_dict: dict[str, object] = raw_params if isinstance(raw_params, dict) else {} litellm_model = str(litellm_params_dict.get("model") or "") # Determine the wildcard pattern to expand. diff --git a/litellm/proxy/auth/network.py b/litellm/proxy/auth/network.py index 230dde072ca..98f15a20a19 100644 --- a/litellm/proxy/auth/network.py +++ b/litellm/proxy/auth/network.py @@ -2,7 +2,7 @@ from __future__ import annotations import ipaddress from collections.abc import Sequence -from typing import Any, Final +from typing import Final from fastapi import Request from pydantic import Field @@ -24,7 +24,7 @@ class TrustedProxyConfig(LiteLLMBaseModel): trusted_proxy_cidrs: Sequence[str] = Field(default_factory=tuple) -def normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs") -> list[str]: +def normalize_cidr_ranges(configured_ranges: object, *, setting_name: str = "trusted_proxy_cidrs") -> list[str]: if not configured_ranges: return [] if isinstance(configured_ranges, str): @@ -40,7 +40,7 @@ def normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str = "truste def parse_trusted_proxy_ranges( - configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs" + configured_ranges: object, *, setting_name: str = "trusted_proxy_cidrs" ) -> list[TrustedProxyNetwork]: networks: Final[list[TrustedProxyNetwork]] = [] for cidr in normalize_cidr_ranges(configured_ranges, setting_name=setting_name): diff --git a/litellm/proxy/auth/trusted_proxy_utils.py b/litellm/proxy/auth/trusted_proxy_utils.py index 54b2e1124e2..55cf3e98444 100644 --- a/litellm/proxy/auth/trusted_proxy_utils.py +++ b/litellm/proxy/auth/trusted_proxy_utils.py @@ -12,7 +12,7 @@ from litellm.proxy.auth.network import ( TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges" -def _get_proxy_general_settings() -> dict[str, Any]: +def _get_proxy_general_settings() -> dict[str, object]: try: from litellm.proxy.proxy_server import general_settings diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py index 593f8b797a5..c4615a3cebf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py @@ -11,7 +11,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import json import sys -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Final, Literal from fastapi import HTTPException @@ -70,7 +70,7 @@ class AporiaGuardrail(CustomGuardrail): return new_messages async def prepare_aporia_request(self, new_messages: list[dict], response_string: str | None = None) -> dict: - data: Final[dict[str, Any]] = {} + data: Final[dict[str, object]] = {} if new_messages is not None: data["messages"] = new_messages if response_string is not None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index 597722eb1fb..b69028594b2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -60,7 +60,9 @@ class AzureGuardrailBase: self.api_base = api_base self.api_version: str | None = kwargs.get("api_version") - async def _post_to_content_safety(self, endpoint_path: str, request_body: dict[str, object]) -> dict[str, Any]: + async def _post_to_content_safety( + self, endpoint_path: str, request_body: dict[str, object] + ) -> Mapping[str, object]: """POST to an Azure Content Safety endpoint with standard auth headers. Args: @@ -85,7 +87,7 @@ class AzureGuardrailBase: json=request_body, timeout=self.timeout, ) - response_json: Final[dict[str, Any]] = response.json() + response_json: Final[dict[str, object]] = response.json() verbose_proxy_logger.debug("Azure Content Safety response [%s]: %s", endpoint_path, response_json) return response_json diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index f900d14bdc1..c5104254c4e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2037,7 +2037,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): @staticmethod def _sanitize_invoke_checks_response_for_logging( response: BedrockGuardrailChecksResponse, - ) -> dict[str, Any]: + ) -> dict[str, object]: """Strip PII location offsets from a checks response before it is logged.""" sanitized: Final[dict[str, Any]] = copy.deepcopy(dict(response)) sensitive: Final = (sanitized.get("results") or {}).get("sensitiveInformation") or {} diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index 1f63851b216..1e24b2aecbc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -949,7 +949,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _finalize_inspection( self, - inspect_response: dict[str, Any], + inspect_response: dict[str, object], request_data: dict, context: _ScanContext, start_time: datetime, @@ -1189,7 +1189,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return any(key in payload for key in cls._DECISION_FIELDS) @classmethod - def _unwrap_verdict_envelope(cls, inspect_response: dict[str, Any]) -> dict[str, Any]: + def _unwrap_verdict_envelope(cls, inspect_response: dict[str, object]) -> dict[str, Any]: """Return the dict that actually holds is_safe / action / rules. Cisco AI Defense returns the verdict at different nesting depths diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 405fd779d24..3188febf59a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -42,6 +42,7 @@ from litellm.types.utils import ( GenericGuardrailAPIInputs, GuardrailStatus, GuardrailTracingDetail, + Message, ModelResponse, ModelResponseStream, ) @@ -1516,7 +1517,7 @@ class ContentFilterGuardrail(CustomGuardrail): @staticmethod def _describe_image_response_content(response: ModelResponse) -> str | None: choice = response.choices[0] - message = getattr(choice, "message", None) + message: Final[Message | None] = getattr(choice, "message", None) if message and getattr(message, "content", None): return message.content return None diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py index d34861838c4..f82c93d54e1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py @@ -21,7 +21,7 @@ import os import re import time from datetime import datetime, timezone -from typing import Any, Final +from typing import Final import pytest from fastapi import HTTPException @@ -60,7 +60,7 @@ def _run(checker, text: str) -> dict: return {"decision": "ALLOW", "score": 0.0, "matched_topic": None} except HTTPException as e: if e.status_code == 400: - detail: Final[dict[str, Any]] = e.detail if isinstance(e.detail, dict) else {} + detail: Final[dict[str, object]] = e.detail if isinstance(e.detail, dict) else {} return { "decision": "BLOCK", "score": detail.get("score", 1.0), 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 7e0688f89f7..b01487bffcf 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 @@ -438,7 +438,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): collect_response_item(item, slots) return tuple(slots) - async def _restore_responses_api_response(self, response: Any, slots: Sequence[Slot], data: MutableRequest) -> Any: + async def _restore_responses_api_response( + self, response: object, slots: Sequence[Slot], data: MutableRequest + ) -> object: """Puts the original values back into a Responses API reply.""" await rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data))) return response @@ -567,7 +569,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async def _restore_tool_call_window( self, - tool_call: Any, + tool_call: object, choice_index: int, carries: CarryWindows, session_id: str, @@ -628,8 +630,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): delta.tool_calls = [*existing, *continuations] async def _flush_trailing( - self, last_chunk: Any, carries: CarryWindows, session_id: str - ) -> AsyncGenerator[Any, None]: + self, last_chunk: object, carries: CarryWindows, session_id: str + ) -> AsyncGenerator[object, None]: """Empties every window still holding text, one chunk per window. This is the net for a stream that ended with no finish_reason at all; a stream diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py index 1e246922fca..9124a98ac36 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py @@ -10,7 +10,7 @@ Permission logic: - end_user_id + mcp_servers → allow only those servers """ -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Final, Literal from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -230,7 +230,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): return tool_name.split("-", 1)[0] @staticmethod - def _get_tool_name_from_definition(tool: Any) -> str | None: + def _get_tool_name_from_definition(tool: object) -> str | None: """ Extract tool name from a definition dict. diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 6d5d39877d8..9ee53c05904 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -59,7 +59,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.presidio import ( PresidioAnalyzeRequest, PresidioAnalyzeResponseItem, ) -from litellm.types.utils import GuardrailStatus, StreamingChoices +from litellm.types.utils import GuardrailStatus, Message, StreamingChoices from litellm.utils import ( EmbeddingResponse, ImageResponse, @@ -1331,7 +1331,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_config: Final = self.get_presidio_settings_from_request_data(request_data or {}) for choice in response.choices: - message = getattr(choice, "message", None) + message: Message | None = getattr(choice, "message", None) if message is None: continue diff --git a/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py b/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py index a6fe9bd7d77..a63695a3af2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py @@ -124,7 +124,7 @@ class SingulrGuardrail(CustomGuardrail): ) @classmethod - def _resolve_metadata_value(cls, request_data: Mapping[str, Any], key: str) -> str | None: + def _resolve_metadata_value(cls, request_data: Mapping[str, object], key: str) -> str | None: for container in cls._metadata_containers(request_data=request_data): value = container.get(key) if value: @@ -132,7 +132,7 @@ class SingulrGuardrail(CustomGuardrail): return None @classmethod - def _resolve_user_role_from_request_data(cls, request_data: Mapping[str, Any]) -> str | None: + def _resolve_user_role_from_request_data(cls, request_data: Mapping[str, object]) -> str | None: for container in cls._metadata_containers(request_data=request_data): auth = container.get("user_api_key_auth") if isinstance(auth, UserAPIKeyAuth) and auth.user_role: @@ -140,7 +140,7 @@ class SingulrGuardrail(CustomGuardrail): return None @classmethod - def _build_metadata(cls, request_data: Mapping[str, Any]) -> Mapping[str, str] | None: + def _build_metadata(cls, request_data: Mapping[str, object]) -> Mapping[str, str] | None: fields: Final = ( "user_api_key_alias", "user_api_key_user_id", diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py index 0f87adf5c93..f8e6f98f893 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py @@ -259,7 +259,7 @@ def _route_has_translation(request_data: dict) -> bool: return any(call_type in mappings for call_type in get_call_types_for_route(route) or ()) -def _request_structured_messages(request_data: dict) -> list[dict[str, Any]] | None: +def _request_structured_messages(request_data: dict) -> list[dict[str, object]] | None: messages: Final = request_data.get("messages") if messages: return messages if isinstance(messages, list) else None @@ -1083,7 +1083,7 @@ class StraikerGuardrail(CustomGuardrail): if last_failure is None or not last_failure.retryable: return parsed, last_failure if attempt < attempts - 1: - backoff = min(self.initial_backoff * (2**attempt), self.max_backoff) + backoff: float = min(self.initial_backoff * (2**attempt), self.max_backoff) await asyncio.sleep(random.uniform(0, backoff)) return None, last_failure or _WebhookFailure("unknown error", is_unreachable=True) diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index b7f3726d017..0c1e22463ff 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -1,5 +1,5 @@ # litellm/proxy/guardrails/guardrail_initializers.py -from typing import Any, Final +from typing import Final import litellm from litellm.integrations.custom_guardrail import CustomGuardrail @@ -218,7 +218,7 @@ def initialize_tool_permission(litellm_params: LitellmParams, guardrail: Guardra ToolPermissionGuardrail, ) - rules: list[dict[str, Any]] | None = None + rules: list[dict[str, object]] | None = None if litellm_params.rules: rules = [] for rule in litellm_params.rules: diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index 7926c9a6cfb..f35a7a45c4e 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -1,4 +1,4 @@ -from typing import Any, Final, cast +from typing import Final, cast import litellm from litellm import Router @@ -69,7 +69,7 @@ def _populate_router_guardrail_list(guardrail_list: list[Guardrail]) -> None: for guardrail in guardrail_list: guardrail_id = guardrail.get("guardrail_id") guardrail_name = guardrail.get("guardrail_name") - litellm_params: Any = guardrail.get("litellm_params", {}) + litellm_params: object = guardrail.get("litellm_params", {}) # Get the callback instance from the registry callback = None diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index b707e72792c..9f5bc0966f8 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -524,7 +524,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): ) from litellm.proxy.proxy_server import llm_router - fetch_kwargs: Final[dict[str, Any]] = { + fetch_kwargs: Final[dict[str, object]] = { "custom_llm_provider": custom_llm_provider, } diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2bfaf57f0fc..363d525baf3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3038,7 +3038,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Resolve the agent_id from either the API key or request metadata. Key-level agent_id takes precedence over metadata/header-supplied agent_id. """ - key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None) + key_agent_id: Final[str | None] = getattr(user_api_key_dict, "agent_id", None) if key_agent_id: return key_agent_id metadata: Final = data.get("metadata") or {} @@ -4225,7 +4225,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations def _get_total_tokens_from_usage( - self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"] + self, usage: object | None, rate_limit_type: Literal["output", "input", "total"] ) -> int: """ Get total tokens from response usage for rate limiting. diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index f937b439042..e7160436638 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -2,7 +2,7 @@ import asyncio import traceback from collections.abc import Callable, Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Protocol, cast +from typing import TYPE_CHECKING, Final, Protocol, cast import litellm from litellm._logging import verbose_proxy_logger @@ -291,7 +291,7 @@ class _ProxyDBLogger(CustomLogger): async def _PROXY_track_cost_callback( self, kwargs, # kwargs to completion - completion_response: litellm.ModelResponse | Any | None, # response from completion + completion_response: litellm.ModelResponse | object | None, # response from completion start_time=None, end_time=None, # start/end time for completion ): diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 152438d0573..452c653eef0 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -562,7 +562,7 @@ def _strip_client_message_redaction_opt_out(data: dict[str, object]) -> None: def _strip_client_callback_credentials( - data: dict[str, Any], # mutable-ok: strips in place on the request body the pre-call pipeline threads through + data: dict[str, object], # mutable-ok: strips in place on the request body the pre-call pipeline threads through ) -> None: """Drop callback credentials and destinations supplied by the caller. @@ -626,7 +626,7 @@ def _strip_client_pricing_overrides(data: dict[str, object]) -> None: def _strip_router_reserved_metadata( - data: dict[str, Any], # mutable-ok: strips in place on the request body the pre-call pipeline threads through + data: dict[str, object], # mutable-ok: strips in place on the request body the pre-call pipeline threads through ) -> None: """Drop the router-owned fallback stamps from any client-supplied metadata bucket.""" for metadata_key in ("metadata", "litellm_metadata"): diff --git a/litellm/proxy/logging_endpoints/callback_logs_endpoints.py b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py index 4a1079871b0..55644cffc1b 100644 --- a/litellm/proxy/logging_endpoints/callback_logs_endpoints.py +++ b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py @@ -88,9 +88,9 @@ class CallbackLogsReplayer: function_id="", ) - metadata: Final[dict[str, Any]] = payload.get("metadata") or {} + metadata: Final[dict[str, object]] = payload.get("metadata") or {} user_api_key_hash: Final = metadata.get("user_api_key_hash") - litellm_metadata: Final[dict[str, Any]] = { + litellm_metadata: Final[dict[str, object]] = { "user_api_key": user_api_key_hash, "user_api_key_hash": user_api_key_hash, "user_api_key_alias": metadata.get("user_api_key_alias"), diff --git a/litellm/proxy/middleware/per_request_root_path_middleware.py b/litellm/proxy/middleware/per_request_root_path_middleware.py index d5df91458d2..3afa268b182 100644 --- a/litellm/proxy/middleware/per_request_root_path_middleware.py +++ b/litellm/proxy/middleware/per_request_root_path_middleware.py @@ -104,7 +104,7 @@ class PerRequestRootPathMiddleware: async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: if scope["type"] in ("http", "websocket"): - path: Final = scope.get("path", "") + path: Final[str] = scope.get("path", "") for prefix in self.root_paths: if path == prefix or path.startswith(prefix + "/"): scope["root_path"] = prefix # rebind-ok: ASGI middleware contract; Router and base_url read it diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index ebdd3e92bb2..f5d9eb11dc6 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -4,7 +4,7 @@ Prometheus Auth Middleware - Pure ASGI implementation import json from collections.abc import MutableMapping -from typing import Any, Final +from typing import Final from fastapi import Request from starlette.routing import get_route_path @@ -52,9 +52,9 @@ class PrometheusAuthMiddleware: # user_api_key_auth reads the request body, which consumes ASGI `receive`. # Buffer those messages and replay them for the inner app; otherwise a # successful auth would forward an exhausted receive and /metrics hangs. - buffered_messages: Final[list[MutableMapping[str, Any]]] = [] + buffered_messages: Final[list[MutableMapping[str, object]]] = [] - async def receive_for_auth() -> MutableMapping[str, Any]: + async def receive_for_auth() -> MutableMapping[str, object]: message: Final = await receive() buffered_messages.append(message) return message @@ -102,7 +102,7 @@ class PrometheusAuthMiddleware: replay_idx = 0 - async def receive_replay() -> MutableMapping[str, Any]: + async def receive_replay() -> MutableMapping[str, object]: nonlocal replay_idx if replay_idx < len(buffered_messages): msg: Final = buffered_messages[replay_idx] diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py index 7d289b80457..7527c504c59 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py @@ -30,7 +30,7 @@ COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS: Final = frozenset(COMPREHEND_MEDICAL_CO class ComprehendMedicalPassthroughLoggingHandler: @staticmethod def _operation_from_response(httpx_response: httpx.Response) -> str: - target: Final = httpx_response.request.headers.get("x-amz-target", "") + target: Final[str] = httpx_response.request.headers.get("x-amz-target", "") return target.split(".")[-1] @staticmethod diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py index d97ddb9a909..9efa371ed7f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py @@ -1,4 +1,5 @@ import re +from collections.abc import Callable from datetime import datetime from typing import TYPE_CHECKING, Any, Final @@ -17,6 +18,7 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrou ) from litellm.types.utils import ( ModelResponse, + ModelResponseStream, TextCompletionResponse, ) @@ -194,7 +196,9 @@ class GeminiPassthroughLoggingHandler: sync_stream=False, logging_obj=litellm_logging_obj, ) - chunk_parsing_logic: Final[Any] = gemini_iterator._common_chunk_parsing_logic + chunk_parsing_logic: Final[Callable[[str], ModelResponseStream | None]] = ( + gemini_iterator._common_chunk_parsing_logic + ) parsed_chunks = [chunk_parsing_logic(chunk) for chunk in all_chunks] else: return None diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/transcribe_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/transcribe_passthrough_logging_handler.py index a8dfb745290..79f6dee90f7 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/transcribe_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/transcribe_passthrough_logging_handler.py @@ -592,7 +592,7 @@ class TranscribePassthroughLoggingHandler: @staticmethod def _operation_from_response(httpx_response: httpx.Response) -> str: headers: Final[Mapping[str, str]] = httpx_response.request.headers - target: Final = headers.get("x-amz-target", "") + target: Final[str] = headers.get("x-amz-target", "") return target.split(".")[-1] @staticmethod diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index be40e225cde..9f0ddb09953 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -1,6 +1,6 @@ import asyncio import re -from collections.abc import Mapping +from collections.abc import Callable, Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, cast from urllib.parse import urlparse @@ -39,6 +39,7 @@ from litellm.types.utils import ( EmbeddingResponse, ImageResponse, ModelResponse, + ModelResponseStream, SpecialEnums, StandardPassThroughResponseObject, TextCompletionResponse, @@ -629,7 +630,7 @@ class VertexPassthroughLoggingHandler: sync_stream=False, logging_obj=litellm_logging_obj, ) - chunk_parsing_logic: Any = vertex_iterator._common_chunk_parsing_logic + chunk_parsing_logic: Callable[..., ModelResponseStream | None] = vertex_iterator._common_chunk_parsing_logic parsed_chunks = [chunk_parsing_logic(chunk) for chunk in all_chunks] elif "rawPredict" in url_route or "streamRawPredict" in url_route: from litellm.llms.anthropic.chat.handler import ModelResponseIterator diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index da3e28e25e4..65a199d38b6 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -367,7 +367,9 @@ class PassThroughEndpointLogging: vertex_ai_live_handler: Final = VertexAILivePassthroughLoggingHandler() # For WebSocket responses, response_body should be a list of messages - websocket_messages: Final[list[dict[str, Any]]] = response_body if isinstance(response_body, list) else [] + websocket_messages: Final[list[dict[str, object]]] = ( + response_body if isinstance(response_body, list) else [] + ) vertex_ai_live_handler_result: Final = vertex_ai_live_handler.vertex_ai_live_passthrough_handler( websocket_messages=websocket_messages, diff --git a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py index 5e9fba0b7cd..59b8d40577a 100644 --- a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py +++ b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py @@ -96,8 +96,8 @@ def parse_upstream_reported_usage(headers: httpx.Headers) -> UpstreamReportedUsa for a target that does not speak this contract (e.g. Anthropic or Vertex, whose cost LiteLLM derives from the response body instead). """ - raw_response_cost: Final = headers.get(UPSTREAM_RESPONSE_COST_HEADER) - raw_total_tokens: Final = headers.get(UPSTREAM_TOTAL_TOKENS_HEADER) + raw_response_cost: Final[str | None] = headers.get(UPSTREAM_RESPONSE_COST_HEADER) + raw_total_tokens: Final[str | None] = headers.get(UPSTREAM_TOTAL_TOKENS_HEADER) if raw_response_cost is None and raw_total_tokens is None: return None return UpstreamReportedUsage( diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 1a741af6654..faf268d3986 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -498,7 +498,7 @@ class ProxyInitializationHelpers: if ciphers is not None: print("\033[1;33mLiteLLM: --ciphers is not applied when using --run_granian.\033[0m\n") - kwargs: Final[dict[str, Any]] = { + kwargs: Final[dict[str, object]] = { "target": "litellm.proxy.proxy_server:app", "address": host, "port": port, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 06bfb467b52..56a35c53a4a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -14798,7 +14798,7 @@ async def _fetch_db_models_for_search( filter for `team_public_model_name` instead and keep the DB cost bounded by `search`. """ - db_where_condition: Final[dict[str, Any]] = { + db_where_condition: Final[dict[str, object]] = { "model_name": {"contains": search_lower, "mode": "insensitive"} if model_name is None else model_name } if db_model_ids_in_router: diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 1a0eb5e7e43..936296cba26 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -426,7 +426,7 @@ async def parse_rag_ingest_request( file_data: tuple[str, bytes, str] | None = None file_url: str | None = None file_id: str | None = None - ingest_options: dict[str, Any] = {} + ingest_options: dict[str, object] = {} if "multipart/form-data" in content_type: # Form upload diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 9d4304e1d47..49b868ac71a 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -68,7 +68,7 @@ async def _refresh_router_search_tools() -> None: verbose_proxy_logger.exception("Search tool router refresh failed after a management write: %s", e) -def _with_loaded_tools_where_undecryptable(db_search_tools: Sequence[dict[str, Any]]) -> list[dict[str, Any]]: +def _with_loaded_tools_where_undecryptable(db_search_tools: Sequence[dict[str, object]]) -> list[dict[str, Any]]: from litellm.proxy.proxy_server import llm_router kept_search_tools: Final = keep_loaded_search_tools_that_do_not_decrypt( diff --git a/litellm/proxy/types_utils/utils.py b/litellm/proxy/types_utils/utils.py index a85a96aecfd..c3fd237c0d8 100644 --- a/litellm/proxy/types_utils/utils.py +++ b/litellm/proxy/types_utils/utils.py @@ -65,7 +65,7 @@ def get_instance_fn(value: str, config_file_path: str | None = None) -> Any: raise e -def _load_instance_from_remote_storage(remote_url: str, config_file_path: str | None = None) -> Any: +def _load_instance_from_remote_storage(remote_url: str, config_file_path: str | None = None) -> object: """ Load custom logger instance from S3 or GCS URL. diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index f01e515d356..e21f706ba60 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -62,7 +62,7 @@ def _index_of_output_item_type(items: Sequence[object], item_type: str) -> int | ) -def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | None) -> tuple[Any, ...]: +def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | None) -> tuple[object, ...]: if item_id is None: return items diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 13b353dbdd0..3a27d21d0ce 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -266,7 +266,7 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _transform_tool_choice( tool_choice: Any, - ) -> str | dict[str, Any] | None: + ) -> str | dict[str, object] | None: """ Transform tool_choice from various formats to OpenAI Chat Completion format. diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index c1dabf8a5c9..4154452af05 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -2861,7 +2861,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs["litellm_metadata"]["proxy_server_request"] = proxy_server_request call_kwargs["proxy_server_request"] = proxy_server_request - async def _stream_and_forward(self, model: str, call_kwargs: dict[str, Any]) -> _MutableJsonObject | None: + async def _stream_and_forward(self, model: str, call_kwargs: dict[str, object]) -> _MutableJsonObject | None: """ Stream ``litellm.aresponses`` and forward every chunk over the WebSocket. diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 0c025506f22..a97c4392c09 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -312,17 +312,17 @@ class ResponsesAPIRequestUtils: def _update_responses_api_response_id_with_model_id( responses_api_response: ResponsesAPIResponse, custom_llm_provider: str | None, - litellm_metadata: dict[str, Any] | None = None, + litellm_metadata: dict[str, object] | None = None, ) -> ResponsesAPIResponse: ... @overload @staticmethod def _update_responses_api_response_id_with_model_id( - responses_api_response: dict[str, Any], + responses_api_response: dict[str, object], custom_llm_provider: str | None, - litellm_metadata: dict[str, Any] | None = None, - ) -> dict[str, Any]: + litellm_metadata: dict[str, object] | None = None, + ) -> dict[str, object]: ... # fmt: on @@ -332,7 +332,7 @@ class ResponsesAPIRequestUtils: responses_api_response: ResponsesAPIResponse | dict[str, Any], custom_llm_provider: str | None, litellm_metadata: dict[str, Any] | None = None, - ) -> ResponsesAPIResponse | dict[str, Any]: + ) -> ResponsesAPIResponse | dict[str, object]: """Update the responses_api_response_id with model_id and custom_llm_provider. Handles both ``ResponsesAPIResponse`` objects and plain dictionaries returned @@ -1004,7 +1004,7 @@ class ResponsesAPIRequestUtils: responses_api_response: ResponsesAPIResponse | dict[str, Any], custom_llm_provider: str | None, litellm_metadata: dict[str, Any] | None = None, - ) -> ResponsesAPIResponse | dict[str, Any]: + ) -> ResponsesAPIResponse | dict[str, object]: """Encode container IDs in the response output with provider/model info. This walks through all output items and encodes any container_id fields diff --git a/litellm/router_strategy/adaptive_router/hooks.py b/litellm/router_strategy/adaptive_router/hooks.py index c4a2eae1ef9..0bb00a08f51 100644 --- a/litellm/router_strategy/adaptive_router/hooks.py +++ b/litellm/router_strategy/adaptive_router/hooks.py @@ -135,7 +135,7 @@ def _recent_tool_results( return results -def _assistant_content_and_tool_calls(response_obj: Any) -> tuple: +def _assistant_content_and_tool_calls(response_obj: Any) -> tuple[object, Sequence[Mapping[str, object]]]: """Return (assistant_text, tool_calls_list) extracted from a ModelResponse-ish object.""" if response_obj is None: return None, [] diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 631b0c3df3d..98601038014 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -207,14 +207,14 @@ class RouterBudgetLimiting(CustomLogger): def _filter_out_deployments_above_budget( self, - potential_deployments: list[dict[str, Any]], + potential_deployments: list[dict[str, object]], healthy_deployments: list[dict[str, Any]], provider_configs: dict[str, GenericBudgetInfo], deployment_configs: dict[str, GenericBudgetInfo], deployment_providers: list[str | None], spend_map: dict[str, float], request_tags: list[str], - ) -> tuple[list[dict[str, Any]], str]: + ) -> tuple[list[dict[str, object]], str]: """ Filter out deployments that have exceeded their budget limit. Follow budget checks are run here: diff --git a/litellm/search/main.py b/litellm/search/main.py index b2dd51799a1..e3db8bfb0ac 100644 --- a/litellm/search/main.py +++ b/litellm/search/main.py @@ -29,7 +29,7 @@ def _build_search_optional_params( search_domain_filter: list[str] | None = None, max_tokens_per_page: int | None = None, country: str | None = None, -) -> dict[str, Any]: +) -> dict[str, object]: """ Helper function to build optional_params dict from Perplexity Search API parameters. @@ -42,7 +42,7 @@ def _build_search_optional_params( Returns: Dict with non-None optional parameters """ - optional_params: Final[dict[str, Any]] = {} + optional_params: Final[dict[str, object]] = {} if max_results is not None: optional_params["max_results"] = max_results diff --git a/litellm/videos/utils.py b/litellm/videos/utils.py index 306d59c3d58..59d96ad5e36 100644 --- a/litellm/videos/utils.py +++ b/litellm/videos/utils.py @@ -75,12 +75,12 @@ class VideoGenerationRequestUtils: cleaned_kwargs: Final = filter_out_litellm_params(kwargs={k: v for k, v in raw_kwargs.items() if v is not None}) - optional_params: Final[dict[str, Any]] = { + optional_params: Final[dict[str, object]] = { **base_params, **cleaned_kwargs, } - merged_extra_body: dict[str, Any] = {} + merged_extra_body: dict[str, object] = {} for extra_body_candidate in (top_level_extra_body, kwargs_extra_body): if isinstance(extra_body_candidate, dict): for key, value in extra_body_candidate.items(): diff --git a/tests/integration/pricing/test_bundled_cost_map_price_wire.py b/tests/integration/pricing/test_bundled_cost_map_price_wire.py new file mode 100644 index 00000000000..c054856e2c8 --- /dev/null +++ b/tests/integration/pricing/test_bundled_cost_map_price_wire.py @@ -0,0 +1,59 @@ +import json +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.openai_wire import answering_model_discovery, posted_targets +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-4o-mini" +_API_KEY: Final = "synthetic-openai-key" +_PROMPT_TOKENS: Final = 1000 +_COMPLETION_TOKENS: Final = 500 +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_PACKAGED_PRICES: Final = Path(__file__).resolve().parents[3] / "litellm" / "model_prices_and_context_window_backup.json" +_REPLY: Final = json.dumps( + { + "id": "chatcmpl-bundled-price", + "object": "chat.completion", + "created": 1, + "model": _BACKEND, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "bundled price control"}, "finish_reason": "stop"} + ], + "usage": { + "prompt_tokens": _PROMPT_TOKENS, + "completion_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, + }, + } +).encode() + + +def _peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + return Reply(body=_REPLY) + + +def test_a_deployment_with_no_configured_price_is_charged_at_the_packaged_cost_map_rates(gateway: Gateway) -> None: + packaged: Final = _JSON_OBJECT.validate_json(_PACKAGED_PRICES.read_bytes())[_BACKEND] + assert isinstance(packaged, dict) + input_rate: Final = packaged["input_cost_per_token"] + output_rate: Final = packaged["output_cost_per_token"] + assert isinstance(input_rate, float) and isinstance(output_rate, float) + assert input_rate > 0 and output_rate > 0 + with wire_server(answering_model_discovery(_peer)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic priced request"}]}, + ) + assert response.status_code == 200, response.text + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx( + _PROMPT_TOKENS * input_rate + _COMPLETION_TOKENS * output_rate, rel=1e-9 + ) + assert posted_targets(wire) == ("/v1/chat/completions",) diff --git a/tests/integration/providers/test_bedrock_anthropic_beta_header_wire.py b/tests/integration/providers/test_bedrock_anthropic_beta_header_wire.py new file mode 100644 index 00000000000..aa15ee623aa --- /dev/null +++ b/tests/integration/providers/test_bedrock_anthropic_beta_header_wire.py @@ -0,0 +1,154 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" +_TOKEN: Final = "synthetic-bedrock-bearer" +_CONTEXT_BETA: Final = "context-1m-2025-08-07" +_THINKING_BETA: Final = "interleaved-thinking-2025-05-14" +_REPLY_TEXT: Final = "beta header control" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_INVOKE_REPLY: Final = json.dumps( + { + "id": "msg_beta_header", + "type": "message", + "role": "assistant", + "model": _BACKEND, + "content": [{"type": "text", "text": _REPLY_TEXT}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 12, "output_tokens": 5}, + } +).encode() +_CONVERSE_REPLY: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _REPLY_TEXT}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 12, "outputTokens": 5, "totalTokens": 17}, + "metrics": {"latencyMs": 1}, + } +).encode() + + +def _invoke_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == f"/model/{_BACKEND}/invoke", request.target + assert request.headers["authorization"] == f"Bearer {_TOKEN}" + return Reply(body=_INVOKE_REPLY) + + +def _converse_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target.endswith("/converse"), request.target + assert request.headers["authorization"] == f"Bearer {_TOKEN}" + return Reply(body=_CONVERSE_REPLY) + + +def _sent_betas(outbound: dict[str, JsonValue]) -> tuple[str, ...] | None: + betas: Final = outbound.get("anthropic_beta") + if betas is None: + return None + assert isinstance(betas, list), outbound + return tuple(sorted(str(beta) for beta in betas)) + + +def _beta_header(value: str | None) -> dict[str, str]: + return {} if value is None else {"anthropic-beta": value} + + +@pytest.mark.parametrize( + ("header", "expected"), + ( + pytest.param(json.dumps([_CONTEXT_BETA, _THINKING_BETA]), (_CONTEXT_BETA, _THINKING_BETA), id="json-array"), + pytest.param(f'[ " {_CONTEXT_BETA} " ]', (_CONTEXT_BETA,), id="json-array-padded-entry"), + pytest.param("[1, 2]", ("1", "2"), id="json-array-of-integers"), + pytest.param("[null]", ("None",), id="json-array-holding-null"), + pytest.param(f"[{_CONTEXT_BETA}]", (f"[{_CONTEXT_BETA}]",), id="bracketed-but-not-json"), + pytest.param(f"{_CONTEXT_BETA}, {_THINKING_BETA}", (_CONTEXT_BETA, _THINKING_BETA), id="comma-separated"), + pytest.param("[]", None, id="empty-json-array"), + pytest.param(None, None, id="header-absent"), + ), +) +def test_bedrock_invoke_chat_sends_the_client_anthropic_beta_header_as_a_body_list( + gateway: Gateway, + header: str | None, + expected: tuple[str, ...] | None, +) -> None: + with wire_server(_invoke_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock/invoke/{_BACKEND}", + api_key=_TOKEN, + aws_region_name="us-east-1", + api_base=wire.url, + aws_bedrock_runtime_endpoint=wire.url, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic beta request"}], "max_tokens": 16}, + headers=_beta_header(header), + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["choices"][0]["message"]["content"] == _REPLY_TEXT, response.text + assert payload["usage"]["prompt_tokens"] == 12 and payload["usage"]["completion_tokens"] == 5, response.text + requests: Final = wire.drain() + assert len(requests) == 1, requests + outbound: Final = _JSON_OBJECT.validate_json(requests[0].body) + assert outbound["messages"] == [ + {"role": "user", "content": [{"type": "text", "text": "synthetic beta request"}]} + ], outbound + assert _sent_betas(outbound) == (None if expected is None else tuple(sorted(expected))), outbound + + +def test_bedrock_invoke_messages_sends_a_comma_separated_anthropic_beta_header_as_a_body_list( + gateway: Gateway, +) -> None: + with wire_server(_invoke_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock/invoke/{_BACKEND}", + api_key=_TOKEN, + aws_region_name="us-east-1", + api_base=wire.url, + aws_bedrock_runtime_endpoint=wire.url, + ) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "messages": [{"role": "user", "content": "synthetic beta request"}], "max_tokens": 16}, + headers=_beta_header(f"{_CONTEXT_BETA}, {_THINKING_BETA}"), + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["content"] == [{"type": "text", "text": _REPLY_TEXT}], response.text + requests: Final = wire.drain() + assert len(requests) == 1, requests + assert _sent_betas(_JSON_OBJECT.validate_json(requests[0].body)) == (_CONTEXT_BETA,), requests[0].body + + +def test_bedrock_converse_chat_sends_a_comma_separated_anthropic_beta_header_as_additional_model_fields( + gateway: Gateway, +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock/converse/{_BACKEND}", + api_key=_TOKEN, + aws_region_name="us-east-1", + api_base=wire.url, + aws_bedrock_runtime_endpoint=wire.url, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic beta request"}], "max_tokens": 16}, + headers=_beta_header(f"{_CONTEXT_BETA}, {_THINKING_BETA}"), + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["choices"][0]["message"]["content"] == _REPLY_TEXT, response.text + requests: Final = wire.drain() + assert len(requests) == 1, requests + outbound: Final = _JSON_OBJECT.validate_json(requests[0].body) + assert outbound["additionalModelRequestFields"] == {"anthropic_beta": [_CONTEXT_BETA]}, outbound diff --git a/tests/integration/providers/test_container_file_routes_wire.py b/tests/integration/providers/test_container_file_routes_wire.py new file mode 100644 index 00000000000..bcf23c595e4 --- /dev/null +++ b/tests/integration/providers/test_container_file_routes_wire.py @@ -0,0 +1,91 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.openai_wire import MODEL_DISCOVERY, answering_model_discovery +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-4o-mini" +_API_KEY: Final = "synthetic-openai-key" +_CONTAINER: Final = "cntr_wire" +_FILE: Final = "cfile_1" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_FILE_OBJECT: Final[dict[str, JsonValue]] = { + "id": _FILE, + "object": "container.file", + "container_id": _CONTAINER, + "created_at": 1, + "bytes": 5, + "path": "/mnt/data/notes.txt", + "source": "user", +} +_REPLIES: Final = { + ("POST", "/v1/containers"): Reply( + body=json.dumps( + {"id": _CONTAINER, "object": "container", "created_at": 1, "status": "running", "name": "wire"} + ).encode() + ), + ("GET", f"/v1/containers/{_CONTAINER}/files"): Reply( + body=json.dumps( + {"object": "list", "data": [_FILE_OBJECT], "first_id": _FILE, "last_id": _FILE, "has_more": False} + ).encode() + ), + ("GET", f"/v1/containers/{_CONTAINER}/files/{_FILE}"): Reply(body=json.dumps(_FILE_OBJECT).encode()), + ("GET", f"/v1/containers/{_CONTAINER}/files/{_FILE}/content"): Reply( + body=b"hello", content_type="application/octet-stream" + ), + ("DELETE", f"/v1/containers/{_CONTAINER}/files/{_FILE}"): Reply( + body=json.dumps({"id": _FILE, "object": "container.file.deleted", "deleted": True}).encode() + ), +} + + +def _peer(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + return _REPLIES[(request.method, request.target.split("?")[0])] + + +def test_container_file_routes_reach_the_openai_paths_the_packaged_endpoint_table_declares( + gateway: Gateway, +) -> None: + with wire_server(answering_model_discovery(_peer)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + created: Final = gateway.request("POST", "/v1/containers", {"model": model, "name": "wire"}) + assert created.status_code == 200, created.text + container: Final = str(_JSON_OBJECT.validate_json(created.content)["id"]) + assert container.startswith("cntr_") and container != _CONTAINER, created.text + files: Final = f"/v1/containers/{container}/files" + + listed: Final = gateway.request("GET", files) + assert listed.status_code == 200, listed.text + assert _JSON_OBJECT.validate_json(listed.content)["data"][0]["id"] == _FILE, listed.text + + paged: Final = gateway.request("GET", files, params={"limit": "2", "order": "desc", "after": "cfile_0"}) + assert paged.status_code == 200, paged.text + + retrieved: Final = gateway.request("GET", f"{files}/{_FILE}") + assert retrieved.status_code == 200, retrieved.text + assert _JSON_OBJECT.validate_json(retrieved.content)["path"] == "/mnt/data/notes.txt", retrieved.text + + content: Final = gateway.request("GET", f"{files}/{_FILE}/content") + assert content.status_code == 200 and content.content == b"hello", content.text + + deleted: Final = gateway.request("DELETE", f"{files}/{_FILE}") + assert deleted.status_code == 200, deleted.text + assert _JSON_OBJECT.validate_json(deleted.content) == { + "id": _FILE, + "object": "container.file.deleted", + "deleted": True, + }, deleted.text + + upstream_files: Final = f"/v1/containers/{_CONTAINER}/files" + calls: Final = [(request.method, request.target) for request in wire.drain()] + assert [call for call in calls if call != MODEL_DISCOVERY] == [ + ("POST", "/v1/containers"), + ("GET", upstream_files), + ("GET", f"{upstream_files}?after=cfile_0&limit=2&order=desc"), + ("GET", f"{upstream_files}/{_FILE}"), + ("GET", f"{upstream_files}/{_FILE}/content"), + ("DELETE", f"{upstream_files}/{_FILE}"), + ] diff --git a/tests/integration/providers/test_gemini_image_config_and_usage_wire.py b/tests/integration/providers/test_gemini_image_config_and_usage_wire.py new file mode 100644 index 00000000000..27b408b27aa --- /dev/null +++ b/tests/integration/providers/test_gemini_image_config_and_usage_wire.py @@ -0,0 +1,244 @@ +import base64 +import json +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-2.5-flash-image" +_API_KEY: Final = "synthetic-gemini-key" +_PNG: Final = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg==" +) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_ONLY_MODALITIES: Final[dict[str, JsonValue]] = {"response_modalities": ["IMAGE", "TEXT"]} +_FULL_USAGE: Final[dict[str, JsonValue]] = { + "promptTokenCount": 263, + "candidatesTokenCount": 1290, + "totalTokenCount": 1553, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 5}, {"modality": "IMAGE", "tokenCount": 258}], + "candidatesTokensDetails": [{"modality": "IMAGE", "tokenCount": 1290}], +} +_FULL_USAGE_REPLY: Final[dict[str, JsonValue]] = { + "total_tokens": 1553, + "input_tokens": 263, + "input_tokens_details": {"image_tokens": 258, "text_tokens": 5}, + "output_tokens": 1290, + "output_tokens_details": {"image_tokens": 1290, "text_tokens": 0}, +} +_COUNTS_ONLY_REPLY: Final[dict[str, JsonValue]] = { + "total_tokens": 30, + "input_tokens": 10, + "input_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + "output_tokens": 20, + "output_tokens_details": {"image_tokens": 20, "text_tokens": 0}, +} + + +def _image_reply(usage: JsonValue, include_usage: bool = True) -> bytes: + candidates: Final = [ + {"content": {"parts": [{"inlineData": {"mimeType": "image/png", "data": base64.b64encode(_PNG).decode()}}]}} + ] + return json.dumps({"candidates": candidates, **({"usageMetadata": usage} if include_usage else {})}).encode() + + +def _peer(reply: bytes): + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == f"/models/{_BACKEND}:generateContent", request.target.split("?")[0] + return Reply(body=reply) + + return respond + + +def _edit(gateway: Gateway, model: str, image_config: str | None) -> httpx.Response: + fields: Final = {"model": model, "prompt": "synthetic edit request"} + return gateway.request_multipart( + "/v1/images/edits", + fields if image_config is None else {**fields, "imageConfig": image_config}, + {"image": ("pixel.png", _PNG, "image/png")}, + ) + + +def _generation_config(request: Request) -> JsonValue: + return _JSON_OBJECT.validate_json(request.body)["generationConfig"] + + +@pytest.mark.parametrize( + ("image_config", "expected"), + ( + pytest.param( + '{"aspectRatio": "1:1"}', + {**_ONLY_MODALITIES, "imageConfig": {"aspectRatio": "1:1"}}, + id="json-object", + ), + pytest.param("null", _ONLY_MODALITIES, id="json-null"), + pytest.param("[1, 2]", _ONLY_MODALITIES, id="json-array"), + pytest.param('"1:1"', _ONLY_MODALITIES, id="json-string"), + pytest.param("7", _ONLY_MODALITIES, id="json-number"), + pytest.param(None, _ONLY_MODALITIES, id="field-absent"), + ), +) +def test_gemini_image_edit_forwards_only_a_json_object_image_config( + gateway: Gateway, + image_config: str | None, + expected: dict[str, JsonValue], +) -> None: + with wire_server(_peer(_image_reply(_FULL_USAGE))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _edit(gateway, model, image_config) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert [image["b64_json"] for image in payload["data"]] == [base64.b64encode(_PNG).decode()], response.text + assert payload["usage"] == _FULL_USAGE_REPLY, response.text + requests: Final = wire.drain() + assert len(requests) == 1, requests + assert _generation_config(requests[0]) == expected, requests[0].body + + +@pytest.mark.parametrize("image_config", (pytest.param("{not json", id="truncated-json"), pytest.param("", id="empty"))) +def test_gemini_image_edit_rejects_an_image_config_string_that_is_not_json( + gateway: Gateway, + image_config: str, +) -> None: + with wire_server(_peer(_image_reply(_FULL_USAGE))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _edit(gateway, model, image_config) + assert response.status_code == 400, response.text + error: Final = _JSON_OBJECT.validate_json(response.content)["error"] + assert isinstance(error, dict), response.text + assert error["type"] == "invalid_request_error" and error["code"] == "400", response.text + assert str(error["message"]).startswith( + "litellm.UnsupportedParamsError: `imageConfig` must be valid JSON when provided as a string." + ), response.text + assert wire.drain() == (), "a rejected imageConfig must not reach the provider" + + +@pytest.mark.parametrize( + ("usage", "include_usage", "expected"), + ( + pytest.param(_FULL_USAGE, True, _FULL_USAGE_REPLY, id="counts-and-modality-details"), + pytest.param( + {"promptTokenCount": 10, "candidatesTokenCount": 20, "totalTokenCount": 30}, + True, + _COUNTS_ONLY_REPLY, + id="counts-only", + ), + pytest.param( + { + "promptTokenCount": 10, + "candidatesTokenCount": 20, + "totalTokenCount": 30, + "promptTokensDetails": "none", + "candidatesTokensDetails": {"modality": "IMAGE"}, + }, + True, + _COUNTS_ONLY_REPLY, + id="details-that-are-not-lists", + ), + pytest.param( + { + "promptTokenCount": 10, + "candidatesTokenCount": 25, + "totalTokenCount": 35, + "candidatesTokensDetails": [ + {"modality": "IMAGE", "tokenCount": 20}, + {"modality": "TEXT", "tokenCount": 2}, + ], + }, + True, + { + "total_tokens": 35, + "input_tokens": 10, + "input_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + "output_tokens": 25, + "output_tokens_details": {"image_tokens": 20, "text_tokens": 5}, + }, + id="output-text-is-the-remainder", + ), + pytest.param( + {}, + True, + { + "total_tokens": 0, + "input_tokens": 0, + "input_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + "output_tokens": 0, + "output_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + }, + id="empty-usage-object", + ), + pytest.param( + None, + False, + { + "total_tokens": 0, + "input_tokens": 0, + "input_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + "output_tokens": 0, + }, + id="usage-absent", + ), + ), +) +def test_gemini_image_edit_reports_the_provider_usage_metadata_as_image_usage( + gateway: Gateway, + usage: JsonValue, + include_usage: bool, + expected: dict[str, JsonValue], +) -> None: + with wire_server(_peer(_image_reply(usage, include_usage))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _edit(gateway, model, None) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["usage"] == expected, response.text + assert len(wire.drain()) == 1 + + +@pytest.mark.parametrize( + ("count", "detail"), + ( + pytest.param( + "ten", + "Input should be a valid integer, unable to parse string as an integer [type=int_parsing", + id="count-is-a-word", + ), + pytest.param(None, "Input should be a valid integer [type=int_type", id="count-is-null"), + ), +) +def test_gemini_image_edit_answers_500_when_a_provider_token_count_is_not_a_number( + gateway: Gateway, + count: JsonValue, + detail: str, +) -> None: + usage: Final[dict[str, JsonValue]] = {"promptTokenCount": count, "candidatesTokenCount": 20, "totalTokenCount": 30} + with wire_server(_peer(_image_reply(usage))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _edit(gateway, model, None) + assert response.status_code == 500, response.text + error: Final = _JSON_OBJECT.validate_json(response.content)["error"] + assert isinstance(error, dict), response.text + assert error["type"] == "internal_server_error", response.text + assert "1 validation error for ImageUsage\ninput_tokens\n" in str(error["message"]), response.text + assert detail in str(error["message"]), response.text + assert len(wire.drain()) >= 1 + + +def test_gemini_image_generation_reports_the_provider_usage_metadata_as_image_usage(gateway: Gateway) -> None: + with wire_server(_peer(_image_reply(_FULL_USAGE))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/images/generations", {"model": model, "prompt": "synthetic generation request"} + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert [image["b64_json"] for image in payload["data"]] == [base64.b64encode(_PNG).decode()], response.text + assert payload["usage"] == _FULL_USAGE_REPLY, response.text + requests: Final = wire.drain() + assert len(requests) == 1, requests + assert _JSON_OBJECT.validate_json(requests[0].body)["contents"] == [ + {"parts": [{"text": "synthetic generation request"}]} + ], requests[0].body diff --git a/tests/integration/providers/test_interactions_openai_bridge_wire.py b/tests/integration/providers/test_interactions_openai_bridge_wire.py new file mode 100644 index 00000000000..cd18a227cd9 --- /dev/null +++ b/tests/integration/providers/test_interactions_openai_bridge_wire.py @@ -0,0 +1,93 @@ +import json +from datetime import UTC, datetime, timedelta +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.openai_wire import answering_model_discovery +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-interactions-bridge" +_API_KEY: Final = "synthetic-openai-key" +_REPLY_TEXT: Final = "interactions bridge control" +_CREATED_AT: Final = 86400 +_WIDEST_UTC_OFFSET: Final = timedelta(hours=14) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _reply(include_usage: bool) -> bytes: + usage: Final = {"usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13}} if include_usage else {} + return json.dumps( + { + "id": "resp_interactions_bridge", + "object": "response", + "created_at": _CREATED_AT, + "status": "completed", + "model": _BACKEND, + "output": [ + { + "type": "message", + "id": "msg_interactions_bridge", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _REPLY_TEXT, "annotations": []}], + } + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + **usage, + } + ).encode() + + +def _peer(include_usage: bool): + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/responses", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + return Reply(body=_reply(include_usage)) + + return respond + + +@pytest.mark.parametrize("route", ("/v1beta/interactions", "/interactions")) +@pytest.mark.parametrize( + ("include_usage", "expected_usage"), + ( + pytest.param(True, {"total_input_tokens": 9, "total_output_tokens": 4}, id="usage-reported"), + pytest.param(False, None, id="usage-absent"), + ), +) +def test_interactions_with_an_openai_model_answers_in_the_interactions_shape( + gateway: Gateway, + route: str, + include_usage: bool, + expected_usage: dict[str, JsonValue] | None, +) -> None: + with wire_server(answering_model_discovery(_peer(include_usage))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + route, + {"model": model, "input": "synthetic interaction", "system_instruction": "synthetic instruction"}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["object"] == "interaction" and payload["status"] == "completed", response.text + assert str(payload["id"]).startswith("resp_"), response.text + assert payload["model"] == model, response.text + assert payload["outputs"] == [{"type": "text", "text": _REPLY_TEXT}], response.text + assert payload["steps"] == [ + {"type": "model_output", "content": [{"type": "text", "text": _REPLY_TEXT}]} + ], response.text + assert payload["usage"] == expected_usage, response.text + created: Final = payload["created"] + assert isinstance(created, str) and created == payload["updated"], response.text + local_created: Final = datetime.fromisoformat(created).replace(tzinfo=UTC) + assert abs(local_created - datetime.fromtimestamp(_CREATED_AT, UTC)) <= _WIDEST_UTC_OFFSET, response.text + requests: Final = tuple(request for request in wire.drain() if request.method == "POST") + assert len(requests) == 1, requests + outbound: Final = _JSON_OBJECT.validate_json(requests[0].body) + assert outbound["model"] == _BACKEND and outbound["input"] == "synthetic interaction", requests[0].body + assert outbound["instructions"] == "synthetic instruction", requests[0].body diff --git a/tests/integration/providers/test_responses_reasoning_item_wire.py b/tests/integration/providers/test_responses_reasoning_item_wire.py new file mode 100644 index 00000000000..7e9b823c5c6 --- /dev/null +++ b/tests/integration/providers/test_responses_reasoning_item_wire.py @@ -0,0 +1,126 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.openai_wire import answering_model_discovery +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-reasoning-item" +_API_KEY: Final = "synthetic-openai-key" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_SUMMARY: Final[list[JsonValue]] = [{"type": "summary_text", "text": "thought"}] +_FIRST_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": "first synthetic turn"} +_LAST_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": "second synthetic turn"} +_REPLY: Final = json.dumps( + { + "id": "resp_reasoning_item", + "object": "response", + "created_at": 1, + "status": "completed", + "model": _BACKEND, + "output": [ + { + "type": "message", + "id": "msg_reasoning_item", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "reasoning item control", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13}, + } +).encode() + + +def _peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/responses", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + return Reply(body=_REPLY) + + +@pytest.mark.parametrize( + ("item", "expected"), + ( + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": _SUMMARY}, + {"type": "reasoning", "id": "rs_1", "summary": _SUMMARY}, + id="well-formed", + ), + pytest.param( + { + "type": "reasoning", + "id": "rs_1", + "summary": [], + "status": None, + "content": None, + "encrypted_content": None, + }, + {"type": "reasoning", "id": "rs_1", "summary": []}, + id="null-optional-fields-dropped", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": _SUMMARY, "status": "completed"}, + {"type": "reasoning", "id": "rs_1", "summary": _SUMMARY, "status": "completed"}, + id="status-kept", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": [], "encrypted_content": "opaque-blob"}, + {"type": "reasoning", "id": "rs_1", "summary": [], "encrypted_content": "opaque-blob"}, + id="encrypted-content-kept", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": [], "vendor_note": "kept"}, + {"type": "reasoning", "id": "rs_1", "summary": [], "vendor_note": "kept"}, + id="unknown-field-kept", + ), + pytest.param( + {"type": "reasoning", "id": 7, "summary": []}, + {"type": "reasoning", "id": 7, "summary": []}, + id="id-is-a-number", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": "not a list"}, + {"type": "reasoning", "id": "rs_1", "summary": "not a list"}, + id="summary-is-text", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": None}, + {"type": "reasoning", "id": "rs_1", "summary": None}, + id="summary-is-null", + ), + pytest.param( + {"type": "reasoning", "summary": []}, + {"type": "reasoning", "summary": []}, + id="id-absent", + ), + pytest.param( + {"type": "reasoning", "id": "rs_1", "summary": [], "status": "sideways"}, + {"type": "reasoning", "id": "rs_1", "summary": [], "status": "sideways"}, + id="status-outside-the-enum", + ), + ), +) +def test_responses_forwards_a_reasoning_input_item_to_openai( + gateway: Gateway, + item: dict[str, JsonValue], + expected: dict[str, JsonValue], +) -> None: + with wire_server(answering_model_discovery(_peer)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": [_FIRST_TURN, item, _LAST_TURN]} + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["status"] == "completed", response.text + assert payload["output"][0]["content"][0]["text"] == "reasoning item control", response.text + requests: Final = tuple(request for request in wire.drain() if request.method == "POST") + assert len(requests) == 1, requests + outbound: Final = _JSON_OBJECT.validate_json(requests[0].body)["input"] + assert isinstance(outbound, list) and len(outbound) == 3, requests[0].body + assert outbound[1] == expected, requests[0].body diff --git a/tests/integration/providers/test_validated_reply_boundaries_under_faults.py b/tests/integration/providers/test_validated_reply_boundaries_under_faults.py new file mode 100644 index 00000000000..0a358e685f6 --- /dev/null +++ b/tests/integration/providers/test_validated_reply_boundaries_under_faults.py @@ -0,0 +1,503 @@ +import base64 +import json +import re +import signal +import threading +import uuid +from collections import Counter +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from itertools import product +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, object_value +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_API_KEY: Final = "synthetic-provider-key" +_ENVIRONMENT_KEY: Final = "synthetic-environment-key" +_XAI_BACKEND: Final = "grok-test" +_GEMINI_BACKEND: Final = "gemini-2.5-flash-image" +_BEDROCK_BACKEND: Final = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" +_OPENAI_BACKEND: Final = "gpt-4o-mini" +_XAI_MODEL: Final = "faults-xai-chat" +_GEMINI_MODEL: Final = "faults-gemini-image" +_BEDROCK_MODEL: Final = "faults-bedrock-invoke" +_OPENAI_MODEL: Final = "faults-openai" +_XAI_CHAT: Final = "/xai/v1/chat/completions" +_OPENAI_CHAT: Final = "/openai/v1/chat/completions" +_GENERATE: Final = f"/models/{_GEMINI_BACKEND}:generateContent" +_INVOKE: Final = f"/model/{_BEDROCK_BACKEND}/invoke" +_OPENAI_MODEL_LISTING: Final = "/openai/v1/models" +_CONTAINERS: Final = "/openai/v1/containers" +_CONTAINER: Final = "cntr_faults" +_CONTAINER_FILES: Final = f"{_CONTAINERS}/{_CONTAINER}/files" +_FILE: Final = "cfile_faults" +_BETAS: Final = ("context-1m-2025-08-07", "interleaved-thinking-2025-05-14") +_IMAGE_CONFIG: Final[dict[str, JsonValue]] = {"aspectRatio": "1:1"} +_SERVER_TOOL_USAGE: Final[dict[str, JsonValue]] = {"web_search_calls": 2, "x_search_calls": 1} +_PROMPT_TOKENS: Final = 1000 +_COMPLETION_TOKENS: Final = 500 +_BURST_ROUNDS: Final = 8 +_GATE_SECONDS: Final = 60 +_PNG: Final = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg==" +) +_IMAGE_USAGE: Final[dict[str, JsonValue]] = { + "total_tokens": 1553, + "input_tokens": 263, + "input_tokens_details": {"image_tokens": 258, "text_tokens": 5}, + "output_tokens": 1290, + "output_tokens_details": {"image_tokens": 1290, "text_tokens": 0}, +} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_PACKAGED_PRICES: Final = Path(__file__).resolve().parents[3] / "litellm" / "model_prices_and_context_window_backup.json" + + +def _chat_reply(model: str, usage: dict[str, JsonValue]) -> Reply: + return Reply( + body=json.dumps( + { + "id": "chatcmpl-faults", + "object": "chat.completion", + "created": 1, + "model": model, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "faults control"}, "finish_reason": "stop"} + ], + "usage": usage, + } + ).encode() + ) + + +_XAI_REPLIES: Final = MappingProxyType( + { + "healthy": _chat_reply( + _XAI_BACKEND, + { + "prompt_tokens": 12, + "completion_tokens": 5, + "total_tokens": 17, + "server_side_tool_usage_details": _SERVER_TOOL_USAGE, + }, + ), + "dropped": Reply(drop_connection=True), + "overloaded": Reply( + status=503, body=json.dumps({"error": {"message": "overloaded", "type": "server_error"}}).encode() + ), + } +) +_REPLIES: Final = MappingProxyType( + { + ("POST", _OPENAI_CHAT): _chat_reply( + _OPENAI_BACKEND, + { + "prompt_tokens": _PROMPT_TOKENS, + "completion_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, + }, + ), + ("POST", _GENERATE): Reply( + body=json.dumps( + { + "candidates": [ + { + "content": { + "parts": [ + {"inlineData": {"mimeType": "image/png", "data": base64.b64encode(_PNG).decode()}} + ] + } + } + ], + "usageMetadata": { + "promptTokenCount": 263, + "candidatesTokenCount": 1290, + "totalTokenCount": 1553, + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 5}, + {"modality": "IMAGE", "tokenCount": 258}, + ], + "candidatesTokensDetails": [{"modality": "IMAGE", "tokenCount": 1290}], + }, + } + ).encode() + ), + ("POST", _INVOKE): Reply( + body=json.dumps( + { + "id": "msg_faults", + "type": "message", + "role": "assistant", + "model": _BEDROCK_BACKEND, + "content": [{"type": "text", "text": "faults control"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 12, "output_tokens": 5}, + } + ).encode() + ), + ("GET", _OPENAI_MODEL_LISTING): Reply(body=json.dumps({"object": "list", "data": []}).encode()), + ("POST", _CONTAINERS): Reply( + body=json.dumps( + {"id": _CONTAINER, "object": "container", "created_at": 1, "status": "running", "name": "faults"} + ).encode() + ), + ("GET", _CONTAINER_FILES): Reply( + body=json.dumps( + { + "object": "list", + "data": [ + { + "id": _FILE, + "object": "container.file", + "container_id": _CONTAINER, + "created_at": 1, + "bytes": 5, + "path": "/mnt/data/notes.txt", + "source": "user", + } + ], + "first_id": _FILE, + "last_id": _FILE, + "has_more": False, + } + ).encode() + ), + } +) + + +def _path(request: Request) -> str: + return request.target.split("?")[0] + + +def _marker(request: Request) -> str: + messages: Final = _JSON_OBJECT.validate_json(request.body)["messages"] + assert isinstance(messages, list), request.body + return str(object_value(messages[-1])["content"]) + + +def _peer(request: Request) -> Reply: + path: Final = _path(request) + if path != _GENERATE: + assert request.headers["authorization"] == f"Bearer {_API_KEY}", path + if path == _XAI_CHAT: + return _XAI_REPLIES[_marker(request).split("-")[0]] + return _REPLIES[(request.method, path)] + + +@dataclass(frozen=True, slots=True) +class _Call: + send: Callable[[Gateway], httpx.Response] + check: Callable[[httpx.Response], None] + + +def _uncached(marker: str) -> str: + return f"{marker}-{uuid.uuid4().hex}" + + +def _chat(model: str, marker: str, gateway: Gateway) -> httpx.Response: + return gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": _uncached(marker)}]} + ) + + +def _beta_chat(gateway: Gateway) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": _BEDROCK_MODEL, + "messages": [{"role": "user", "content": _uncached("synthetic beta request")}], + "max_tokens": 16, + }, + headers={"anthropic-beta": json.dumps(_BETAS)}, + ) + + +def _edit(image_config: str, gateway: Gateway) -> httpx.Response: + return gateway.request_multipart( + "/v1/images/edits", + {"model": _GEMINI_MODEL, "prompt": "synthetic edit request", "imageConfig": image_config}, + {"image": ("pixel.png", _PNG, "image/png")}, + ) + + +def _list_files(container: str, gateway: Gateway) -> httpx.Response: + return gateway.request("GET", f"/v1/containers/{container}/files") + + +def _error(response: httpx.Response) -> dict[str, JsonValue]: + return object_value(_JSON_OBJECT.validate_json(response.content)["error"]) + + +def _check_enriched(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + usage: Final = object_value(_JSON_OBJECT.validate_json(response.content)["usage"]) + assert usage["server_side_tool_usage_details"] == _SERVER_TOOL_USAGE, response.text + assert object_value(usage["prompt_tokens_details"])["web_search_requests"] == 2, response.text + + +def _check_dropped(response: httpx.Response) -> None: + assert response.status_code == 500, response.text + error: Final = _error(response) + assert error["type"] == "internal_server_error" and error["code"] == "500", response.text + assert "XaiException - Server disconnected" in str(error["message"]), response.text + + +def _check_overloaded(response: httpx.Response) -> None: + assert response.status_code == 503, response.text + error: Final = _error(response) + assert error["type"] == "internal_server_error" and error["code"] == "503", response.text + assert str(error["message"]).startswith( + "litellm.ServiceUnavailableError: ServiceUnavailableError: XaiException - " + ), response.text + + +def _check_image(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["usage"] == _IMAGE_USAGE, response.text + + +def _check_rejected_image_config(response: httpx.Response) -> None: + assert response.status_code == 400, response.text + error: Final = _error(response) + assert error["type"] == "invalid_request_error" and error["code"] == "400", response.text + assert str(error["message"]).startswith( + "litellm.UnsupportedParamsError: `imageConfig` must be valid JSON when provided as a string." + ), response.text + + +def _check_beta(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + usage: Final = object_value(_JSON_OBJECT.validate_json(response.content)["usage"]) + assert usage["prompt_tokens"] == 12 and usage["completion_tokens"] == 5, response.text + + +def _check_priced(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + packaged: Final = object_value(_JSON_OBJECT.validate_json(_PACKAGED_PRICES.read_bytes())[_OPENAI_BACKEND]) + input_rate: Final = packaged["input_cost_per_token"] + output_rate: Final = packaged["output_cost_per_token"] + assert isinstance(input_rate, float) and isinstance(output_rate, float), packaged + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx( + _PROMPT_TOKENS * input_rate + _COMPLETION_TOKENS * output_rate, rel=1e-9 + ) + + +def _check_files(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + listed: Final = _JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(listed, list) and [object_value(item)["id"] for item in listed] == [_FILE], response.text + + +_ENRICHED: Final = _Call(partial(_chat, _XAI_MODEL, "healthy"), _check_enriched) +_IMAGE: Final = _Call(partial(_edit, json.dumps(_IMAGE_CONFIG)), _check_image) +_REJECTED_IMAGE_CONFIG: Final = _Call(partial(_edit, "{not json"), _check_rejected_image_config) +_BETA: Final = _Call(_beta_chat, _check_beta) +_PRICED: Final = _Call(partial(_chat, _OPENAI_MODEL, "synthetic priced request"), _check_priced) + + +def _burst(gateway: Gateway, calls: tuple[_Call, ...]) -> None: + with ThreadPoolExecutor(max_workers=len(calls)) as pool: + futures: Final = tuple(pool.submit(call.send, gateway) for call in calls) + responses: Final = tuple(future.result(timeout=60) for future in futures) + for call, response in zip(calls, responses, strict=True): + call.check(response) + + +def _attempt(call: _Call, gateway: Gateway) -> httpx.Response | None: + try: + return call.send(gateway) + except httpx.TransportError: + return None + + +def _faults_config(wire: Wire, directory: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["model_list"] = [ + { + "model_name": _XAI_MODEL, + "litellm_params": {"model": f"xai/{_XAI_BACKEND}", "api_base": f"{wire.url}/xai/v1", "api_key": _API_KEY}, + }, + { + "model_name": _GEMINI_MODEL, + "litellm_params": {"model": f"gemini/{_GEMINI_BACKEND}", "api_base": wire.url, "api_key": _API_KEY}, + }, + { + "model_name": _BEDROCK_MODEL, + "litellm_params": { + "model": f"bedrock/invoke/{_BEDROCK_BACKEND}", + "api_key": _API_KEY, + "aws_region_name": "us-east-1", + "api_base": wire.url, + "aws_bedrock_runtime_endpoint": wire.url, + }, + }, + { + "model_name": _OPENAI_MODEL, + "litellm_params": { + "model": f"openai/{_OPENAI_BACKEND}", + "api_base": f"{wire.url}/openai/v1", + "api_key": _API_KEY, + }, + }, + ] + path: Final = directory / "validated-reply-boundaries-under-faults.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _environment(wire: Wire) -> dict[str, str]: + return { + "OPENAI_API_KEY": _ENVIRONMENT_KEY, + "OPENAI_API_BASE": f"{wire.url}/environment/v1", + "OPENAI_BASE_URL": f"{wire.url}/environment/v1", + } + + +def _sent_betas(request: Request) -> tuple[str, ...]: + betas: Final = _JSON_OBJECT.validate_json(request.body)["anthropic_beta"] + assert isinstance(betas, list), request.body + return tuple(sorted(str(beta) for beta in betas)) + + +def _sent_image_config(request: Request) -> JsonValue: + return object_value(_JSON_OBJECT.validate_json(request.body)["generationConfig"])["imageConfig"] + + +def _assert_outbound(received: tuple[Request, ...], expected: dict[tuple[str, str], int]) -> None: + assert ( + Counter( + (request.method, _path(request)) for request in received if _path(request) != _OPENAI_MODEL_LISTING + ) + == expected + ) + assert {_sent_betas(request) for request in received if _path(request) == _INVOKE} == {tuple(sorted(_BETAS))} + assert all(_sent_image_config(request) == _IMAGE_CONFIG for request in received if _path(request) == _GENERATE) + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def test_xai_usage_enrichment_survives_a_burst_the_provider_partly_drops_and_overloads(gateway: Gateway) -> None: + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"xai/{_XAI_BACKEND}", api_base=f"{wire.url}/xai/v1", api_key=_API_KEY) + kinds: Final = (("healthy", _check_enriched), ("dropped", _check_dropped), ("overloaded", _check_overloaded)) + calls: Final = tuple( + _Call(partial(_chat, model, f"{kind}-{index}"), check) for index, (kind, check) in product(range(10), kinds) + ) + _burst(gateway, calls) + _check_enriched(_chat(model, "healthy-after-the-burst", gateway)) + received: Final = wire.drain() + assert Counter(_marker(request).split("-")[0] for request in received) == { + "healthy": 11, + "dropped": 10, + "overloaded": 10, + } + + +@pytest.mark.timeout(180) +def test_a_worker_killed_mid_burst_leaves_the_sibling_validating_provider_payloads( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + + def gated(request: Request) -> Reply: + if _path(request) != _OPENAI_MODEL_LISTING: + held.put(request.target) + assert release.wait(timeout=_GATE_SECONDS), "The burst was never released" + return _peer(request) + + calls: Final = (_ENRICHED, _IMAGE, _BETA) * _BURST_ROUNDS + with wire_server(gated) as wire: + config: Final = _faults_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, _environment(wire), config=config, workers=2) as owned: + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + with ( + httpx.Client( + base_url=owned.gateway.client.base_url, timeout=2 * _GATE_SECONDS, trust_env=False + ) as outlasting_the_gate, + ThreadPoolExecutor(max_workers=len(calls)) as pool, + ): + patient: Final = Gateway(outlasting_the_gate, owned.gateway.key, owned.gateway.upstream_url) + futures: Final = tuple(pool.submit(_attempt, call, patient) for call in calls) + eventually(held.qsize, lambda size: size == len(calls), seconds=_GATE_SECONDS) + held_by: Final = {pid: _open_upstream_connections(pid, wire.url) for pid in workers} + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + outcomes: Final = tuple(future.result(timeout=60) for future in futures) + assert sum(held_by.values()) == len(calls), held_by + served: Final = tuple( + (call, response) for call, response in zip(calls, outcomes, strict=True) if response is not None + ) + assert len(served) == held_by[survivor_pid] >= len(calls) // 2, (held_by, len(served)) + for call, response in served: + call.check(response) + _burst(owned.gateway, (_ENRICHED, _IMAGE, _BETA, _REJECTED_IMAGE_CONFIG, _PRICED)) + _assert_outbound( + wire.drain(), + { + ("POST", _XAI_CHAT): _BURST_ROUNDS + 1, + ("POST", _GENERATE): _BURST_ROUNDS + 1, + ("POST", _INVOKE): _BURST_ROUNDS + 1, + ("POST", _OPENAI_CHAT): 1, + }, + ) + + +@pytest.mark.timeout(240) +def test_a_restarted_proxy_answers_with_the_same_validated_shapes(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(_peer) as wire: + config: Final = _faults_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, _environment(wire), config=config, workers=2) as first: + container: Final = str(first.gateway.post("/v1/containers", {"model": _OPENAI_MODEL, "name": "faults"})["id"]) + calls: Final = ( + _ENRICHED, + _IMAGE, + _REJECTED_IMAGE_CONFIG, + _BETA, + _PRICED, + _Call(partial(_list_files, container), _check_files), + ) * (_BURST_ROUNDS // 2) + _burst(first.gateway, calls) + with owned_proxy_process(gateway, tmp_path, _environment(wire), config=config, workers=2) as second: + _burst(second.gateway, calls) + _assert_outbound( + wire.drain(), + { + ("POST", _XAI_CHAT): _BURST_ROUNDS, + ("POST", _GENERATE): _BURST_ROUNDS, + ("POST", _INVOKE): _BURST_ROUNDS, + ("POST", _OPENAI_CHAT): _BURST_ROUNDS, + ("POST", _CONTAINERS): 1, + ("GET", _CONTAINER_FILES): _BURST_ROUNDS, + }, + ) diff --git a/tests/integration/providers/test_xai_server_tool_usage_wire.py b/tests/integration/providers/test_xai_server_tool_usage_wire.py new file mode 100644 index 00000000000..1dbe1a05dc8 --- /dev/null +++ b/tests/integration/providers/test_xai_server_tool_usage_wire.py @@ -0,0 +1,124 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "grok-test" +_API_KEY: Final = "synthetic-xai-key" +_REPLY_TEXT: Final = "server tool usage control" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_PLAIN_USAGE: Final[dict[str, JsonValue]] = {"prompt_tokens": 12, "completion_tokens": 5, "total_tokens": 17} +_ABSENT: Final = "absent" + + +def _reply(usage: JsonValue) -> bytes: + return json.dumps( + { + "id": "chatcmpl-xai-usage", + "object": "chat.completion", + "created": 1, + "model": _BACKEND, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": _REPLY_TEXT}, "finish_reason": "stop"} + ], + "usage": usage, + } + ).encode() + + +def _peer(usage: JsonValue): + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + assert _JSON_OBJECT.validate_json(request.body)["messages"] == [ + {"role": "user", "content": "synthetic usage request"} + ], request.body + return Reply(body=_reply(usage)) + + return respond + + +def _usage_with(details: JsonValue) -> dict[str, JsonValue]: + return {**_PLAIN_USAGE, "server_side_tool_usage_details": details} + + +@pytest.mark.parametrize( + ("usage", "expected_details", "expected_web_search_requests"), + ( + pytest.param( + _usage_with({"web_search_calls": 2, "x_search_calls": 1}), + {"web_search_calls": 2, "x_search_calls": 1}, + 2, + id="web-search-calls-counted", + ), + pytest.param(_usage_with({"x_search_calls": 3}), {"x_search_calls": 3}, None, id="no-web-search-calls"), + pytest.param( + _usage_with({"web_search_calls": "many"}), {"web_search_calls": "many"}, None, id="call-count-is-a-word" + ), + pytest.param(_usage_with([1, 2]), [1, 2], None, id="details-are-a-list"), + pytest.param(_usage_with(None), _ABSENT, None, id="details-are-null"), + pytest.param(_PLAIN_USAGE, _ABSENT, None, id="details-absent"), + ), +) +def test_xai_chat_reports_server_side_tool_usage_from_the_provider_reply( + gateway: Gateway, + usage: dict[str, JsonValue], + expected_details: JsonValue, + expected_web_search_requests: int | None, +) -> None: + with wire_server(_peer(usage)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"xai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic usage request"}]}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["choices"][0]["message"]["content"] == _REPLY_TEXT, response.text + reported: Final = payload["usage"] + assert isinstance(reported, dict), response.text + assert {name: reported[name] for name in _PLAIN_USAGE} == _PLAIN_USAGE, response.text + assert reported.get("server_side_tool_usage_details", _ABSENT) == expected_details, response.text + prompt_details: Final = reported.get("prompt_tokens_details") + web_search_requests: Final = ( + prompt_details.get("web_search_requests") if isinstance(prompt_details, dict) else None + ) + assert web_search_requests == expected_web_search_requests, response.text + assert len(wire.drain()) == 1 + + +def test_xai_chat_reports_zero_usage_when_the_provider_usage_is_null(gateway: Gateway) -> None: + with wire_server(_peer(None)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"xai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic usage request"}]}, + ) + assert response.status_code == 200, response.text + reported: Final = _JSON_OBJECT.validate_json(response.content)["usage"] + assert isinstance(reported, dict), response.text + assert {name: reported[name] for name in _PLAIN_USAGE} == dict.fromkeys(_PLAIN_USAGE, 0), response.text + assert "server_side_tool_usage_details" not in reported, response.text + assert len(wire.drain()) == 1 + + +@pytest.mark.parametrize("usage", (pytest.param([1, 2], id="usage-is-a-list"), pytest.param("lots", id="usage-is-text"))) +def test_xai_chat_answers_500_when_the_provider_usage_is_not_an_object(gateway: Gateway, usage: JsonValue) -> None: + with wire_server(_peer(usage)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"xai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic usage request"}]}, + ) + assert response.status_code == 500, response.text + error: Final = _JSON_OBJECT.validate_json(response.content)["error"] + assert isinstance(error, dict), response.text + assert error["type"] == "internal_server_error", response.text + assert "XaiException - Invalid response object" in str(error["message"]), response.text + assert len(wire.drain()) >= 1 diff --git a/tests/unit/integrations/gitlab/test_init.py b/tests/unit/integrations/gitlab/test_init.py new file mode 100644 index 00000000000..b8eab935016 --- /dev/null +++ b/tests/unit/integrations/gitlab/test_init.py @@ -0,0 +1,56 @@ +import httpx +import pytest +import respx + +from litellm.integrations.gitlab import prompt_initializer_registry +from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec + +_GITLAB_CONFIG = { + "project": "group/repo", + "access_token": "glpat-test", + "base_url": "https://gitlab.example.test/api/v4", + "prompts_path": "prompts", +} + + +def _prompt_spec(litellm_params: PromptLiteLLMParams) -> PromptSpec: + return PromptSpec(prompt_id="chat/greet", litellm_params=litellm_params) + + +def test_gitlab_initializer_loads_the_prompt_file_at_the_per_prompt_git_ref(respx_mock: respx.MockRouter) -> None: + route = respx_mock.get(host="gitlab.example.test").mock( + return_value=httpx.Response( + 200, text="---\nmodel: gpt-4o\n---\nHello {{name}}", headers={"content-type": "text/plain"} + ) + ) + litellm_params = PromptLiteLLMParams(prompt_integration="gitlab", gitlab_config=_GITLAB_CONFIG, git_ref="v1.2.0") + + manager = prompt_initializer_registry["gitlab"](litellm_params, _prompt_spec(litellm_params)) + fetches_by_initializer = route.call_count + + assert manager.get_prompt_template("chat/greet", {"name": "Ada"}) == ( + "Hello Ada", + {"model": "gpt-4o", "temperature": None, "max_tokens": None}, + ) + assert (fetches_by_initializer, route.call_count) == (1, 1) + request = route.calls.last.request + assert str(request.url) == ( + "https://gitlab.example.test/api/v4/projects/group%2Frepo/repository/files/prompts%2Fchat%2Fgreet.prompt/raw" + "?ref=v1.2.0" + ) + assert request.headers["Private-Token"] == "glpat-test" + + +@pytest.mark.parametrize( + "integration_params", + [ + pytest.param({}, id="config-absent"), + pytest.param({"gitlab_config": None}, id="config-none"), + pytest.param({"gitlab_config": {}}, id="config-empty"), + ], +) +def test_gitlab_initializer_rejects_prompt_without_gitlab_config(integration_params: dict[str, object]) -> None: + litellm_params = PromptLiteLLMParams(prompt_integration="gitlab", git_ref="v1.2.0", **integration_params) + + with pytest.raises(ValueError, match="gitlab_config is required for gitlab prompt integration"): + prompt_initializer_registry["gitlab"](litellm_params, _prompt_spec(litellm_params)) diff --git a/tests/unit/interactions/test_litellm_responses_bridge.py b/tests/unit/interactions/test_litellm_responses_bridge.py index 3abd0a6ca98..58f2756e7fb 100644 --- a/tests/unit/interactions/test_litellm_responses_bridge.py +++ b/tests/unit/interactions/test_litellm_responses_bridge.py @@ -5,11 +5,17 @@ Inherits from BaseInteractionsTest to run the same test suite against the litellm_responses bridge provider, which calls litellm.responses() internally. """ +from collections.abc import Mapping +from datetime import datetime +from typing import Final + +import pytest from litellm.interactions.litellm_responses_transformation.transformation import ( LiteLLMResponsesInteractionsConfig, ) from litellm.types.interactions import Turn +from litellm.types.llms.openai import ResponsesAPIResponse class TestBridgeInputTransformation: @@ -78,3 +84,61 @@ class TestBridgeInputTransformation: [{"type": "user_input", "content": [image_part]}] ) assert transformed == [{"role": "user", "content": [image_part]}] + + +def _responses_api_response(status: str, usage: Mapping[str, int] | None) -> ResponsesAPIResponse: + return ResponsesAPIResponse( + id="resp_123", + created_at=1700000000, + model="gpt-4o", + object="response", + status=status, + output=[ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Hello there", "annotations": []}], + } + ], + usage=usage, + ) + + +class TestBridgeResponseTransformation: + @pytest.mark.parametrize( + ("status", "usage", "model", "expected_model", "expected_usage"), + [ + ( + "completed", + {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7}, + None, + "gpt-4o", + {"total_input_tokens": 3, "total_output_tokens": 4}, + ), + ("in_progress", None, "bridge-model", "bridge-model", None), + ], + ) + def test_responses_response_becomes_interaction( + self, + status: str, + usage: Mapping[str, int] | None, + model: str | None, + expected_model: str, + expected_usage: Mapping[str, int] | None, + ): + interaction: Final = LiteLLMResponsesInteractionsConfig.transform_responses_response_to_interactions_response( + _responses_api_response(status, usage), model + ) + + assert interaction.id == "resp_123" + assert interaction.object == "interaction" + assert interaction.status == status + assert interaction.model == expected_model + assert interaction.outputs == [{"type": "text", "text": "Hello there"}] + assert interaction.steps == [{"type": "model_output", "content": [{"type": "text", "text": "Hello there"}]}] + assert interaction.usage == expected_usage + assert interaction.created is not None + assert datetime.fromisoformat(interaction.created).timestamp() == 1700000000 + assert interaction.updated == interaction.created diff --git a/tests/unit/llms/azure/completion/__init__.py b/tests/unit/llms/azure/completion/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/azure/completion/test_handler.py b/tests/unit/llms/azure/completion/test_handler.py new file mode 100644 index 00000000000..9208938fcc2 --- /dev/null +++ b/tests/unit/llms/azure/completion/test_handler.py @@ -0,0 +1,52 @@ +from typing import Final + +import pytest +import respx +from httpx import Response + +import litellm +from litellm import atext_completion, text_completion + +COMPLETIONS_URL: Final = ( + "https://example-resource.openai.azure.com/openai/deployments/gpt-35-turbo-instruct/completions" +) +AZURE_TEXT_CALL: Final = { + "model": "azure_text/gpt-35-turbo-instruct", + "prompt": "hello", + "api_base": "https://example-resource.openai.azure.com", + "api_version": "2024-02-01", + "api_key": "azure-test-key", +} + + +@pytest.fixture +def rejected_completions_endpoint(): + return respx.post(url__startswith=COMPLETIONS_URL).mock( + return_value=Response( + 400, + json={"error": {"message": "bad prompt", "code": "invalid_prompt"}}, + headers={"x-request-id": "req-1"}, + ) + ) + + +@respx.mock +def test_completion_surfaces_the_status_and_headers_of_a_rejected_request(rejected_completions_endpoint): + with pytest.raises(litellm.BadRequestError) as rejected: + text_completion(**AZURE_TEXT_CALL) + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" + + +@respx.mock +@pytest.mark.parametrize("stream", [False, True]) +async def test_acompletion_surfaces_the_status_and_headers_of_a_rejected_request( + rejected_completions_endpoint, monkeypatch, stream: bool +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with pytest.raises(litellm.BadRequestError) as rejected: + await atext_completion(**AZURE_TEXT_CALL, stream=stream) + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" diff --git a/tests/unit/llms/azure_ai/agents/__init__.py b/tests/unit/llms/azure_ai/agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/azure_ai/agents/test_transformation.py b/tests/unit/llms/azure_ai/agents/test_transformation.py new file mode 100644 index 00000000000..47767ed997a --- /dev/null +++ b/tests/unit/llms/azure_ai/agents/test_transformation.py @@ -0,0 +1,67 @@ +from typing import Final + +import pytest + +from litellm.llms.azure_ai.agents.transformation import AzureAIAgentsConfig + +MESSAGES: Final = [ + {"role": "system", "content": "be brief"}, + {"role": "user", "content": [{"type": "text", "text": "hello "}, {"type": "text", "text": "world"}]}, +] +FLATTENED_MESSAGES: Final = [ + {"role": "system", "content": "be brief"}, + {"role": "user", "content": "hello world"}, +] + + +@pytest.mark.parametrize( + ("optional_params", "expected"), + [ + pytest.param( + {}, + { + "agent_id": "asst_from_model", + "messages": FLATTENED_MESSAGES, + "api_version": AzureAIAgentsConfig.DEFAULT_API_VERSION, + }, + id="agent-id-from-model-and-default-api-version", + ), + pytest.param( + { + "agent_id": "asst_override", + "api_version": "2025-05-15-preview", + "thread_id": "thread_9", + "instructions": "use the search tool", + }, + { + "agent_id": "asst_override", + "messages": FLATTENED_MESSAGES, + "api_version": "2025-05-15-preview", + "thread_id": "thread_9", + "instructions": "use the search tool", + }, + id="caller-agent-id-api-version-thread-and-instructions", + ), + pytest.param( + {"assistant_id": "asst_legacy_name"}, + { + "agent_id": "asst_legacy_name", + "messages": FLATTENED_MESSAGES, + "api_version": AzureAIAgentsConfig.DEFAULT_API_VERSION, + }, + id="assistant-id-names-the-agent", + ), + ], +) +def test_transform_request_builds_the_agent_run_payload( + optional_params: dict[str, object], expected: dict[str, object] +) -> None: + payload = AzureAIAgentsConfig().transform_request( + model="azure_ai/agents/asst_from_model", + messages=MESSAGES, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert payload == expected diff --git a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py index dcb6aec8091..bfcdd8da3d8 100644 --- a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py +++ b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py @@ -9,7 +9,10 @@ import json from unittest.mock import MagicMock, patch import pytest +import respx +from httpx import Response +import litellm from litellm.llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion from litellm.types.utils import ModelResponse @@ -248,3 +251,26 @@ class TestAzureAnthropicChatCompletion: mock_get_client.assert_called_once_with(params={"timeout": timeout}) assert mock_client.post.call_args.kwargs["timeout"] == timeout assert result is not None + + +@respx.mock +def test_completion_surfaces_the_status_and_headers_of_a_rejected_request(): + respx.post("https://example-resource.services.ai.azure.com/anthropic/v1/messages").mock( + return_value=Response( + 400, + json={"type": "error", "error": {"type": "invalid_request_error", "message": "max_tokens too large"}}, + headers={"x-request-id": "req-1"}, + ) + ) + + with pytest.raises(litellm.BadRequestError) as rejected: + litellm.completion( + model="azure_ai/claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + api_base="https://example-resource.services.ai.azure.com/anthropic", + api_key="azure-test-key", + ) + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" + assert "max_tokens too large" in str(rejected.value) diff --git a/tests/unit/llms/base_llm/guardrail_translation/__init__.py b/tests/unit/llms/base_llm/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/guardrail_translation/test_utils.py b/tests/unit/llms/base_llm/guardrail_translation/test_utils.py new file mode 100644 index 00000000000..65208a2bd3b --- /dev/null +++ b/tests/unit/llms/base_llm/guardrail_translation/test_utils.py @@ -0,0 +1,34 @@ +from typing import Final + +import pytest + +import litellm +from litellm.llms.base_llm.guardrail_translation.utils import ( + effective_skip_system_message_for_guardrail, + effective_skip_tool_message_for_guardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail + + +@pytest.mark.parametrize( + ("per_guardrail", "global_default", "expected"), + [ + (True, False, True), + (False, True, False), + (None, True, True), + (None, False, False), + ], +) +def test_a_per_guardrail_skip_flag_wins_over_the_global_setting( + monkeypatch: pytest.MonkeyPatch, per_guardrail: bool | None, global_default: bool, expected: bool +): + monkeypatch.setattr(litellm, "skip_system_message_in_guardrail", global_default) + monkeypatch.setattr(litellm, "skip_tool_message_in_guardrail", global_default) + guardrail: Final = LakeraAIGuardrail( + api_key="lakera-test-key", + skip_system_message_in_guardrail=per_guardrail, + skip_tool_message_in_guardrail=per_guardrail, + ) + + assert effective_skip_system_message_for_guardrail(guardrail) is expected + assert effective_skip_tool_message_for_guardrail(guardrail) is expected diff --git a/tests/unit/llms/bedrock/test_bedrock_common_utils.py b/tests/unit/llms/bedrock/test_bedrock_common_utils.py index d4c182dc952..09354155221 100644 --- a/tests/unit/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/unit/llms/bedrock/test_bedrock_common_utils.py @@ -1133,3 +1133,21 @@ def test_build_bedrock_stream_error_resolves_status_from_the_exception_type( assert error.status_code == expected_status assert error.message == expected_message + + +@pytest.mark.parametrize( + ("header_value", "expected"), + [ + ( + '["interleaved-thinking-2025-05-14", "claude-code-20250219"]', + ["interleaved-thinking-2025-05-14", "claude-code-20250219"], + ), + (' [" context-1m-2025-08-07 "] ', ["context-1m-2025-08-07"]), + ("[]", []), + ("[not-json]", ["[not-json]"]), + ], +) +def test_get_anthropic_beta_from_headers_reads_a_json_array_header(header_value: str, expected: list[str]): + from litellm.llms.bedrock.common_utils import get_anthropic_beta_from_headers + + assert get_anthropic_beta_from_headers({"anthropic-beta": header_value}) == expected diff --git a/tests/unit/llms/custom_httpx/test_container_handler.py b/tests/unit/llms/custom_httpx/test_container_handler.py index a1b5a66696d..1b09165087a 100644 --- a/tests/unit/llms/custom_httpx/test_container_handler.py +++ b/tests/unit/llms/custom_httpx/test_container_handler.py @@ -1,3 +1,4 @@ +from typing import Final from unittest.mock import MagicMock import httpx @@ -7,6 +8,7 @@ import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.container_handler import generic_container_handler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.containers.main import DeleteContainerFileResponse from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager @@ -100,3 +102,25 @@ def test_json_endpoint_still_raises_provider_error_message(): assert exc_info.value.status_code == 404 assert exc_info.value.message == "File not found." + + +def test_json_endpoint_sends_the_configured_route_and_parses_its_response_model(): + handler: Final = HTTPHandler() + handler.client = httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response( + 200, + json={ + "id": request.url.path, + "object": "container.file.deleted", + "deleted": request.method == "DELETE", + }, + ) + ) + ) + + response: Final = _handle("delete_container_file", handler) + + assert type(response) is DeleteContainerFileResponse + assert response.id.endswith("/containers/cntr_real/files/cfile_nonexistent") + assert response.deleted is True diff --git a/tests/unit/llms/gemini/test_image_usage_transformation.py b/tests/unit/llms/gemini/test_image_usage_transformation.py new file mode 100644 index 00000000000..54e0668d2d9 --- /dev/null +++ b/tests/unit/llms/gemini/test_image_usage_transformation.py @@ -0,0 +1,47 @@ +from typing import Final + +import pytest + +from litellm.llms.gemini.image_usage_transformation import transform_gemini_image_usage + +PROMPT_DETAILS: Final = [{"modality": "TEXT", "tokenCount": 30}, {"modality": "IMAGE", "tokenCount": 5}] + + +@pytest.mark.parametrize( + ("candidates_details", "expected_output_details"), + [ + ({}, {"image_tokens": 1716, "text_tokens": 0}), + ( + {"candidatesTokensDetails": [{"modality": "IMAGE", "tokenCount": 1120}]}, + {"image_tokens": 1120, "text_tokens": 596}, + ), + ( + {"candidates_tokens_details": [{"modality": "TEXT", "token_count": 16}]}, + {"image_tokens": 0, "text_tokens": 1716}, + ), + ], +) +def test_transform_gemini_image_usage_reports_image_and_chat_style_counts( + candidates_details: dict[str, object], expected_output_details: dict[str, int] +): + usage: Final = transform_gemini_image_usage( + { + "promptTokenCount": 35, + "candidatesTokenCount": 1716, + "totalTokenCount": 1751, + "promptTokensDetails": PROMPT_DETAILS, + **candidates_details, + } + ) + + assert usage.model_dump() == { + "input_tokens": 35, + "input_tokens_details": {"image_tokens": 5, "text_tokens": 30}, + "output_tokens": 1716, + "total_tokens": 1751, + "prompt_tokens": 35, + "prompt_tokens_details": {"image_tokens": 5, "text_tokens": 30}, + "completion_tokens": 1716, + "completion_tokens_details": expected_output_details, + "output_tokens_details": expected_output_details, + } diff --git a/tests/unit/llms/litellm_proxy/skills/test_prompt_injection.py b/tests/unit/llms/litellm_proxy/skills/test_prompt_injection.py new file mode 100644 index 00000000000..431c5ed9332 --- /dev/null +++ b/tests/unit/llms/litellm_proxy/skills/test_prompt_injection.py @@ -0,0 +1,83 @@ +from typing import Final + +import pytest + +from litellm.llms.litellm_proxy.skills.prompt_injection import SkillPromptInjectionHandler +from litellm.proxy._types import LiteLLM_SkillsTable + +NO_ARGUMENTS_SCHEMA: Final = {"type": "object", "properties": {}, "required": []} +SQL_ARGUMENTS_SCHEMA: Final = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], +} + + +@pytest.mark.parametrize( + ("skill", "expected"), + [ + pytest.param( + LiteLLM_SkillsTable( + skill_id="translate-file v2", + display_title="Document Translator", + description="Converts files between languages", + instructions="Translate the uploaded document", + ), + { + "name": "translate_file_v2", + "description": "Translate the uploaded document", + "input_schema": NO_ARGUMENTS_SCHEMA, + }, + id="instructions-describe-the-tool-and-the-id-becomes-a-function-name", + ), + pytest.param( + LiteLLM_SkillsTable( + skill_id="warehouse_sql", + display_title="Warehouse SQL Analyst", + description="Runs SQL against the inventory database", + metadata={"parameters": SQL_ARGUMENTS_SCHEMA}, + ), + { + "name": "warehouse_sql", + "description": "Runs SQL against the inventory database", + "input_schema": SQL_ARGUMENTS_SCHEMA, + }, + id="metadata-parameters-become-the-input-schema", + ), + pytest.param( + LiteLLM_SkillsTable( + skill_id="trip-planner", + display_title="Trip Planner", + metadata={"parameters": "not a schema"}, + ), + {"name": "trip_planner", "description": "Trip Planner", "input_schema": NO_ARGUMENTS_SCHEMA}, + id="non-dict-parameters-keep-the-no-arguments-schema", + ), + pytest.param( + LiteLLM_SkillsTable(skill_id="bare"), + {"name": "bare", "description": "Skill: bare", "input_schema": NO_ARGUMENTS_SCHEMA}, + id="a-skill-with-no-text-is-described-by-its-id", + ), + ], +) +def test_convert_skill_to_anthropic_tool_builds_the_messages_api_tool( + skill: LiteLLM_SkillsTable, expected: dict[str, object] +) -> None: + assert SkillPromptInjectionHandler().convert_skill_to_anthropic_tool(skill) == expected + + +@pytest.mark.parametrize( + ("instructions", "expected_description"), + [ + ("x" * 1024, "x" * 1024), + ("x" * 1025, "x" * 1021 + "..."), + ], +) +def test_convert_skill_to_anthropic_tool_caps_the_description_at_1024_characters( + instructions: str, expected_description: str +) -> None: + tool = SkillPromptInjectionHandler().convert_skill_to_anthropic_tool( + LiteLLM_SkillsTable(skill_id="long-instructions", instructions=instructions) + ) + + assert tool["description"] == expected_description diff --git a/tests/unit/llms/openai/completion/test_completion_handler.py b/tests/unit/llms/openai/completion/test_completion_handler.py index 329956605ab..5c0f633430e 100644 --- a/tests/unit/llms/openai/completion/test_completion_handler.py +++ b/tests/unit/llms/openai/completion/test_completion_handler.py @@ -88,3 +88,62 @@ async def test_acompletion_forwards_client_headers_to_provider( request_headers = mock_completions_endpoint.calls.last.request.headers assert request_headers["x-mycorp-llmcall-id"] == "abc-123" + + +@pytest.fixture +def rejected_completions_endpoint(): + return respx.post("https://api.openai.com/v1/completions").mock( + return_value=Response( + 400, + json={"error": {"message": "bad prompt", "type": "invalid_request_error"}}, + headers={"x-request-id": "req-1"}, + ) + ) + + +@respx.mock +def test_completion_surfaces_the_status_and_headers_of_a_rejected_request(rejected_completions_endpoint): + with pytest.raises(litellm.BadRequestError) as rejected: + text_completion(model="gpt-3.5-turbo-instruct", prompt="hello") + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" + + +@respx.mock +def test_streaming_completion_surfaces_the_status_and_headers_of_a_rejected_request(rejected_completions_endpoint): + with pytest.raises(litellm.BadRequestError) as rejected: + list(text_completion(model="gpt-3.5-turbo-instruct", prompt="hello", stream=True)) + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" + + +@respx.mock +async def test_acompletion_surfaces_the_status_and_headers_of_a_rejected_request( + rejected_completions_endpoint, monkeypatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with pytest.raises(litellm.BadRequestError) as rejected: + await atext_completion(model="gpt-3.5-turbo-instruct", prompt="hello") + + assert rejected.value.status_code == 400 + assert rejected.value.litellm_response_headers["x-request-id"] == "req-1" + + +@respx.mock +async def test_async_streaming_completion_reports_an_error_event_sent_mid_stream(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + respx.post("https://api.openai.com/v1/completions").mock( + return_value=Response( + 200, + text='data: {"error": {"message": "upstream overloaded", "type": "server_error"}}\n\n', + headers={"content-type": "text/event-stream"}, + ) + ) + with pytest.raises(litellm.InternalServerError) as failed: + async for _ in await atext_completion(model="gpt-3.5-turbo-instruct", prompt="hello", stream=True): + pass + + assert failed.value.status_code == 500 + assert "upstream overloaded" in str(failed.value) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_transformation.py b/tests/unit/llms/openai/responses/test_openai_responses_transformation.py index 6fbf2c225e7..f33f61b3e3d 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_transformation.py @@ -1017,6 +1017,32 @@ class TestOpenAIResponsesAPIConfig: assert len(request["input"]) == len(history) - 1 assert all(item.get("type") != "reasoning" for item in request["input"]) + @pytest.mark.parametrize( + ("reasoning_item", "expected"), + [ + ( + {"id": "rs_1", "type": "reasoning", "summary": [], "status": None, "note": None}, + {"id": "rs_1", "type": "reasoning", "summary": []}, + ), + ( + {"id": "rs_1", "type": "reasoning", "summary": "not a list", "status": None, "note": None}, + {"id": "rs_1", "type": "reasoning", "summary": "not a list", "note": None}, + ), + ], + ) + def test_a_reasoning_input_item_loses_its_null_status_whether_or_not_it_fits_the_openai_model( + self, reasoning_item: dict[str, object], expected: dict[str, object] + ): + request: Final = self.config.transform_responses_api_request( + model="gpt-5.6", + input=[reasoning_item, {"role": "user", "content": "And Berlin?"}], + response_api_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert request["input"] == [expected, {"role": "user", "content": "And Berlin?"}] + class TestAzureResponsesAPIConfig: def setup_method(self): diff --git a/tests/unit/llms/vertex_ai/agent_engine/test_transformation.py b/tests/unit/llms/vertex_ai/agent_engine/test_transformation.py index 19616682c59..06aaca44e8a 100644 --- a/tests/unit/llms/vertex_ai/agent_engine/test_transformation.py +++ b/tests/unit/llms/vertex_ai/agent_engine/test_transformation.py @@ -5,13 +5,22 @@ Tests the request transformation and streaming chunk parsing without making real """ +import json +from datetime import datetime +from typing import Final + +import httpx import pytest +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.vertex_ai.agent_engine.sse_iterator import ( VertexAgentEngineResponseIterator, ) from litellm.llms.vertex_ai.agent_engine.transformation import VertexAgentEngineConfig +from litellm.types.utils import LlmProviders +from litellm.utils import CustomStreamWrapper class TestVertexAgentEngineTransformRequest: @@ -125,3 +134,39 @@ class TestVertexAgentEngineChunkParser: assert result.choices[0].delta.content == "Partial response..." assert result.choices[0].finish_reason is None assert result.usage is None + + +async def test_async_stream_wrapper_without_a_client_posts_through_the_cached_vertex_ai_client(): + api_base: Final = "https://us-central1-aiplatform.googleapis.com/v1/reasoningEngines/123:streamQuery" + sent_requests: Final[list[httpx.Request]] = [] + + def respond(request: httpx.Request) -> httpx.Response: + sent_requests.append(request) + return httpx.Response(200, text='{"content": {"parts": [{"text": "hi"}], "role": "model"}}') + + cached_client: Final = get_async_httpx_client(llm_provider=LlmProviders.VERTEX_AI, params={}) + cached_client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + messages: Final = [{"role": "user", "content": "hi"}] + + wrapper: Final = await VertexAgentEngineConfig().get_async_custom_stream_wrapper( + model="agent_engine/123", + custom_llm_provider="vertex_ai", + logging_obj=Logging( + model="agent_engine/123", + messages=messages, + stream=True, + call_type="acompletion", + start_time=datetime(2026, 1, 1), + litellm_call_id="call-1", + function_id="fn-1", + ), + api_base=api_base, + headers={}, + data={"class_method": "stream_query"}, + messages=messages, + litellm_params={}, + ) + + assert type(wrapper) is CustomStreamWrapper + assert [str(request.url) for request in sent_requests] == [api_base] + assert json.loads(sent_requests[0].content) == {"class_method": "stream_query"} diff --git a/tests/unit/llms/xai/test_xai_chat_transformation.py b/tests/unit/llms/xai/test_xai_chat_transformation.py index 704d0061103..601b384aae4 100644 --- a/tests/unit/llms/xai/test_xai_chat_transformation.py +++ b/tests/unit/llms/xai/test_xai_chat_transformation.py @@ -1,6 +1,8 @@ +from typing import Final from unittest.mock import Mock import httpx +import pytest import litellm from litellm.llms.xai.chat.transformation import ( @@ -204,6 +206,56 @@ class TestXAIChatWebSearchBilling: assert response.usage.prompt_tokens_details is None assert getattr(response.usage, "server_side_tool_usage_details", None) is None + @pytest.mark.parametrize( + ("tool_details", "expected_web_search_requests"), + [(_TOOL_DETAILS, 3), (None, None)], + ) + def test_transform_response_reads_tool_usage_details_from_the_response_body( + self, + tool_details: dict[str, int] | None, + expected_web_search_requests: int | None, + ): + raw_response: Final = httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-xai", + "object": "chat.completion", + "created": 0, + "model": "grok-4", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 100, + "completion_tokens": 20, + "total_tokens": 120, + "server_side_tool_usage_details": tool_details, + }, + }, + ) + + response: Final = XAIChatConfig().transform_response( + model="grok-4", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert getattr(response.usage, "server_side_tool_usage_details", None) == tool_details + assert ( + getattr(response.usage.prompt_tokens_details, "web_search_requests", None) == expected_web_search_requests + ) + assert response.usage.total_tokens == 120 + class TestXAIReportedCost: """xAI reports what it charged; the transformation moves it to where litellm bills from. diff --git a/tests/unit/proxy/auth/test_handle_jwt.py b/tests/unit/proxy/auth/test_handle_jwt.py index 640b3d8053d..39c2466033d 100644 --- a/tests/unit/proxy/auth/test_handle_jwt.py +++ b/tests/unit/proxy/auth/test_handle_jwt.py @@ -13,6 +13,7 @@ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( DEFAULT_JWKS_STALE_TTL, JWTLiteLLMRoleMap, + LiteLLMRoutes, LiteLLM_JWTAuth, LiteLLM_ModelTable, LiteLLM_TeamMembership, @@ -38,6 +39,7 @@ from litellm.proxy.auth.handle_jwt import ( NoMatchingJWTPublicKeyError, ) from litellm.proxy.auth.model_access_denied import ModelAccessDeniedHTTPException +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.types.agents import AgentResponse @@ -8225,3 +8227,28 @@ async def test_delegated_jwt_uses_granting_team_policy_before_route_authorizatio assert result["managed_agent_context"] == context if team_claim != "granting-team": assert any(call.kwargs.get("check_db_only") is True for call in load_team.call_args_list) + + +@pytest.mark.asyncio +async def test_check_admin_access_names_the_route_and_the_expanded_allow_list_when_denied(): + handler: Final = JWTHandler() + handler.update_environment( + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + litellm_jwtauth=LiteLLM_JWTAuth(admin_allowed_routes=["info_routes", "/custom/admin/route"]), + ) + + with pytest.raises(Exception, match="Admin not allowed to access this route") as denied: + await JWTAuthManager.check_admin_access( + jwt_handler=handler, + scopes=["litellm_proxy_admin"], + route="/key/generate", + user_id="admin-user", + org_id=None, + api_key="jwt-token", + ) + + assert str(denied.value) == ( + "Admin not allowed to access this route. Route=/key/generate, " + f"Allowed Routes={[*LiteLLMRoutes.info_routes.value, '/custom/admin/route']}" + ) diff --git a/tests/unit/proxy/auth/test_litellm_license.py b/tests/unit/proxy/auth/test_litellm_license.py index 83e26968f97..4b84dce3600 100644 --- a/tests/unit/proxy/auth/test_litellm_license.py +++ b/tests/unit/proxy/auth/test_litellm_license.py @@ -1,10 +1,14 @@ import asyncio import json +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx +import pytest from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey +from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.proxy.auth.litellm_license import LicenseCheck @@ -126,3 +130,30 @@ def test_valid_signed_wildcard_license_lifts_the_limit() -> None: assert license_check.verify_license_without_api_request(public_key=named_public_key, license_key=named_key) is True assert license_check.grants_feature("auto_router") is False assert license_check.auto_router_capability_limit() == 1 + + +@pytest.mark.parametrize( + ("reply", "premium"), + [ + ({"verify": True}, True), + ({"verify": False}, False), + (["verify", True], False), + ({"verify": "true"}, False), + ({"verified": True}, False), + ], +) +def test_is_premium_follows_the_license_server_reply_for_an_unsigned_license( + monkeypatch: pytest.MonkeyPatch, reply: object, premium: bool +) -> None: + monkeypatch.setenv("LITELLM_LICENSE", "license-the-public-key-did-not-sign") + requested: Final[list[str]] = [] + + def license_server(request: httpx.Request) -> httpx.Response: + requested.append(str(request.url)) + return httpx.Response(200, json=reply) + + license_check: Final = LicenseCheck() + license_check.http_handler = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(license_server))) + + assert license_check.is_premium() is premium + assert set(requested) == {"https://license.litellm.ai/verify_license/license-the-public-key-did-not-sign"} diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/test_aporia_ai.py b/tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/test_aporia_ai.py new file mode 100644 index 00000000000..f1a942d5825 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/aporia_ai/test_aporia_ai.py @@ -0,0 +1,87 @@ +import json + +import httpx +import pytest +import respx +from fastapi import HTTPException + +from litellm.proxy.guardrails.guardrail_hooks.aporia_ai.aporia_ai import AporiaGuardrail + +pytestmark = pytest.mark.usefixtures("httpx_transport") + +_API_BASE = "https://aporia.example.test/project-1" +_USER_MESSAGE = {"role": "user", "content": "what is my balance"} + + +def _guardrail() -> AporiaGuardrail: + return AporiaGuardrail( + api_key="aporia-key", + api_base=_API_BASE, + guardrail_name="aporia-guard", + event_hook="during_call", + default_on=True, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("new_messages", "response_string", "expected"), + [ + pytest.param( + [_USER_MESSAGE], + "your balance is 10", + {"messages": [_USER_MESSAGE], "response": "your balance is 10", "validation_target": "both"}, + id="prompt-and-response", + ), + pytest.param( + [_USER_MESSAGE], + None, + {"messages": [_USER_MESSAGE], "validation_target": "prompt"}, + id="prompt-only", + ), + pytest.param( + [], + "your balance is 10", + {"messages": [], "response": "your balance is 10", "validation_target": "response"}, + id="response-only", + ), + pytest.param([], None, {"messages": []}, id="nothing-to-validate"), + ], +) +async def test_prepare_aporia_request_sets_validation_target_from_supplied_content( + new_messages: list[dict], response_string: str | None, expected: dict[str, object] +) -> None: + request = await _guardrail().prepare_aporia_request(new_messages=new_messages, response_string=response_string) + + assert request == expected + + +@pytest.mark.asyncio +async def test_make_aporia_api_request_posts_prepared_body_to_validate_endpoint(respx_mock: respx.MockRouter) -> None: + route = respx_mock.post(f"{_API_BASE}/validate").respond(json={"action": "passthrough"}) + + await _guardrail().make_aporia_api_request( + request_data={}, new_messages=[_USER_MESSAGE], response_string="your balance is 10" + ) + + sent = route.calls.last.request + assert json.loads(sent.content) == { + "messages": [_USER_MESSAGE], + "response": "your balance is 10", + "validation_target": "both", + } + assert sent.headers["X-APORIA-API-KEY"] == "aporia-key" + + +@pytest.mark.asyncio +async def test_make_aporia_api_request_raises_400_when_aporia_blocks(respx_mock: respx.MockRouter) -> None: + verdict = {"action": "block", "revised_response": "blocked by policy"} + respx_mock.post(f"{_API_BASE}/validate").mock(return_value=httpx.Response(200, json=verdict)) + + with pytest.raises(HTTPException) as blocked: + await _guardrail().make_aporia_api_request( + request_data={}, new_messages=[_USER_MESSAGE], response_string="your balance is 10" + ) + + assert blocked.value.status_code == 400 + assert blocked.value.detail == {"error": "Violated guardrail policy", "aporia_ai_response": verdict} diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/unit/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index d0ad068aeb4..f7fb7e1fed0 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -12,6 +12,7 @@ import pytest from fastapi import HTTPException +from litellm import Router from litellm.constants import ( CONTENT_FILTER_STREAMING_HOLDBACK_CHARS, CONTENT_FILTER_STREAMING_SCAN_CONTEXT_CHARS, @@ -3468,3 +3469,57 @@ class TestContentFilterToolCallArguments: request_data={}, input_type="response", ) + + +def _guardrail_describing_images_as(description: str) -> ContentFilterGuardrail: + router: Final = Router( + model_list=[ + { + "model_name": "vision", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-test", + "mock_response": description, + }, + } + ] + ) + return ContentFilterGuardrail( + guardrail_name="image-filter", + blocked_words=[BlockedWord(keyword="project-titan", action=ContentFilterAction.BLOCK)], + llm_router=router, + image_model="vision", + ) + + +@pytest.mark.asyncio +async def test_image_whose_description_contains_blocked_keyword_is_rejected(): + description: Final = "A badge showing project-titan on a desk" + guardrail: Final = _guardrail_describing_images_as(description) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"images": ["https://example.com/badge.png"]}, + request_data={}, + input_type="request", + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == { + "error": f"Content blocked: keyword 'project-titan' detected (Image description): {description}", + "keyword": "project-titan", + "description": None, + } + + +@pytest.mark.asyncio +async def test_image_whose_description_is_clean_passes_through_unchanged(): + guardrail: Final = _guardrail_describing_images_as("A plain wooden desk") + + result: Final = await guardrail.apply_guardrail( + inputs={"images": ["https://example.com/desk.png"]}, + request_data={}, + input_type="request", + ) + + assert result == {"images": ["https://example.com/desk.png"], "texts": []} diff --git a/tests/unit/proxy/guardrails/test_guardrail_initializers.py b/tests/unit/proxy/guardrails/test_guardrail_initializers.py new file mode 100644 index 00000000000..52af817626b --- /dev/null +++ b/tests/unit/proxy/guardrails/test_guardrail_initializers.py @@ -0,0 +1,84 @@ +import pytest +from fastapi import HTTPException + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ToolPermissionGuardrail +from litellm.proxy.guardrails.guardrail_initializers import initialize_tool_permission +from litellm.types.guardrails import LitellmParams + +_RULES = [ + {"id": "allow_bash", "tool_name": "^Bash$", "decision": "allow"}, + {"id": "deny_read", "tool_name": "^Read$", "decision": "deny"}, +] + + +def _tool(name: str) -> dict[str, object]: + return {"type": "function", "function": {"name": name, "parameters": {"type": "object", "properties": {}}}} + + +def _initialize(**tool_permission_params: object) -> ToolPermissionGuardrail: + litellm_params = LitellmParams( + guardrail="tool_permission", mode="pre_call", default_on=True, **tool_permission_params + ) + return initialize_tool_permission(litellm_params, {"guardrail_name": "tool-guard"}) + + +async def _pre_call(guardrail: ToolPermissionGuardrail, tool_names: list[str]) -> dict: + return await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=DualCache(), + data={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "list the files"}], + "tools": [_tool(name) for name in tool_names], + }, + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_initialize_tool_permission_rewrite_mode_strips_only_the_tool_a_configured_rule_denies() -> None: + guardrail = _initialize(rules=_RULES, default_action="allow", on_disallowed_action="rewrite") + + data = await _pre_call(guardrail, ["Bash", "Read", "Grep"]) + + assert data["tools"] == [_tool("Bash"), _tool("Grep")] + assert litellm.callbacks == [guardrail] + + +@pytest.mark.asyncio +async def test_initialize_tool_permission_block_mode_rejects_request_naming_the_matching_rule() -> None: + guardrail = _initialize(rules=_RULES, default_action="allow") + + with pytest.raises(HTTPException) as blocked: + await _pre_call(guardrail, ["Bash", "Read"]) + + assert blocked.value.status_code == 400 + assert blocked.value.detail == { + "error": "Violated guardrail policy", + "detection_message": "Tool 'Read' denied by rule 'deny_read'", + } + + +@pytest.mark.asyncio +async def test_initialize_tool_permission_without_rules_denies_every_tool_by_default_action() -> None: + guardrail = _initialize() + + with pytest.raises(HTTPException) as blocked: + await _pre_call(guardrail, ["Bash"]) + + assert blocked.value.detail == { + "error": "Violated guardrail policy", + "detection_message": "Tool 'Bash' denied by default action", + } + + +@pytest.mark.asyncio +async def test_initialize_tool_permission_without_rules_and_allow_default_keeps_every_tool() -> None: + guardrail = _initialize(default_action="allow") + + data = await _pre_call(guardrail, ["Bash", "Read"]) + + assert data["tools"] == [_tool("Bash"), _tool("Read")] diff --git a/tests/unit/proxy/guardrails/test_init_guardrails.py b/tests/unit/proxy/guardrails/test_init_guardrails.py index fcd7e537937..2cac5d60068 100644 --- a/tests/unit/proxy/guardrails/test_init_guardrails.py +++ b/tests/unit/proxy/guardrails/test_init_guardrails.py @@ -229,8 +229,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): initialized = [ callback for callback in litellm.callbacks - if isinstance(callback, _OPTIONAL_PresidioPIIMasking) - and callback.guardrail_name == "test_presidio_chunk_size" + if isinstance(callback, _OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" ] assert initialized, "presidio guardrail was not registered as a callback" assert initialized[-1].presidio_analyze_chunk_size_bytes == 250_000 @@ -526,3 +525,45 @@ async def test_presidio_initialized_output_dispatch( ) assert response.choices[0].message.content == expected assert len(selected) == expected_calls + + +def test_init_guardrails_v2_publishes_initialized_guardrail_to_the_proxy_router( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm + from litellm.proxy import proxy_server + from litellm.proxy.guardrails import guardrail_registry + from litellm.proxy.guardrails.guardrail_hooks.aporia_ai import AporiaGuardrail + + router: Final = litellm.Router(model_list=[]) + monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", InMemoryGuardrailHandler()) + monkeypatch.setattr(proxy_server, "llm_router", router) + + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "aporia-guard", + "guardrail_id": "aporia-1", + "litellm_params": { + "guardrail": "aporia", + "mode": "during_call", + "default_on": True, + "api_key": "aporia-key", + "api_base": "https://aporia.example.test/project-1", + }, + } + ] + ) + + (registered,) = (callback for callback in litellm.callbacks if isinstance(callback, AporiaGuardrail)) + assert router.get_available_guardrail("aporia-guard") == { + "guardrail_name": "aporia-guard", + "litellm_params": { + "guardrail": "aporia", + "mode": "during_call", + "api_key": "aporia-key", + "api_base": "https://aporia.example.test/project-1", + }, + "callback": registered, + "id": "aporia-1", + } diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py index b89ae530d6f..52c664a65a5 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py @@ -1,5 +1,6 @@ import json from datetime import datetime +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -446,3 +447,35 @@ class TestGeminiPassthroughLoggingHandler: assert result["kwargs"]["response_cost"] == pytest.approx(expected_cost) assert result["kwargs"]["custom_llm_provider"] == "gemini" assert mock_logging_obj.model_call_details["custom_llm_provider"] == "gemini" + + +def test_build_complete_streaming_response_assembles_gemini_sse_chunks(): + logging_obj: Final = LiteLLMLoggingObj( + model="gemini-2.5-flash", + messages=[{"role": "user", "content": "Hello"}], + stream=True, + call_type="pass_through_endpoint", + start_time=datetime(2026, 1, 1), + litellm_call_id="gemini-stream-call-id", + function_id="gemini-stream-function-id", + ) + logging_obj.update_environment_variables(litellm_params={}, optional_params={}, model="gemini-2.5-flash") + first_chunk: Final = {"candidates": [{"content": {"parts": [{"text": "Hello"}], "role": "model"}, "index": 0}]} + last_chunk: Final = { + "candidates": [ + {"content": {"parts": [{"text": " there!"}], "role": "model"}, "finishReason": "STOP", "index": 0} + ], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 8, "totalTokenCount": 18}, + } + + response: Final = GeminiPassthroughLoggingHandler._build_complete_streaming_response( + all_chunks=[f"data: {json.dumps(first_chunk)}", f"data: {json.dumps(last_chunk)}"], + litellm_logging_obj=logging_obj, + model="gemini-2.5-flash", + url_route="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:streamGenerateContent", + ) + + assert isinstance(response, litellm.ModelResponse) + assert response.choices[0].message.content == "Hello there!" + assert response.choices[0].finish_reason == "stop" + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (10, 8, 18) diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py index f9df5a120f5..cc1246f4ede 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py @@ -1,4 +1,6 @@ +import json from datetime import datetime +from typing import Final import httpx import pytest @@ -93,3 +95,35 @@ def test_interactions_usage_object_is_read_into_prompt_and_completion_tokens(): assert response.usage.completion_tokens == 29 assert response.usage.completion_tokens_details.text_tokens == 9 assert result["kwargs"]["custom_llm_provider"] == "vertex_ai" + + +def test_build_complete_streaming_response_assembles_stream_generate_content_sse_chunks(): + logging_obj: Final = Logging( + model="gemini-2.5-flash", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="pass_through_endpoint", + start_time=datetime(2026, 1, 1), + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.update_environment_variables(litellm_params={}, optional_params={}, model="gemini-2.5-flash") + first_chunk: Final = {"candidates": [{"content": {"parts": [{"text": "Hello"}], "role": "model"}, "index": 0}]} + last_chunk: Final = { + "candidates": [ + {"content": {"parts": [{"text": " there!"}], "role": "model"}, "finishReason": "STOP", "index": 0} + ], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 8, "totalTokenCount": 18}, + } + + response: Final = VertexPassthroughLoggingHandler._build_complete_streaming_response( + all_chunks=[f"data: {json.dumps(first_chunk)}", f"data: {json.dumps(last_chunk)}"], + litellm_logging_obj=logging_obj, + model="gemini-2.5-flash", + url_route="/v1/projects/p/locations/us-central1/publishers/google/models/gemini-2.5-flash:streamGenerateContent", + ) + + assert isinstance(response, ModelResponse) + assert response.choices[0].message.content == "Hello there!" + assert response.choices[0].finish_reason == "stop" + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (10, 8, 18) diff --git a/tests/unit/proxy/pass_through_endpoints/test_success_handler.py b/tests/unit/proxy/pass_through_endpoints/test_success_handler.py new file mode 100644 index 00000000000..67b9fbc12db --- /dev/null +++ b/tests/unit/proxy/pass_through_endpoints/test_success_handler.py @@ -0,0 +1,78 @@ +from datetime import datetime, timezone +from typing import Final + +import httpx +import pytest + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging + +pytestmark: Final = pytest.mark.usefixtures("local_model_cost_map") + +_LIVE_ROUTE: Final = "/vertex_ai/live" +_LIVE_MODEL: Final = "gemini-live-2.5-flash" +_START: Final = datetime(2025, 1, 2, 3, 4, 5, tzinfo=timezone.utc) +_END: Final = datetime(2025, 1, 2, 3, 4, 9, tzinfo=timezone.utc) +_FIRST_TURN_USAGE: Final = {"promptTokenCount": 10, "candidatesTokenCount": 4, "totalTokenCount": 14} +_SECOND_TURN_USAGE: Final = {"promptTokenCount": 30, "candidatesTokenCount": 6, "totalTokenCount": 36} + + +def _logging_obj() -> LiteLLMLoggingObj: + return LiteLLMLoggingObj( + model="unknown", + messages=[{"role": "user", "content": "WebSocket connection"}], + stream=True, + call_type="pass_through_endpoint", + start_time=_START, + litellm_call_id="call-live", + function_id="websocket_passthrough", + ) + + +def _normalize_live_session(response_body: dict | list[dict[str, object]] | None) -> dict: + return PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=httpx.Response(200, request=httpx.Request("GET", f"https://proxy.example.test{_LIVE_ROUTE}")), + response_body=response_body, + request_body={}, + logging_obj=_logging_obj(), + url_route=_LIVE_ROUTE, + result="websocket_connection_successful", + start_time=_START, + end_time=_END, + cache_hit=False, + model=_LIVE_MODEL, + ) + + +def test_vertex_ai_live_route_sums_usage_of_every_turn_in_the_websocket_frames() -> None: + normalized = _normalize_live_session( + [ + {"setupComplete": {}}, + {"serverContent": {"modelTurn": {"parts": [{"text": "hello"}]}}}, + {"usageMetadata": _FIRST_TURN_USAGE}, + {"serverContent": {"modelTurn": {"parts": [{"text": "goodbye"}]}}}, + {"usageMetadata": _SECOND_TURN_USAGE}, + ] + ) + + response = normalized["standard_logging_response_object"] + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (40, 10, 50) + assert response.model == _LIVE_MODEL + assert normalized["kwargs"] == {"model": _LIVE_MODEL, "custom_llm_provider": "vertex_ai"} + + +@pytest.mark.parametrize( + "response_body", + [ + pytest.param({"usageMetadata": _FIRST_TURN_USAGE}, id="single-json-object"), + pytest.param(None, id="no-body"), + pytest.param([{"setupComplete": {}}], id="frames-without-usage"), + ], +) +def test_vertex_ai_live_route_without_usage_frames_yields_no_logging_response( + response_body: dict | list[dict[str, object]] | None, +) -> None: + normalized = _normalize_live_session(response_body) + + assert normalized["standard_logging_response_object"] is None + assert normalized["kwargs"] == {"model": _LIVE_MODEL}