diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index a29f0a1b43a..c648f7faea0 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -46,6 +46,8 @@ from litellm.proxy._types import ( UserAPIKeyAuth, WebhookEvent, ) +from litellm.repositories.table_repositories import InvitationLinkRepository +from litellm.repositories.user_repository import UserRepository from litellm.secret_managers.main import get_secret_bool from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL @@ -844,7 +846,7 @@ class BaseEmailLogger(CustomLogger): ) return None - user_row = await prisma_client.db.litellm_usertable.find_unique( + user_row = await UserRepository(prisma_client).table.find_unique( where={"user_id": user_id} ) @@ -929,7 +931,7 @@ class BaseEmailLogger(CustomLogger): try: # Try to get existing invitation existing_invitations = ( - await prisma_client.db.litellm_invitationlink.find_many( + await InvitationLinkRepository(prisma_client).table.find_many( where={"user_id": user_id}, order={"created_at": "desc"}, ) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 35fc3510414..985a9b6de38 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -15,6 +15,10 @@ from litellm.constants import ( MANAGED_OBJECT_STALENESS_CUTOFF_DAYS, MAX_OBJECTS_PER_POLL_CYCLE, ) +from litellm.repositories.table_repositories import ManagedObjectRepository +from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository if TYPE_CHECKING: from prisma import models as prisma_models @@ -58,25 +62,19 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable - return table + return ManagedObjectRepository(prisma_client).table def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]": - table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable - return table + return UserRepository(prisma_client).table def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]": - table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = ( - prisma_client.db.litellm_verificationtoken - ) - return table + return VerificationTokenRepository(prisma_client).table def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]": - table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable - return table + return TeamRepository(prisma_client).table class CheckBatchCost: diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index cdeea0d3d4b..738b9ebe1e0 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -16,6 +16,7 @@ from litellm.constants import ( MAX_OBJECTS_PER_POLL_CYCLE, STALE_OBJECT_CLEANUP_BATCH_SIZE, ) +from litellm.repositories.table_repositories import ManagedObjectRepository from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN @@ -43,8 +44,7 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable - return table + return ManagedObjectRepository(prisma_client).table class CheckResponsesCost: diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index e0a94612646..011bf95defd 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1647,7 +1647,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): owner_filter: Final = build_owner_filter(user_api_key_dict) if owner_filter is None: - return FileListPage(**build_list_page([])) + return FileListPage.model_validate(build_list_page([])) if after: cursor_row = await _managed_file_table(self.prisma_client).find_first( @@ -1686,7 +1686,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): cursor_id = chunk[-1].unified_file_id chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE) - return FileListPage(**build_list_page(matches[:page_size], has_more=len(matches) > page_size)) + return FileListPage.model_validate(build_list_page(matches[:page_size], has_more=len(matches) > page_size)) def _is_batch_polling_enabled(self) -> bool: """ diff --git a/litellm/a2a_protocol/providers/base.py b/litellm/a2a_protocol/providers/base.py index 5a5eff8cf35..6196cc789fc 100644 --- a/litellm/a2a_protocol/providers/base.py +++ b/litellm/a2a_protocol/providers/base.py @@ -4,7 +4,6 @@ Base configuration for A2A protocol providers. from abc import ABC, abstractmethod from collections.abc import AsyncIterator -from typing import Any class BaseA2AProviderConfig(ABC): @@ -19,10 +18,10 @@ class BaseA2AProviderConfig(ABC): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Handle non-streaming A2A request. @@ -40,10 +39,10 @@ class BaseA2AProviderConfig(ABC): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming A2A request. diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py index 2b37c0c4906..22dc3aa508a 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py @@ -3,7 +3,7 @@ Bedrock AgentCore A2A provider configuration. """ from collections.abc import AsyncIterator -from typing import Any, Final +from typing import Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.bedrock_agentcore.handler import ( @@ -23,10 +23,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: """Handle non-streaming request to AgentCore A2A agent.""" litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: @@ -43,10 +43,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """Handle streaming request to AgentCore A2A agent.""" litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py index da5eb522187..aecc887fbfe 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py @@ -7,7 +7,7 @@ completion bridge that would otherwise strip the envelope. import json from collections.abc import AsyncIterator, Mapping -from typing import Any, Final +from typing import Final from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( @@ -30,9 +30,9 @@ class BedrockAgentCoreA2AHandler: async def handle_non_streaming( request_id: str, params: Mapping[str, object], - litellm_params: dict[str, Any], + litellm_params: dict[str, object], agent_extra_headers: dict[str, str] | None = None, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Handle non-streaming A2A request to AgentCore. @@ -77,9 +77,9 @@ class BedrockAgentCoreA2AHandler: async def handle_streaming( request_id: str, params: Mapping[str, object], - litellm_params: dict[str, Any], + litellm_params: dict[str, object], agent_extra_headers: dict[str, str] | None = None, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming A2A request to AgentCore. diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py index c486f1f6d95..5d4e0d5fb39 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py @@ -219,7 +219,7 @@ class BedrockAgentCoreA2ATransformation: return url, signed_headers, signed_body @staticmethod - async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, Any]]: + async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, object]]: """ Parse SSE events from an httpx streaming response. diff --git a/litellm/a2a_protocol/providers/langflow/config.py b/litellm/a2a_protocol/providers/langflow/config.py index 54d403f88c0..b8819dbe4bc 100644 --- a/litellm/a2a_protocol/providers/langflow/config.py +++ b/litellm/a2a_protocol/providers/langflow/config.py @@ -1,5 +1,4 @@ from collections.abc import AsyncIterator -from typing import Any from litellm.a2a_protocol.litellm_completion_bridge.handler import ( A2A_USER_API_KEY_HASH_PARAM, @@ -16,10 +15,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( @@ -39,10 +38,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py index 20404e3702b..f4c758af2d0 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py @@ -2,8 +2,7 @@ Pydantic AI provider configuration. """ -from collections.abc import AsyncIterator -from typing import Any +from collections.abc import AsyncIterator, Mapping from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.pydantic_ai_agents.handler import PydanticAIHandler @@ -20,9 +19,12 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + timeout: float = 60.0, + agent_extra_headers: Mapping[str, str] | None = None, + **kwargs: object, ) -> dict[str, object]: """Handle non-streaming request to Pydantic AI agent.""" if api_base is None: @@ -31,14 +33,14 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): request_id=request_id, params=params, api_base=api_base, - timeout=kwargs.get("timeout", 60.0), - agent_extra_headers=kwargs.get("agent_extra_headers"), + timeout=timeout, + agent_extra_headers=agent_extra_headers, ) async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, ) -> AsyncIterator[dict[str, object]]: diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py index c083c0267f7..f3ceb899a92 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py @@ -5,8 +5,8 @@ Pydantic AI agents follow A2A protocol but don't support streaming natively. This handler provides fake streaming by converting non-streaming responses into streaming chunks. """ -from collections.abc import AsyncIterator -from typing import Any, Final +from collections.abc import AsyncIterator, Mapping +from typing import Final from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( @@ -26,11 +26,11 @@ class PydanticAIHandler: @staticmethod async def handle_non_streaming( request_id: str, - params: dict[str, Any], + params: Mapping[str, object], api_base: str | None = None, timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, - ) -> dict[str, Any]: + agent_extra_headers: Mapping[str, str] | None = None, + ) -> dict[str, object]: """ Handle non-streaming request to Pydantic AI agent. @@ -63,13 +63,13 @@ class PydanticAIHandler: @staticmethod async def handle_streaming( request_id: str, - params: dict[str, Any], + params: Mapping[str, object], api_base: str | None = None, timeout: float = 60.0, chunk_size: int = 50, delay_ms: int = 10, - agent_extra_headers: dict[str, str] | None = None, - ) -> AsyncIterator[dict[str, Any]]: + agent_extra_headers: Mapping[str, str] | None = None, + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming request to Pydantic AI agent with fake streaming. diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py index 024e8c179c2..35fe6e76f4d 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py @@ -7,7 +7,7 @@ This module provides fake streaming by converting non-streaming responses into s import asyncio from collections.abc import AsyncIterator, Mapping, Sequence -from typing import Any, Final, Protocol, cast, runtime_checkable +from typing import Final, Protocol, runtime_checkable from uuid import uuid4 from pydantic import TypeAdapter @@ -100,7 +100,7 @@ class PydanticAITransformation: request_id: str, max_attempts: int = 30, poll_interval: float = 0.5, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Poll for task completion using tasks/get method. @@ -156,7 +156,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a request to Pydantic AI agent and return the raw task response. @@ -200,7 +200,7 @@ class PydanticAITransformation: # Send request to Pydantic AI agent using shared async HTTP client client: Final = get_async_httpx_client( - llm_provider=cast(Any, "pydantic_ai_agent"), + llm_provider="pydantic_ai_agent", params={"timeout": timeout}, ) response: Final = await client.post( @@ -242,7 +242,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a non-streaming A2A request to Pydantic AI agent and wait for completion. @@ -278,7 +278,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a request to Pydantic AI agent and return the raw task response. diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index 44873edf271..a1a6c3a5957 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -3,11 +3,11 @@ A2A provider configuration for IBM watsonx Orchestrate (WXO). """ from collections.abc import AsyncIterator -from typing import Any, Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import ( WatsonxOrchestrateHandler, + WXOLitellmParams, ) @@ -17,12 +17,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + litellm_params: WXOLitellmParams | None = None, + **kwargs: object, ) -> dict[str, object]: """Handle a non-streaming A2A request via WXO runs API.""" - litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " @@ -37,12 +38,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + litellm_params: WXOLitellmParams | None = None, + **kwargs: object, ) -> AsyncIterator[dict[str, object]]: """Handle a streaming A2A request via WXO streaming runs API.""" - litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index c66b07c321c..0a2871d25cc 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -7,7 +7,7 @@ import hashlib import json import time from collections.abc import AsyncIterator -from typing import Any, Final, NamedTuple, Protocol +from typing import Final, NamedTuple, Protocol import httpx from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -228,7 +228,7 @@ class WatsonxOrchestrateHandler: return run_data @staticmethod - async def _accumulate_wxo_sse_text(response: Any) -> str: + async def _accumulate_wxo_sse_text(response: _SSELineSource) -> str: source: Final[_WXOView] = {"sse_source": response} accumulated_text = "" async for line in source["sse_source"].aiter_lines(): diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index cec25634bb8..d82e58ecc91 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -12,7 +12,7 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and from __future__ import annotations -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final import litellm @@ -39,10 +39,10 @@ def resolve_embedding_router( def build_router_embedding_metadata( - request_metadata: dict[str, Any] | None, -) -> dict[str, Any]: + request_metadata: Mapping[str, object] | None, +) -> Mapping[str, object]: """Forward the caller's full metadata, flagged as a semantic-cache embedding.""" - metadata: Final[dict[str, Any]] = dict(request_metadata or {}) + metadata: Final = dict(request_metadata or {}) metadata["semantic-cache-embedding"] = True return metadata diff --git a/litellm/files/main.py b/litellm/files/main.py index 723784795b0..9057211256d 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -331,7 +331,7 @@ async def afile_retrieve( else: response = init_response - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) except Exception as e: raise e diff --git a/litellm/files/streaming.py b/litellm/files/streaming.py index 5d23ebf32ae..460050872f4 100644 --- a/litellm/files/streaming.py +++ b/litellm/files/streaming.py @@ -34,7 +34,7 @@ class FileContentStreamingResponse: self.custom_llm_provider = custom_llm_provider self.logging_obj = logging_obj self.standard_logging_object: StandardLoggingPayload | None = None - self._hidden_params: dict[str, Any] = {} + self._hidden_params: dict[str, object] = {} self._logging_completed = False self._close_completed = False self._start_time = ( @@ -121,7 +121,7 @@ class FileContentStreamingResponse: return response def _sync_hidden_params(self) -> None: - litellm_params: dict[str, Any] = {} + litellm_params: dict[str, object] = {} if self.logging_obj is not None: litellm_params = self.logging_obj.model_call_details.get("litellm_params", {}) or {} diff --git a/litellm/harness/context.py b/litellm/harness/context.py index eaae9aafe1f..f75ee7e49f1 100644 --- a/litellm/harness/context.py +++ b/litellm/harness/context.py @@ -4,7 +4,7 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, TypeAlias from pydantic import BaseModel @@ -42,7 +42,7 @@ class SessionContext: api_base: str | None = None endpoint: ModelEndpoint | None = None instructions: str | None = None - tools: Sequence[Callable[..., Any]] = () + tools: Sequence[Callable[..., object]] = () skills: Sequence[str] = () disable_tools: Sequence[str] = () permissions: PermissionMode = "full" @@ -50,7 +50,7 @@ class SessionContext: output: type[BaseModel] | None = None max_turns: int | None = None timeout: float | None = None - metadata: Mapping[str, Any] = field(default_factory=dict) + metadata: Mapping[str, object] = field(default_factory=dict) options: HarnessOptions | None = None # Set by the handler after each turn. final_text: str = "" diff --git a/litellm/harness/handlers/cli_handler.py b/litellm/harness/handlers/cli_handler.py index 0fda16c34fe..d217d1e5bbb 100644 --- a/litellm/harness/handlers/cli_handler.py +++ b/litellm/harness/handlers/cli_handler.py @@ -13,7 +13,7 @@ import asyncio import os from collections import deque from collections.abc import AsyncIterator, Sequence -from typing import Any, Final +from typing import Final from litellm._logging import verbose_logger from litellm.constants import HARNESS_STDERR_TAIL_LINES, HARNESS_STREAM_READ_CHUNK_BYTES @@ -125,7 +125,7 @@ class CLIHarnessHandler(BaseHarnessHandler): self._proc = proc tail: Final[deque[str]] = deque(maxlen=HARNESS_STDERR_TAIL_LINES) # mutable-ok: bounded stderr ring buffer stderr_task = asyncio.ensure_future(drain_stderr(proc.stderr, tail)) - state: Any = self.config.create_stream_state() + state: Final[object] = self.config.create_stream_state() exit_code: int | None = None try: await send_stdin(proc, request.stdin) diff --git a/litellm/harness/options.py b/litellm/harness/options.py index 18865359ba8..5d7c8ae0807 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -9,7 +9,7 @@ from typing import Any, Literal @dataclass(frozen=True) class ClaudeCodeOptions: - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @@ -17,20 +17,20 @@ class ClaudeCodeOptions: class CodexOptions: reasoning_effort: Literal["low", "medium", "high", "xhigh"] | None = None web_search: bool = False - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @dataclass(frozen=True) class OpenCodeOptions: agent: str = "build" - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @dataclass(frozen=True) class DeepAgentsOptions: - subagents: Sequence[Any] = () + subagents: Sequence[object] = () recursion_limit: int | None = None diff --git a/litellm/harness/runtime.py b/litellm/harness/runtime.py index da90b28a751..dc24a1f0ef7 100644 --- a/litellm/harness/runtime.py +++ b/litellm/harness/runtime.py @@ -1005,24 +1005,24 @@ def aagent( `await litellm.aagent(...)` returns a Result. With stream=True it returns an async iterator of events instead: `async for event in litellm.aagent(..., stream=True)`. """ - kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to arun_agent/astream_agent - "sandbox": sandbox, - "model": model, - "api_key": api_key, - "api_base": api_base, - "instructions": instructions, - "tools": tools, - "skills": skills, - "disable_tools": disable_tools, - "permissions": permissions, - "on_approval": on_approval, - "output": output, - "max_turns": max_turns, - "timeout": timeout, - "metadata": metadata, - "options": options, - "install": install, - } - if stream: - return astream_agent(harness, prompt, **kwargs) - return arun_agent(harness, prompt, **kwargs) + call: Final = astream_agent if stream else arun_agent + return call( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 9787800833b..18336a3d06e 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -8,7 +8,7 @@ import threading from collections.abc import AsyncIterator, Callable, Coroutine, Mapping, Sequence from concurrent.futures import Future from typing import ( - Any, + Final, TypeVar, ) @@ -417,24 +417,24 @@ def agent( Prefix the model with `litellm_proxy/` to route every model call through your LiteLLM AI Gateway. """ - kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to _run/_stream - "sandbox": sandbox, - "model": model, - "api_key": api_key, - "api_base": api_base, - "instructions": instructions, - "tools": tools, - "skills": skills, - "disable_tools": disable_tools, - "permissions": permissions, - "on_approval": on_approval, - "output": output, - "max_turns": max_turns, - "timeout": timeout, - "metadata": metadata, - "options": options, - "install": install, - } - if stream: - return _stream(harness, prompt, **kwargs) - return _run(harness, prompt, **kwargs) + call: Final = _stream if stream else _run + return call( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) diff --git a/litellm/harness/types.py b/litellm/harness/types.py index 475568a9979..6b0118f4349 100644 --- a/litellm/harness/types.py +++ b/litellm/harness/types.py @@ -7,7 +7,7 @@ import json from collections.abc import Mapping from dataclasses import dataclass, field from enum import Enum -from typing import Any, Literal +from typing import Literal from pydantic import BaseModel @@ -71,7 +71,7 @@ class ToolCall: id: str name: str native_name: str - input: Mapping[str, Any] + input: Mapping[str, object] builtin: bool = True @@ -100,7 +100,7 @@ class Approval: """A request to run a tool. The turn waits until allow() or deny() is called.""" tool: str - input: Mapping[str, Any] + input: Mapping[str, object] _decision: asyncio.Future[tuple[bool, str]] = field( default_factory=lambda: asyncio.get_event_loop().create_future(), compare=False, diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 0271cf1e03c..3f89ac6978f 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -2,6 +2,7 @@ import json from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly, TypedDict, override from litellm._logging import verbose_logger @@ -29,6 +30,8 @@ from litellm.integrations._types.open_inference import ( ToolCallAttributes, ) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class ArizeOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @@ -609,9 +612,7 @@ def _coerce_response_obj_for_attrs(response_obj): text: Final = getattr(response_obj, "text", None) if isinstance(text, str) and text: try: - parsed: Final = json.loads(text) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(text)) except Exception: pass return response_obj @@ -1062,9 +1063,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): inner = candidate.get("response") if isinstance(inner, str): try: - parsed = json.loads(inner) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(inner)) except Exception: continue if isinstance(inner, dict): @@ -1078,9 +1077,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): return original if isinstance(original, str): try: - parsed = json.loads(original) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(original)) except Exception: return None return None diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 22bb50cd739..9565a64e69c 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -15,6 +15,7 @@ from types import MappingProxyType from typing import Any, Final, Literal import httpx +from pydantic import ConfigDict, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger @@ -57,6 +58,7 @@ from litellm.types.utils import ( _EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({}) _EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""} _MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024 +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset( {"agent", "assistant", "developer", "function", "model", "system", "tool", "user"} ) @@ -282,8 +284,10 @@ def _to_dd_arguments(raw_arguments: object) -> dict[str, object] | str: return raw_arguments if isinstance(raw_arguments, dict) else str(raw_arguments) if len(raw_arguments) > _MAX_PARSED_TOOL_ARGUMENT_CHARS: return raw_arguments - parsed: Final = safe_json_loads(raw_arguments) - return parsed if isinstance(parsed, dict) else raw_arguments + try: + return _JSON_OBJECT.validate_python(safe_json_loads(raw_arguments)) + except ValidationError: + return raw_arguments def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]: diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index 21906b0d996..9bc324256fb 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -4,12 +4,12 @@ import json import os import re import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime, timezone, tzinfo from typing import Any, Final, Protocol, cast import httpx -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from typing_extensions import ReadOnly, TypedDict import litellm @@ -35,6 +35,11 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai" # unavailable, invalid credentials) cannot leak memory unboundedly. GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000 +_UNTYPED_VALUE: Final = TypeAdapter(object) +_ZERO_ARGUMENT_CALLABLE: Final[TypeAdapter[Callable[[], object]]] = TypeAdapter( + Callable[[], object], config=ConfigDict(hide_input_in_errors=True) +) + class _GalileoLoginBody(TypedDict): """Decoded body of the Galileo login response.""" @@ -473,9 +478,9 @@ class GalileoObserve(CustomLogger): if isinstance(value, str): return value - def _json_default(obj: Any) -> object: + def _json_default(obj: object) -> object: if hasattr(obj, "model_dump"): - return obj.model_dump() + return _ZERO_ARGUMENT_CALLABLE.validate_python(getattr(obj, "model_dump", None))() return str(obj) return json.dumps(value, default=_json_default) @@ -496,7 +501,7 @@ class GalileoObserve(CustomLogger): if hasattr(message, "json"): message_json: Final[object] = message.json() if isinstance(message_json, str): - return json.loads(message_json) + return _UNTYPED_VALUE.validate_python(json.loads(message_json)) return message_json return message return None diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 8d588896b2f..956e754944a 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -8,6 +8,8 @@ from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypedDict, cast +from pydantic import ConfigDict, TypeAdapter + import litellm from litellm._logging import verbose_logger from litellm.integrations._types.open_inference import ( @@ -95,6 +97,8 @@ class _ResponseWithUsageView(TypedDict, total=False): usage: "_UsageCompletionTokensView | None" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + # Cap on credential-scoped providers held at once; each one owns an exporter thread. _MAX_DYNAMIC_TRACER_PROVIDERS: Final = 256 @@ -1633,7 +1637,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): attributes = self.config.attributes if attributes is None and self.callback_name in (None, "otel"): otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} - raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None + raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: attributes = _build_metric_attribute_filter(raw) ( @@ -2919,7 +2923,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): import json try: - _parsed: Final[Mapping[str, object]] = json.loads(_raw_response) + _parsed: Final = _JSON_OBJECT.validate_python(json.loads(_raw_response)) for param, val in _parsed.items(): self.safe_set_attribute( span=span, diff --git a/litellm/integrations/opik/opik_payload_builder/api.py b/litellm/integrations/opik/opik_payload_builder/api.py index 4334c81eb6e..d16bc71c416 100644 --- a/litellm/integrations/opik/opik_payload_builder/api.py +++ b/litellm/integrations/opik/opik_payload_builder/api.py @@ -1,12 +1,17 @@ """Public API for Opik payload building.""" +from collections.abc import Mapping from datetime import datetime from typing import Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.integrations.opik import utils from . import extractors, payload_builders, types +_STANDARD_LOGGING_FIELDS: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + def build_opik_payload( kwargs: dict[str, Any], @@ -36,12 +41,14 @@ def build_opik_payload( - First element is TracePayload if creating a new trace, None if attaching to existing - Second element is always SpanPayload """ - standard_logging_object: Final = kwargs["standard_logging_object"] + standard_logging_object: Final = _STANDARD_LOGGING_FIELDS.validate_python(kwargs["standard_logging_object"]) # Extract litellm params and metadata litellm_params: Final = kwargs.get("litellm_params", {}) or {} litellm_metadata: Final = litellm_params.get("metadata", {}) or {} - standard_logging_metadata: Final = standard_logging_object.get("metadata", {}) or {} + standard_logging_metadata: Final = _STANDARD_LOGGING_FIELDS.validate_python( + standard_logging_object.get("metadata", {}) or {} + ) # Extract and merge Opik metadata opik_metadata: Final = extractors.extract_opik_metadata(litellm_metadata, standard_logging_metadata) diff --git a/litellm/integrations/opik/opik_payload_builder/extractors.py b/litellm/integrations/opik/opik_payload_builder/extractors.py index 4dd3d40fae3..1ceb9aa763f 100644 --- a/litellm/integrations/opik/opik_payload_builder/extractors.py +++ b/litellm/integrations/opik/opik_payload_builder/extractors.py @@ -4,8 +4,12 @@ import json from collections.abc import Mapping from typing import Any, Final +from pydantic import TypeAdapter + from litellm import _logging +_DECODED_JSON: Final = TypeAdapter(object) + def normalize_provider_name(provider: str | None) -> str | None: """ @@ -149,7 +153,7 @@ def apply_proxy_header_overrides( thread_id = value elif param_key == "tags": try: - parsed_tags: object = json.loads(value) + parsed_tags = _DECODED_JSON.validate_python(json.loads(value)) if isinstance(parsed_tags, list): tags.extend(parsed_tags) except (json.JSONDecodeError, TypeError): diff --git a/litellm/integrations/opik/opik_payload_builder/types.py b/litellm/integrations/opik/opik_payload_builder/types.py index 546ce55f840..2f418b2ef9e 100644 --- a/litellm/integrations/opik/opik_payload_builder/types.py +++ b/litellm/integrations/opik/opik_payload_builder/types.py @@ -1,7 +1,8 @@ """Type definitions for Opik payload building.""" +from collections.abc import Mapping from dataclasses import dataclass -from typing import Any, Final, Literal +from typing import Final, Literal @dataclass @@ -13,9 +14,9 @@ class TracePayload: name: str start_time: str end_time: str - input: Any - output: Any - metadata: dict[str, Any] + input: object + output: object + metadata: Mapping[str, object] tags: list[str] thread_id: str | None = None @@ -32,9 +33,9 @@ class SpanPayload: model: str start_time: str end_time: str - input: Any - output: Any - metadata: dict[str, Any] + input: object + output: object + metadata: Mapping[str, object] tags: list[str] usage: dict[str, int] parent_span_id: str | None = None diff --git a/litellm/integrations/opik/utils.py b/litellm/integrations/opik/utils.py index 2caacd1f871..79b6199423e 100644 --- a/litellm/integrations/opik/utils.py +++ b/litellm/integrations/opik/utils.py @@ -2,7 +2,8 @@ import configparser import os import time import uuid -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final CONFIG_FILE_PATH_DEFAULT: Final[str] = "~/.opik.config" @@ -93,14 +94,14 @@ def create_usage_object(usage): return usage_dict -def _remove_nulls(x: dict[str, Any]) -> dict[str, Any]: +def _remove_nulls(x: Mapping[str, object]) -> dict[str, object]: """Remove None values from dict.""" return {k: v for k, v in x.items() if v is not None} def get_traces_and_spans_from_payload( - payload: list[dict[str, Any]], -) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + payload: Sequence[Mapping[str, object]], +) -> tuple[list[dict[str, object]], list[dict[str, object]]]: """ Separate traces and spans from payload. diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index ddb8e127408..6e9038c0cad 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -2,7 +2,7 @@ from enum import Enum from functools import lru_cache -from typing import Annotated, Any, Final +from typing import Annotated, Final from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator from pydantic.fields import FieldInfo @@ -315,7 +315,7 @@ class OpenTelemetryV2Config(BaseSettings): mode="before", ) @classmethod - def _split_csv(cls, value: Any) -> Any: + def _split_csv(cls, value: object) -> object: """Accept a comma-separated string for list fields. Env vars are strings, but these fields are lists. Pydantic-settings would diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index e1623f4697f..76b23467679 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -307,7 +307,7 @@ class GenAIMetricRecorder: return common_attrs - def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes: + def _bounded_attributes(self, kwargs: Mapping[str, object]) -> MetricAttributes: """The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`. The cap runs BEFORE the operator's include/exclude filter so the filter can @@ -322,7 +322,7 @@ class GenAIMetricRecorder: attributes = None if self._callback_name in (None, "otel"): otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} - raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None + raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: attributes = _build_metric_attribute_filter(raw) # A bad filter (include_list + exclude_list both set, an unfilterable name) diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index f504292cb64..08249d33561 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -595,7 +595,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): and not (self.s3_drop_on_terminal_error and _is_terminal(response)) and attempt < max_retries - 1 ): - wait_time = 2**attempt # 1s, 2s + wait_time = 1 << attempt # 1s, 2s verbose_logger.log( logging.DEBUG if _in_flush.get() else logging.WARNING, "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", @@ -897,7 +897,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): and not (self.s3_drop_on_terminal_error and _is_terminal(response)) and attempt < max_retries - 1 ): - wait_time = 2**attempt # 1s, 2s + wait_time = 1 << attempt # 1s, 2s verbose_logger.warning( "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", response.status_code, diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index 80530622e36..f40a166e67d 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -6,12 +6,12 @@ It searches the vector store for relevant context, runs the request's pre-call g over that context, and appends it to the messages. """ -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from itertools import chain from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args -from pydantic import TypeAdapter, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError from typing_extensions import assert_never import litellm @@ -44,10 +44,12 @@ else: LiteLLMLoggingObj = Any SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures" +_PROVIDER_FIELDS_ATTRIBUTE: Final = "provider_specific_fields" _DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate" _FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode) _OBJECT_ADAPTER: Final = TypeAdapter(object) _STR_KEYED_ADAPTER: Final = TypeAdapter(dict[str, object]) +_ITERABLE_ADAPTER: Final = TypeAdapter(Iterable[object], config=ConfigDict(hide_input_in_errors=True)) _GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA: Final = frozenset( {"guardrails", "guardrail_config", "policies", "include_guardrail_response"} ) @@ -475,7 +477,7 @@ class VectorStorePreCallHook(CustomLogger): async def async_post_call_streaming_deployment_hook( self, request_data: dict, - response_chunk: Any, + response_chunk: object, call_type: CallTypes | None, ) -> object | None: """ @@ -496,15 +498,17 @@ class VectorStorePreCallHook(CustomLogger): return response_chunk # Add search results to streaming chunk - if hasattr(response_chunk, "choices") and response_chunk.choices: - for choice in response_chunk.choices: - if hasattr(choice, "delta") and choice.delta: - provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {} + choices: Final[object] = getattr(response_chunk, "choices", None) + if choices: + for choice in _ITERABLE_ADAPTER.validate_python(choices): + delta: object = getattr(choice, "delta", None) + if delta: + provider_fields = getattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, None) or {} if search_results: provider_fields["search_results"] = search_results if search_failures: provider_fields[SEARCH_FAILURES_FIELD] = search_failures - choice.delta.provider_specific_fields = provider_fields + setattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, provider_fields) # Return modified chunk return response_chunk diff --git a/litellm/llms/a2a/chat/transformation.py b/litellm/llms/a2a/chat/transformation.py index af4c6f69944..6171b341ca4 100644 --- a/litellm/llms/a2a/chat/transformation.py +++ b/litellm/llms/a2a/chat/transformation.py @@ -3,8 +3,8 @@ A2A Protocol Transformation for LiteLLM """ import uuid -from collections.abc import Iterator, Mapping -from typing import TYPE_CHECKING, Any, Final +from collections.abc import AsyncIterator, Iterator, Mapping +from typing import TYPE_CHECKING, Final import httpx @@ -72,9 +72,9 @@ class A2AConfig(BaseConfig): agent_name: str, api_base: str | None, api_key: str | None, - headers: dict[str, Any] | None, - optional_params: dict[str, Any], - ) -> tuple[str | None, str | None, dict[str, Any] | None]: + headers: dict[str, object] | None, + optional_params: dict[str, object], + ) -> tuple[str | None, str | None, dict[str, object] | None]: """ Resolve agent configuration from the registry for a registered agent. @@ -376,7 +376,7 @@ class A2AConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Iterator | Any, + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, json_mode: bool | None = False, ) -> BaseModelResponseIterator: @@ -397,7 +397,7 @@ class A2AConfig(BaseConfig): json_mode=json_mode, ) - def _openai_message_to_a2a_message(self, message: dict[str, Any]) -> dict[str, Any]: + def _openai_message_to_a2a_message(self, message: Mapping[str, object]) -> dict[str, object]: """ Convert OpenAI message to A2A message format. diff --git a/litellm/llms/aiohttp_openai/chat/transformation.py b/litellm/llms/aiohttp_openai/chat/transformation.py index a06c670e3f1..dba7a7405ab 100644 --- a/litellm/llms/aiohttp_openai/chat/transformation.py +++ b/litellm/llms/aiohttp_openai/chat/transformation.py @@ -7,9 +7,11 @@ https://github.com/BerriAI/litellm/issues/6592 New config to ensure we introduce this without causing breaking changes for users """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final from aiohttp import ClientResponse +from pydantic import ConfigDict, TypeAdapter from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig from litellm.types.llms.openai import AllMessageValues @@ -23,6 +25,10 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECTS: Final = TypeAdapter( + Iterable[Mapping[str, object]], config=ConfigDict(strict=True, hide_input_in_errors=True) +) + class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): def get_complete_url( @@ -73,7 +79,9 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): ) -> ModelResponse: _json_response: Final = await raw_response.json() model_response.id = _json_response.get("id") - model_response.choices = [Choices(**choice) for choice in _json_response.get("choices")] + model_response.choices = [ + Choices.model_validate(choice) for choice in _JSON_OBJECTS.validate_python(_json_response.get("choices")) + ] model_response.created = _json_response.get("created") model_response.model = _json_response.get("model") model_response.object = _json_response.get("object") diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index 920916726f2..401c09b580b 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -27,7 +27,7 @@ class AnthropicTokenCounter(BaseTokenCounter): self, model_to_use: str, messages: list[dict[str, Any]] | None, - contents: list[dict[str, Any]] | None, + contents: list[dict[str, object]] | None, deployment: dict[str, Any] | None = None, request_model: str = "", tools: list[dict[str, Any]] | None = None, diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index da21f130f1a..082902e8897 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -19,6 +19,7 @@ from typing import Final, cast import httpx from openai.types.file_deleted import FileDeleted +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.litellm_core_utils.url_utils import encode_url_path_segment @@ -47,6 +48,8 @@ ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com" ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14" ANTHROPIC_MESSAGE_BATCH_ID_PREFIX: Final = "msgbatch_" +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AnthropicFilesConfig(BaseFilesConfig): """ @@ -211,7 +214,7 @@ class AnthropicFilesConfig(BaseFilesConfig): "created_at": "2025-01-01T00:00:00Z" } """ - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) return self._parse_anthropic_file(response_json) def transform_retrieve_file_request( @@ -230,7 +233,7 @@ class AnthropicFilesConfig(BaseFilesConfig): logging_obj: LiteLLMLoggingObj, litellm_params: dict, ) -> OpenAIFileObject: - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) return self._parse_anthropic_file(response_json) def transform_delete_file_request( @@ -249,13 +252,9 @@ class AnthropicFilesConfig(BaseFilesConfig): logging_obj: LiteLLMLoggingObj, litellm_params: dict, ) -> FileDeleted: - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) file_id: Final = response_json.get("id", "") - return FileDeleted( - id=file_id, - deleted=True, - object="file", - ) + return FileDeleted.model_validate({"id": file_id, "deleted": True, "object": "file"}) def transform_list_files_request( self, diff --git a/litellm/llms/azure/audio_transcription/transformation.py b/litellm/llms/azure/audio_transcription/transformation.py index 44623ec778a..17597da3895 100644 --- a/litellm/llms/azure/audio_transcription/transformation.py +++ b/litellm/llms/azure/audio_transcription/transformation.py @@ -5,10 +5,12 @@ Maps OpenAI-compatible audio transcription calls to Azure Speech REST recognition for short audio. """ -from typing import Any, Final +from collections.abc import Mapping +from typing import Final from urllib.parse import urlencode, urlparse import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( @@ -23,6 +25,8 @@ from litellm.types.llms.openai import ( ) from litellm.types.utils import FileTypes, TranscriptionResponse +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AzureSpeechAudioTranscriptionException(BaseLLMException): pass @@ -127,7 +131,8 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): raw_response: httpx.Response, ) -> TranscriptionResponse: response_json: Final = raw_response.json() - recognition_status: Final = response_json.get("RecognitionStatus") + payload: Final = _JSON_OBJECT.validate_python(response_json) + recognition_status: Final = payload.get("RecognitionStatus") if recognition_status is not None and recognition_status != "Success": raise AzureSpeechAudioTranscriptionException( message=(f"Azure AI Speech transcription failed with RecognitionStatus={recognition_status}."), @@ -135,7 +140,7 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): headers=raw_response.headers, ) - text: Final = self._extract_text(response_json) + text: Final = self._extract_text(payload) response: Final = TranscriptionResponse(text=text) response._hidden_params = response_json return response @@ -194,9 +199,10 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return "detailed" return "simple" - def _extract_text(self, response_json: dict[str, Any]) -> str: - if isinstance(response_json.get("DisplayText"), str): - return response_json["DisplayText"] + def _extract_text(self, response_json: Mapping[str, object]) -> str: + display_text: Final = response_json.get("DisplayText") + if isinstance(display_text, str): + return display_text nbest: Final = response_json.get("NBest") if isinstance(nbest, list) and nbest: diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py index 36d5a56db0d..6d6e10ce1dc 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py @@ -30,7 +30,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig): messages: list[dict[str, Any]], api_key: str, api_base: str, - litellm_params: dict[str, Any] | None = None, + litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = None, tools: list[dict[str, Any]] | None = None, system: object = None, diff --git a/litellm/llms/base_llm/harness/transformation.py b/litellm/llms/base_llm/harness/transformation.py index 643f808d4d8..d6db3d2ecb1 100644 --- a/litellm/llms/base_llm/harness/transformation.py +++ b/litellm/llms/base_llm/harness/transformation.py @@ -18,7 +18,7 @@ from __future__ import annotations from abc import ABC, abstractmethod from collections.abc import Mapping, Sequence from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, ClassVar, Generic, TypeVar +from typing import TYPE_CHECKING, ClassVar, Generic, TypeVar from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.types import Capabilities, Event, Harness @@ -133,7 +133,7 @@ class BaseCLIHarnessConfig(BaseHarnessConfig[OptionsT], Generic[OptionsT, Stream """Fresh per-turn parser state.""" @abstractmethod - def transform_stream_line(self, line: Mapping[str, Any], state: StreamStateT) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: StreamStateT) -> Sequence[Event]: """One decoded JSON line from stdout to zero or more events. Pure.""" @abstractmethod diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index fd4c3dc1659..e87c5c35fed 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -1,6 +1,6 @@ from collections.abc import Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final, Literal from openai.types.batch import BatchRequestCounts from openai.types.batch import Metadata as OpenAIBatchMetadata @@ -17,7 +17,12 @@ if TYPE_CHECKING: # AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses. # Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response` # so create / retrieve return consistent statuses. -_BEDROCK_MIJ_STATUS_TO_OPENAI: Final = { +_BEDROCK_MIJ_STATUS_TO_OPENAI: Final[ + Mapping[ + str, + Literal["validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"], + ] +] = { "Submitted": "validating", "Validating": "validating", "Scheduled": "validating", @@ -92,7 +97,7 @@ def _record_counts_from_response(response: Mapping[str, object]) -> BatchRequest ) -def _to_epoch(value: Any) -> int | None: +def _to_epoch(value: object) -> int | None: if value is None: return None if isinstance(value, (int, float)): @@ -349,10 +354,7 @@ class BedrockBatchesHandler: ) bedrock_status: Final = str(response.get("status", "")) - openai_status: Final = cast( - Any, - _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress"), - ) + openai_status: Final = _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress") input_uri: Final = response.get("inputDataConfig", {}).get("s3InputDataConfig", {}).get("s3Uri", "") output_prefix: Final = response.get("outputDataConfig", {}).get("s3OutputDataConfig", {}).get("s3Uri", "") diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index 01e25f4671e..f66dcf130b8 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -25,6 +25,7 @@ import base64 from typing import TYPE_CHECKING, Any, Final import httpx +from httpx._types import RequestFiles from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.bedrock.common_utils import BedrockError @@ -165,7 +166,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> tuple[dict, Any]: + ) -> tuple[dict, RequestFiles]: """ Transform OpenAI-style request to Bedrock Stability request format. diff --git a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py index e1b06791c9d..12ace32f43b 100644 --- a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py @@ -80,9 +80,7 @@ class AmazonTitanImageGenerationConfig: non_default_params: dict, optional_params: dict, ): - from typing import Any - - image_generation_config: Final[dict[str, Any]] = {} + image_generation_config: Final[dict[str, object]] = {} for k, v in non_default_params.items(): if k == "size" and v is not None: width, height = v.split("x") @@ -106,11 +104,9 @@ class AmazonTitanImageGenerationConfig: text: str, optional_params: dict, ) -> AmazonTitanImageGenerationRequestBody: - from typing import Any - image_generation_config = optional_params.pop("imageGenerationConfig", {}) negative_text: Final = optional_params.pop("negativeText", None) - text_to_image_params: Final[dict[str, Any]] = {"text": text} + text_to_image_params: Final = AmazonTitanTextToImageParams(text=text) if negative_text: text_to_image_params["negativeText"] = negative_text task_type: Final = optional_params.pop("taskType", "TEXT_IMAGE") @@ -121,7 +117,7 @@ class AmazonTitanImageGenerationConfig: } return AmazonTitanImageGenerationRequestBody( taskType=task_type, - textToImageParams=AmazonTitanTextToImageParams(**text_to_image_params), + textToImageParams=text_to_image_params, imageGenerationConfig=AmazonNovaCanvasImageGenerationConfig(**image_generation_config), ) diff --git a/litellm/llms/bedrock/rerank/handler.py b/litellm/llms/bedrock/rerank/handler.py index 8847381cbc9..4a8308bb46a 100644 --- a/litellm/llms/bedrock/rerank/handler.py +++ b/litellm/llms/bedrock/rerank/handler.py @@ -2,6 +2,7 @@ import json from typing import TYPE_CHECKING, Any, Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging @@ -24,6 +25,8 @@ if TYPE_CHECKING: else: AWSPreparedRequest = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class BedrockRerankHandler(BaseAWSLLM): async def arerank( @@ -55,13 +58,13 @@ class BedrockRerankHandler(BaseAWSLLM): except httpx.TimeoutException: raise BedrockError(status_code=408, message="Timeout error occurred.") - return BedrockRerankConfig()._transform_response(response.json()) + return BedrockRerankConfig()._transform_response(_JSON_DICT.validate_python(response.json())) def rerank( self, model: str, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], optional_params: dict, logging_obj: LitellmLogging, top_n: int | None = None, @@ -136,7 +139,7 @@ class BedrockRerankHandler(BaseAWSLLM): api_key="", ) - response_json: Final = response.json() + response_json: Final = _JSON_DICT.validate_python(response.json()) return BedrockRerankConfig()._transform_response(response_json) diff --git a/litellm/llms/bedrock/search/transformation.py b/litellm/llms/bedrock/search/transformation.py index 19ba7d5673b..565fc12aa1e 100644 --- a/litellm/llms/bedrock/search/transformation.py +++ b/litellm/llms/bedrock/search/transformation.py @@ -37,6 +37,7 @@ from collections.abc import Iterator, Mapping, Sequence from typing import Final import httpx +from pydantic import TypeAdapter from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.search.transformation import ( @@ -81,6 +82,8 @@ _SSE_EVENT_SEPARATOR: Final = re.compile(r"\r?\n[ \t]*\r?\n") _SSE_LINE_PREFIXES: Final = ("event:", "data:", ":", "id:", "retry:") +_JSON_VALUE: Final = TypeAdapter(object) + def _gateway_host_match(api_base: str) -> re.Match[str] | None: return _GATEWAY_HOST_PATTERN.fullmatch(httpx.URL(api_base).host) @@ -128,7 +131,7 @@ def _parse_result_items(raw_text: object) -> tuple[Mapping[str, object], ...]: if not isinstance(raw_text, str): return () try: - parsed: Final = json.loads(raw_text) + parsed: Final = _JSON_VALUE.validate_python(json.loads(raw_text)) except json.JSONDecodeError: return () return _result_items(parsed) @@ -147,7 +150,7 @@ def _iter_sse_events(text: str) -> Iterator[Mapping[str, object]]: if not payload: continue try: - parsed = json.loads(payload) + parsed = _JSON_VALUE.validate_python(json.loads(payload)) except json.JSONDecodeError: continue if isinstance(parsed, dict): diff --git a/litellm/llms/black_forest_labs/image_generation/transformation.py b/litellm/llms/black_forest_labs/image_generation/transformation.py index e5c2bdf59e9..9d586fba26a 100644 --- a/litellm/llms/black_forest_labs/image_generation/transformation.py +++ b/litellm/llms/black_forest_labs/image_generation/transformation.py @@ -8,9 +8,11 @@ API Reference: https://docs.bfl.ai/ """ import time +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -36,6 +38,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): """ @@ -218,7 +222,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): https://docs.bfl.ai/flux_models/flux_1_1_pro """ # Build request body with prompt - request_body: Final[dict[str, Any]] = { + request_body: Final[dict[str, object]] = { "prompt": prompt, } @@ -275,7 +279,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): message=f"Error parsing BFL response: {e}", ) - result: Final = response_data.get("result", {}) + result: Final = _JSON_OBJECT.validate_python(response_data).get("result", {}) if not model_response.data: model_response.data = [] diff --git a/litellm/llms/clarifai/chat/transformation.py b/litellm/llms/clarifai/chat/transformation.py index a0946254de0..cea6e6c7440 100644 --- a/litellm/llms/clarifai/chat/transformation.py +++ b/litellm/llms/clarifai/chat/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.openai.common_utils import OpenAIError @@ -20,6 +22,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ClarifaiConfig(OpenAIGPTConfig): """ @@ -111,7 +115,7 @@ class ClarifaiConfig(OpenAIGPTConfig): headers=raw_response.headers, ) from e - response: Final = ModelResponse(**completion_response) + response: Final = ModelResponse(**_JSON_OBJECT.validate_python(completion_response)) if response.model is not None: response.model = "clarifai/" + model diff --git a/litellm/llms/claude_code/harness/transformation.py b/litellm/llms/claude_code/harness/transformation.py index ea3cdd67ffb..f75d86c7e63 100644 --- a/litellm/llms/claude_code/harness/transformation.py +++ b/litellm/llms/claude_code/harness/transformation.py @@ -16,6 +16,8 @@ from dataclasses import dataclass, field from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.options import ClaudeCodeOptions from litellm.harness.types import ( @@ -48,6 +50,7 @@ if TYPE_CHECKING: CLAUDE_BINARY: Final = "claude" SYNTHETIC_MODEL: Final = "" +_BLOCK: Final = TypeAdapter(Mapping[object, object], config=ConfigDict(hide_input_in_errors=True)) BASE_COMMAND: Final = ("-p", "--output-format", "stream-json", "--verbose", "--input-format", "text") @@ -157,7 +160,7 @@ def _stringify_block(block: object) -> str: return json.dumps(block, ensure_ascii=False) -def _message_blocks(event: Mapping[str, object]) -> Sequence[Any]: +def _message_blocks(event: Mapping[str, object]) -> Sequence[object]: message: Final = event.get("message") content: Final = message.get("content") if isinstance(message, Mapping) else None if isinstance(content, str): @@ -211,8 +214,7 @@ def _user_events(event: Mapping[str, object]) -> Sequence[Event]: output=stringify_tool_output(block.get("content")), is_error=bool(block.get("is_error", False)), ) - for block in _message_blocks(event) - if _is_tool_result(block) + for block in map(_BLOCK.validate_python, filter(_is_tool_result, _message_blocks(event))) ) ) diff --git a/litellm/llms/codex/harness/transformation.py b/litellm/llms/codex/harness/transformation.py index ab17cf3d869..86a5f8274e6 100644 --- a/litellm/llms/codex/harness/transformation.py +++ b/litellm/llms/codex/harness/transformation.py @@ -17,6 +17,8 @@ from dataclasses import dataclass, field from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.options import CodexOptions @@ -61,6 +63,7 @@ MANAGED_CONFIG_KEYS: Final = frozenset( ) _BARE_TOML_KEY: Final = re.compile(r"^[A-Za-z0-9_-]+$") _TOOL_ITEM_TYPES: Final = frozenset({"command_execution", "file_change", "web_search", "mcp_tool_call"}) +_CHANGES: Final = TypeAdapter(tuple[Mapping[object, object], ...], config=ConfigDict(hide_input_in_errors=True)) @dataclass @@ -92,7 +95,7 @@ def _tool_input(item: Mapping[str, Any]) -> tuple[str, str, Mapping[str, Any], b return name, tool, tool_args, False -def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: +def _tool_output(item: Mapping[str, object]) -> tuple[str, bool]: """(output text, is_error) for a completed tool-like item.""" item_type = item.get("type") status = item.get("status") @@ -101,7 +104,8 @@ def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: is_error = status == "failed" or (exit_code is not None and exit_code != 0) return str(item.get("aggregated_output") or ""), is_error if item_type == "file_change": - lines = (f"{c.get('kind', '')} {c.get('path', '')}".strip() for c in item.get("changes") or ()) + changes: Final = _CHANGES.validate_python(item.get("changes") or ()) + lines: Final = (f"{c.get('kind', '')} {c.get('path', '')}".strip() for c in changes) return "\n".join(lines), status == "failed" if item_type == "web_search": return "", status == "failed" @@ -118,7 +122,7 @@ def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: def _tool_item_events( - item_id: str, item: Mapping[str, Any], completed: bool, state: CodexStreamState + item_id: str, item: Mapping[str, object], completed: bool, state: CodexStreamState ) -> Iterator[Event]: if item_id not in state.started: state.started.add(item_id) @@ -129,7 +133,7 @@ def _tool_item_events( yield ToolResult(id=item_id, output=output, is_error=is_error) -def _item_events(event_type: str, item: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: +def _item_events(event_type: str, item: Mapping[str, object], state: CodexStreamState) -> Sequence[Event]: item_type = item.get("type") item_id = str(item.get("id") or "") completed = event_type == "item.completed" @@ -181,7 +185,7 @@ def _config_override(key: object, value: object) -> str: return f"{dotted}={toml_value(value)}" -def config_overrides(config: Mapping[str, Any]) -> Sequence[str]: +def config_overrides(config: Mapping[str, object]) -> Sequence[str]: """`-c` override strings for CodexOptions.config, rejecting managed keys.""" overrides: Final = (_config_override(key, value) for key, value in config.items()) return list(overrides) # mutable-ok: public helper; tests compare to a list @@ -310,7 +314,7 @@ class CodexHarnessConfig(BaseCLIHarnessConfig): def create_stream_state(self) -> CodexStreamState: return CodexStreamState() - def transform_stream_line(self, line: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: CodexStreamState) -> Sequence[Event]: """turn.completed usage is ignored on purpose: the session endpoint accounts it.""" event_type = line.get("type") if event_type == "thread.started": diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index a8e755406d8..b6d94d4853a 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -1,7 +1,8 @@ from collections.abc import Mapping -from typing import Any, Final +from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -12,6 +13,8 @@ from litellm.types.rerank import OptionalRerankParams, RerankRequest, RerankResp from ..common_utils import CohereError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class CohereRerankConfig(BaseRerankConfig): """ @@ -51,7 +54,7 @@ class CohereRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -148,7 +151,7 @@ class CohereRerankConfig(BaseRerankConfig): except Exception: raise CohereError(message=raw_response.text, status_code=raw_response.status_code) - return RerankResponse(**raw_response_json) + return RerankResponse.model_validate(_JSON_OBJECT.validate_python(raw_response_json)) def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CohereError(message=error_message, status_code=status_code) diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py index 4432c151a64..567dbfd308e 100644 --- a/litellm/llms/cometapi/image_generation/transformation.py +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -20,6 +22,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class CometAPIImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://api.cometapi.com" @@ -155,7 +160,8 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): # CometAPI returns OpenAI-compatible format # Expected format: {"created": timestamp, "data": [{"url": "...", "b64_json": "..."}]} if "data" in response_data: - for image_data in response_data["data"]: + payload: Final = _JSON_OBJECT.validate_python(response_data) + for image_data in _JSON_OBJECTS.validate_python(payload["data"]): image_obj = ImageObject( b64_json=image_data.get("b64_json"), url=image_data.get("url"), diff --git a/litellm/llms/dashscope/image_generation/transformation.py b/litellm/llms/dashscope/image_generation/transformation.py index ffa60a3d9bc..6a172e1e074 100644 --- a/litellm/llms/dashscope/image_generation/transformation.py +++ b/litellm/llms/dashscope/image_generation/transformation.py @@ -23,9 +23,11 @@ Response format: } """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -45,6 +47,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + DEFAULT_API_BASE: Final = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation" CHAT_COMPATIBLE_MODE_PATH: Final = "/compatible-mode/v1" @@ -192,9 +197,10 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): # DashScope can return API-level errors in a 200 response body. # Example: {"code": "InvalidParameter", "message": "Size not supported"} - if "code" in response_data and "output" not in response_data: + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if "code" in response_object and "output" not in response_object: raise self.get_error_class( - error_message=str(response_data.get("message", response_data)), + error_message=str(response_object.get("message", response_object)), status_code=raw_response.status_code, headers=raw_response.headers, ) @@ -202,9 +208,11 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): if not model_response.data: model_response.data = [] - choices: Final = response_data.get("output", {}).get("choices", []) + output: Final = _JSON_OBJECT.validate_python(response_object.get("output", {})) + choices: Final = _JSON_OBJECTS.validate_python(output.get("choices", [])) for choice in choices: - content_list = choice.get("message", {}).get("content", []) + message = _JSON_OBJECT.validate_python(choice.get("message", {})) + content_list = _JSON_OBJECTS.validate_python(message.get("content", [])) for content_item in content_list: image_url = content_item.get("image") if image_url: diff --git a/litellm/llms/dashscope/rerank/transformation.py b/litellm/llms/dashscope/rerank/transformation.py index 14ad756ec9c..4ec7d697f02 100644 --- a/litellm/llms/dashscope/rerank/transformation.py +++ b/litellm/llms/dashscope/rerank/transformation.py @@ -26,10 +26,11 @@ as supported only for gte-rerank-v2 / qwen3-vl-rerank. Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api """ -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Iterable, Mapping +from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -48,6 +49,11 @@ from ..common_utils import DashScopeError, resolve_dashscope_family_rerank_api_b DEFAULT_RERANK_URL: Final = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class DashScopeRerankConfig(BaseRerankConfig): """ @@ -117,7 +123,7 @@ class DashScopeRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -197,7 +203,8 @@ class DashScopeRerankConfig(BaseRerankConfig): message=response_json.get("message", str(response_json)), ) - results: Final = response_json.get("results") + payload: Final = _JSON_OBJECT.validate_python(response_json) + results: Final = payload.get("results") if results is None: raise DashScopeError( status_code=raw_response.status_code, @@ -210,7 +217,7 @@ class DashScopeRerankConfig(BaseRerankConfig): # "document": {"text": "..."} # which already matches LiteLLM's RerankResponseDocument shape. transformed_results: Final[list[dict]] = [] - for r in results: + for r in _JSON_OBJECTS.validate_python(results): item: dict[str, object] = { "index": r["index"], "relevance_score": r["relevance_score"], @@ -223,14 +230,14 @@ class DashScopeRerankConfig(BaseRerankConfig): item["document"] = {"text": doc} transformed_results.append(item) - usage: Final = response_json.get("usage") or {} - total_tokens: Final = usage.get("total_tokens") + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage") or {}) + total_tokens: Final = _OPTIONAL_INT.validate_python(usage.get("total_tokens")) billed_units: Final = RerankBilledUnits(total_tokens=total_tokens) tokens: Final = RerankTokens(input_tokens=total_tokens) meta: Final = RerankResponseMeta(billed_units=billed_units, tokens=tokens) return RerankResponse( - id=response_json.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=transformed_results, meta=meta, ) diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py index 14f87f0753b..1f8728230bc 100644 --- a/litellm/llms/deepagents/harness/transformation.py +++ b/litellm/llms/deepagents/harness/transformation.py @@ -173,7 +173,7 @@ def stream_events( ] -def tool_call_event(call: Mapping[str, Any]) -> ToolCall: +def tool_call_event(call: Mapping[str, object]) -> ToolCall: native = str(call.get("name") or "") args = call.get("args") return ToolCall( @@ -185,7 +185,7 @@ def tool_call_event(call: Mapping[str, Any]) -> ToolCall: ) -def _node_messages(update: Mapping[Any, Any]) -> Iterator[object]: +def _node_messages(update: Mapping[object, object]) -> Iterator[object]: for node, delta in update.items(): if node not in _EVENT_NODES or not isinstance(delta, Mapping): continue @@ -231,7 +231,7 @@ def interrupts_in( return list(items) # mutable-ok: list return; callers/tests compare to lists -def final_ai_text(messages: Sequence[Any]) -> str: +def final_ai_text(messages: Sequence[object]) -> str: for message in reversed(messages): if getattr(message, "type", None) == "ai": text = content_text(getattr(message, "content", "")) diff --git a/litellm/llms/deepgram/audio_transcription/transformation.py b/litellm/llms/deepgram/audio_transcription/transformation.py index 95bdd406c3f..c5c9fe12423 100644 --- a/litellm/llms/deepgram/audio_transcription/transformation.py +++ b/litellm/llms/deepgram/audio_transcription/transformation.py @@ -2,10 +2,12 @@ Translates from OpenAI's `/v1/audio/transcriptions` to Deepgram's `/v1/listen` """ +from collections.abc import Iterable, Mapping from typing import Final from urllib.parse import urlencode from httpx import Headers, Response +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -22,6 +24,9 @@ from ...base_llm.audio_transcription.transformation import ( ) from ..common_utils import DeepgramException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: @@ -105,7 +110,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response["task"] = "transcribe" # Use detected_language if available, otherwise default to "en" - detected_language: Final = first_channel.get("detected_language") + detected_language: Final = _JSON_OBJECT.validate_python(first_channel).get("detected_language") response["language"] = detected_language if detected_language else "en" response["duration"] = response_json["metadata"]["duration"] @@ -114,7 +119,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): if "words" in first_alternative: response["words"] = [ {"word": word["word"], "start": word["start"], "end": word["end"]} - for word in first_alternative["words"] + for word in _JSON_OBJECTS.validate_python(first_alternative["words"]) ] # Store full response in hidden params diff --git a/litellm/llms/e2b/sandbox/transformation.py b/litellm/llms/e2b/sandbox/transformation.py index 4928ca0c092..d2702b116dc 100644 --- a/litellm/llms/e2b/sandbox/transformation.py +++ b/litellm/llms/e2b/sandbox/transformation.py @@ -8,9 +8,11 @@ Talks to e2b's REST API directly over httpx (no e2b SDK dependency): """ import json +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.sandbox.transformation import ( SANDBOX_MAX_OUTPUT_BYTES, @@ -32,6 +34,12 @@ JUPYTER_PORT: Final = 49999 DEFAULT_SANDBOX_TIMEOUT: Final = 300 MAX_OUTPUT_BYTES: Final = SANDBOX_MAX_OUTPUT_BYTES +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_MESSAGE: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter( + Mapping[str, object] | None, config=ConfigDict(hide_input_in_errors=True) +) +_TEXT: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class E2BSandboxConfig(BaseSandboxConfig): def _http(self, client: AsyncHTTPHandler | None) -> AsyncHTTPHandler: @@ -73,12 +81,14 @@ class E2BSandboxConfig(BaseSandboxConfig): headers={"X-API-Key": key, "Content-Type": "application/json"}, json=body, ) - data: Final = response.json() + data: Final = _JSON_OBJECT.validate_python(response.json()) - handle: Final = ContainerHandle( - id=data["sandboxID"], - provider="e2b", - domain=data.get("domain") or E2B_DEFAULT_DOMAIN, + handle: Final = ContainerHandle.model_validate( + { + "id": data["sandboxID"], + "provider": "e2b", + "domain": data.get("domain") or E2B_DEFAULT_DOMAIN, + } ) handle._hidden_params = { "envd_access_token": data.get("envdAccessToken"), @@ -158,7 +168,7 @@ class E2BSandboxConfig(BaseSandboxConfig): def _parse_lines(lines: list[str]) -> CodeExecutionResult: def _try_parse(stripped: str): try: - return json.loads(stripped) + return _MESSAGE.validate_python(json.loads(stripped)) except json.JSONDecodeError: return None @@ -178,10 +188,12 @@ class E2BSandboxConfig(BaseSandboxConfig): None, ) - return CodeExecutionResult( - stdout="".join(m.get("text", "") for m in of_type("stdout")), - stderr="".join(m.get("text", "") for m in of_type("stderr")), - results=[{k: v for k, v in m.items() if k != "type"} for m in of_type("result")], - error=error, - execution_count=execution_count, + return CodeExecutionResult.model_validate( + { + "stdout": "".join(_TEXT.validate_python(m.get("text", "")) for m in of_type("stdout")), + "stderr": "".join(_TEXT.validate_python(m.get("text", "")) for m in of_type("stderr")), + "results": [{k: v for k, v in m.items() if k != "type"} for m in of_type("result")], + "error": error, + "execution_count": execution_count, + } ) diff --git a/litellm/llms/elevenlabs/audio_transcription/transformation.py b/litellm/llms/elevenlabs/audio_transcription/transformation.py index a2b25760aac..0dbf61917e2 100644 --- a/litellm/llms/elevenlabs/audio_transcription/transformation.py +++ b/litellm/llms/elevenlabs/audio_transcription/transformation.py @@ -2,9 +2,11 @@ Translates from OpenAI's `/v1/audio/transcriptions` to ElevenLabs's `/v1/speech-to-text` """ +from collections.abc import Iterable, Mapping from typing import Final from httpx import Headers, Response +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -22,6 +24,9 @@ from ...base_llm.audio_transcription.transformation import ( ) from ..common_utils import ElevenLabsException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): @property @@ -115,21 +120,22 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): """ try: response_json: Final = raw_response.json() + response_object: Final = _JSON_OBJECT.validate_python(response_json) # Extract the main transcript text - text: Final = response_json.get("text", "") + text: Final = response_object.get("text", "") # Create TranscriptionResponse object response: Final = TranscriptionResponse(text=text) # Add additional metadata matching OpenAI format response["task"] = "transcribe" - response["language"] = response_json.get("language_code", "unknown") + response["language"] = response_object.get("language_code", "unknown") # Map ElevenLabs words to OpenAI format - if "words" in response_json: + if "words" in response_object: response["words"] = [] - for word_data in response_json["words"]: + for word_data in _JSON_OBJECTS.validate_python(response_object["words"]): # Only include actual words, skip spacing and audio events if word_data.get("type") == "word": response["words"].append( diff --git a/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py index 7c63e1077f1..d4d2941db43 100644 --- a/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIFluxProV11UltraConfig(FalAIBaseConfig): """ @@ -228,16 +232,17 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig): if not model_response.data: model_response.data = [] - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) model_response.data.extend(fal_images_to_image_objects(images)) # Add additional metadata from Flux Pro response if hasattr(model_response, "_hidden_params"): - if "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] - if "timings" in response_data: - model_response._hidden_params["timings"] = response_data["timings"] - if "has_nsfw_concepts" in response_data: - model_response._hidden_params["has_nsfw_concepts"] = response_data["has_nsfw_concepts"] + if "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] + if "timings" in response_object: + model_response._hidden_params["timings"] = response_object["timings"] + if "has_nsfw_concepts" in response_object: + model_response._hidden_params["has_nsfw_concepts"] = response_object["has_nsfw_concepts"] return model_response 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 ad1852a622b..970addbd5a2 100644 --- a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIIdeogramV3Config(FalAIBaseConfig): """ @@ -169,7 +173,8 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): if not model_response.data: model_response.data = [] - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) if isinstance(images, list): for image_entry in images: if isinstance(image_entry, dict): @@ -184,7 +189,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): ) ) - if hasattr(model_response, "_hidden_params") and "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] + if hasattr(model_response, "_hidden_params") and "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] return model_response diff --git a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py index 934ce420d53..038a06c5749 100644 --- a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIRecraftV3Config(FalAIBaseConfig): """ @@ -89,7 +93,7 @@ class FalAIRecraftV3Config(FalAIBaseConfig): return optional_params - def _map_image_size(self, size: str) -> Any: + def _map_image_size(self, size: str) -> str | Mapping[str, int]: """ Map OpenAI size format to Recraft v3 image_size format. @@ -203,7 +207,7 @@ class FalAIRecraftV3Config(FalAIBaseConfig): model_response.data = [] # Handle Recraft v3 response format - images: Final = response_data.get("images", []) + images: Final = _JSON_OBJECT.validate_python(response_data).get("images", []) if isinstance(images, list): for image_data in images: if isinstance(image_data, dict): diff --git a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py index 79d8800773b..659011ae537 100644 --- a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py +++ b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIStableDiffusionConfig(FalAIBaseConfig): """ @@ -124,7 +128,7 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): return optional_params - def _map_image_size(self, size: str) -> Any: + def _map_image_size(self, size: str) -> str | Mapping[str, int]: """ Map OpenAI size format to Stable Diffusion image_size format. @@ -243,7 +247,8 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): model_response.data = [] # Handle Stable Diffusion response format - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) if isinstance(images, list): for image_data in images: if isinstance(image_data, dict): @@ -264,11 +269,11 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): # Add additional metadata from Stable Diffusion response if hasattr(model_response, "_hidden_params"): - if "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] - if "timings" in response_data: - model_response._hidden_params["timings"] = response_data["timings"] - if "has_nsfw_concepts" in response_data: - model_response._hidden_params["has_nsfw_concepts"] = response_data["has_nsfw_concepts"] + if "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] + if "timings" in response_object: + model_response._hidden_params["timings"] = response_object["timings"] + if "has_nsfw_concepts" in response_object: + model_response._hidden_params["has_nsfw_concepts"] = response_object["has_nsfw_concepts"] return model_response diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index 509dbd5ff24..3d9d813c6e9 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -4,10 +4,11 @@ Fireworks AI Rerank API transformation Reference: https://docs.fireworks.ai/inference-api-reference/rerank """ -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Iterable, Mapping, Sequence +from typing import Final import httpx +from pydantic import BaseModel, ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -23,6 +24,26 @@ from litellm.types.rerank import ( ) +class _FireworksAIUsageFields(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + total_tokens: int | None = 0 + prompt_tokens: int | None = 0 + completion_tokens: int | None = 0 + + +class _FireworksAIResultFields(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + index: int | float | str + relevance_score: int | float | str + + +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str) + + class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): """ Fireworks AI Rerank API configuration @@ -59,7 +80,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: Sequence[str | Mapping[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -204,23 +225,18 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): # } # Extract usage information - usage: Final = raw_response_json.get("usage", {}) - _billed_units: Final = RerankBilledUnits(search_units=usage.get("total_tokens", 0)) - _tokens: Final = RerankTokens( - input_tokens=usage.get("prompt_tokens", 0), - output_tokens=usage.get("completion_tokens", 0), - ) - rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + response_json: Final = _JSON_OBJECT.validate_python(raw_response_json) + usage: Final = _JSON_OBJECT.validate_python(response_json.get("usage", {})) # Extract results - Fireworks AI uses "data" instead of "results" - _results: Final[list[dict] | None] = raw_response_json.get("data") or raw_response_json.get("results") + _results: Final = response_json.get("data") or response_json.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") rerank_results: Final[list[RerankResponseResult]] = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): # Validate required fields exist if not all(key in result for key in ["index", "relevance_score"]): raise ValueError(f"Missing required fields in the result={result}") @@ -239,9 +255,10 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): document = RerankResponseDocument(text=str(text)) # Create typed result + fields = _FireworksAIResultFields.model_validate(result) rerank_result = RerankResponseResult( - index=int(result["index"]), - relevance_score=float(result["relevance_score"]), + index=int(fields.index), + relevance_score=float(fields.relevance_score), ) # Only add document if it exists @@ -250,7 +267,15 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): rerank_results.append(rerank_result) - response_id: Final = raw_response_json.get("id") or str(uuid.uuid4()) + usage_fields: Final = _FireworksAIUsageFields.model_validate(usage) + _billed_units: Final = RerankBilledUnits(search_units=usage_fields.total_tokens) + _tokens: Final = RerankTokens( + input_tokens=usage_fields.prompt_tokens, + output_tokens=usage_fields.completion_tokens, + ) + rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + + response_id: Final = _STR.validate_python(response_json.get("id") or str(uuid.uuid4())) return RerankResponse( id=response_id, diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py index bb8d7455031..4d13d0d2e6b 100644 --- a/litellm/llms/gemini/image_generation/transformation.py +++ b/litellm/llms/gemini/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -31,6 +33,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GoogleImageGenConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta" @@ -216,13 +221,15 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): # Extract usage metadata for Gemini models if "usageMetadata" in response_data: - model_response.usage = transform_gemini_image_usage(response_data["usageMetadata"]) + model_response.usage = transform_gemini_image_usage( + _JSON_DICT.validate_python(response_data["usageMetadata"]) + ) web_search_requests: Final = get_gemini_image_web_search_requests(response_data) if web_search_requests and model_response.usage is not None: setattr(model_response.usage, "web_search_requests", web_search_requests) else: # Original Imagen format - predictions with generated images - predictions: Final = response_data.get("predictions", []) + predictions: Final = _JSON_OBJECTS.validate_python(response_data.get("predictions", [])) for prediction in predictions: # Google AI returns base64 encoded images in the prediction model_response.data.append( diff --git a/litellm/llms/gigachat/embedding/transformation.py b/litellm/llms/gigachat/embedding/transformation.py index 927f5e944b6..321f5525433 100644 --- a/litellm/llms/gigachat/embedding/transformation.py +++ b/litellm/llms/gigachat/embedding/transformation.py @@ -8,9 +8,11 @@ API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/res from __future__ import annotations import types +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm import LlmProviders from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -22,6 +24,8 @@ from litellm.types.utils import EmbeddingResponse from ..authenticator import get_access_token +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GigaChatEmbeddingError(BaseLLMException): """GigaChat Embedding API error.""" @@ -165,7 +169,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): "total_tokens": total_tokens, } - return EmbeddingResponse(**response_json) + return EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) def validate_environment( self, diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index 7ea7a89b4ca..c1acf6fccaa 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -11,6 +11,7 @@ import os from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._logging import verbose_logger from litellm.exceptions import AuthenticationError @@ -33,6 +34,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_RESPONSE_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): """ @@ -154,7 +157,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): logging_obj.post_call(original_response=raw_response.text) # GitHub Copilot returns standard OpenAI-compatible embedding response - response_json: Final = raw_response.json() + response_json: Final = _RESPONSE_OBJECT.validate_python(raw_response.json()) return convert_to_model_response_object( response_object=response_json, diff --git a/litellm/llms/jina_ai/embedding/transformation.py b/litellm/llms/jina_ai/embedding/transformation.py index 260d9e6e494..e184ad2628b 100644 --- a/litellm/llms/jina_ai/embedding/transformation.py +++ b/litellm/llms/jina_ai/embedding/transformation.py @@ -7,9 +7,11 @@ Docs - https://jina.ai/embeddings/ """ import types +from collections.abc import Mapping from typing import Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter from litellm import LlmProviders from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -22,6 +24,8 @@ from litellm.utils import is_base64_encoded from ..common_utils import JinaAIError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class JinaAIEmbeddingConfig(BaseEmbeddingConfig): """ @@ -139,7 +143,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): additional_args={"complete_input_dict": request_data}, original_response=response_json, ) - return EmbeddingResponse(**response_json) + return EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) def validate_environment( self, diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 2a81c38fe34..f4a3a9f0f09 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -6,10 +6,11 @@ Why separate file? Make it easy to see how transformation works Docs - https://jina.ai/reranker """ -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final from httpx import URL, Response +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj @@ -23,6 +24,12 @@ from litellm.types.rerank import ( ) from litellm.types.utils import ModelInfo +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_BILLED_UNITS: Final = TypeAdapter(RerankBilledUnits) +_TOKENS: Final = TypeAdapter(RerankTokens) +_STR: Final = TypeAdapter(str) + class JinaAIRerankConfig(BaseRerankConfig): def get_supported_cohere_rerank_params(self, model: str) -> list: @@ -100,13 +107,11 @@ class JinaAIRerankConfig(BaseRerankConfig): logging_obj.post_call(original_response=raw_response.text) - _json_response: Final = raw_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(raw_response.json()) - _billed_units: Final = RerankBilledUnits(**_json_response.get("usage", {})) - _tokens: Final = RerankTokens(**_json_response.get("usage", {})) - rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + usage: Final = _JSON_OBJECT.validate_python(_json_response.get("usage", {})) - _results: Final[list[dict] | None] = _json_response.get("results") + _results: Final = _json_response.get("results") if _results is None: raise ValueError(f"No results found in the response={_json_response}") @@ -115,7 +120,7 @@ class JinaAIRerankConfig(BaseRerankConfig): # Jina AI returns: {"index": 0, "relevance_score": 0.72, "document": "hello"} # LiteLLM expects: {"index": 0, "relevance_score": 0.72, "document": {"text": "hello"}} transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result = { "index": result["index"], "relevance_score": result["relevance_score"], @@ -128,8 +133,12 @@ class JinaAIRerankConfig(BaseRerankConfig): transformed_result["document"] = result["document"] transformed_results.append(transformed_result) + _billed_units: Final = _BILLED_UNITS.validate_python(usage) + _tokens: Final = _TOKENS.validate_python(usage) + rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + return RerankResponse( - id=_json_response.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(_json_response.get("id") or str(uuid.uuid4())), results=transformed_results, meta=rerank_meta, ) # Return response diff --git a/litellm/llms/manus/files/transformation.py b/litellm/llms/manus/files/transformation.py index 8c7b02f8d92..94b0e62199d 100644 --- a/litellm/llms/manus/files/transformation.py +++ b/litellm/llms/manus/files/transformation.py @@ -11,10 +11,12 @@ Reference: https://open.manus.im/docs/openai-compatibility#file-management """ import time -from typing import Any, Final +from collections.abc import Iterable, Mapping +from typing import Final import httpx from openai.types.file_deleted import FileDeleted +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -39,6 +41,10 @@ from litellm.types.utils import LlmProviders MANUS_API_BASE: Final = "https://api.manus.im" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_TEXT: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ManusFilesConfig(BaseFilesConfig): """ @@ -337,8 +343,7 @@ class ManusFilesConfig(BaseFilesConfig): litellm_params: dict, ) -> FileDeleted: """Transform delete file response.""" - response_json: Final = raw_response.json() - return FileDeleted(**response_json) + return FileDeleted.model_validate(_JSON_OBJECT.validate_python(raw_response.json())) def transform_list_files_request( self, @@ -366,19 +371,20 @@ class ManusFilesConfig(BaseFilesConfig): litellm_params: dict, ) -> list[OpenAIFileObject]: """Transform list files response.""" - response_json: Final = raw_response.json() - files_data: Final = response_json.get("data", []) + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) + files_data: Final = _JSON_OBJECTS.validate_python(response_json.get("data", [])) return [self._parse_file_dict(f) for f in files_data] - def _parse_file_dict(self, file_dict: dict[str, Any]) -> OpenAIFileObject: + def _parse_file_dict(self, file_dict: Mapping[str, object]) -> OpenAIFileObject: """Parse a file dict into OpenAIFileObject.""" created_at_str: Final = file_dict.get("created_at", "") if created_at_str: + created_at_text: Final = _TEXT.validate_python(created_at_str) try: created_at = int( time.mktime( time.strptime( - created_at_str.replace("Z", "+00:00")[:19], + created_at_text.replace("Z", "+00:00")[:19], "%Y-%m-%dT%H:%M:%S", ) ) @@ -388,15 +394,17 @@ class ManusFilesConfig(BaseFilesConfig): else: created_at = int(time.time()) - return OpenAIFileObject( - id=file_dict.get("id", ""), - bytes=file_dict.get("bytes", 0), - created_at=created_at, - filename=file_dict.get("filename", ""), - object="file", - purpose=file_dict.get("purpose", "assistants"), - status=file_dict.get("status", "uploaded"), - status_details=file_dict.get("status_details"), + return OpenAIFileObject.model_validate( + { + "id": file_dict.get("id", ""), + "bytes": file_dict.get("bytes", 0), + "created_at": created_at, + "filename": file_dict.get("filename", ""), + "object": "file", + "purpose": file_dict.get("purpose", "assistants"), + "status": file_dict.get("status", "uploaded"), + "status_details": file_dict.get("status_details"), + } ) def transform_file_content_request( diff --git a/litellm/llms/minimax/text_to_speech/transformation.py b/litellm/llms/minimax/text_to_speech/transformation.py index e38a8a2c3a3..1cabc43eb1c 100644 --- a/litellm/llms/minimax/text_to_speech/transformation.py +++ b/litellm/llms/minimax/text_to_speech/transformation.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final import httpx from httpx import Headers +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -26,6 +27,9 @@ else: LiteLLMLoggingObj = Any HttpxBinaryResponseContent = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class MinimaxException(BaseLLMException): """Custom exception for MiniMax API errors""" @@ -299,7 +303,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): try: # Parse JSON response - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) # MiniMax API response format check # The API can return different structures: @@ -320,7 +324,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): # Extract audio data # MiniMax returns audio in "data" field - data: Final = response_json.get("data", {}) + data: Final = _JSON_OBJECT.validate_python(response_json.get("data", {})) # Check if response contains a URL (output_format='url') audio_url: Final = data.get("audio_url", None) @@ -334,7 +338,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): ) # Get hex-encoded audio data - audio_hex: Final = data.get("audio", "") or response_json.get("audio_file", "") + audio_hex: Final = _STR.validate_python(data.get("audio", "") or response_json.get("audio_file", "") or "") if not audio_hex: raise MinimaxException( diff --git a/litellm/llms/modelscope/image_generation/transformation.py b/litellm/llms/modelscope/image_generation/transformation.py index 3a8a37307d6..08cdbacdad0 100644 --- a/litellm/llms/modelscope/image_generation/transformation.py +++ b/litellm/llms/modelscope/image_generation/transformation.py @@ -6,9 +6,11 @@ Handles transformation between OpenAI-compatible format and ModelScope API forma API Reference: https://modelscope.cn/docs/model-service/API-Inference/intro """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Final import httpx +from pydantic import ConfigDict, TypeAdapter from typing_extensions import override from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -29,6 +31,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = object +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): """ @@ -177,8 +182,10 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): ) # Check for errors in response - if "error" in response_data: - error_msg: Final = response_data["error"].get("message", str(response_data["error"])) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if "error" in response_object: + error: Final = _JSON_OBJECT.validate_python(response_object["error"]) + error_msg: Final = error.get("message", str(error)) raise self.get_error_class( error_message=f"ModelScope error: {error_msg}", status_code=raw_response.status_code, @@ -186,7 +193,7 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): ) # Extract images from response - data_list: Final = response_data.get("data", []) + data_list: Final = _JSON_OBJECTS.validate_python(response_object.get("data", [])) if not model_response.data: model_response.data = [] diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index 15cfdb6bece..c8e1a331c8c 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -1,7 +1,8 @@ from collections.abc import Mapping -from typing import Any, Final, Literal +from typing import Final, Literal import httpx +from pydantic import ConfigDict, TypeAdapter from typing_extensions import Required, TypedDict import litellm @@ -44,6 +45,14 @@ class NvidiaNimRerankResponse(TypedDict): rankings: Required[list[NvidiaNimRankingResult]] +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_NUMBER: Final[TypeAdapter[bool | int | float]] = TypeAdapter( + bool | int | float, config=ConfigDict(strict=True, hide_input_in_errors=True) +) +_BILLED_UNITS: Final = TypeAdapter(RerankBilledUnits) +_STR: Final = TypeAdapter(str) + + class NvidiaNimRerankConfig(BaseRerankConfig): """ Reference: https://docs.api.nvidia.com/nim/reference/nvidia-llama-3_2-nv-rerankqa-1b-v2-infer @@ -115,7 +124,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -133,7 +142,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): Nvidia NIM specific params (passed through as-is from non_default_params): - truncate: How to truncate input if too long (NONE, END) """ - optional_nvidia_nim_rerank_params: Final[dict[str, Any]] = { + optional_nvidia_nim_rerank_params: Final[dict[str, object]] = { "query": query, "documents": documents, } @@ -327,15 +336,18 @@ class NvidiaNimRerankConfig(BaseRerankConfig): # Construct metadata with billed_units # Nvidia NIM uses "usage" field with "total_tokens" - usage: Final = raw_response_json.get("usage", {}) - total_tokens: Final = usage.get("total_tokens", 0) + payload: Final = _JSON_OBJECT.validate_python(raw_response_json) + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage", {})) + total_tokens: Final = _NUMBER.validate_python(usage.get("total_tokens", 0)) - billed_units: Final[RerankBilledUnits] = {"total_tokens": total_tokens if total_tokens > 0 else len(results)} + billed_units: Final = _BILLED_UNITS.validate_python( + {"total_tokens": total_tokens if total_tokens > 0 else len(results)} + ) meta: Final[RerankResponseMeta] = {"billed_units": billed_units} return RerankResponse( - id=raw_response_json.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=results, meta=meta, ) diff --git a/litellm/llms/nvidia_riva/audio_transcription/transformation.py b/litellm/llms/nvidia_riva/audio_transcription/transformation.py index cc78f5b99ef..262dfb7bd78 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/transformation.py +++ b/litellm/llms/nvidia_riva/audio_transcription/transformation.py @@ -102,7 +102,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): if endpointing_config is not None: recognition_config["endpointing_config"] = endpointing_config - request_payload: Final[dict[str, Any]] = { + request_payload: Final[dict[str, object]] = { "recognition_config": recognition_config, "response_format": optional_params.get("response_format") or "json", "timestamp_granularities": optional_params.get("timestamp_granularities"), @@ -135,7 +135,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): # gRPC auth is constructed in the handler, not via HTTP headers. return headers - def _build_recognition_config_dict(self, model: str, optional_params: dict) -> dict[str, Any]: + def _build_recognition_config_dict(self, model: str, optional_params: dict) -> dict[str, object]: """ Build the Riva ``RecognitionConfig`` shape as a plain dict. @@ -159,7 +159,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): "profanity_filter": optional_params.get("profanity_filter", False), } - def _build_endpointing_config_dict(self, optional_params: dict) -> dict[str, Any] | None: + def _build_endpointing_config_dict(self, optional_params: dict) -> dict[str, object] | None: """ Translate an OpenAI-style ``chunking_strategy`` into Riva's ``EndpointingConfig`` shape, or pass through an explicit @@ -177,7 +177,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return None if isinstance(chunking, dict) and chunking.get("type") == "server_vad": - config: Final[dict[str, Any]] = {} + config: Final[dict[str, object]] = {} if "threshold" in chunking: threshold: Final = float(chunking["threshold"]) config["start_threshold"] = threshold @@ -245,7 +245,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response["task"] = "transcribe" if response_format == "verbose_json": - words: Final[list[dict[str, Any]]] = [] + words: Final[list[dict[str, object]]] = [] if timestamp_granularities and "word" in timestamp_granularities: for item in final_results: for word in item.get("words", []) or []: diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index dec43717387..edadc293771 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -22,9 +22,11 @@ Supported models: Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -52,6 +54,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + # OCI sends up to 96 texts per embedText request (Cohere limit). OCI_EMBED_BATCH_LIMIT: Final = 96 @@ -275,7 +279,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): ) try: - parsed: Final = OCIEmbedResponse(**json_response) + parsed: Final = OCIEmbedResponse.model_validate(_JSON_OBJECT.validate_python(json_response)) except Exception as e: raise OCIError( status_code=500, diff --git a/litellm/llms/openai/responses/count_tokens/handler.py b/litellm/llms/openai/responses/count_tokens/handler.py index 66782a83a19..0d7a553d633 100644 --- a/litellm/llms/openai/responses/count_tokens/handler.py +++ b/litellm/llms/openai/responses/count_tokens/handler.py @@ -5,6 +5,7 @@ Uses httpx for HTTP requests to OpenAI's /v1/responses/input_tokens endpoint. """ import json +from collections.abc import Sequence from typing import Any, Final import httpx @@ -26,11 +27,11 @@ class OpenAICountTokensHandler(OpenAICountTokensConfig): async def handle_count_tokens_request( self, model: str, - input: str | list[Any], + input: str | Sequence[object], api_key: str, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, - tools: list[dict[str, Any]] | None = None, + tools: list[dict[str, object]] | None = None, instructions: str | None = None, ) -> dict[str, Any]: """ diff --git a/litellm/llms/opencode/harness/transformation.py b/litellm/llms/opencode/harness/transformation.py index af5fa1ae71d..158af0fae0a 100644 --- a/litellm/llms/opencode/harness/transformation.py +++ b/litellm/llms/opencode/harness/transformation.py @@ -165,7 +165,7 @@ def _as_dict(value: object) -> Mapping[str, Any]: return value if isinstance(value, dict) else MappingProxyType({}) -def _tool_events(part: Mapping[str, Any]) -> Sequence[Event]: +def _tool_events(part: Mapping[str, object]) -> Sequence[Event]: native = str(part.get("tool") or "") call_id = str(part.get("callID") or part.get("id") or "") state = _as_dict(part.get("state")) @@ -195,7 +195,7 @@ def _error_message(error: object) -> str: return str(error.get("name") or "opencode reported an error") -def validate_user_config(config: Mapping[str, Any]) -> None: +def validate_user_config(config: Mapping[str, object]) -> None: """Reject OpenCodeOptions.config keys LiteLLM manages (or that bypass permissions).""" for key in config: if key in MANAGED_CONFIG_KEYS: @@ -241,7 +241,7 @@ def build_opencode_config( user_config: Mapping[str, Any] | None = None, instructions_path: str | None = None, skills_path: str | None = None, -) -> Mapping[str, Any]: +) -> Mapping[str, object]: """The full opencode config: user config underneath, LiteLLM-managed keys on top.""" user: Final = user_config or MappingProxyType({}) validate_user_config(user) @@ -379,7 +379,7 @@ class OpenCodeHarnessConfig(BaseCLIHarnessConfig): def create_stream_state(self) -> OpenCodeStreamState: return OpenCodeStreamState() - def transform_stream_line(self, line: Mapping[str, Any], state: OpenCodeStreamState) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: OpenCodeStreamState) -> Sequence[Event]: """step_finish token counts are ignored on purpose: the session endpoint accounts usage.""" session_id = line.get("sessionID") if session_id and state.session_id is None: diff --git a/litellm/llms/openrouter/image_generation/transformation.py b/litellm/llms/openrouter/image_generation/transformation.py index 67d90d027ec..cfca853e280 100644 --- a/litellm/llms/openrouter/image_generation/transformation.py +++ b/litellm/llms/openrouter/image_generation/transformation.py @@ -27,9 +27,11 @@ Response format: } """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -55,6 +57,11 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): """ @@ -355,15 +362,16 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): model_response.data = [] try: - choices: Final = response_json.get("choices", []) + response_object: Final = _JSON_DICT.validate_python(response_json) + choices: Final = _JSON_OBJECTS.validate_python(response_object.get("choices", [])) for choice in choices: - message = choice.get("message", {}) - images = message.get("images", []) + message = _JSON_OBJECT.validate_python(choice.get("message", {})) + images = _JSON_OBJECTS.validate_python(message.get("images", [])) for image_data in images: - image_url_obj = image_data.get("image_url", {}) - image_url = image_url_obj.get("url") + image_url_obj = _JSON_OBJECT.validate_python(image_data.get("image_url", {})) + image_url = _STR.validate_python(image_url_obj.get("url") or "") if image_url: if image_url.startswith("data:"): @@ -389,7 +397,7 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): ) # Extract and set usage and cost information - self._set_usage_and_cost(model_response, response_json, model) + self._set_usage_and_cost(model_response, response_object, model) return model_response diff --git a/litellm/llms/ovhcloud/audio_transcription/transformation.py b/litellm/llms/ovhcloud/audio_transcription/transformation.py index 6fc56ebb61f..086c1afbcd4 100644 --- a/litellm/llms/ovhcloud/audio_transcription/transformation.py +++ b/litellm/llms/ovhcloud/audio_transcription/transformation.py @@ -5,9 +5,11 @@ Our unified API follows the OpenAI standard. More information on our website: https://endpoints.ai.cloud.ovh.net """ +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( @@ -24,6 +26,8 @@ from litellm.types.utils import FileTypes, TranscriptionResponse from ..utils import OVHCloudException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: @@ -145,7 +149,8 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): headers=raw_response.headers, ) - text: Final = response_json.get("text") or response_json.get("transcript") or "" + payload: Final = _JSON_OBJECT.validate_python(response_json) + text: Final = payload.get("text") or payload.get("transcript") or "" response: Final = TranscriptionResponse(text=text) # OVHCloud field migration (deadline: 2026-05-11): @@ -153,9 +158,7 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): # Prefer `seconds`, fall back to `duration`, normalize to `duration` # so downstream consumers see a consistent key. duration: Final = ( - response_json["seconds"] - if "seconds" in response_json and response_json["seconds"] is not None - else response_json.get("duration") + payload["seconds"] if "seconds" in payload and payload["seconds"] is not None else payload.get("duration") ) if duration is not None: response_json["duration"] = duration diff --git a/litellm/llms/recraft/image_generation/transformation.py b/litellm/llms/recraft/image_generation/transformation.py index f65bf1e7292..0fcfafb79fb 100644 --- a/litellm/llms/recraft/image_generation/transformation.py +++ b/litellm/llms/recraft/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -21,6 +23,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class RecraftImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://external.api.recraft.ai" @@ -141,7 +146,8 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): if not model_response.data: model_response.data = [] - for image_data in response_data["data"]: + payload: Final = _JSON_OBJECT.validate_python(response_data) + for image_data in _JSON_OBJECTS.validate_python(payload["data"]): model_response.data.append( ImageObject( url=image_data.get("url", None), diff --git a/litellm/llms/replicate/chat/transformation.py b/litellm/llms/replicate/chat/transformation.py index f7e09b7bec0..e8dd8816997 100644 --- a/litellm/llms/replicate/chat/transformation.py +++ b/litellm/llms/replicate/chat/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH @@ -26,6 +28,8 @@ if TYPE_CHECKING: else: LoggingClass = Any +_TEXTS: Final = TypeAdapter(Iterable[str], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ReplicateConfig(BaseConfig): """ @@ -253,7 +257,7 @@ class ReplicateConfig(BaseConfig): message=f"LiteLLM Error - prediction not succeeded - {raw_response_json}", headers=raw_response.headers, ) - outputs: Final = raw_response_json.get("output", []) + outputs: Final = _TEXTS.validate_python(raw_response_json.get("output", [])) response_str = "".join(outputs) if len(response_str) == 0: # edge case, where result from replicate is empty response_str = " " diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index 8e8f7ea61aa..0d86752fc5f 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -1,9 +1,10 @@ import functools import json -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm import verbose_logger @@ -11,6 +12,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.utils import GenericStreamingChunk as GChunk from litellm.types.utils import StreamingChatCompletionChunk +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + def _load_sagemaker_response_stream_shape(): try: @@ -18,7 +21,7 @@ def _load_sagemaker_response_stream_shape(): from botocore.model import ServiceModel loader: Final = Loader() - service_dict: Final = loader.load_service_model("sagemaker-runtime", "service-2") + service_dict: Final = _JSON_OBJECT.validate_python(loader.load_service_model("sagemaker-runtime", "service-2")) return ServiceModel(service_dict).shape_for("InvokeEndpointWithResponseStreamOutput") except Exception as e: verbose_logger.warning( diff --git a/litellm/llms/snowflake/embedding/transformation.py b/litellm/llms/snowflake/embedding/transformation.py index 75f35a8f379..6aa66de1db5 100644 --- a/litellm/llms/snowflake/embedding/transformation.py +++ b/litellm/llms/snowflake/embedding/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +12,8 @@ from litellm.types.utils import EmbeddingResponse from ..utils import SnowflakeBaseConfig, SnowflakeException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): """ @@ -53,7 +57,7 @@ class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): # convert embeddings to 1d array for item in response_json["data"]: item["embedding"] = item["embedding"][0] - returned_response: Final = EmbeddingResponse(**response_json) + returned_response: Final = EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) returned_response.model = "snowflake/" + (returned_response.model or "") diff --git a/litellm/llms/stability/image_generation/transformation.py b/litellm/llms/stability/image_generation/transformation.py index 656ffe395c8..1856523ce49 100644 --- a/litellm/llms/stability/image_generation/transformation.py +++ b/litellm/llms/stability/image_generation/transformation.py @@ -6,9 +6,11 @@ Handles transformation between OpenAI-compatible format and Stability AI API for API Reference: https://platform.stability.ai/docs/api-reference """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -33,6 +35,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class StabilityImageGenerationConfig(BaseImageGenerationConfig): """ @@ -234,7 +238,8 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): ) # Check finish_reason - finish_reason: Final = response_data.get("finish_reason", "") + payload: Final = _JSON_OBJECT.validate_python(response_data) + finish_reason: Final = payload.get("finish_reason", "") if finish_reason == "CONTENT_FILTERED": raise self.get_error_class( error_message="Content was filtered by Stability AI safety systems", @@ -246,7 +251,7 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): model_response.data = [] # Extract image from response - image_b64: Final = response_data.get("image") + image_b64: Final = payload.get("image") if image_b64: model_response.data.append( ImageObject( diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index d5ed7da3815..c7976456d87 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -26,6 +26,7 @@ from litellm.secret_managers.main import get_secret_str _UrlEncodableParams: Final = TypeAdapter(dict[str, str | int | float | bool]) _StrList: Final = TypeAdapter(list[str]) _StrFrozenSet: Final = TypeAdapter(frozenset[str]) +_DecodedJson: Final = TypeAdapter(object) _TINYFISH_PARAMS_KEY: Final = "_tinyfish_params" _TINYFISH_DOCS_URL: Final = "https://docs.tinyfish.ai/search-api" @@ -218,7 +219,7 @@ class TinyfishSearchConfig(BaseSearchConfig): ) try: - raw_json: Final[object] = raw_response.json() # any-ok: httpx Response.json() -> Any + raw_json: Final = _DecodedJson.validate_python(raw_response.json()) except json.JSONDecodeError: raise self._wrap_error( error_message=f"Expected JSON response, got: {raw_response.text[:200]}", @@ -276,7 +277,7 @@ class TinyfishSearchConfig(BaseSearchConfig): # for other envelope shapes (CDN HTML pages, other JSON envelopes, plain text). inner_message = error_message try: - body: Final[object] = json.loads(error_message) # any-ok: json.loads -> Any + body: Final = _DecodedJson.validate_python(json.loads(error_message)) if isinstance(body, dict): error_obj: Final[object] = body.get("error") # any-ok: untyped dict if isinstance(error_obj, dict): diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index b8079e52c97..2fff0854f10 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -4,7 +4,9 @@ Re rank api LiteLLM supports the re rank API format, no paramter transformation occurs """ -from typing import Any, Final +from typing import Final + +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base import BaseLLM @@ -15,6 +17,8 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.llms.together_ai.rerank.transformation import TogetherAIRerankConfig from litellm.types.rerank import RerankRequest, RerankResponse +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + def _rerank_url(api_base: str) -> str: return f"{api_base.rstrip('/')}/rerank" @@ -27,7 +31,7 @@ class TogetherAIRerank(BaseLLM): api_key: str, api_base: str, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], top_n: int | None = None, rank_fields: list[str] | None = None, return_documents: bool | None = True, @@ -66,13 +70,13 @@ class TogetherAIRerank(BaseLLM): if response.status_code != 200: raise Exception(response.text) - _json_response: Final = response.json() + _json_response: Final = _JSON_DICT.validate_python(response.json()) return TogetherAIRerankConfig()._transform_response(_json_response) async def async_rerank( # New async method self, - request_data_dict: dict[str, Any], + request_data_dict: dict[str, object], api_key: str, api_base: str, ) -> RerankResponse: @@ -91,6 +95,6 @@ class TogetherAIRerank(BaseLLM): if response.status_code != 200: raise Exception(response.text) - _json_response: Final = response.json() + _json_response: Final = _JSON_DICT.validate_python(response.json()) return TogetherAIRerankConfig()._transform_response(_json_response) diff --git a/litellm/llms/vertex_ai/image_generation/image_generation_handler.py b/litellm/llms/vertex_ai/image_generation/image_generation_handler.py index 6a5bb484540..7ba7b060991 100644 --- a/litellm/llms/vertex_ai/image_generation/image_generation_handler.py +++ b/litellm/llms/vertex_ai/image_generation/image_generation_handler.py @@ -1,8 +1,10 @@ import json +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx from openai.types.image import Image +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.custom_httpx.http_handler import ( @@ -17,11 +19,13 @@ from litellm.types.utils import ImageResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +_PREDICTIONS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class VertexImageGeneration(VertexLLM): def process_image_generation_response( self, - json_response: dict[str, Any], + json_response: Mapping[str, object], model_response: ImageResponse, model: str | None = None, ) -> ImageResponse: @@ -32,12 +36,11 @@ class VertexImageGeneration(VertexLLM): model=model, ) - predictions: Final = json_response["predictions"] + predictions: Final = _PREDICTIONS.validate_python(json_response["predictions"]) response_data: Final[list[Image]] = [] for prediction in predictions: - bytes_base64_encoded = prediction["bytesBase64Encoded"] - image_object = Image(b64_json=bytes_base64_encoded) + image_object = Image.model_validate({"b64_json": prediction["bytesBase64Encoded"]}) response_data.append(image_object) model_response.data = response_data diff --git a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py index 17ccf16837b..9f3126fd4ff 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py @@ -1,7 +1,9 @@ import os +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -31,6 +33,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): """ @@ -321,10 +326,11 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): ) ) - if usage_metadata := response_data.get("usageMetadata", None): - model_response.usage = self._transform_image_usage(usage_metadata) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if usage_metadata := response_object.get("usageMetadata", None): + model_response.usage = self._transform_image_usage(_JSON_DICT.validate_python(usage_metadata)) - web_search_requests: Final = get_gemini_image_web_search_requests(response_data) + web_search_requests: Final = get_gemini_image_web_search_requests(response_object) if web_search_requests and model_response.usage is not None: setattr(model_response.usage, "web_search_requests", web_search_requests) diff --git a/litellm/llms/vertex_ai/text_to_speech/transformation.py b/litellm/llms/vertex_ai/text_to_speech/transformation.py index a7b079fb89c..a78145fdcf4 100644 --- a/litellm/llms/vertex_ai/text_to_speech/transformation.py +++ b/litellm/llms/vertex_ai/text_to_speech/transformation.py @@ -6,11 +6,12 @@ Reference: https://cloud.google.com/text-to-speech/docs/reference/rest/v1/text/s """ import base64 -from collections.abc import Coroutine +from collections.abc import Coroutine, Mapping from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypeAlias, Union import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.exceptions import UnsupportedParamsError @@ -45,6 +46,9 @@ else: _LyriaVoice: TypeAlias = str | dict | None +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(hide_input_in_errors=True)) + class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): """ @@ -465,14 +469,14 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): from litellm.types.llms.openai import HttpxBinaryResponseContent # Parse JSON response - _json_response: Final = raw_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(raw_response.json()) # Get base64-encoded audio content response_content: Final = _json_response.get("audioContent") if not response_content: raise ValueError("No audioContent in Vertex AI TTS response") - binary_data: Final = base64.b64decode(response_content) + binary_data: Final = base64.b64decode(_STR.validate_python(response_content)) media_type: Final = speech_media_type_from_audio_bytes(binary_data) response: Final = httpx.Response( status_code=200, diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index b48efe229b3..3acb2f2ed58 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -4,10 +4,11 @@ Transformation logic for Voyage AI's /v1/rerank endpoint. Docs - https://docs.voyageai.com/docs/reranker """ -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj @@ -23,6 +24,11 @@ from litellm.types.utils import ModelInfo from ..embedding.transformation import VoyageError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class VoyageRerankConfig(BaseRerankConfig): def get_supported_cohere_rerank_params(self, model: str) -> list: @@ -103,13 +109,14 @@ class VoyageRerankConfig(BaseRerankConfig): ) # Voyage AI returns results in "data" key, not "results" - _results: Final[list[dict] | None] = _json_response.get("data") + payload: Final = _JSON_OBJECT.validate_python(_json_response) + _results: Final = payload.get("data") if _results is None: raise ValueError(f"No results found in the response={_json_response}") # Transform to LiteLLM format transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result: dict[str, object] = { "index": result["index"], "relevance_score": result["relevance_score"], @@ -121,14 +128,14 @@ class VoyageRerankConfig(BaseRerankConfig): transformed_result["document"] = result["document"] transformed_results.append(transformed_result) - usage: Final = _json_response.get("usage", {}) - total_tokens: Final = usage.get("total_tokens", 0) + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage", {})) + total_tokens: Final = _OPTIONAL_INT.validate_python(usage.get("total_tokens", 0)) _billed_units: Final = RerankBilledUnits(total_tokens=total_tokens) _tokens: Final = RerankTokens(input_tokens=total_tokens, output_tokens=0) rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) return RerankResponse( - id=_json_response.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=transformed_results, meta=rerank_meta, ) diff --git a/litellm/llms/watsonx/audio_transcription/transformation.py b/litellm/llms/watsonx/audio_transcription/transformation.py index 2169b9bf49a..77ef03b8e10 100644 --- a/litellm/llms/watsonx/audio_transcription/transformation.py +++ b/litellm/llms/watsonx/audio_transcription/transformation.py @@ -4,9 +4,11 @@ Translates from OpenAI's `/v1/audio/transcriptions` to IBM WatsonX's `/ml/v1/aud WatsonX follows the OpenAI spec for audio transcription. """ +from collections.abc import Mapping from typing import Final from httpx import Response +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -25,6 +27,8 @@ from ...openai.transcriptions.whisper_transformation import ( ) from ..common_utils import IBMWatsonXMixin +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTranscriptionConfig): """ @@ -174,8 +178,9 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran # Extract only valid fields for TranscriptionResponse.__init__() # TranscriptionResponse only accepts 'text' and 'usage' in __init__() - text: Final = raw_response_json.get("text") - usage: Final = raw_response_json.get("usage") + response_object: Final = _JSON_OBJECT.validate_python(raw_response_json) + text: Final = response_object.get("text") + usage: Final = response_object.get("usage") # Create response with only valid fields response_kwargs: Final = {} @@ -187,14 +192,14 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran if not response_kwargs: raise ValueError( "Invalid response format. Received response does not match the expected format. Got: ", - raw_response_json, + response_object, ) response: Final = TranscriptionResponse(**response_kwargs) # Add other fields using dictionary-style assignment (like duration, task, etc.) # Skip fields that TranscriptionResponse doesn't accept in __init__() - for key, value in raw_response_json.items(): + for key, value in response_object.items(): if key not in [ "text", "usage", diff --git a/litellm/llms/watsonx/embed/transformation.py b/litellm/llms/watsonx/embed/transformation.py index 80cd28a058d..609eeba90c0 100644 --- a/litellm/llms/watsonx/embed/transformation.py +++ b/litellm/llms/watsonx/embed/transformation.py @@ -2,9 +2,11 @@ Translates from OpenAI's `/v1/embeddings` to IBM's `/text/embeddings` route. """ +from collections.abc import Iterable, Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.embedding.transformation import ( BaseEmbeddingConfig, @@ -16,6 +18,9 @@ from litellm.types.utils import EmbeddingResponse, Usage from ..common_utils import IBMWatsonXMixin, _get_api_params +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_TOKEN_COUNT: Final = TypeAdapter(int) + class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): def get_supported_openai_params(self, model: str) -> list: @@ -95,7 +100,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): json_resp: Final = raw_response.json() if model_response is None: model_response = EmbeddingResponse(model=json_resp.get("model_id", None)) - results: Final = json_resp.get("results", []) + results: Final = _JSON_OBJECTS.validate_python(json_resp.get("results", [])) embedding_response: Final = [] for idx, result in enumerate(results): embedding_response.append( @@ -107,7 +112,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): ) model_response.object = "list" model_response.data = embedding_response - input_tokens: Final = json_resp.get("input_token_count", 0) + input_tokens: Final = _TOKEN_COUNT.validate_python(json_resp.get("input_token_count", 0) or 0) setattr( model_response, "usage", diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index ff95f14951a..c2cfd6597f3 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -5,10 +5,11 @@ Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank """ import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig @@ -24,6 +25,11 @@ from litellm.types.rerank import ( from ..common_utils import IBMWatsonXMixin, _generate_watsonx_token, _get_api_params +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): """ @@ -171,13 +177,14 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): headers=raw_response.headers, ) - _results: Final[list[dict] | None] = raw_response_json.get("results") + payload: Final = _JSON_OBJECT.validate_python(raw_response_json) + _results: Final = payload.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result: dict[str, object] = { "index": result["index"], "relevance_score": result["score"], @@ -191,11 +198,11 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): transformed_results.append(transformed_result) - response_id: Final = raw_response_json.get("id") or str(uuid.uuid4()) + response_id: Final = _STR.validate_python(payload.get("id") or str(uuid.uuid4())) # Extract usage information _tokens: Final = RerankTokens( - input_tokens=raw_response_json.get("input_token_count", 0), + input_tokens=_OPTIONAL_INT.validate_python(payload.get("input_token_count", 0)), ) rerank_meta: Final = RerankResponseMeta(tokens=_tokens) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index b85ff2ce50e..a2aeef3f67f 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -8,7 +8,7 @@ from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast from fastapi import HTTPException -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly from litellm._logging import verbose_proxy_logger @@ -162,6 +162,8 @@ _CLIENT_FORWARDED_AUTH_TYPES: Final["frozenset[str]"] = frozenset({"true_passthr # Minted token material that must never survive a client rotation on a persisted row. _MINTED_TOKEN_CREDENTIAL_FIELDS: Final["frozenset[str]"] = frozenset({"access_token", "refresh_token", "expires_in"}) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + def _bind_submitted_oauth_client( credentials: Mapping[str, object], issuer: str | None, url: str | None @@ -1916,7 +1918,7 @@ def mcp_oauth_token_identity(server: object) -> tuple[object, ...]: creds: Final = getattr(server, "credentials", None) if isinstance(creds, str): try: - parsed: dict[str, object] | None = json.loads(creds) + parsed: Mapping[str, object] | None = _JSON_OBJECT.validate_python(json.loads(creds)) except ValueError: parsed = None else: @@ -2382,13 +2384,10 @@ def _decode_user_env_vars(stored: str) -> dict[str, str]: "re-enter them rather than silently forwarding ciphertext" ) return {} - parsed: dict[str, object] | None try: - parsed = json.loads(decrypted) + parsed: Final = _JSON_OBJECT.validate_python(json.loads(decrypted)) except (ValueError, TypeError): return {} - if not isinstance(parsed, dict): - return {} return {str(k): str(v) for k, v in parsed.items()} diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index c88c6f2570a..dca48babc81 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -19,7 +19,7 @@ from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse, StreamingResponse -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm._logging import verbose_proxy_logger @@ -78,6 +78,9 @@ _PASCAL_TO_WIRE: Final[Mapping[str, str]] = { } +_DECODED_JSON: Final = TypeAdapter(object) + + def _sse_event(payload: object) -> str: """Frame a JSON-RPC object as a single A2A SSE event (``data: \\n\\n``).""" return f"data: {json.dumps(payload)}\n\n" @@ -91,7 +94,7 @@ def _to_jsonrpc_object(chunk: object) -> object: """ if isinstance(chunk, (str, bytes, bytearray)): try: - return json.loads(chunk) + return _DECODED_JSON.validate_python(json.loads(chunk)) except (json.JSONDecodeError, UnicodeDecodeError): return chunk if hasattr(chunk, "model_dump"): diff --git a/litellm/proxy/client/chat.py b/litellm/proxy/client/chat.py index a330d057490..cf48bee6958 100644 --- a/litellm/proxy/client/chat.py +++ b/litellm/proxy/client/chat.py @@ -3,9 +3,14 @@ from collections.abc import Iterator from typing import Any, Final import requests +from pydantic import ConfigDict, TypeAdapter from .exceptions import UnauthorizedError +_SSE_LINE: Final[TypeAdapter[bytes | bytearray]] = TypeAdapter( + bytes | bytearray, config=ConfigDict(arbitrary_types_allowed=True, strict=True, hide_input_in_errors=True) +) + class ChatClient: def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 600): @@ -172,7 +177,7 @@ class ChatClient: # Parse SSE stream for line in response.iter_lines(): if line: - line = line.decode("utf-8") + line = _SSE_LINE.validate_python(line).decode("utf-8") if line.startswith("data: "): data_str = line[6:] # Remove 'data: ' prefix if data_str.strip() == "[DONE]": diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 98af32fa7aa..684449acd0c 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -8,6 +8,7 @@ from urllib.parse import urlencode import click import requests +from pydantic import ConfigDict, TypeAdapter from rich.console import Console from rich.table import Table from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never @@ -121,6 +122,8 @@ class CliAuthResult(TypedDict): _TeamMapping: Final = TypeVar("_TeamMapping", bound=Mapping[str, object]) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + KEYRING_INSTALL_HINT: Final = "pip install 'litellm[cli]'" KEYRING_ENABLE_HINT: Final = "keyring --enable (or unset PYTHON_KEYRING_BACKEND)" @@ -485,7 +488,7 @@ def prompt_team_selection_fallback( def _response_error_detail(response: requests.Response) -> str | None: try: - body: Final[dict[str, object] | list[object] | str | int | float | bool | None] = response.json() + body: Final = _JSON_OBJECT.validate_python(response.json()) except ValueError: return None detail: Final = body.get("detail") if isinstance(body, dict) else None diff --git a/litellm/proxy/client/cli/commands/debug.py b/litellm/proxy/client/cli/commands/debug.py index 4e3914143d8..2aca60ca810 100644 --- a/litellm/proxy/client/cli/commands/debug.py +++ b/litellm/proxy/client/cli/commands/debug.py @@ -115,6 +115,7 @@ class RequestResponsePayload(BaseModel): _SESSION_PAGE: Final = TypeAdapter(SessionLogsPage) _PAYLOAD: Final[TypeAdapter[RequestResponsePayload | None]] = TypeAdapter(RequestResponsePayload | None) _JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +_BACKTICK_RUNS: Final = TypeAdapter(tuple[str, ...]) _SESSION_PAGE_SIZE: Final = 100 _TRANSPORT_BODY_CHARS: Final = 500 @@ -192,7 +193,7 @@ def _fmt_json(value: JsonValue, max_chars: int) -> str: def _fenced(text: str, info: str = "") -> tuple[str, str, str]: - longest_run: Final = max((len(run) for run in re.findall(r"`+", text)), default=0) + longest_run: Final = max((len(run) for run in _BACKTICK_RUNS.validate_python(re.findall(r"`+", text))), default=0) fence: Final = "`" * max(3, longest_run + 1) return (f"{fence}{info}", text, fence) diff --git a/litellm/proxy/db/model_insights_tasks.py b/litellm/proxy/db/model_insights_tasks.py index 865965dcf75..e117afb881e 100644 --- a/litellm/proxy/db/model_insights_tasks.py +++ b/litellm/proxy/db/model_insights_tasks.py @@ -1,14 +1,20 @@ import json +from collections.abc import Mapping from functools import lru_cache from pathlib import Path from typing import Final +from pydantic import ConfigDict, TypeAdapter + from litellm.types.model_insights import ModelInsightTask _TASKS_FILE: Final = Path(__file__).resolve().parent.parent / "model_insights_tasks.json" +_TASK_ENTRIES: Final = TypeAdapter( + Mapping[str, Mapping[str, object]], config=ConfigDict(strict=True, hide_input_in_errors=True) +) @lru_cache(maxsize=1) def load_model_insight_tasks() -> dict[str, ModelInsightTask]: - raw: Final = json.loads(_TASKS_FILE.read_text()) - return {name: ModelInsightTask(task_type=name, **entry) for name, entry in raw.items()} + raw: Final = _TASK_ENTRIES.validate_python(json.loads(_TASKS_FILE.read_text())) + return {name: ModelInsightTask.model_validate(dict(task_type=name, **entry)) for name, entry in raw.items()} diff --git a/litellm/proxy/fine_tuning_endpoints/endpoints.py b/litellm/proxy/fine_tuning_endpoints/endpoints.py index a13ad00713d..e09f1ec8ba8 100644 --- a/litellm/proxy/fine_tuning_endpoints/endpoints.py +++ b/litellm/proxy/fine_tuning_endpoints/endpoints.py @@ -426,7 +426,7 @@ async def list_fine_tuning_jobs( route_type=CallTypes.alist_fine_tuning_jobs.value, ) - response: Any | None = None + response: object = None if target_model_names and isinstance(target_model_names, str): target_model_names_list: Final = target_model_names.split(",") if len(target_model_names_list) != 1: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 00f0ef1756f..2a0b528a50b 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -8,7 +8,7 @@ import time import traceback from collections.abc import Iterable, Mapping from datetime import datetime, timedelta, timezone -from typing import Any, Final, Literal, TypedDict, cast +from typing import Final, Literal, TypedDict import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response, status @@ -1431,7 +1431,7 @@ async def shared_health_check_status_endpoint( ) -def _read_license_data() -> dict[str, Any] | None: +def _read_license_data() -> EnterpriseLicenseData | None: from litellm.proxy.proxy_server import _license_check, premium_user_data license_data: EnterpriseLicenseData | None = premium_user_data or _license_check.airgapped_license_data @@ -1453,10 +1453,10 @@ def _read_license_data() -> dict[str, Any] | None: if license_data is None: return None - return cast(dict[str, Any], license_data) + return license_data -def _read_allowed_features(license_data: dict[str, Any]) -> list: +def _read_allowed_features(license_data: Mapping[str, object]) -> list: raw_allowed_features: Final = license_data.get("allowed_features") if isinstance(raw_allowed_features, list): return list(raw_allowed_features) @@ -1707,7 +1707,7 @@ def _show_env_credential_login_warning() -> bool: async def _get_health_readiness_details( response: Response | None = None, -) -> dict[str, Any]: +) -> dict[str, object]: """ Detailed health payload for authenticated diagnostics. """ @@ -1726,7 +1726,7 @@ async def _get_health_readiness_details( success_callback_names = litellm.success_callback # check Cache - cache_type: Any = None + cache_type: object = None if litellm.cache is not None: from litellm.caching.caching import RedisSemanticCache @@ -1735,7 +1735,7 @@ async def _get_health_readiness_details( if isinstance(litellm.cache.cache, RedisSemanticCache): # ping the cache # TODO: @ishaan-jaff - we should probably not ping the cache on every /health/readiness check - index_info: Any + index_info: object try: index_info = await litellm.cache.cache._index_info() except Exception as e: diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 9806a824c63..5db0fce7d92 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -54,6 +54,8 @@ from litellm.repositories.autorouter_session_repository import AutoRouterSession from litellm.repositories.base_repository import SupportsModelDump from litellm.repositories.daily_activity_sql import build_where_clause from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository from litellm.router_strategy.complexity_router import ComplexityRouter from litellm.router_utils.auto_router_model_naming import ( StrategyRouterDependencyRole, @@ -187,15 +189,15 @@ def _team_table(prisma_client: "PrismaClient") -> _TeamTable: def _verification_tokens(prisma_client: "PrismaClient") -> _VerificationTokenTable: - return prisma_client.db.litellm_verificationtoken + return VerificationTokenRepository(prisma_client).table def _team_rows(prisma_client: "PrismaClient") -> _TeamRowsTable: - return prisma_client.db.litellm_teamtable + return TeamRepository(prisma_client).table def _user_rows(prisma_client: "PrismaClient") -> _UserRowsTable: - return prisma_client.db.litellm_usertable + return UserRepository(prisma_client).table def _shadow_eval_jobs(prisma_client: "PrismaClient") -> _ShadowEvalJobTable: diff --git a/litellm/proxy/management_endpoints/callback_management_endpoints.py b/litellm/proxy/management_endpoints/callback_management_endpoints.py index 4f9f46af08f..d910b475b57 100644 --- a/litellm/proxy/management_endpoints/callback_management_endpoints.py +++ b/litellm/proxy/management_endpoints/callback_management_endpoints.py @@ -7,12 +7,15 @@ import os from typing import Final from fastapi import APIRouter, Depends +from pydantic import TypeAdapter from litellm.litellm_core_utils.logging_callback_manager import CallbacksByType from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router: Final = APIRouter() +_DECODED_JSON: Final = TypeAdapter(object) + @router.get( "/callbacks/list", @@ -51,6 +54,6 @@ async def get_callback_configs(): ) with open(config_path, "r") as f: - configs: Final = json.load(f) + configs: Final = _DECODED_JSON.validate_python(json.load(f)) return configs diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 323b9e434a8..037b9082cd4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -355,9 +355,12 @@ _KEY_METADATA_REQUEST_FIELDS: Final = frozenset( ) +_DECODED_JSON: Final = TypeAdapter(object) + + def _decode_json_string_column(column: str, value: object) -> object: if column in _KEY_UPDATE_JSON_STRING_COLUMNS and isinstance(value, str): - return json.loads(value) + return _DECODED_JSON.validate_python(json.loads(value)) return value diff --git a/litellm/proxy/management_endpoints/management_v1/budgets.py b/litellm/proxy/management_endpoints/management_v1/budgets.py index 106dbfaf7b7..f1760b857c5 100644 --- a/litellm/proxy/management_endpoints/management_v1/budgets.py +++ b/litellm/proxy/management_endpoints/management_v1/budgets.py @@ -85,8 +85,7 @@ class PrismaBudgetListExecutor: async def count(self, where: tuple[Predicate, ...]) -> int: clauses, params = where_sql(where) sql: Final = f"SELECT COUNT(*) AS count FROM {BUDGET_TABLE}" + (f" WHERE {clauses}" if clauses else "") - rows: Final = await self.prisma_client.db.query_raw(sql, *params) - counted: Final = _ROW_COUNTS.validate_python(rows) + counted: Final = _ROW_COUNTS.validate_python(await self.prisma_client.db.query_raw(sql, *params)) return counted[0].count if counted else 0 async def find_many(self, plan: QueryPlan) -> Sequence[BudgetListItem]: @@ -97,8 +96,7 @@ class PrismaBudgetListExecutor: + f" ORDER BY {order_by_sql(plan.order)}" + f" LIMIT ${len(params) + 1} OFFSET ${len(params) + 2}" ) - rows: Final = await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip) - return _BUDGET_ROWS.validate_python(rows) + return _BUDGET_ROWS.validate_python(await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip)) def _serialize(row: BudgetListItem) -> BudgetListItem: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 2aed0e2afb8..e5bea71ac7c 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -533,6 +533,8 @@ if MCP_AVAILABLE: except Exception as e: verbose_proxy_logger.debug("Failed to write temporary MCP server to Redis cache: %s", e) + _CACHED_VALUE: Final = TypeAdapter(object) + @with_service_target(MCP_SERVERS_TARGET) async def _get_temporary_mcp_server_from_redis( server_id: str, @@ -550,8 +552,8 @@ if MCP_AVAILABLE: return None try: - cached_server: Final = await cache_backend.async_get_cache( - key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}" + cached_server: Final = _CACHED_VALUE.validate_python( + await cache_backend.async_get_cache(key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}") ) except Exception as e: verbose_proxy_logger.debug("Failed reading temporary MCP server from Redis cache: %s", e) diff --git a/litellm/proxy/management_endpoints/model_insights_endpoints.py b/litellm/proxy/management_endpoints/model_insights_endpoints.py index b787aac2d9f..53e8dfcc71e 100644 --- a/litellm/proxy/management_endpoints/model_insights_endpoints.py +++ b/litellm/proxy/management_endpoints/model_insights_endpoints.py @@ -115,7 +115,7 @@ def _deployment_filter(rows: list[_GroupedModel]) -> list[dict[str, str]]: def _daily_metric(row: _GroupedDaily) -> ModelInsightDailyMetric: - return ModelInsightDailyMetric(date=row.date, **_metric(row).model_dump()) + return ModelInsightDailyMetric.model_validate({**_metric(row).model_dump(), "date": row.date}) def _daily_total(row: _GroupedDate) -> ModelInsightDailyTotal: @@ -146,12 +146,14 @@ def _summarize_tasks(rows: list[_GroupedTask], metric: ModelInsightsMetric) -> l } grand: Final = sum(totals.values()) return [ - ModelInsightTaskSummary( - **(catalog.get(task) or _UNCATEGORIZED_TASK).model_copy(update={"task_type": task}).model_dump(), - value=value, - share=value / grand * 100 if grand else 0.0, - leader=leaders[task].model_group, - provider=leaders[task].custom_llm_provider, + ModelInsightTaskSummary.model_validate( + { + **(catalog.get(task) or _UNCATEGORIZED_TASK).model_copy(update={"task_type": task}).model_dump(), + "value": value, + "share": value / grand * 100 if grand else 0.0, + "leader": leaders[task].model_group, + "provider": leaders[task].custom_llm_provider, + } ) for task, value in sorted(totals.items(), key=lambda item: item[1], reverse=True) ] diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 11a0075ffd0..253013a2ebf 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -2547,8 +2547,8 @@ async def add_new_model( enforced=bool(general_settings.get(ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING, False)), ) - clean_model_info: Final = ModelInfo( - **without_server_derived_pricing(model_params.model_info.model_dump(exclude_none=True)) + clean_model_info: Final = ModelInfo.model_validate( + dict(without_server_derived_pricing(model_params.model_info.model_dump(exclude_none=True))) ) model_params.model_info = ( # rebind-ok: downstream team-model handling mutates this same object clean_model_info.model_copy(update=MappingProxyType({"member_auto_router": True})) @@ -3194,6 +3194,9 @@ def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: return unique_models +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + + def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: """A DB row's model_info column arrives as a dict or as its JSON string depending on the query path, and every consumer needs the mapping. Single owner of that parse: @@ -3204,10 +3207,9 @@ def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: if not isinstance(model_info, str): return None try: - parsed: Final = json.loads(model_info) + return _JSON_OBJECT.validate_python(json.loads(model_info)) except (TypeError, ValueError): return None - return parsed if isinstance(parsed, Mapping) else None def _expects_liveness_on_this_pod(model_info: object) -> bool: diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index c8ae7af41db..26abe35b49a 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -27,7 +27,7 @@ from typing import ( import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, status -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger @@ -306,7 +306,7 @@ async def _verify_org_access( ) -_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _BUDGET_SETTABLE_FIELDS: Final = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"} _ORG_COLUMN_FIELDS: Final = frozenset({"organization_alias", "models"}) _ORG_METADATA_FIELDS: Final = tuple( @@ -720,7 +720,7 @@ async def update_organization( ) # Transform UI payload to expected format - raw_data: Final[dict[str, object]] = await request.json() + raw_data: Final = _STR_OBJECT_DICT_ADAPTER.validate_python(await request.json()) raw_data_with_flat_budget_fields: Final = handle_nested_budget_structure_in_organization_update_request(raw_data) # Create validated data model @@ -767,7 +767,9 @@ async def update_organization( # Merge metadata from existing organization with updated metadata if updated_organization_row_json.get("metadata") is not None: existing_metadata: Final = existing_organization_row.metadata or {} - updated_metadata: Final[dict[str, object]] = updated_organization_row_json.get("metadata", {}) + updated_metadata: Final = _STR_OBJECT_DICT_ADAPTER.validate_python( + updated_organization_row_json.get("metadata", {}) + ) merged_metadata: Final[Mapping[str, object]] = _update_dictionary( existing_dict=cast( # cast-ok: prisma de-serializes a Json column to the plain python dict it stores "dict[str, object]", existing_metadata @@ -786,12 +788,16 @@ async def update_organization( ) budget_fields: Final = { - k: v for k, v in data.model_dump().items() if k in _BUDGET_SETTABLE_FIELDS and k in data.model_fields_set + k: v + for k, v in _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump()).items() + if k in _BUDGET_SETTABLE_FIELDS and k in data.model_fields_set } if budget_fields and existing_organization_row.budget_id: await update_budget( - budget_obj=BudgetNewRequest(budget_id=existing_organization_row.budget_id, **budget_fields), + budget_obj=BudgetNewRequest.model_validate( + {"budget_id": existing_organization_row.budget_id, **budget_fields} + ), user_api_key_dict=user_api_key_dict, ) diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index 69356922ea1..c3c499eeec1 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -12,12 +12,12 @@ All /policy management endpoints import copy import json import os -from collections.abc import AsyncGenerator, AsyncIterator +from collections.abc import AsyncGenerator, AsyncIterator, Mapping from typing import TYPE_CHECKING, Final, Literal, cast from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict import litellm @@ -449,6 +449,23 @@ async def validate_policy( return result +class _LoadedPolicy(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + inherit: str | None = None + scope: PolicyScopeResponse = Field(default_factory=PolicyScopeResponse) + guardrails: PolicyGuardrailsResponse = Field(default_factory=PolicyGuardrailsResponse) + resolved_guardrails: tuple[str, ...] = () + inheritance_chain: tuple[str, ...] = () + + +class _LoadedPolicies(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + policies: Mapping[str, _LoadedPolicy] = Field(default_factory=dict) + total_count: int = 0 + + @router.get( "/policy/list", tags=["policy management"], @@ -472,19 +489,19 @@ async def list_policies( """ from litellm.proxy.policy_engine.init_policies import get_policies_summary - summary: Final = get_policies_summary() + summary: Final = _LoadedPolicies.model_validate(get_policies_summary()) return PolicyListResponse( policies={ name: PolicySummaryItem( - inherit=data.get("inherit"), - scope=PolicyScopeResponse(**data.get("scope", {})), - guardrails=PolicyGuardrailsResponse(**data.get("guardrails", {})), - resolved_guardrails=data.get("resolved_guardrails", []), - inheritance_chain=data.get("inheritance_chain", []), + inherit=policy.inherit, + scope=policy.scope, + guardrails=policy.guardrails, + resolved_guardrails=list(policy.resolved_guardrails), + inheritance_chain=list(policy.inheritance_chain), ) - for name, data in summary.get("policies", {}).items() + for name, policy in summary.policies.items() }, - total_count=summary.get("total_count", 0), + total_count=summary.total_count, ) diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index a557e3a6082..d8ea96e6a46 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -116,7 +116,7 @@ async def get_router_settings( config: Final = await proxy_config.get_config() router_settings_from_config: Final = config.get("router_settings", {}) - current_values: Final[dict[str, Any]] = {} + current_values: Final[dict[str, object]] = {} if llm_router is not None: # Router exposes routing groups as private `_routing_groups`; the # generic `hasattr` loop below would miss them. diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index b1094684389..9a1cf701142 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -17,6 +17,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Final, Protocol, TypedDict, overload from fastapi import APIRouter, Depends, HTTPException, Query +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view @@ -58,6 +59,8 @@ if TYPE_CHECKING: router: Final = APIRouter() +_DECODED_JSON: Final = TypeAdapter(object) + class _TagRecord(Protocol): tag_name: str @@ -523,7 +526,7 @@ async def info_tag( model_info: object = {} if tag_record.model_info: if isinstance(tag_record.model_info, str): - model_info = json.loads(tag_record.model_info) + model_info = _DECODED_JSON.validate_python(json.loads(tag_record.model_info)) else: model_info = tag_record.model_info @@ -646,7 +649,7 @@ async def list_tags( model_info: object = {} if tag_record.model_info: if isinstance(tag_record.model_info, str): - model_info = json.loads(tag_record.model_info) + model_info = _DECODED_JSON.validate_python(json.loads(tag_record.model_info)) else: model_info = tag_record.model_info diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index fe976c861e5..e452afccef9 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1721,7 +1721,7 @@ async def new_team( complete_team_data_dict = complete_team_data.model_dump(exclude_none=True) # Serialize router_settings to JSON (matching key creation pattern) - router_settings_value: Final = getattr(data, "router_settings", None) + router_settings_value: Final = data.router_settings router_settings_json: Final = ( safe_dumps(router_settings_value) if router_settings_value is not None else safe_dumps({}) ) diff --git a/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py b/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py index 2381a5cc2db..3d408accba3 100644 --- a/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py +++ b/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py @@ -18,11 +18,11 @@ class FileContentStreamingHandler: *, custom_llm_provider: str, file_id: str, - data: dict[str, Any], + data: dict[str, object], should_route: bool, original_file_id: str | None, - credentials: dict[str, Any] | None, - ) -> tuple[str, str, dict[str, Any]]: + credentials: dict[str, object] | None, + ) -> tuple[str, str, dict[str, object]]: """ Resolve the provider, file ID, and request payload to use for streaming. diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index f5e62da5962..dfec77a6c16 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -1652,7 +1652,8 @@ async def delete_file( def _as_file_list_page(response: object) -> object: if not isinstance(response, list): return response - return FileListPage(**build_list_page(_LISTED_FILES_ADAPTER.validate_python(response))) + page: Final = build_list_page(_LISTED_FILES_ADAPTER.validate_python(response)) + return FileListPage.model_validate(page) @router.get( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index 7cec3bac207..dfb731972ac 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -5,6 +5,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger @@ -48,6 +49,8 @@ else: PassThroughEndpointLogging = Any EndpointType = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AnthropicPassthroughLoggingHandler: @staticmethod @@ -343,11 +346,9 @@ class AnthropicPassthroughLoggingHandler: if not line.startswith("data:"): continue try: - data = json.loads(line[len("data:") :].strip()) + data = _JSON_OBJECT.validate_python(json.loads(line[len("data:") :].strip())) except (json.JSONDecodeError, ValueError): continue - if not isinstance(data, dict): - continue etype = data.get("type") if etype == "message_delta": return False 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 24c1b865db6..be40e225cde 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 @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast from urllib.parse import urlparse import httpx -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger @@ -55,6 +55,7 @@ else: _VERTEX_INTERACTIONS_PATH: Final = re.compile(r"/projects/[^/]+/locations/[^/]+/interactions/?$") _INTERACTIONS_RESPONSE_BODY: Final = TypeAdapter(dict[str, object]) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) def _interactions_model( @@ -99,7 +100,7 @@ class VertexPassthroughLoggingHandler: litellm_model_response: Final = ModelResponse( model=model, usage=InteractionsUsageObjectTransformation.transform_interactions_usage_object( - cast(Mapping[str, Any], usage_object) + _JSON_OBJECT.validate_python(usage_object) ), ) logging_obj.custom_llm_provider = custom_llm_provider @@ -349,7 +350,7 @@ class VertexPassthroughLoggingHandler: model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) - _json_response: Final[dict[str, object]] = httpx_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(httpx_response.json()) litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse() if VertexPassthroughLoggingHandler._is_audio_predict_response( diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index f55eb4f7863..efe7a5002e0 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -421,7 +421,7 @@ class PolicyRegistry: if policy_request.condition is not None: data["condition"] = json.dumps(policy_request.condition.model_dump()) if policy_request.pipeline is not None: - validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) + validated_pipeline: Final = GuardrailPipeline.model_validate(policy_request.pipeline) data["pipeline"] = json.dumps(validated_pipeline.model_dump()) created_policy: Final = await _policy_table(prisma_client).create(data=data) @@ -496,7 +496,7 @@ class PolicyRegistry: if policy_request.condition is not None: update_data["condition"] = json.dumps(policy_request.condition.model_dump()) if policy_request.pipeline is not None: - validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) + validated_pipeline: Final = GuardrailPipeline.model_validate(policy_request.pipeline) update_data["pipeline"] = json.dumps(validated_pipeline.model_dump()) updated_policy: Final = await _policy_table(prisma_client).update( diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 7814903e975..da9a3033187 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -496,10 +496,11 @@ _AUTOROUTER_PRESETS_ADAPTER: Final = TypeAdapter(dict[str, AutoRouterPresetRecor def _load_bundled_autorouter_presets() -> Mapping[str, AutoRouterPresetRecord]: - raw: Final = json.loads( - files("litellm.proxy.public_endpoints").joinpath("autorouter_presets.json").read_text(encoding="utf-8") + return _AUTOROUTER_PRESETS_ADAPTER.validate_python( + json.loads( + files("litellm.proxy.public_endpoints").joinpath("autorouter_presets.json").read_text(encoding="utf-8") + ) ) - return _AUTOROUTER_PRESETS_ADAPTER.validate_python(raw) async def _fetch_remote_autorouter_presets(url: str) -> Mapping[str, AutoRouterPresetRecord]: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 2dd1c9419f1..f9235a33b42 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -11,6 +11,7 @@ from types import MappingProxyType from typing import Final, NoReturn, SupportsFloat, SupportsIndex, SupportsInt, cast from fastapi import HTTPException, status +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._internal_context import with_service_target @@ -71,6 +72,9 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = { "Project": Litellm_EntityType.PROJECT.value, } +_CACHED_VALUE: Final = TypeAdapter(object) +_WINDOW_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class _CounterReservationUnavailable(Exception): def __init__( @@ -762,8 +766,8 @@ async def _get_team_member_budget_counter( else: default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id") if isinstance(default_budget_id, str): - default_budget: Final = await user_api_key_cache.async_get_cache( - key=f"team_member_default_budget:{default_budget_id}", + default_budget: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=f"team_member_default_budget:{default_budget_id}") ) default_cap: Final = _to_float(_get_value(default_budget, "max_budget")) if default_cap is not None and default_cap > 0: @@ -799,8 +803,8 @@ async def _get_org_budget_counter( if org_id is None: return None - org_table: Final = await user_api_key_cache.async_get_cache( - key=f"org_id:{org_id}:with_budget", + org_table: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=f"org_id:{org_id}:with_budget") ) if org_table is None: return None @@ -833,7 +837,9 @@ async def _get_project_budget_counter( return None source_cache_key: Final = project_cache_key(valid_token.project_id) - project_object: Final = await user_api_key_cache.async_get_cache(key=source_cache_key) + project_object: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=source_cache_key) + ) if project_object is None: return None @@ -901,10 +907,9 @@ def _coerce_window(window: object) -> Mapping[str, object]: return window if isinstance(window, str): try: - parsed: Final[object] = json.loads(window) + return _WINDOW_OBJECT.validate_python(json.loads(window)) except Exception: return {} - return parsed if isinstance(parsed, Mapping) else {} model_dump: Final = getattr(window, "model_dump", None) if not callable(model_dump): return {} diff --git a/litellm/proxy/spend_tracking/spend_log_error_logger.py b/litellm/proxy/spend_tracking/spend_log_error_logger.py index 03037ba854f..28cd535127f 100644 --- a/litellm/proxy/spend_tracking/spend_log_error_logger.py +++ b/litellm/proxy/spend_tracking/spend_log_error_logger.py @@ -25,7 +25,7 @@ troubleshoot. The UI suppression follows the same gate. import logging import os -from typing import Any, Final +from typing import Final from litellm._logging import verbose_proxy_logger from litellm.secret_managers.main import str_to_bool @@ -58,7 +58,7 @@ def should_suppress_spend_log_tracebacks() -> bool: def spend_log_error( message: str, - *args: Any, + *args: object, exc: BaseException | None = None, ) -> None: """Log a spend-tracking error, with the traceback gated on the env var. diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index da1bc0a1feb..fb799546b37 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -23,6 +23,7 @@ from litellm.constants import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, + header_value, httpxSpecialProvider, ) from litellm.rag.ingestion.file_parsers import extract_text_from_pdf @@ -79,7 +80,7 @@ class BaseRAGIngestion(ABC): from litellm.litellm_core_utils.credential_accessor import CredentialAccessor credential_name: Final = self.vector_store_config.get("litellm_credential_name") - if credential_name and litellm.credential_list: + if isinstance(credential_name, str) and credential_name and litellm.credential_list: credential_values: Final = CredentialAccessor.get_credential_values(credential_name) if not credential_values: return @@ -125,7 +126,8 @@ class BaseRAGIngestion(ABC): response.raise_for_status() file_content = response.content filename = file_url.split("/")[-1] or "document" - content_type = response.headers.get("content-type", "application/octet-stream") + content_type_header: Final = header_value(response.headers, "content-type") + content_type = "application/octet-stream" if content_type_header is None else content_type_header return filename, file_content, content_type, None if file_id: diff --git a/litellm/rag/ingestion/gemini_ingestion.py b/litellm/rag/ingestion/gemini_ingestion.py index 1cf5db549e4..a015f3cf614 100644 --- a/litellm/rag/ingestion/gemini_ingestion.py +++ b/litellm/rag/ingestion/gemini_ingestion.py @@ -9,6 +9,8 @@ from __future__ import annotations from typing import TYPE_CHECKING, Final, cast +from pydantic import BaseModel, ConfigDict + from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -23,6 +25,13 @@ if TYPE_CHECKING: from litellm.types.rag import RAGIngestOptions +class _WhiteSpaceConfig(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + max_tokens_per_chunk: object = 800 + max_overlap_tokens: object = 400 + + class GeminiRAGIngestion(BaseRAGIngestion): """ Gemini-specific RAG ingestion using File Search API. @@ -234,12 +243,13 @@ class GeminiRAGIngestion(BaseRAGIngestion): # Add chunking configuration if provided chunking_strategy: Final = self.chunking_strategy if chunking_strategy and isinstance(chunking_strategy, dict): - white_space_config: Final = chunking_strategy.get("white_space_config") + white_space_config: Final[object] = chunking_strategy.get("white_space_config") if white_space_config: + white_space: Final = _WhiteSpaceConfig.model_validate(white_space_config) request_body["chunkingConfig"] = { "whiteSpaceConfig": { - "maxTokensPerChunk": white_space_config.get("max_tokens_per_chunk", 800), - "maxOverlapTokens": white_space_config.get("max_overlap_tokens", 400), + "maxTokensPerChunk": white_space.max_tokens_per_chunk, + "maxOverlapTokens": white_space.max_overlap_tokens, } } diff --git a/litellm/rag/ingestion/openai_ingestion.py b/litellm/rag/ingestion/openai_ingestion.py index 925a49d0f7a..6c7c619237d 100644 --- a/litellm/rag/ingestion/openai_ingestion.py +++ b/litellm/rag/ingestion/openai_ingestion.py @@ -7,7 +7,7 @@ so this implementation skips the embedding step and directly uploads files. from __future__ import annotations -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final import litellm from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion @@ -106,7 +106,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion): vector_store_id=vector_store_id, file_id=existing_file_id, custom_llm_provider="openai", - chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy), + chunking_strategy=self.chunking_strategy, api_key=api_key, api_base=api_base, ) @@ -134,7 +134,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion): vector_store_id=vector_store_id, file_id=result_file_id, custom_llm_provider="openai", - chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy), + chunking_strategy=self.chunking_strategy, api_key=api_key, api_base=api_base, ) diff --git a/litellm/rag/ingestion/s3_vectors_ingestion.py b/litellm/rag/ingestion/s3_vectors_ingestion.py index e2aa5555eec..e3eafd680de 100644 --- a/litellm/rag/ingestion/s3_vectors_ingestion.py +++ b/litellm/rag/ingestion/s3_vectors_ingestion.py @@ -18,7 +18,7 @@ from __future__ import annotations import hashlib import uuid from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, TypedDict +from typing import TYPE_CHECKING, Final, TypedDict import litellm from litellm._logging import verbose_logger @@ -194,7 +194,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): url: str, data: str | None = None, headers: dict[str, str] | None = None, - ) -> Any: + ) -> httpx.Response: """ Helper to sign and execute AWS API requests using httpx + SigV4. diff --git a/litellm/rag/main.py b/litellm/rag/main.py index 1a5301f0579..5ab39bec080 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -18,6 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._internal_context import is_internal_call @@ -71,6 +72,13 @@ _SEARCH_ARGS_SET_BY_PIPELINE: Final = frozenset( ) +_VECTOR_STORE_OPTIONS: Final = TypeAdapter(Mapping[object, object], config=ConfigDict(hide_input_in_errors=True)) + + +def _vector_store_provider(ingest_options: Mapping[str, object]) -> object: + return _VECTOR_STORE_OPTIONS.validate_python(ingest_options.get("vector_store", {})).get("custom_llm_provider") + + def get_ingestion_class(provider: str) -> type[BaseRAGIngestion]: """ Get the ingestion class for a given provider. @@ -137,7 +145,7 @@ async def _execute_ingest_pipeline( @client async def aingest( - ingest_options: dict[str, Any], + ingest_options: Mapping[str, object], file_data: tuple[str, bytes, str] | None = None, file: dict[str, str] | None = None, file_url: str | None = None, @@ -197,7 +205,7 @@ async def aingest( except Exception as e: raise litellm.exception_type( model=None, - custom_llm_provider=ingest_options.get("vector_store", {}).get("custom_llm_provider"), + custom_llm_provider=_vector_store_provider(ingest_options), original_exception=e, completion_kwargs=local_vars, extra_kwargs=kwargs, @@ -226,7 +234,7 @@ def _suppressed_sub_call_billing() -> Iterator[None]: async def _execute_query_pipeline( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -276,12 +284,16 @@ async def _execute_query_pipeline( search_provider: Final = retrieval_config.get("custom_llm_provider", "openai") try: - search_cost = sum( - vector_store_search_cost( - model=search_provider if "/" in search_provider else None, - custom_llm_provider=search_provider, - response=search_response, + search_cost = ( + sum( + vector_store_search_cost( + model=search_provider if "/" in search_provider else None, + custom_llm_provider=search_provider, + response=search_response, + ) ) + if isinstance(search_provider, str) + else 0.0 ) except Exception: # noqa: BLE001 - cost accounting must never break the query path search_cost = 0.0 @@ -354,7 +366,7 @@ async def _execute_query_pipeline( async def aquery( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -403,7 +415,7 @@ async def aquery( def query( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -450,7 +462,7 @@ def query( @client def ingest( - ingest_options: dict[str, Any], + ingest_options: Mapping[str, object], file_data: tuple[str, bytes, str] | None = None, file: dict[str, str] | None = None, file_url: str | None = None, @@ -517,7 +529,7 @@ def ingest( except Exception as e: raise litellm.exception_type( model=None, - custom_llm_provider=ingest_options.get("vector_store", {}).get("custom_llm_provider"), + custom_llm_provider=_vector_store_provider(ingest_options), original_exception=e, completion_kwargs=local_vars, extra_kwargs=kwargs, diff --git a/litellm/repositories/autorouter_session_repository.py b/litellm/repositories/autorouter_session_repository.py index 82e2c091728..640fd26ffe3 100644 --- a/litellm/repositories/autorouter_session_repository.py +++ b/litellm/repositories/autorouter_session_repository.py @@ -2,7 +2,7 @@ Repository for the auto-router per-session rollup (LiteLLM_AutoRouterSession). """ -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.autorouter_session import LiteLLM_AutoRouterSession from litellm.repositories.base_repository import BaseRepository @@ -12,10 +12,21 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _AutoRouterSessionDb(Protocol): + @property + def litellm_autoroutersession(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _AutoRouterSessionDb: ... + + class AutoRouterSessionRepository(BaseRepository[LiteLLM_AutoRouterSession]): @property def table(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: - return self.prisma_client.db.litellm_autoroutersession + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_autoroutersession @property def model_class(self) -> type[LiteLLM_AutoRouterSession]: diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index cbacb6e8f90..6541db4c00b 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -4,20 +4,24 @@ Model repository for database operations on LiteLLM_ProxyModelTable. import json from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final + +from pydantic import ConfigDict, TypeAdapter from litellm.models.model import LiteLLM_ProxyModelTable from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import PrismaTableRepository if TYPE_CHECKING: from prisma import models as prisma_models +_LITELLM_PARAMS: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class _ProxyModelTableRepository(PrismaTableRepository["prisma_models.LiteLLM_ProxyModelTable"]): table_name = "litellm_proxymodeltable" @@ -60,22 +64,26 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): decrypted[key] = value return decrypted - def _to_model(self, record: Any) -> LiteLLM_ProxyModelTable | None: + def _to_model(self, record: DbRecord | None) -> LiteLLM_ProxyModelTable | None: """Convert a database record to a Model with decryption.""" if record is None: return None - data: Final = record.dict() if hasattr(record, "dict") else dict(record) + data: Final = dict(record_to_dict(record)) - if isinstance(data.get("litellm_params"), str): - data["litellm_params"] = json.loads(data["litellm_params"]) - if isinstance(data.get("model_info"), str): - data["model_info"] = json.loads(data["model_info"]) + litellm_params: Final = data.get("litellm_params") + if isinstance(litellm_params, str): + data["litellm_params"] = json.loads(litellm_params) + model_info: Final = data.get("model_info") + if isinstance(model_info, str): + data["model_info"] = json.loads(model_info) if data.get("litellm_params"): - data["litellm_params"] = self._decrypt_litellm_params(data["litellm_params"]) + data["litellm_params"] = self._decrypt_litellm_params( + _LITELLM_PARAMS.validate_python(data["litellm_params"]) + ) - return LiteLLM_ProxyModelTable(**data) + return LiteLLM_ProxyModelTable.model_validate(data) async def find_by_id(self, model_id: str, id_field: str = "model_id") -> LiteLLM_ProxyModelTable | None: return await super().find_by_id(model_id, id_field) diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index b732d2ff94c..99d6aeed88b 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -2,7 +2,7 @@ ObjectPermission repository for database operations on LiteLLM_ObjectPermissionTable. """ -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.repositories.base_repository import BaseRepository @@ -12,6 +12,19 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _ObjectPermissionDb(Protocol): + @property + def litellm_objectpermissiontable(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _ObjectPermissionDb: ... + + @property + def writer_db(self) -> _ObjectPermissionDb: ... + + class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): """Repository for object permission database operations.""" @@ -21,7 +34,8 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: - database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + client: Final[_PrismaClientView] = self.prisma_client + database: Final = client.writer_db if self._use_writer else client.db return database.litellm_objectpermissiontable @property @@ -48,7 +62,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable: """Create a new object permission record.""" - data: Final[dict[str, Any]] = {} + data: Final[dict[str, object]] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: @@ -90,7 +104,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """Update an object permission record.""" - data: Final[dict[str, Any]] = {} + data: Final[dict[str, object]] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 905e813f35e..0ea5009a5b6 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -3,7 +3,7 @@ Project repository for database operations on LiteLLM_ProjectTable. """ from collections.abc import Mapping -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.project import LiteLLM_ProjectTable from litellm.repositories.base_repository import BaseRepository @@ -13,12 +13,23 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _ProjectDb(Protocol): + @property + def litellm_projecttable(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _ProjectDb: ... + + class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): """Repository for project database operations.""" @property def table(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: - return self.prisma_client.db.litellm_projecttable + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_projecttable @property def model_class(self) -> type[LiteLLM_ProjectTable]: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index b201dbc566b..5ed1574c6bc 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -5,7 +5,7 @@ User repository for database operations on LiteLLM_UserTable. import json from collections.abc import Mapping, Sequence from itertools import chain -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter @@ -37,6 +37,19 @@ ORDER BY p.user_id _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...]) +class _UserDb(Protocol): + @property + def litellm_usertable(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _UserDb: ... + + @property + def writer_db(self) -> _UserDb: ... + + class UserRepository(BaseRepository[LiteLLM_UserTable]): """Repository for user database operations.""" @@ -46,7 +59,8 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: - database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + client: Final[_PrismaClientView] = self.prisma_client + database: Final = client.writer_db if self._use_writer else client.db return database.litellm_usertable @property diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 9aa881fc4c9..95c981502e7 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -123,6 +123,9 @@ class AdaptiveRouter: prefs = self.model_to_prefs.get(model) or _default_prefs() self._cells[(rt, model)] = initial_cell(prefs, rt) + def cell(self, request_type: RequestType, model: str) -> BanditCell: + return self._cells[(request_type, model)] + async def load_state_from_db(self, prisma_client: object) -> None: """Add each row's persisted delta to a freshly computed cold-start prior. diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index fc0865654c7..43200271f9f 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -2945,7 +2945,7 @@ class ComplexityRouter(CustomLogger): raise ValueError(f"No candidate models left for tier {tier_key} after routing-plugin filtering") return self._pick_from_tier_value(context.candidate_models, tier_key) - def _ensure_adaptive_router(self) -> Any | None: + def _ensure_adaptive_router(self) -> AdaptiveRouter | None: if not self.config.adaptive: return None if self.adaptive_router is not None: @@ -3087,7 +3087,7 @@ class ComplexityRouter(CustomLogger): pools: Final = self._tier_pools() classified_candidates: Final = _allowed(tuple(pools.get(_tier_name(classified_tier), ())), fit_filter) cold_start_candidates: Final = tuple( - model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0 + model for model in classified_candidates if adaptive.cell(request_type, model).total_samples == 0 ) if cold_start_candidates: chosen_model: Final = random.choice(cold_start_candidates) @@ -3106,7 +3106,7 @@ class ComplexityRouter(CustomLogger): "candidates": [ { "model": model, - "total_samples": adaptive._cells[(request_type, model)].total_samples, + "total_samples": adaptive.cell(request_type, model).total_samples, } for model in cold_start_candidates ], @@ -3123,7 +3123,7 @@ class ComplexityRouter(CustomLogger): best_score = float("-inf") candidate_scores: Final[list[dict[str, object]]] = [] for model in self._adaptive_candidate_models(classified_tier, hard_floor, hard_ceiling, fit_filter): - cell = adaptive._cells[(request_type, model)] + cell = adaptive.cell(request_type, model) quality_sample = thompson_sample(cell) cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs) if self.config.adaptive_eligible == "classified_tier": diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index 1d7656f253e..200de29b664 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -34,7 +34,7 @@ from typing import Final, get_args import litellm from litellm.types.llms.openai import REASONING_EFFORT -REASONING_EFFORT_ADVERTISEMENT_ORDER: Final = get_args(REASONING_EFFORT) +REASONING_EFFORT_ADVERTISEMENT_ORDER: Final[tuple[str, ...]] = get_args(REASONING_EFFORT) _EMPTY_ENTRY: Final[Mapping[str, object]] = MappingProxyType({}) _EFFORT_FLAGS: Final = ( diff --git a/litellm/secret_managers/aws_secret_manager.py b/litellm/secret_managers/aws_secret_manager.py index 1aed44e9a31..b3a9fbf31c8 100644 --- a/litellm/secret_managers/aws_secret_manager.py +++ b/litellm/secret_managers/aws_secret_manager.py @@ -14,9 +14,13 @@ import os import re from typing import Any, Final +from pydantic import TypeAdapter + import litellm from litellm.proxy._types import KeyManagementSystem +_PARSED_LITERAL: Final = TypeAdapter(object) + def validate_environment(): if "AWS_REGION_NAME" not in os.environ: @@ -107,7 +111,7 @@ class AWSKeyManagementService_V2: if isinstance(secret, str): secret = secret.strip() try: - secret_value_as_bool: Final = ast.literal_eval(secret) + secret_value_as_bool: Final = _PARSED_LITERAL.validate_python(ast.literal_eval(secret)) if isinstance(secret_value_as_bool, bool): return secret_value_as_bool except Exception: diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index fbed4fb8e75..275c029b0b7 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -6,7 +6,7 @@ import traceback from typing import Final import httpx -from pydantic import BaseModel, ValidationError +from pydantic import BaseModel, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger @@ -19,6 +19,8 @@ from litellm.secret_managers.get_azure_ad_token_provider import ( oidc_cache: Final = DualCache() +_PARSED_LITERAL: Final = TypeAdapter(object) + _OIDC_TOKEN_EXPIRY_MARGIN_SECONDS: Final = 60 @@ -348,7 +350,7 @@ def get_secret( secret = os.getenv(secret_name) try: if isinstance(secret, str): - secret_value_as_bool = ast.literal_eval(secret) + secret_value_as_bool = _PARSED_LITERAL.validate_python(ast.literal_eval(secret)) if isinstance(secret_value_as_bool, bool): return secret_value_as_bool else: diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index e78d2aa5f6a..c0cb3796934 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -10,6 +10,8 @@ from typing import ( get_args, ) +from pydantic import ConfigDict, TypeAdapter + from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import remove_items_at_indices from litellm.repositories.table_repositories import ( @@ -30,6 +32,8 @@ if TYPE_CHECKING: else: PrismaClient = Any +_DB_ROW: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class VectorStoreIndexRegistry: def __init__(self, vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = []): @@ -95,7 +99,9 @@ class VectorStoreIndexRegistry: ) for vector_store in _vector_stores_from_db: _dict_vector_store = dict(vector_store) - _litellm_managed_vector_store = LiteLLM_ManagedVectorStoreIndex(**_dict_vector_store) + _litellm_managed_vector_store = LiteLLM_ManagedVectorStoreIndex.model_validate( + _DB_ROW.validate_python(_dict_vector_store) + ) vector_stores_from_db.append(_litellm_managed_vector_store) return vector_stores_from_db diff --git a/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py b/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py index db561fd1dc2..8d5e15b4b71 100644 --- a/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py +++ b/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py @@ -2,10 +2,15 @@ Tests for Pydantic AI agents header forwarding via agent_extra_headers. """ -from unittest.mock import AsyncMock, MagicMock, patch +from typing import Final +from unittest.mock import ANY, AsyncMock, MagicMock, patch +import httpx import pytest +import respx +import litellm +from litellm.a2a_protocol.providers.pydantic_ai_agents.config import PydanticAIProviderConfig from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( PydanticAITransformation, ) @@ -198,3 +203,67 @@ async def test_provider_config_threads_agent_extra_headers(): sent_headers = mock_client.post.await_args.kwargs["headers"] assert sent_headers["x-trace-id"] == "abc-123" assert sent_headers["Content-Type"] == "application/json" + + +COMPLETED_TASK: Final = { + "jsonrpc": "2.0", + "id": "req-5", + "result": { + "id": "task-5", + "kind": "task", + "status": {"state": "completed"}, + "history": [], + "artifacts": [{"artifactId": "a-5", "parts": [{"kind": "text", "text": "ok"}]}], + }, +} + + +def _user_params() -> dict[str, object]: + return {"message": {"role": "user", "parts": [{"kind": "text", "text": "hello"}], "messageId": "msg-user-5"}} + + +@pytest.fixture +def agent_route(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + return respx_mock.post("http://agent.test/").mock(return_value=httpx.Response(200, json=COMPLETED_TASK)) + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_defaults_to_sixty_second_timeout(agent_route: respx.Route) -> None: + response: Final = await PydanticAIProviderConfig().handle_non_streaming( + "req-5", _user_params(), "http://agent.test" + ) + + request: Final = agent_route.calls.last.request + assert response == { + "jsonrpc": "2.0", + "id": "req-5", + "result": {"kind": "message", "role": "agent", "parts": [{"kind": "text", "text": "ok"}], "messageId": ANY}, + } + assert request.extensions["timeout"] == {"connect": 60.0, "read": 60.0, "write": 60.0, "pool": 60.0} + assert request.headers["content-type"] == "application/json" + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_forwards_timeout_and_ignores_unrelated_keywords( + agent_route: respx.Route, +) -> None: + response: Final = await PydanticAIProviderConfig().handle_non_streaming( + request_id="req-5", + params=_user_params(), + api_base="http://agent.test", + timeout=5.5, + agent_extra_headers={"x-trace-id": "abc-123"}, + litellm_params={"custom_llm_provider": "pydantic_ai_agents"}, + ) + + request: Final = agent_route.calls.last.request + assert response["id"] == "req-5" + assert request.extensions["timeout"] == {"connect": 5.5, "read": 5.5, "write": 5.5, "pool": 5.5} + assert request.headers["x-trace-id"] == "abc-123" + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_requires_api_base() -> None: + with pytest.raises(ValueError, match="api_base is required for PydanticAIProviderConfig"): + await PydanticAIProviderConfig().handle_non_streaming("req-5", _user_params()) diff --git a/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py b/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py new file mode 100644 index 00000000000..38b896723f7 --- /dev/null +++ b/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py @@ -0,0 +1,161 @@ +import json +import re +from typing import Final +from unittest.mock import ANY + +import httpx +import pytest +import respx + +import litellm +from litellm.a2a_protocol.providers.watsonx_orchestrate.config import WatsonxOrchestrateA2AConfig +from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import WXOLitellmParams + +MISSING_LITELLM_PARAMS: Final = re.escape( + "litellm_params is required for WatsonxOrchestrateA2AConfig " + "(must contain cp4d_host, instance_id, wxo_agent_id, api_key)" +) +MISSING_HOST: Final = re.escape("'cp4d_host' is required in litellm_params for WXO agents") +RUNS_URL: Final = "https://wxo-config.test/orchestrate/cpd/instances/inst-1/v1/orchestrate/runs" +LITELLM_PARAMS: Final[WXOLitellmParams] = { + "cp4d_host": "https://wxo-config.test", + "instance_id": "inst-1", + "wxo_agent_id": "agent-1", + "api_key": "config-test-key", + "auth_mode": "ibm_cloud", +} +EXPECTED_RUN_BODY: Final = { + "agent_id": "agent-1", + "message": {"role": "user", "content": [{"response_type": "text", "text": "hello"}]}, +} + + +def _a2a_params() -> dict[str, object]: + return {"message": {"role": "user", "parts": [{"kind": "text", "text": "hello"}], "messageId": "m-1"}} + + +def _mock_wxo_route(respx_mock: respx.MockRouter, url: str) -> respx.Route: + respx_mock.post("https://iam.cloud.ibm.com/identity/token").mock( + return_value=httpx.Response(200, json={"access_token": "tok", "expires_in": 0}) + ) + return respx_mock.post(url).mock( + return_value=httpx.Response(200, json={"status": "completed", "results": "wxo says hi"}) + ) + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +@pytest.mark.parametrize("litellm_params", [None, {}]) +async def test_handle_non_streaming_rejects_empty_litellm_params(litellm_params: WXOLitellmParams | None) -> None: + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await WatsonxOrchestrateA2AConfig().handle_non_streaming("req-1", _a2a_params(), litellm_params=litellm_params) + + +async def test_handle_non_streaming_requires_litellm_params_when_only_other_keywords_are_given() -> None: + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await WatsonxOrchestrateA2AConfig().handle_non_streaming( + "req-1", _a2a_params(), "https://ignored.test", agent_extra_headers={"x-tenant-id": "acme"} + ) + + +async def test_handle_non_streaming_hands_litellm_params_to_the_wxo_handler() -> None: + with pytest.raises(ValueError, match=MISSING_HOST): + await WatsonxOrchestrateA2AConfig().handle_non_streaming( + "req-1", _a2a_params(), litellm_params={"instance_id": "inst-1"} + ) + + +@pytest.mark.usefixtures("httpx_transport") +async def test_handle_non_streaming_runs_the_agent_and_ignores_unrelated_keywords( + respx_mock: respx.MockRouter, +) -> None: + runs_route: Final = _mock_wxo_route(respx_mock, RUNS_URL) + + response: Final = await WatsonxOrchestrateA2AConfig().handle_non_streaming( + request_id="req-1", + params=_a2a_params(), + api_base="https://ignored.test", + litellm_params=LITELLM_PARAMS, + agent_extra_headers={"x-tenant-id": "acme"}, + ) + + run_request: Final = runs_route.calls.last.request + assert response == { + "jsonrpc": "2.0", + "id": "req-1", + "result": { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": "wxo says hi"}], + "messageId": ANY, + }, + } + assert json.loads(run_request.content) == EXPECTED_RUN_BODY + assert run_request.headers["authorization"] == "Bearer tok" + assert "x-tenant-id" not in run_request.headers + + +@pytest.mark.parametrize("litellm_params", [None, {}]) +async def test_handle_streaming_rejects_empty_litellm_params(litellm_params: WXOLitellmParams | None) -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), litellm_params=litellm_params + ) + + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await anext(stream) + + +async def test_handle_streaming_requires_litellm_params_when_only_other_keywords_are_given() -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), "https://ignored.test", agent_extra_headers={"x-tenant-id": "acme"} + ) + + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await anext(stream) + + +async def test_handle_streaming_hands_litellm_params_to_the_wxo_handler() -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), litellm_params={"instance_id": "inst-1"} + ) + + with pytest.raises(ValueError, match=MISSING_HOST): + await anext(stream) + + +@pytest.mark.usefixtures("httpx_transport") +async def test_handle_streaming_runs_the_agent_and_ignores_unrelated_keywords(respx_mock: respx.MockRouter) -> None: + stream_route: Final = _mock_wxo_route(respx_mock, f"{RUNS_URL}/stream") + + chunks: Final = [ + chunk + async for chunk in WatsonxOrchestrateA2AConfig().handle_streaming( + request_id="req-1", + params=_a2a_params(), + api_base="https://ignored.test", + litellm_params=LITELLM_PARAMS, + agent_extra_headers={"x-tenant-id": "acme"}, + ) + ] + + stream_request: Final = stream_route.calls.last.request + assert [chunk["id"] for chunk in chunks] == ["req-1", "req-1", "req-1", "req-1"] + assert chunks[2]["result"] == { + "contextId": ANY, + "kind": "artifact-update", + "taskId": ANY, + "artifact": {"artifactId": ANY, "parts": [{"kind": "text", "text": "wxo says hi"}]}, + } + assert chunks[3]["result"] == { + "contextId": ANY, + "final": True, + "kind": "status-update", + "status": {"state": "completed"}, + "taskId": ANY, + } + assert json.loads(stream_request.content) == EXPECTED_RUN_BODY + assert stream_request.headers["authorization"] == "Bearer tok" + assert "x-tenant-id" not in stream_request.headers diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 52e44ca5448..6cfdb75fd01 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -2,6 +2,8 @@ import asyncio import json import os import unittest.mock as mock +from types import SimpleNamespace +from typing import Final from unittest.mock import patch import pytest @@ -19,7 +21,12 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import ( ) from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER -from litellm.proxy._types import CallInfo, Litellm_EntityType, WebhookEvent +from litellm.proxy._types import ( + CallInfo, + LiteLLM_UserTable, + Litellm_EntityType, + WebhookEvent, +) from litellm.constants import EMAIL_BUDGET_ALERT_TTL @@ -1417,3 +1424,69 @@ async def test_budget_alert_release_failure_does_not_propagate(base_email_logger ) mock_cache.async_delete_cache.assert_awaited_once() + + +class _RecordingEmailLogger(BaseEmailLogger): + def __init__(self): + super().__init__() + self.recipients = [] + + async def send_email(self, from_email, to_email, subject, html_body): + self.recipients.append(to_email) + + +class _UserTable: + def __init__(self, rows): + self._rows = rows + + async def find_unique(self, where): + return self._rows.get(where["user_id"]) + + +def _prisma_client_with_users(*rows): + table: Final = _UserTable({row.user_id: row for row in rows}) + return SimpleNamespace(db=SimpleNamespace(litellm_usertable=table)) + + +def _key_created_event(user_id): + return SendKeyCreatedEmailEvent( + user_id=user_id, + user_email=None, + virtual_key="sk-test", + max_budget=None, + spend=0.0, + event_group=Litellm_EntityType.USER, + event="key_created", + event_message="Key Created", + ) + + +@pytest.mark.asyncio +async def test_key_created_email_goes_to_the_address_stored_for_the_user(monkeypatch): + stored_user: Final = LiteLLM_UserTable( + user_id="user-1", user_email="stored@example.com" + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + _prisma_client_with_users(stored_user), + ) + email_logger: Final = _RecordingEmailLogger() + + await email_logger.send_key_created_email(_key_created_event("user-1")) + + assert email_logger.recipients == [["stored@example.com"]] + + +@pytest.mark.asyncio +async def test_key_created_email_is_refused_for_a_user_the_database_does_not_know( + monkeypatch, +): + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", _prisma_client_with_users() + ) + email_logger: Final = _RecordingEmailLogger() + + with pytest.raises(ValueError, match="User email not found for user_id: user-2"): + await email_logger.send_key_created_email(_key_created_event("user-2")) + + assert email_logger.recipients == [] diff --git a/tests/unit/files/test_main.py b/tests/unit/files/test_main.py index 2704bdfb6ff..cb70b39f4d5 100644 --- a/tests/unit/files/test_main.py +++ b/tests/unit/files/test_main.py @@ -3,9 +3,12 @@ from urllib.parse import parse_qs, urlparse import httpx import pytest +import respx +from pydantic import ValidationError import litellm from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.types.llms.openai import OpenAIFileObject NATIVE_VERTEX_ROWS: Final = ( b'{"request": {"contents": [{"role": "user", "parts": [{"text": "Who won the 2024 Tour de France?"}]}],' @@ -69,3 +72,52 @@ def test_create_file_passthrough_kwarg_ships_native_rows_byte_for_byte_under_the assert upload.read() == NATIVE_VERTEX_ROWS assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/") assert file_object.id == f"gs://my-bucket/{object_name}" + + +OPENAI_FILES_API_BASE: Final = "https://files.test/v1" +PROVIDER_FILE: Final = { + "id": "file-abc123", + "bytes": 120, + "created_at": 1700000000, + "filename": "batch.jsonl", + "object": "file", + "purpose": "batch", +} +UNSET_OPTIONAL_FIELDS: Final = {"status": None, "expires_at": None, "status_details": None} + + +@pytest.mark.parametrize( + "provider_extras, expected", + [ + ({}, {**PROVIDER_FILE, **UNSET_OPTIONAL_FIELDS}), + ( + {"status": "processed", "expires_at": 1800000000, "status_details": "ok", "provider_only": {"a": [1]}}, + {**PROVIDER_FILE, "status": "processed", "expires_at": 1800000000, "status_details": "ok"}, + ), + ], + ids=["required-fields-only", "optional-and-unknown-fields"], +) +@respx.mock +async def test_afile_retrieve_returns_the_provider_file_as_an_openai_file_object(provider_extras, expected): + respx.get(f"{OPENAI_FILES_API_BASE}/files/file-abc123").respond(200, json={**PROVIDER_FILE, **provider_extras}) + + file_object: Final = await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="sk-test", api_base=OPENAI_FILES_API_BASE + ) + + assert type(file_object) is OpenAIFileObject + assert file_object.model_dump() == expected + + +@respx.mock +async def test_afile_retrieve_rejects_a_provider_file_without_its_size(): + provider_file: Final = {key: value for key, value in PROVIDER_FILE.items() if key != "bytes"} + respx.get(f"{OPENAI_FILES_API_BASE}/files/file-abc123").respond(200, json=provider_file) + + with pytest.raises(ValidationError) as exc_info: + await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="sk-test", api_base=OPENAI_FILES_API_BASE + ) + + assert exc_info.value.title == "OpenAIFileObject" + assert [error["loc"] for error in exc_info.value.errors()] == [("bytes",)] diff --git a/tests/unit/integrations/arize/test_arize_utils.py b/tests/unit/integrations/arize/test_arize_utils.py index 167b083e147..0d833b90148 100644 --- a/tests/unit/integrations/arize/test_arize_utils.py +++ b/tests/unit/integrations/arize/test_arize_utils.py @@ -5,6 +5,7 @@ from typing import Optional import asyncio +import httpx import pytest import litellm @@ -13,6 +14,7 @@ from litellm.integrations._types.open_inference import ( SpanAttributes, ToolCallAttributes, ) +from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs, _parse_passthrough_response from litellm.integrations.arize.arize import ArizeLogger from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import Choices, StandardCallbackDynamicParams @@ -1513,3 +1515,42 @@ def test_arize_mcp_emitter_is_inert_without_a_standard_logging_object(): written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} assert SpanAttributes.TOOL_NAME not in written + + +@pytest.mark.parametrize( + ("text", "expected"), + [ + ('{"id": "msg_1", "usage": {"input_tokens": 3}}', {"id": "msg_1", "usage": {"input_tokens": 3}}), + ("{}", {}), + ], +) +def test_coerce_response_obj_decodes_a_json_object_body(text: str, expected: dict[str, object]): + assert _coerce_response_obj_for_attrs(httpx.Response(200, text=text)) == expected + + +@pytest.mark.parametrize("text", ["[1, 2]", '"text"', "7", "null", "true", "not json"]) +def test_coerce_response_obj_keeps_the_response_when_the_body_is_not_a_json_object(text: str): + response = httpx.Response(200, text=text) + + assert _coerce_response_obj_for_attrs(response) is response + + +@pytest.mark.parametrize( + ("raw", "coerced", "kwargs", "expected"), + [ + (None, {"response": '{"id": "wrapped"}'}, {}, {"id": "wrapped"}), + (None, {"response": "[1, 2]"}, {}, None), + ({"id": "raw"}, {"response": "[1, 2]"}, {}, {"id": "raw"}), + ({"id": "raw"}, {"response": "not json"}, {}, {"id": "raw"}), + (None, {"response": "7"}, {"original_response": '{"id": "original"}'}, {"id": "original"}), + (None, None, {"original_response": '{"id": "original"}'}, {"id": "original"}), + (None, None, {"original_response": "[1, 2]"}, None), + (None, None, {"original_response": '"text"'}, None), + (None, None, {"original_response": "null"}, None), + (None, None, {"original_response": "not json"}, None), + ], +) +def test_parse_passthrough_response_reads_only_json_objects_from_text( + raw: object, coerced: object, kwargs: dict[str, object], expected: dict[str, object] | None +): + assert _parse_passthrough_response(raw, coerced, kwargs) == expected diff --git a/tests/unit/integrations/datadog/test_datadog_llm_obs.py b/tests/unit/integrations/datadog/test_datadog_llm_obs.py index 77d2518696c..a7826e11483 100644 --- a/tests/unit/integrations/datadog/test_datadog_llm_obs.py +++ b/tests/unit/integrations/datadog/test_datadog_llm_obs.py @@ -1139,3 +1139,31 @@ def test_reasoning_content_survives_the_mapping(logger: DataDogLLMObsLogger) -> ) assert payload["meta"]["output"]["messages"][0]["reasoning_content"] == "thinking" + + +@pytest.mark.parametrize( + ("raw_arguments", "shipped"), + [ + ('{"city": "Paris", "days": [1, 2]}', {"city": "Paris", "days": [1, 2]}), + ("{}", {}), + ("[1, 2]", "[1, 2]"), + ("null", "null"), + ("true", "true"), + ('"text"', '"text"'), + ("1.5", "1.5"), + ("", ""), + ], +) +def test_tool_arguments_ship_as_an_object_only_when_they_decode_to_one( + logger: DataDogLLMObsLogger, raw_arguments: str, shipped: object +) -> None: + payload = build( + logger, + response_message={ + "role": "assistant", + "content": None, + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": raw_arguments}}], + }, + ) + + assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == shipped diff --git a/tests/unit/integrations/opik/opik_payload_builder/__init__.py b/tests/unit/integrations/opik/opik_payload_builder/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/integrations/opik/opik_payload_builder/test_api.py b/tests/unit/integrations/opik/opik_payload_builder/test_api.py new file mode 100644 index 00000000000..d2e20909d3c --- /dev/null +++ b/tests/unit/integrations/opik/opik_payload_builder/test_api.py @@ -0,0 +1,230 @@ +from collections import OrderedDict +from collections.abc import Mapping +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Final +from uuid import UUID + +import pytest +from pydantic import ValidationError + +from litellm.integrations.opik.opik_payload_builder import build_opik_payload +from litellm.integrations.opik.opik_payload_builder.types import SpanPayload, TracePayload +from litellm.types.utils import ModelResponse, Usage + +_START: Final = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) +_END: Final = datetime(2026, 1, 2, 3, 4, 6, tzinfo=timezone.utc) +_MESSAGES: Final = [{"role": "user", "content": "hi"}] +_RESPONSE: Final = {"id": "chatcmpl-1", "choices": []} +_HIDDEN_PARAMS: Final = {"model_id": "deployment-1"} +_LOGGING_METADATA: Final = { + "user_api_key_alias": "team-key", + "requester_metadata": {"opik": {"thread_id": "thread-1"}}, +} +_STANDARD_LOGGING_OBJECT: Final = { + "call_type": "acompletion", + "status": "success", + "model": "gpt-4o", + "metadata": _LOGGING_METADATA, + "messages": _MESSAGES, + "response": _RESPONSE, + "hidden_params": _HIDDEN_PARAMS, + "trace_id": "not-forwarded-to-opik", +} +_RESPONSE_OBJ: Final = ModelResponse( + model="gpt-4o", + created=1767323045, + usage=Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3), +) +_OPIK_FIELDS_OF_THE_LOGGING_OBJECT: Final = { + "type": "acompletion", + "status": "success", + "model": "gpt-4o", + "hidden_params": _HIDDEN_PARAMS, +} +_EXISTING_TRACE: Final = {"metadata": {"opik": {"current_span_data": {"trace_id": "trace-1", "id": "span-0"}}}} + + +def _payloads_attached_to_trace_1(standard_logging_object: object) -> tuple[TracePayload | None, SpanPayload]: + return build_opik_payload( + kwargs={"standard_logging_object": standard_logging_object, "litellm_params": _EXISTING_TRACE}, + response_obj=_RESPONSE_OBJ, + start_time=_START, + end_time=_END, + project_name="default-project", + ) + + +def _span_attached_to_trace_1(span_id: str, metadata: Mapping[str, object]) -> SpanPayload: + return SpanPayload( + id=span_id, + project_name="default-project", + trace_id="trace-1", + parent_span_id="span-0", + name="gpt-4o_chat.completion_1767323045", + type="llm", + model="gpt-4o", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=[], + usage={"completion_tokens": 2, "prompt_tokens": 1, "total_tokens": 3}, + ) + + +def test_build_opik_payload_creates_a_trace_and_its_span_from_the_standard_logging_object() -> None: + trace, span = build_opik_payload( + kwargs={ + "standard_logging_object": _STANDARD_LOGGING_OBJECT, + "custom_llm_provider": "openai", + "response_cost": 0.25, + }, + response_obj=_RESPONSE_OBJ, + start_time=_START, + end_time=_END, + project_name="default-project", + ) + metadata: Final = { + "thread_id": "thread-1", + "created_from": "litellm", + **_LOGGING_METADATA, + **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT, + "cost": {"total_tokens": 0.25, "currency": "USD"}, + } + + assert trace is not None + assert (trace, span) == ( + TracePayload( + project_name="default-project", + id=trace.id, + name="chat.completion", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=["openai"], + thread_id="thread-1", + ), + SpanPayload( + id=span.id, + project_name="default-project", + trace_id=trace.id, + name="gpt-4o_chat.completion_1767323045", + type="llm", + model="gpt-4o", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=["openai"], + usage={"completion_tokens": 2, "prompt_tokens": 1, "total_tokens": 3}, + provider="openai", + total_cost=0.25, + ), + ) + assert [UUID(trace.id).version, UUID(span.id).version, trace.id != span.id] == [7, 7, True] + assert [span.input is _MESSAGES, span.output is _RESPONSE, span.metadata["hidden_params"] is _HIDDEN_PARAMS] == [ + True, + True, + True, + ] + assert list(span.metadata) == [ + "thread_id", + "created_from", + "user_api_key_alias", + "requester_metadata", + "type", + "status", + "model", + "hidden_params", + "cost", + ] + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + _STANDARD_LOGGING_OBJECT, + OrderedDict(_STANDARD_LOGGING_OBJECT), + MappingProxyType(_STANDARD_LOGGING_OBJECT), + {**_STANDARD_LOGGING_OBJECT, "metadata": MappingProxyType(_LOGGING_METADATA)}, + ], +) +def test_build_opik_payload_reads_any_string_keyed_mapping_as_the_standard_logging_object( + standard_logging_object: object, +) -> None: + trace, span = _payloads_attached_to_trace_1(standard_logging_object) + + assert (trace, span) == ( + None, + _span_attached_to_trace_1( + span.id, + { + "thread_id": "thread-1", + "created_from": "litellm", + **_LOGGING_METADATA, + **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT, + }, + ), + ) + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + {key: value for key, value in _STANDARD_LOGGING_OBJECT.items() if key != "metadata"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": None}, + {**_STANDARD_LOGGING_OBJECT, "metadata": {}}, + {**_STANDARD_LOGGING_OBJECT, "metadata": ""}, + ], +) +def test_build_opik_payload_without_standard_logging_metadata_keeps_only_the_logging_object_fields( + standard_logging_object: object, +) -> None: + trace, span = _payloads_attached_to_trace_1(standard_logging_object) + + assert (trace, span) == ( + None, + _span_attached_to_trace_1(span.id, {"created_from": "litellm", **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT}), + ) + + +def test_build_opik_payload_of_an_empty_standard_logging_object_has_empty_input_and_output() -> None: + _, span = _payloads_attached_to_trace_1({}) + + assert (span.input, span.output, span.metadata) == ({}, {}, {"created_from": "litellm"}) + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + None, + "standard_logging_object", + ["messages", "response"], + list(_STANDARD_LOGGING_OBJECT.items()), + 7, + {**_STANDARD_LOGGING_OBJECT, 7: "keys must be strings"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": "metadata"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": ["user_api_key_alias"]}, + {**_STANDARD_LOGGING_OBJECT, "metadata": 7}, + {**_STANDARD_LOGGING_OBJECT, "metadata": {7: "keys must be strings"}}, + ], +) +def test_build_opik_payload_rejects_a_standard_logging_object_that_is_not_a_string_keyed_mapping( + standard_logging_object: object, +) -> None: + with pytest.raises(ValidationError) as raised: + _payloads_attached_to_trace_1(standard_logging_object) + + assert "input_value" not in str(raised.value) + + +def test_build_opik_payload_without_a_standard_logging_object_raises_a_key_error() -> None: + with pytest.raises(KeyError, match="standard_logging_object"): + build_opik_payload( + kwargs={}, response_obj=_RESPONSE_OBJ, start_time=_START, end_time=_END, project_name="default-project" + ) diff --git a/tests/unit/integrations/opik/test_opik_extractors.py b/tests/unit/integrations/opik/test_opik_extractors.py index 6f85a1c6090..4ef9f500b54 100644 --- a/tests/unit/integrations/opik/test_opik_extractors.py +++ b/tests/unit/integrations/opik/test_opik_extractors.py @@ -1,4 +1,7 @@ +import pytest + from litellm.integrations.opik.opik_payload_builder.extractors import ( + apply_proxy_header_overrides, extract_opik_metadata, ) @@ -82,3 +85,32 @@ def test_extract_opik_metadata_requester_metadata_overrides_all_other_sources(): "workspace": "requester-workspace", "thread_id": "requester-thread", } + + +@pytest.mark.parametrize( + ("opik_tags_header", "expected_tags"), + [ + ('["from-header", "second"]', ["from-request", "from-header", "second"]), + ('["text", 7, null, {"nested": [1]}]', ["from-request", "text", 7, None, {"nested": [1]}]), + ("[]", ["from-request"]), + ('{"not": "a list"}', ["from-request"]), + ('"not-a-list"', ["from-request"]), + ("null", ["from-request"]), + ("not json", ["from-request"]), + ], +) +def test_opik_tags_header_adds_tags_only_when_it_is_a_json_list(opik_tags_header: str, expected_tags: list[object]): + overrides = apply_proxy_header_overrides("project", ["from-request"], None, {"opik_tags": opik_tags_header}) + + assert overrides == ("project", expected_tags, None) + + +def test_opik_headers_override_the_project_name_and_thread_id(): + overrides = apply_proxy_header_overrides( + "project", + ["from-request"], + "thread-from-request", + {"opik_project_name": "header-project", "opik_thread_id": "header-thread", "opik_tags": "", "x-other": "1"}, + ) + + assert overrides == ("header-project", ["from-request"], "header-thread") diff --git a/tests/unit/integrations/test_galileo.py b/tests/unit/integrations/test_galileo.py index d0709b966d4..d067b31a651 100644 --- a/tests/unit/integrations/test_galileo.py +++ b/tests/unit/integrations/test_galileo.py @@ -1,7 +1,11 @@ +from dataclasses import dataclass from datetime import datetime, timezone +from decimal import Decimal +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +from pydantic import BaseModel, ValidationError from litellm.integrations.galileo import GalileoObserve @@ -905,3 +909,78 @@ async def test_galileo_async_log_success_appends_and_flushes(galileo_v2_env): assert "/ingest/traces/" in flushed_url["url"] assert logger.in_memory_records == [] + + +class RealtimeEvent(BaseModel): + type: str + + +@dataclass(frozen=True) +class ThirdPartyEvent: + model_dump: object + + +@pytest.mark.parametrize( + ("event", "expected_output"), + [ + (RealtimeEvent(type="response.done"), '[{"type": "response.done"}]'), + (ThirdPartyEvent(model_dump=lambda: {"type": "custom"}), '[{"type": "custom"}]'), + (ThirdPartyEvent(model_dump=lambda: [Decimal("2.5")]), '[["2.5"]]'), + (Decimal("1.5"), '["1.5"]'), + ], +) +def test_galileo_realtime_output_serializes_events_that_json_cannot_encode( + galileo_v2_env: None, event: object, expected_output: str +) -> None: + assert GalileoObserve().get_output_str_from_response([event], {"call_type": "_arealtime"}) == expected_output + + +@pytest.mark.parametrize("model_dump", [None, "not callable"]) +def test_galileo_realtime_output_with_a_model_dump_that_cannot_be_called_raises_a_validation_error( + galileo_v2_env: None, model_dump: object +) -> None: + with pytest.raises(ValidationError) as raised: + GalileoObserve().get_output_str_from_response( + [ThirdPartyEvent(model_dump=model_dump)], {"call_type": "_arealtime"} + ) + + assert "input_value" not in str(raised.value) + + +@dataclass(frozen=True) +class ThirdPartyMessage: + json: object + + +def _chat_response_carrying(message: object) -> ModelResponse: + response: Final = ModelResponse(choices=[Choices(message=Message(content="replaced", role="assistant"))]) + response.choices[0].message = message + return response + + +@pytest.mark.parametrize( + ("message", "expected_output"), + [ + ( + ThirdPartyMessage(json=lambda: '{"role":"assistant","content":"hi"}'), + '{"role": "assistant", "content": "hi"}', + ), + ( + ThirdPartyMessage(json=lambda: {"role": "assistant", "content": "hi"}), + '{"role": "assistant", "content": "hi"}', + ), + (ThirdPartyMessage(json=lambda: '"just text"'), "just text"), + (ThirdPartyMessage(json=lambda: None), ""), + ({"role": "assistant", "content": "hi"}, '{"role": "assistant", "content": "hi"}'), + ("plain reply", "plain reply"), + (None, ""), + ], +) +def test_galileo_output_str_of_a_chat_response_whose_message_is_not_a_litellm_message( + galileo_v2_env: None, message: object, expected_output: str +) -> None: + output: Final = GalileoObserve().get_output_str_from_response( + _chat_response_carrying(message), {"call_type": "acompletion"} + ) + + assert output == expected_output diff --git a/tests/unit/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py index 175bd95c263..2ce3561d441 100644 --- a/tests/unit/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -6762,3 +6762,37 @@ class TestOpenTelemetryNonInferenceUsage(unittest.TestCase): self.assertEqual( self._time_per_output_token_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 1 ) + + +def _raw_response_span_attributes(original_response: str) -> dict[str, object]: + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider() + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + span = tracer_provider.get_tracer(__name__).start_span("raw_gen_ai_request") + + OpenTelemetry(tracer_provider=tracer_provider).set_raw_request_attributes( + span, + {"litellm_params": {"custom_llm_provider": "vertex_ai"}, "original_response": original_response}, + None, + ) + span.end() + + return dict(span_exporter.get_finished_spans()[0].attributes or {}) + + +@pytest.mark.parametrize( + ("original_response", "expected"), + [ + ('{"id": "r1", "model": "m"}', {"llm.vertex_ai.id": "r1", "llm.vertex_ai.model": "m"}), + ("{}", {}), + ("not json", {"llm.vertex_ai.stringified_raw_response": "not json"}), + ("[1, 2]", {}), + ('"text"', {}), + ("7", {}), + ("null", {}), + ], +) +def test_set_raw_request_attributes_stamps_only_json_object_responses( + original_response: str, expected: dict[str, object] +): + assert _raw_response_span_attributes(original_response) == expected diff --git a/tests/unit/integrations/test_opik_utils.py b/tests/unit/integrations/test_opik_utils.py index a4250acf1dc..376c70a4ed9 100644 --- a/tests/unit/integrations/test_opik_utils.py +++ b/tests/unit/integrations/test_opik_utils.py @@ -4,7 +4,7 @@ import uuid from datetime import datetime, timezone from unittest.mock import patch -from litellm.integrations.opik.utils import create_uuid7 +from litellm.integrations.opik.utils import create_uuid7, get_traces_and_spans_from_payload def _timestamp_ms(uuid_str: str) -> int: @@ -27,3 +27,22 @@ def test_create_uuid7_encodes_timestamp_in_milliseconds(): value = create_uuid7() assert _timestamp_ms(value) == int(fixed.timestamp() * 1000) + + +def test_queued_opik_payloads_are_split_into_traces_and_spans_without_their_null_fields(): + traces, spans = get_traces_and_spans_from_payload( + [ + {"id": "trace-1", "name": "chat.completion", "thread_id": None, "input": {"kept": None}}, + {"id": "span-1", "type": "llm", "parent_span_id": None, "tags": [], "total_cost": 0}, + {"id": "span-2", "type": None}, + ] + ) + + assert (traces, spans) == ( + [{"id": "trace-1", "name": "chat.completion", "input": {"kept": None}}], + [{"id": "span-1", "type": "llm", "tags": [], "total_cost": 0}, {"id": "span-2"}], + ) + + +def test_an_empty_opik_queue_has_no_traces_and_no_spans(): + assert get_traces_and_spans_from_payload([]) == ([], []) diff --git a/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index a0766ac3d58..726ce2f7f1d 100644 --- a/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -1,5 +1,5 @@ import logging -from collections.abc import Iterator, Mapping +from collections.abc import Callable, Iterable, Iterator, Mapping from dataclasses import dataclass, field from types import MappingProxyType from typing import Literal, Protocol @@ -33,6 +33,7 @@ from litellm.types.utils import ( ) from litellm.types.vector_stores import ( VectorStoreResultContent, + VectorStoreSearchFailure, VectorStoreSearchResponse, VectorStoreSearchResult, ) @@ -457,6 +458,166 @@ async def test_a_failing_vector_store_is_reported_on_the_streaming_chunk(registr ) +@dataclass +class ThirdPartyMessage: + provider_specific_fields: dict[str, object] | None = None + + +@dataclass +class ThirdPartyChoice: + message: ThirdPartyMessage | None = None + delta: ThirdPartyMessage | None = None + + +@dataclass +class ThirdPartyResponse: + choices: object + + +_SEARCH_FAILURES = ( + VectorStoreSearchFailure(vector_store_id="vs-broken", custom_llm_provider="bedrock", error="search timed out"), +) + + +def _logging_obj_with_search_failures() -> FakeLoggingObj: + logging_obj = FakeLoggingObj({}) + logging_obj.model_call_details["vector_store_search_failures"] = _SEARCH_FAILURES + return logging_obj + + +@pytest.mark.asyncio +async def test_search_failures_join_the_provider_fields_the_message_already_carries() -> None: + existing_fields: dict[str, object] = {"citations": ["doc-1"]} + response = ModelResponse( + choices=[Choices(message=Message(content="an answer", provider_specific_fields=existing_fields))] + ) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": _logging_obj_with_search_failures()}, + response=response, + call_type=CallTypes.acompletion, + ) + + assert returned is response + assert _first_message(response).provider_specific_fields is existing_fields + assert existing_fields == {"citations": ["doc-1"], "vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +async def test_search_failures_join_the_provider_fields_the_streaming_delta_already_carries() -> None: + existing_fields: dict[str, object] = {"citations": ["doc-1"]} + chunk = ModelResponseStream( + choices=[StreamingChoices(delta=Delta(content="an answer", provider_specific_fields=existing_fields))] + ) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert chunk.choices[0].delta.provider_specific_fields is existing_fields + assert existing_fields == {"citations": ["doc-1"], "vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("as_choices", [list, tuple, iter]) +async def test_a_chunk_that_only_looks_like_a_chat_completion_chunk_is_annotated_too( + as_choices: Callable[[list[ThirdPartyChoice]], Iterable[ThirdPartyChoice]], +) -> None: + delta = ThirdPartyMessage() + chunk = ThirdPartyResponse(choices=as_choices([ThirdPartyChoice(delta=None), ThirdPartyChoice(delta=delta)])) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert delta.provider_specific_fields == {"vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("choices", [None, [], ()]) +async def test_a_response_without_choices_is_returned_untouched( + choices: object, warnings: list[logging.LogRecord] +) -> None: + response = ThirdPartyResponse(choices=choices) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": _logging_obj_with_search_failures()}, + response=response, + call_type=CallTypes.acompletion, + ) + + assert returned is response + assert warnings == [] + + +@pytest.mark.asyncio +async def test_a_chunk_whose_choices_cannot_be_iterated_is_logged_and_passed_through( + warnings: list[logging.LogRecord], +) -> None: + chunk = ThirdPartyResponse(choices=7) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert [record.levelname for record in warnings] == ["ERROR"] + assert warnings[0].getMessage().startswith("Error adding search results to streaming chunk: ") + assert "input_value" not in warnings[0].getMessage() + + +@pytest.mark.asyncio +async def test_the_search_receives_the_requests_own_metadata_object(registry_with: RegisterStores) -> None: + registry_with("vs-router") + router = RecordingRouter() + metadata = {"user_api_key_team_id": "team-a"} + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-router"], + FakeLoggingObj(metadata), + ) + + assert [call["metadata"] is metadata for call in router.calls] == [True] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("litellm_params", [None, "metadata", ["metadata"], {"other": "value"}]) +async def test_a_request_without_metadata_in_its_litellm_params_searches_with_empty_metadata( + registry_with: RegisterStores, litellm_params: object +) -> None: + registry_with("vs-router") + router = RecordingRouter() + logging_obj = FakeLoggingObj({}) + logging_obj.model_call_details["litellm_params"] = litellm_params + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-router"], + logging_obj, + ) + + assert [call["metadata"] for call in router.calls] == [{}] + + @pytest.mark.asyncio async def test_error_mode_fails_the_request_instead_of_answering_without_the_knowledge_base( registry_with: RegisterStores, diff --git a/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py b/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py index 6440825e135..a419e73a2e1 100644 --- a/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py +++ b/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock import pytest +from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator from litellm.llms.a2a.chat.transformation import A2AConfig from litellm.types.utils import ModelResponse @@ -85,3 +86,38 @@ def test_transform_request_tags_the_message_with_its_kind(optional_params: dict) ) assert request["params"]["message"]["kind"] == "message" + + +def test_get_model_response_iterator_parses_the_agent_stream(): + iterator = A2AConfig().get_model_response_iterator( + streaming_response=iter( + [ + '{"jsonrpc":"2.0","id":"1","result":{"kind":"task","status":{"state":"completed"},' + '"artifacts":[{"parts":[{"kind":"text","text":"7"}]}]}}' + ] + ), + sync_stream=True, + ) + + chunk = next(iterator) + + assert isinstance(iterator, A2AModelResponseIterator) + assert chunk["text"] == "7" + assert chunk["finish_reason"] == "stop" + + +def test_resolve_agent_config_from_registry_returns_the_explicit_headers_object_untouched(): + headers: dict[str, object] = {"X-Test": "value"} + optional_params: dict[str, object] = {"stream": True} + + resolved = A2AConfig.resolve_agent_config_from_registry( + agent_name="test-agent", + api_base="http://explicit.example", + api_key="explicit-key", + headers=headers, + optional_params=optional_params, + ) + + assert resolved[2] is headers + assert resolved[:2] == ("http://explicit.example", "explicit-key") + assert optional_params == {"stream": True} diff --git a/tests/unit/llms/aiohttp_openai/__init__.py b/tests/unit/llms/aiohttp_openai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/aiohttp_openai/chat/__init__.py b/tests/unit/llms/aiohttp_openai/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/aiohttp_openai/chat/test_transformation.py b/tests/unit/llms/aiohttp_openai/chat/test_transformation.py new file mode 100644 index 00000000000..253f9dbe746 --- /dev/null +++ b/tests/unit/llms/aiohttp_openai/chat/test_transformation.py @@ -0,0 +1,101 @@ +from unittest.mock import AsyncMock, Mock + +import pytest +from aiohttp import ClientResponse +from pydantic import ValidationError + +from litellm.llms.aiohttp_openai.chat.transformation import AiohttpOpenAIChatConfig +from litellm.types.utils import ModelResponse + + +async def _transform(body: object) -> ModelResponse: + raw_response = Mock(spec=ClientResponse) + raw_response.json = AsyncMock(return_value=body) + return await AiohttpOpenAIChatConfig().transform_response( + model="gpt-4o", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +async def test_transform_response_copies_the_openai_body_onto_the_model_response(): + response = await _transform( + { + "id": "chatcmpl-1", + "created": 1700000000, + "model": "gpt-4o-2024", + "object": "chat.completion", + "system_fingerprint": "fp_1", + "choices": [ + {"index": 0, "finish_reason": "length", "message": {"role": "assistant", "content": "Hi"}}, + { + "index": 1, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + }, + ], + } + ) + + assert response.id == "chatcmpl-1" + assert response.created == 1700000000 + assert response.model == "gpt-4o-2024" + assert response.object == "chat.completion" + assert response.system_fingerprint == "fp_1" + assert [(choice.index, choice.finish_reason, choice.message.content) for choice in response.choices] == [ + (0, "length", "Hi"), + (1, "tool_calls", None), + ] + assert response.choices[1].message.tool_calls[0].function.name == "lookup" + + +async def test_transform_response_fills_choice_defaults_for_an_empty_choice(): + response = await _transform({"choices": [{}]}) + + assert [(choice.index, choice.finish_reason, choice.message.role) for choice in response.choices] == [ + (0, "stop", "assistant") + ] + assert response.id is None + + +async def test_transform_response_returns_no_choices_for_an_empty_choices_list(): + response = await _transform({"id": "chatcmpl-1", "choices": []}) + + assert response.choices == [] + assert response.id == "chatcmpl-1" + + +@pytest.mark.parametrize( + "body", + [ + {"id": "chatcmpl-1"}, + {"choices": None}, + {"choices": 7}, + {"choices": "secret-completion"}, + {"choices": {"message": "secret-completion"}}, + {"choices": ["secret-completion"]}, + {"choices": [{"index": 0}, ["secret-completion"]]}, + ], +) +async def test_transform_response_rejects_choices_that_are_not_a_list_of_objects_without_echoing_them(body: object): + with pytest.raises(ValidationError) as exc_info: + await _transform(body) + + assert "secret-completion" not in str(exc_info.value) + + +async def test_transform_response_rejects_a_choice_whose_message_is_not_an_object(): + with pytest.raises(ValidationError, match="validation error for Choices"): + await _transform({"choices": [{"message": "text"}]}) diff --git a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py index f385c2f2211..e9736d47537 100644 --- a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py +++ b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py @@ -12,6 +12,8 @@ import time import httpx import pytest +from openai.types.file_deleted import FileDeleted +from pydantic import ValidationError from unittest.mock import Mock, patch from litellm.llms.anthropic.files.transformation import ( @@ -196,6 +198,45 @@ class TestAnthropicFilesConfig: assert result.purpose == "messages" assert result.status == "uploaded" + def test_create_file_response_maps_the_anthropic_file_onto_an_openai_file(self) -> None: + result = self.config.transform_create_file_response( + model=None, + raw_response=httpx.Response( + 200, + json={ + "id": "file-abc123", + "type": "file", + "filename": "document.pdf", + "mime_type": "application/pdf", + "size_bytes": 12345, + "created_at": "2025-01-15T10:30:00Z", + }, + ), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == OpenAIFileObject( + id="file-abc123", + bytes=12345, + created_at=1736937000, + filename="document.pdf", + object="file", + purpose="messages", + status="uploaded", + status_details=None, + ) + + @pytest.mark.parametrize("body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7"]) + def test_create_file_response_rejects_a_body_that_is_not_a_json_object(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_create_file_response( + model=None, + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_retrieve_file_request(self): url, params = self.config.transform_retrieve_file_request( file_id="file-abc123", @@ -245,6 +286,43 @@ class TestAnthropicFilesConfig: assert result.id == "file-abc123" assert result.bytes == 5000 + def test_retrieve_file_response_maps_the_anthropic_file_onto_an_openai_file(self) -> None: + result = self.config.transform_retrieve_file_response( + raw_response=httpx.Response( + 200, + json={ + "id": "file-abc123", + "type": "file", + "filename": "document.pdf", + "mime_type": "application/pdf", + "size_bytes": 5000, + "created_at": "2025-06-01T12:00:00Z", + }, + ), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == OpenAIFileObject( + id="file-abc123", + bytes=5000, + created_at=1748779200, + filename="document.pdf", + object="file", + purpose="messages", + status="uploaded", + status_details=None, + ) + + @pytest.mark.parametrize("body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7"]) + def test_retrieve_file_response_rejects_a_body_that_is_not_a_json_object(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_retrieve_file_response( + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_delete_file_request(self): url, params = self.config.transform_delete_file_request( file_id="file-abc123", @@ -271,6 +349,35 @@ class TestAnthropicFilesConfig: assert result.deleted is True assert result.object == "file" + @pytest.mark.parametrize( + ("payload", "expected_id"), + [ + ({"id": "file-abc123", "type": "file_deleted"}, "file-abc123"), + ({"id": "file-abc123", "unknown": [1, {"nested": None}]}, "file-abc123"), + ({"type": "error", "error": {"type": "not_found_error", "message": "File not found"}}, ""), + ({}, ""), + ], + ) + def test_delete_file_response_reports_the_id_anthropic_returned(self, payload: object, expected_id: str) -> None: + result = self.config.transform_delete_file_response( + raw_response=httpx.Response(200, json=payload), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == FileDeleted(id=expected_id, deleted=True, object="file") + + @pytest.mark.parametrize( + "body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7", b'{"id": null}', b'{"id": 7}'] + ) + def test_delete_file_response_rejects_a_body_without_a_string_id(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_delete_file_response( + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_list_files_request(self): url, params = self.config.transform_list_files_request( purpose=None, diff --git a/tests/unit/llms/azure/test_azure_speech_audio_transcription.py b/tests/unit/llms/azure/test_azure_speech_audio_transcription.py index b447645bae8..4330f05e79c 100644 --- a/tests/unit/llms/azure/test_azure_speech_audio_transcription.py +++ b/tests/unit/llms/azure/test_azure_speech_audio_transcription.py @@ -3,6 +3,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.azure.audio_transcription.transformation import ( @@ -226,3 +227,22 @@ def test_azure_speech_transcription_routes_through_provider_config(monkeypatch): assert audio_handler.call_args.kwargs["custom_llm_provider"] == "azure" +def test_azure_speech_audio_transcription_response_keeps_the_raw_payload_as_hidden_params(): + payload = {"RecognitionStatus": "Success", "NBest": [{"Lexical": "hello world", "Confidence": 0.9}], "Offset": 3} + + response = AzureSpeechAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + assert response.text == "hello world" + assert response._hidden_params == payload + + +@pytest.mark.parametrize("payload", [7, "spoken secret", [{"DisplayText": "spoken secret"}]]) +def test_azure_speech_audio_transcription_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + AzureSpeechAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + assert "spoken secret" not in str(exc_info.value) diff --git a/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py b/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py new file mode 100644 index 00000000000..6ca70324162 --- /dev/null +++ b/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py @@ -0,0 +1,14 @@ +from litellm.llms.bedrock.image_edit.stability_transformation import BedrockStabilityImageEditConfig + + +def test_transform_image_edit_request_returns_the_json_body_and_no_files(): + request = BedrockStabilityImageEditConfig().transform_image_edit_request( + model="stability.stable-image-inpaint-v1:0", + prompt="add a red hat", + image=b"\x89PNG-bytes", + image_edit_optional_request_params={}, + litellm_params={}, + headers={}, + ) + + assert request == ({"output_format": "png", "prompt": "add a red hat", "image": "iVBORy1ieXRlcw=="}, {}) diff --git a/tests/unit/llms/bedrock/image_generation/__init__.py b/tests/unit/llms/bedrock/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py b/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py new file mode 100644 index 00000000000..eb8b9de6721 --- /dev/null +++ b/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py @@ -0,0 +1,57 @@ +import pytest + +from litellm.llms.bedrock.image_generation.amazon_titan_transformation import AmazonTitanImageGenerationConfig + + +@pytest.mark.parametrize( + ("non_default_params", "expected"), + [ + ( + {"size": "1024x512", "n": 2, "quality": "hd"}, + {"imageGenerationConfig": {"width": 1024, "height": 512, "numberOfImages": 2, "quality": "premium"}}, + ), + ({"quality": "low"}, {"imageGenerationConfig": {"quality": "standard"}}), + ({"quality": "auto", "size": None, "n": None}, {}), + ], +) +def test_map_openai_params_builds_the_image_generation_config( + non_default_params: dict[str, object], expected: dict[str, object] +): + optional_params = AmazonTitanImageGenerationConfig.map_openai_params( + non_default_params=non_default_params, optional_params={} + ) + + assert optional_params == expected + + +@pytest.mark.parametrize( + ("optional_params", "expected"), + [ + ( + {"imageGenerationConfig": {"width": 512, "numberOfImages": 2}, "negativeText": "blurry"}, + { + "taskType": "TEXT_IMAGE", + "textToImageParams": {"text": "a cat", "negativeText": "blurry"}, + "imageGenerationConfig": {"width": 512, "numberOfImages": 2}, + }, + ), + ( + {"taskType": "COLOR_GUIDED_GENERATION", "negativeText": ""}, + { + "taskType": "COLOR_GUIDED_GENERATION", + "textToImageParams": {"text": "a cat"}, + "imageGenerationConfig": {}, + }, + ), + ], +) +def test_transform_request_body_builds_the_titan_request( + optional_params: dict[str, object], expected: dict[str, object] +): + request_body = AmazonTitanImageGenerationConfig.transform_request_body( + text="a cat", optional_params=optional_params + ) + + assert request_body == expected + assert list(request_body["textToImageParams"]) == list(expected["textToImageParams"]) + assert optional_params == {} diff --git a/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py b/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py index c40830b238f..495e0f56c84 100644 --- a/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py +++ b/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py @@ -8,7 +8,9 @@ forward_client_headers_to_llm_api were not being passed to Bedrock rerank provid import json from unittest.mock import AsyncMock, MagicMock, Mock, patch +import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.bedrock.base_aws_llm import Boto3CredentialsInfo @@ -506,3 +508,60 @@ async def test_bedrock_rerank_records_llm_api_duration(): assert response._hidden_params["litellm_overhead_time_ms"] is not None assert response._hidden_params["_response_ms"] >= response._hidden_params["litellm_overhead_time_ms"] + + +RERANK_MODEL_ARN = "arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0" +NON_OBJECT_BODIES = [7, "sensitive-document", [{"index": 0, "relevanceScore": 0.9, "id": "sensitive-document"}]] + + +def _rerank_through_transport(payload: object, *, is_async: bool): + transport = httpx.MockTransport(lambda request: httpx.Response(200, json=payload)) + client = ( + AsyncHTTPHandler(transport=transport) if is_async else HTTPHandler(client=httpx.Client(transport=transport)) + ) + return BedrockRerankHandler().rerank( + model=RERANK_MODEL_ARN, + query=test_query, + documents=test_documents, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "example-secret", + "aws_region_name": "us-west-2", + }, + logging_obj=Mock(), + _is_async=is_async, + client=client, + ) + + +def _assert_is_the_bedrock_ranking(response: litellm.RerankResponse) -> None: + assert response.results == [ + {"index": 2, "relevance_score": 0.95}, + {"index": 0, "relevance_score": 0.1}, + {"index": 1, "relevance_score": 0.05}, + ] + assert response.meta == {"billed_units": {"search_units": 1}, "tokens": {}} + + +def test_bedrock_rerank_maps_the_upstream_ranking(): + _assert_is_the_bedrock_ranking(_rerank_through_transport(bedrock_rerank_response, is_async=False)) + + +async def test_bedrock_arerank_maps_the_upstream_ranking(): + _assert_is_the_bedrock_ranking(await _rerank_through_transport(bedrock_rerank_response, is_async=True)) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +def test_bedrock_rerank_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _rerank_through_transport(payload, is_async=False) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +async def test_bedrock_arerank_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + await _rerank_through_transport(payload, is_async=True) + + assert "sensitive-document" not in str(exc_info.value) diff --git a/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py b/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py index 20bf65ee385..4526c428f1a 100644 --- a/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py +++ b/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py @@ -9,10 +9,12 @@ the sharded CI (coverage collection runs against this tree). import json import os +import httpx import pytest from unittest.mock import AsyncMock, patch, MagicMock import litellm +from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.llms.bedrock.search.transformation import ( AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION, AgentCoreSearchConfig, @@ -635,3 +637,44 @@ class TestAgentCoreSearchEdgeCases: monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map()) assert search_provider_cost_per_query(model="agentcore/search", custom_llm_provider="agentcore") == (0.0, 0.0) + + +def _transform_text(text: str) -> SearchResponse: + return AgentCoreSearchConfig().transform_search_response( + raw_response=httpx.Response(200, text=text), + logging_obj=MagicMock(), + ) + + +def _as_tuples(response: SearchResponse) -> list[tuple[str, str, str, str | None, str | None]]: + return [(r.title, r.url, r.snippet, r.date, r.last_updated) for r in response.results] + + +EXPECTED_MCP_RESULTS = [ + ("Test Result 1", "https://example.com/1", "Snippet for result 1", "2026-06-16", None), + ("Test Result 2", "https://example.com/2", "Snippet for result 2", None, None), +] + + +@pytest.mark.parametrize( + "text", + [ + json.dumps(_mcp_response_body()), + f"event: message\ndata: {json.dumps(_mcp_response_body())}\n\n", + f"data: 5\n\ndata: [1]\n\ndata: {json.dumps(_mcp_response_body())}\n\n", + json.dumps({"result": {"content": [{"type": "text", "text": json.dumps({"results": MCP_RESULTS})}]}}), + json.dumps({"result": {"content": [{"type": "text", "text": "prose"}], "structuredContent": MCP_RESULTS}}), + ], +) +def test_transform_search_response_reads_results_from_a_real_http_response(text: str): + assert _as_tuples(_transform_text(text)) == EXPECTED_MCP_RESULTS + + +@pytest.mark.parametrize( + "block_text", + ["null", "5", "true", '"text"', "{}", '{"results": null}', '{"results": "text"}', '["scalar", 5, null]'], +) +def test_transform_search_response_ignores_text_blocks_without_result_objects(block_text: str): + body = {"result": {"content": [{"type": "text", "text": block_text}]}} + + assert _transform_text(json.dumps(body)).results == [] diff --git a/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py b/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py index d6e2c4a3e06..c947967417d 100644 --- a/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py +++ b/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py @@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.black_forest_labs.image_generation.transformation import ( @@ -355,3 +356,48 @@ class TestBlackForestLabsImageGenerationTransformation: config = get_black_forest_labs_image_generation_config("flux-pro-1.1") assert isinstance(config, BlackForestLabsImageGenerationConfig) + + +def _transform(payload: object) -> ImageResponse: + return BlackForestLabsImageGenerationConfig().transform_image_generation_response( + model="flux-pro-1.1", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("result", "expected_urls"), + [ + ({"sample": "https://bfl.example/a.png"}, ["https://bfl.example/a.png"]), + ( + ["https://bfl.example/a.png", {"url": "https://bfl.example/b.png"}, {"seed": 1}, 7], + ["https://bfl.example/a.png", "https://bfl.example/b.png"], + ), + ], +) +def test_transform_image_generation_response_reads_urls_from_the_result(result: object, expected_urls: list[str]): + response = _transform({"status": "Ready", "result": result}) + + assert [image.url for image in response.data] == expected_urls + + +@pytest.mark.parametrize("payload", [{}, {"result": None}, {"result": "https://bfl.example/a.png"}, {"result": []}]) +def test_transform_image_generation_response_without_a_url_is_a_provider_error(payload: dict[str, object]): + with pytest.raises(BlackForestLabsError) as exc_info: + _transform(payload) + + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize("payload", [7, "https://bfl.example/a.png", [{"sample": "https://bfl.example/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "bfl.example" not in str(exc_info.value) diff --git a/tests/unit/llms/clarifai/__init__.py b/tests/unit/llms/clarifai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/clarifai/chat/__init__.py b/tests/unit/llms/clarifai/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/clarifai/chat/test_transformation.py b/tests/unit/llms/clarifai/chat/test_transformation.py new file mode 100644 index 00000000000..e81b2a9558d --- /dev/null +++ b/tests/unit/llms/clarifai/chat/test_transformation.py @@ -0,0 +1,70 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.clarifai.chat.transformation import ClarifaiConfig +from litellm.llms.openai.common_utils import OpenAIError +from litellm.types.utils import ModelResponse + + +def _transform(raw_response: httpx.Response) -> ModelResponse: + return ClarifaiConfig().transform_response( + model="user.app.model", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_response_builds_the_model_response_from_the_body(): + response = _transform( + httpx.Response( + 200, + json={ + "id": "chatcmpl-1", + "created": 1700000000, + "model": "upstream-model", + "system_fingerprint": "fp_1", + "choices": [{"index": 0, "finish_reason": "length", "message": {"role": "assistant", "content": "Hi"}}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + "vendor_field": {"kept": True}, + }, + ) + ) + + assert response.id == "chatcmpl-1" + assert response.created == 1700000000 + assert response.model == "clarifai/user.app.model" + assert response.system_fingerprint == "fp_1" + assert [(choice.finish_reason, choice.message.content) for choice in response.choices] == [("length", "Hi")] + assert response.usage.model_dump()["total_tokens"] == 5 + assert response.vendor_field == {"kept": True} + + +def test_transform_response_keeps_a_missing_model_unset(): + response = _transform(httpx.Response(200, json={"choices": [{"message": {"content": "Hi"}}]})) + + assert response.model is None + assert response.choices[0].message.content == "Hi" + + +@pytest.mark.parametrize("body", [b'["prompt-text"]', b'"prompt-text"', b"7", b"null"]) +def test_transform_response_rejects_a_body_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "prompt-text" not in str(exc_info.value) + + +def test_transform_response_reports_an_unparseable_body_as_an_openai_error(): + with pytest.raises(OpenAIError, match="Failed to parse Clarifai response") as exc_info: + _transform(httpx.Response(502, content=b"bad gateway")) + + assert exc_info.value.status_code == 502 diff --git a/tests/unit/llms/claude_code/harness/test_transformation.py b/tests/unit/llms/claude_code/harness/test_transformation.py index 6348b4f6a3e..1ae4497ef18 100644 --- a/tests/unit/llms/claude_code/harness/test_transformation.py +++ b/tests/unit/llms/claude_code/harness/test_transformation.py @@ -318,6 +318,29 @@ def test_parse_skips_subagent_messages_and_maps_errors(): ] +def test_user_message_yields_only_its_tool_result_objects(): + line = { + "type": "user", + "message": { + "content": [ + "plain text", + None, + ["tool_result"], + {"type": "text", "text": "not a result"}, + {"type": "tool_result", "tool_use_id": "t1", "content": "ok"}, + {"type": "tool_result", "content": [{"type": "text", "text": "a"}, "b"], "is_error": 1}, + {"type": "tool_result", "tool_use_id": 7, "content": None, "extra": {"kept": "out"}}, + ] + }, + } + + assert ClaudeCodeHarnessConfig().transform_stream_line(line, ClaudeCodeStreamState()) == [ + ToolResult(id="t1", output="ok", is_error=False), + ToolResult(id="", output="a\nb", is_error=True), + ToolResult(id="7", output="", is_error=False), + ] + + def test_parse_thinking_and_mcp_tools(): state = ClaudeCodeStreamState() msg = { diff --git a/tests/unit/llms/codex/harness/test_transformation.py b/tests/unit/llms/codex/harness/test_transformation.py index 6371f74fba6..6df322db1ae 100644 --- a/tests/unit/llms/codex/harness/test_transformation.py +++ b/tests/unit/llms/codex/harness/test_transformation.py @@ -11,7 +11,7 @@ from pathlib import Path from typing import Optional import pytest -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from litellm.harness.context import SessionContext from litellm.harness.errors import HarnessError, HarnessInstallFailed, OptionsMismatch @@ -272,6 +272,36 @@ def test_parse_file_change_and_mcp_and_web_search(): assert events[0].name == "web_search" and events[0].input == {"query": "litellm"} +@pytest.mark.parametrize( + ("changes", "output"), + [ + ([{"path": "a.txt", "kind": "add"}, {"path": "b.txt", "kind": "update", "diff": "@@"}], "add a.txt\nupdate b.txt"), + ([{"path": "only-path.txt"}, {"kind": "delete"}, {}], "only-path.txt\ndelete\n"), + ([], ""), + (None, ""), + ], +) +def test_completed_file_change_lists_each_change(changes: object, output: str): + state = CodexStreamState(started={"i1"}) + item = {"id": "i1", "type": "file_change", "changes": changes, "status": "completed"} + + assert parse_event({"type": "item.completed", "item": item}, state) == [ + ToolResult(id="i1", output=output, is_error=False) + ] + + +@pytest.mark.parametrize( + "changes", + [["a.txt"], [{"path": "a.txt"}, None], "a.txt", {"a.txt": {"kind": "add"}}, 7], +) +def test_completed_file_change_rejects_changes_that_are_not_a_list_of_objects(changes: object): + state = CodexStreamState(started={"i1"}) + item = {"id": "i1", "type": "file_change", "changes": changes, "status": "completed"} + + with pytest.raises(ValidationError): + parse_event({"type": "item.completed", "item": item}, state) + + def test_parse_failed_command_is_error_and_unknown_events_ignored(): state = CodexStreamState() item = { diff --git a/tests/unit/llms/cohere/rerank/test_transformation.py b/tests/unit/llms/cohere/rerank/test_transformation.py new file mode 100644 index 00000000000..ac5acf08d0d --- /dev/null +++ b/tests/unit/llms/cohere/rerank/test_transformation.py @@ -0,0 +1,72 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.cohere.rerank.transformation import CohereRerankConfig +from litellm.types.rerank import RerankResponse + + +def _transform(payload: object) -> RerankResponse: + return CohereRerankConfig().transform_rerank_response( + model="rerank-english-v3.0", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=Mock(), + ) + + +def test_transform_rerank_response_keeps_the_cohere_payload(): + response = _transform( + { + "id": "rerank-1", + "results": [{"index": 1, "relevance_score": 0.9, "document": {"text": "sensitive-document"}}], + "meta": {"billed_units": {"search_units": 1}, "tokens": {"input_tokens": 4}}, + "warnings": ["ignored"], + } + ) + + assert response.model_dump() == { + "id": "rerank-1", + "results": [{"index": 1, "relevance_score": 0.9, "document": {"text": "sensitive-document"}}], + "meta": {"billed_units": {"search_units": 1}, "tokens": {"input_tokens": 4}}, + } + + +def test_transform_rerank_response_of_an_empty_object_has_no_results(): + assert _transform({}).model_dump() == {"id": None, "results": None, "meta": None} + + +@pytest.mark.parametrize("payload", [7, "sensitive-document", [{"id": "sensitive-document"}]]) +def test_transform_rerank_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", [{"id": 7}, {"results": "not-a-list"}, {"results": [{"index": 0}]}, {"meta": []}]) +def test_transform_rerank_response_rejects_malformed_fields(payload: dict[str, object]): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_map_cohere_rerank_params_returns_every_cohere_param(): + params = CohereRerankConfig().map_cohere_rerank_params( + non_default_params=None, + model="rerank-english-v3.0", + drop_params=False, + query="capital of france", + documents=["Paris", {"text": "Berlin"}], + top_n=1, + ) + + assert params == { + "query": "capital of france", + "documents": ["Paris", {"text": "Berlin"}], + "top_n": 1, + "rank_fields": None, + "return_documents": True, + "max_chunks_per_doc": None, + } diff --git a/tests/unit/llms/cometapi/image_generation/__init__.py b/tests/unit/llms/cometapi/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/cometapi/image_generation/test_transformation.py b/tests/unit/llms/cometapi/image_generation/test_transformation.py new file mode 100644 index 00000000000..0a14b955339 --- /dev/null +++ b/tests/unit/llms/cometapi/image_generation/test_transformation.py @@ -0,0 +1,54 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.cometapi.image_generation.transformation import CometAPIImageGenerationConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return CometAPIImageGenerationConfig().transform_image_generation_response( + model="dall-e-3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_each_image_object(): + response = _transform({"created": 1, "data": [{"url": "https://img.cometapi.com/a.png"}, {"b64_json": "QUJD"}, {}]}) + + assert [(image.url, image.b64_json) for image in response.data] == [ + ("https://img.cometapi.com/a.png", None), + (None, "QUJD"), + (None, None), + ] + + +@pytest.mark.parametrize("payload", [{}, {"created": 1}, {"data": []}, {"data": ""}, {"data": {}}, []]) +def test_transform_image_generation_response_without_images_has_no_data(payload: object): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["data", "https://img.cometapi.com/a.png"], + {"data": None}, + {"data": "https://img.cometapi.com/a.png"}, + {"data": {"url": "https://img.cometapi.com/a.png"}}, + {"data": ["https://img.cometapi.com/a.png"]}, + {"data": [{"url": "https://img.cometapi.com/a.png"}, 7]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "img.cometapi.com" not in str(exc_info.value) diff --git a/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py b/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py index 4466e5b8767..a913c79af31 100644 --- a/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py +++ b/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.dashscope.common_utils import DashScopeError @@ -347,3 +348,89 @@ class TestProviderConfigManagerDispatch: present_version_params=[], ) assert isinstance(cfg, DashScopeRerankConfig) + + +def _transform(payload: object, status_code: int = 200) -> RerankResponse: + return DashScopeRerankConfig().transform_rerank_response( + model="qwen3-rerank", + raw_response=httpx.Response(status_code, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + ("usage", "expected_total_tokens"), + [ + ({"total_tokens": 79}, 79), + ({}, None), + (None, None), + ], +) +def test_transform_rerank_response_reads_total_tokens_from_usage(usage: object, expected_total_tokens: int | None): + response = _transform({"id": "rerank-1", "results": [{"index": 0, "relevance_score": 0.5}], "usage": usage}) + + assert response.meta == { + "billed_units": {"total_tokens": expected_total_tokens}, + "tokens": {"input_tokens": expected_total_tokens}, + } + + +def test_transform_rerank_response_keeps_provider_id_and_drops_unknown_result_fields(): + response = _transform( + { + "id": "rerank-1", + "results": [{"index": 2, "relevance_score": 0.25, "document": {"text": "doc", "extra": 1}, "extra": 2}], + } + ) + + assert response.id == "rerank-1" + assert response.results == [{"index": 2, "relevance_score": 0.25, "document": {"text": "doc"}}] + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"index": 0, "relevance_score": 0.5}], "usage": "seventy nine"}, + {"results": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"results": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_result_without_index_raises_key_error(): + with pytest.raises(KeyError, match="index"): + _transform({"results": [{"relevance_score": 0.5}]}) + + +def test_transform_rerank_response_error_envelope_raises_with_the_provider_status_and_message(): + with pytest.raises(DashScopeError) as exc_info: + _transform({"code": "Throttling", "message": "slow down"}, status_code=429) + + assert exc_info.value.status_code == 429 + assert exc_info.value.message == "slow down" + + +def test_transform_rerank_response_without_results_raises_dashscope_error_naming_the_body(): + with pytest.raises(DashScopeError) as exc_info: + _transform({"id": "rerank-1"}, status_code=502) + + assert exc_info.value.status_code == 502 + assert exc_info.value.message == "No results in DashScope rerank response: {'id': 'rerank-1'}" + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"results": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py b/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py index 7c1f5256deb..366413df1c8 100644 --- a/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py +++ b/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py @@ -3,6 +3,7 @@ import os import pathlib from unittest.mock import MagicMock +import httpx import pytest @@ -525,3 +526,82 @@ def test_reconstruct_diarized_transcript_multiple_speaker_changes(): assert "Hello" in result assert "back" in result assert "Thanks" in result + + +def _deepgram_payload(alternative: dict[str, object], channel_fields: dict[str, object]) -> dict[str, object]: + return { + "metadata": {"duration": 2.5}, + "results": {"channels": [{"alternatives": [alternative], **channel_fields}]}, + } + + +def _transform_deepgram_response(payload: object) -> TranscriptionResponse: + return DeepgramAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + +@pytest.mark.parametrize( + ("channel_fields", "expected_language"), + [ + ({}, "en"), + ({"detected_language": None}, "en"), + ({"detected_language": ""}, "en"), + ({"detected_language": "fr"}, "fr"), + ({"detected_language": "de", "language_confidence": 0.98}, "de"), + ({"detected_language": ["es", "en"]}, ["es", "en"]), + ], +) +def test_transform_response_reports_the_detected_language_or_english( + channel_fields: dict[str, object], expected_language: object +): + response = _transform_deepgram_response(_deepgram_payload({"transcript": "bonjour"}, channel_fields)) + + assert response.text == "bonjour" + assert response["language"] == expected_language + assert response["duration"] == 2.5 + assert "words" not in response + + +@pytest.mark.parametrize( + ("words", "expected"), + [ + ([], []), + ("", []), + ({}, []), + ( + [{"word": "hello", "start": 0.0, "end": 0.5, "confidence": 0.9}], + [{"word": "hello", "start": 0.0, "end": 0.5}], + ), + ( + [{"word": None, "start": "0.1", "end": [2]}, {"word": "b", "start": 1, "end": 2}], + [{"word": None, "start": "0.1", "end": [2]}, {"word": "b", "start": 1, "end": 2}], + ), + ], +) +def test_transform_response_maps_words_to_openai_word_timestamps(words: object, expected: list[dict[str, object]]): + payload = _deepgram_payload({"transcript": "hello", "words": words}, {}) + + response = _transform_deepgram_response(payload) + + assert response["words"] == expected + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + "words", + [ + "not a list", + ["not an object"], + [{"word": "hello", "start": 0.0, "end": 0.5}, None], + [{"word": "hello", "start": 0.0, "end": 0.5}, ["nested"]], + [{"word": "hello", "start": 0.0}], + ], +) +def test_transform_response_wraps_malformed_words_with_the_raw_body(words: object): + raw_response = httpx.Response(200, json=_deepgram_payload({"transcript": "hello", "words": words}, {})) + + with pytest.raises(ValueError, match="Error transforming Deepgram response: ") as exc_info: + DeepgramAudioTranscriptionConfig().transform_audio_transcription_response(raw_response) + + assert str(exc_info.value).endswith(f"\nResponse: {raw_response.text}") diff --git a/tests/unit/llms/e2b/__init__.py b/tests/unit/llms/e2b/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/e2b/sandbox/__init__.py b/tests/unit/llms/e2b/sandbox/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/e2b/sandbox/test_transformation.py b/tests/unit/llms/e2b/sandbox/test_transformation.py new file mode 100644 index 00000000000..44c0626bc27 --- /dev/null +++ b/tests/unit/llms/e2b/sandbox/test_transformation.py @@ -0,0 +1,121 @@ +import json + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig + + +def _client_answering(body: bytes) -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(lambda request: httpx.Response(200, content=body))) + + +def _lines(*messages: object) -> list[str]: + return [json.dumps(message) for message in messages] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("created", "expected_domain", "expected_tokens"), + [ + ( + {"sandboxID": "sbx_1", "domain": "eu.e2b.app", "envdAccessToken": "envd", "trafficAccessToken": "traffic"}, + "eu.e2b.app", + ("envd", "traffic"), + ), + ({"sandboxID": "sbx_1"}, "e2b.app", (None, None)), + ({"sandboxID": "sbx_1", "domain": None, "envdAccessToken": "envd"}, "e2b.app", ("envd", None)), + ({"sandboxID": "sbx_1", "domain": ""}, "e2b.app", (None, None)), + ], +) +async def test_acreate_sandbox_builds_the_handle_from_the_create_response( + created: dict[str, object], expected_domain: str, expected_tokens: tuple[str | None, str | None] +): + handle = await E2BSandboxConfig().acreate_sandbox( + api_key="e2b_key", client=_client_answering(json.dumps(created).encode()) + ) + + assert (handle.id, handle.provider, handle.domain) == ("sbx_1", "e2b", expected_domain) + assert handle._hidden_params == { + "envd_access_token": expected_tokens[0], + "traffic_access_token": expected_tokens[1], + "api_key": "e2b_key", + "api_base": "https://api.e2b.app", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [b'["secret-token"]', b'"secret-token"', b"7", b"null"]) +async def test_acreate_sandbox_rejects_a_create_response_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + await E2BSandboxConfig().acreate_sandbox(api_key="e2b_key", client=_client_answering(body)) + + assert "secret-token" not in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_acreate_sandbox_requires_a_sandbox_id(): + with pytest.raises(KeyError, match="sandboxID"): + await E2BSandboxConfig().acreate_sandbox(api_key="e2b_key", client=_client_answering(b'{"domain": "e2b.app"}')) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("created", [{"sandboxID": 5}, {"sandboxID": "sbx_1", "domain": 5}]) +async def test_acreate_sandbox_rejects_non_string_handle_fields(created: dict[str, object]): + with pytest.raises(ValidationError, match="ContainerHandle"): + await E2BSandboxConfig().acreate_sandbox( + api_key="e2b_key", client=_client_answering(json.dumps(created).encode()) + ) + + +def test_parse_lines_maps_every_message_type_onto_the_result(): + result = E2BSandboxConfig._parse_lines( + [ + *_lines( + {"type": "stdout", "text": "1\n", "timestamp": 1}, + {"type": "stderr", "text": "warn\n"}, + {"type": "stdout"}, + {"type": "result", "png": "BASE64", "is_main_result": True}, + {"type": "error", "name": "ValueError", "value": "bad", "traceback": "tb", "ignored": 1}, + {"type": "number_of_executions", "execution_count": 3}, + {"type": "stdout", "text": "2\n"}, + None, + ), + "", + "not json", + ] + ) + + assert result.model_dump() == { + "stdout": "1\n2\n", + "stderr": "warn\n", + "results": [{"png": "BASE64", "is_main_result": True}], + "error": {"name": "ValueError", "value": "bad", "traceback": "tb"}, + "execution_count": 3, + "object": "code_execution", + } + + +def test_parse_lines_rejects_an_execution_count_that_is_not_a_number(): + with pytest.raises(ValidationError, match="execution_count"): + E2BSandboxConfig._parse_lines(_lines({"type": "number_of_executions", "execution_count": "many"})) + + +@pytest.mark.parametrize( + "message", + [ + ["secret-output"], + "secret-output", + 7, + False, + {"type": "stdout", "text": ["secret-output"]}, + {"type": "stderr", "text": {"secret-output": 1}}, + ], +) +def test_parse_lines_rejects_malformed_messages_without_echoing_them(message: object): + with pytest.raises(ValidationError) as exc_info: + E2BSandboxConfig._parse_lines(_lines({"type": "stdout", "text": "ok"}, message)) + + assert "secret-output" not in str(exc_info.value) diff --git a/tests/unit/llms/elevenlabs/audio_transcription/__init__.py b/tests/unit/llms/elevenlabs/audio_transcription/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py b/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py new file mode 100644 index 00000000000..52a10f0a327 --- /dev/null +++ b/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py @@ -0,0 +1,100 @@ +import httpx +import pytest + +from litellm.llms.elevenlabs.audio_transcription.transformation import ElevenLabsAudioTranscriptionConfig +from litellm.types.utils import TranscriptionResponse + + +def _transform(payload: object) -> TranscriptionResponse: + return ElevenLabsAudioTranscriptionConfig().transform_audio_transcription_response( + raw_response=httpx.Response(200, json=payload) + ) + + +def test_transform_audio_transcription_response_keeps_only_spoken_words(): + payload = { + "language_code": "en", + "text": "Hello world", + "words": [ + {"type": "word", "text": "Hello", "start": 0.0, "end": 0.4, "speaker_id": "speaker_0"}, + {"type": "spacing", "text": " ", "start": 0.4, "end": 0.5}, + {"type": "audio_event", "text": "(laughter)", "start": 0.5, "end": 0.9}, + {"type": "word", "text": "world", "start": 0.9, "end": 1.3}, + ], + } + + response = _transform(payload) + + assert response.text == "Hello world" + assert response["task"] == "transcribe" + assert response["language"] == "en" + assert response["words"] == [ + {"word": "Hello", "start": 0.0, "end": 0.4}, + {"word": "world", "start": 0.9, "end": 1.3}, + ] + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + ("word", "expected"), + [ + ({"type": "word"}, [{"word": "", "start": 0, "end": 0}]), + ({"type": "word", "text": None, "start": None, "end": None}, [{"word": None, "start": None, "end": None}]), + ({"type": "word", "text": 7, "start": "0.1", "end": [2]}, [{"word": 7, "start": "0.1", "end": [2]}]), + ({"text": "untyped"}, []), + ({"type": None, "text": "untyped"}, []), + ({}, []), + ], +) +def test_transform_audio_transcription_response_maps_one_word( + word: dict[str, object], expected: list[dict[str, object]] +): + assert _transform({"text": "t", "words": [word]})["words"] == expected + + +@pytest.mark.parametrize("words", [[], "", {}]) +def test_transform_audio_transcription_response_with_empty_words_has_empty_word_list(words: object): + assert _transform({"text": "t", "words": words})["words"] == [] + + +@pytest.mark.parametrize( + ("payload", "expected_text", "expected_language"), + [ + ({}, "", "unknown"), + ({"text": None, "language_code": None}, None, None), + ({"text": "bonjour", "language_code": "fr"}, "bonjour", "fr"), + ({"text": "hola", "language_code": ["es"]}, "hola", ["es"]), + ], +) +def test_transform_audio_transcription_response_without_words_key_has_no_word_list( + payload: dict[str, object], expected_text: str | None, expected_language: object +): + response = _transform(payload) + + assert response.text == expected_text + assert response["language"] == expected_language + assert "words" not in response + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "plain text", + 7, + {"text": ["not", "text"]}, + {"text": "t", "words": None}, + {"text": "t", "words": 7}, + {"text": "t", "words": "not a list"}, + {"text": "t", "words": ["not an object"]}, + {"text": "t", "words": [{"type": "word", "text": "ok"}, None]}, + ], +) +def test_transform_audio_transcription_response_wraps_malformed_payloads_with_the_raw_body(payload: object): + raw_response = httpx.Response(200, json=payload) + + with pytest.raises(ValueError, match="Error transforming ElevenLabs response: ") as exc_info: + ElevenLabsAudioTranscriptionConfig().transform_audio_transcription_response(raw_response=raw_response) + + assert str(exc_info.value).endswith(f"\nResponse: {raw_response.text}") diff --git a/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py new file mode 100644 index 00000000000..960d927851e --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py @@ -0,0 +1,56 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIFluxProV11UltraConfig().transform_image_generation_response( + model="fal-ai/flux-pro/v1.1-ultra", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_images_and_metadata(): + response = _transform( + { + "images": [{"url": "https://fal.media/a.png", "width": 2752, "height": 1536}, "https://fal.media/b.png"], + "seed": 42, + "timings": {"inference": 2.5}, + "has_nsfw_concepts": [False, False], + } + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert response.data[0].provider_specific_fields == {"width": 2752, "height": 1536} + assert response._hidden_params["seed"] == 42 + assert response._hidden_params["timings"] == {"inference": 2.5} + assert response._hidden_params["has_nsfw_concepts"] == [False, False] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": []}]) +def test_transform_image_generation_response_without_images_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + assert "timings" not in response._hidden_params + assert "has_nsfw_concepts" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) diff --git a/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py new file mode 100644 index 00000000000..716440004cd --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py @@ -0,0 +1,71 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.ideogram_v3_transformation import FalAIIdeogramV3Config +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIIdeogramV3Config().transform_image_generation_response( + model="fal-ai/ideogram/v3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_files_and_seed(): + response = _transform( + {"images": [{"url": "https://fal.media/a.png", "file_name": "a.png"}, "https://fal.media/b.png"], "seed": 42} + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert [image.b64_json for image in response.data] == [None, None] + assert response._hidden_params["seed"] == 42 + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + (" 1536x1024 ", "landscape_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "square_hd"), + ("axb", "square_hd"), + ({"width": 640, "height": 480, "unit": "px"}, {"width": 640, "height": 480}), + ({"width": "640"}, {"width": "640"}), + (512, 512), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: object, expected: object): + optional_params = FalAIIdeogramV3Config().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "url"}, + optional_params={}, + model="fal-ai/ideogram/v3", + drop_params=False, + ) + + assert optional_params == {"image_size": expected, "num_images": 2} diff --git a/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py new file mode 100644 index 00000000000..09591fe24d5 --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py @@ -0,0 +1,64 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.recraft_v3_transformation import FalAIRecraftV3Config +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIRecraftV3Config().transform_image_generation_response( + model="fal-ai/recraft/v3/text-to-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_file_objects_and_bare_urls(): + response = _transform( + {"images": [{"url": "https://fal.media/a.png", "content_type": "image/png"}, "https://fal.media/b.png", 7]} + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert [image.b64_json for image in response.data] == [None, None] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + ("576x1024", "portrait_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "square_hd"), + ("axb", "square_hd"), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: str, expected: object): + optional_params = FalAIRecraftV3Config().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "url"}, + optional_params={}, + model="fal-ai/recraft/v3/text-to-image", + drop_params=False, + ) + + assert optional_params == {"image_size": expected} diff --git a/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py new file mode 100644 index 00000000000..8f112374efd --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py @@ -0,0 +1,76 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.stable_diffusion_transformation import FalAIStableDiffusionConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIStableDiffusionConfig().transform_image_generation_response( + model="fal-ai/stable-diffusion-v35-medium", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_images_and_metadata(): + response = _transform( + { + "images": [{"url": "https://fal.media/a.png", "width": 1024}, "https://fal.media/b.png", 7], + "seed": 42, + "timings": {"inference": 2.5}, + "has_nsfw_concepts": [False, False], + } + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert response._hidden_params["seed"] == 42 + assert response._hidden_params["timings"] == {"inference": 2.5} + assert response._hidden_params["has_nsfw_concepts"] == [False, False] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + assert "timings" not in response._hidden_params + assert "has_nsfw_concepts" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + ("1024x576", "landscape_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "landscape_4_3"), + ("axb", "landscape_4_3"), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: str, expected: object): + optional_params = FalAIStableDiffusionConfig().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "b64_json"}, + optional_params={}, + model="fal-ai/stable-diffusion-v35-medium", + drop_params=False, + ) + + assert optional_params == {"image_size": expected, "num_images": 2, "output_format": "jpeg"} diff --git a/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py b/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py index a03b7708238..a3d1879f97c 100644 --- a/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py +++ b/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py @@ -8,6 +8,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.fireworks_ai.rerank.transformation import FireworksAIRerankConfig from litellm.types.rerank import RerankResponse @@ -341,3 +342,68 @@ class TestFireworksAIRerankTransform: assert headers["Authorization"] == "Bearer test-api-key" assert headers["Content-Type"] == "application/json" + + +def _transform(payload: object) -> RerankResponse: + return FireworksAIRerankConfig().transform_rerank_response( + model="fireworks_ai/fireworks/qwen3-reranker-8b", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + ("result", "expected"), + [ + ({"index": "1", "relevance_score": "0.5"}, {"index": 1, "relevance_score": 0.5}), + ({"index": 2, "relevance_score": 1}, {"index": 2, "relevance_score": 1.0}), + ], +) +def test_transform_rerank_response_converts_index_and_score_with_int_and_float( + result: dict[str, object], expected: dict[str, object] +): + assert _transform({"id": "rerank-1", "data": [result]}).results == [expected] + + +@pytest.mark.parametrize("usage_fields", [{}, {"usage": {}}]) +def test_transform_rerank_response_without_usage_counters_reports_zero_tokens(usage_fields: dict[str, object]): + response = _transform({"id": "rerank-1", "data": [{"index": 0, "relevance_score": 0.5}], **usage_fields}) + + assert response.meta == {"billed_units": {"search_units": 0}, "tokens": {"input_tokens": 0, "output_tokens": 0}} + + +def test_transform_rerank_response_keeps_provider_id_and_falls_back_to_results_key(): + response = _transform({"id": "rerank-1", "data": [], "results": [{"index": 3, "relevance_score": 0.25}]}) + + assert response.id == "rerank-1" + assert response.results == [{"index": 3, "relevance_score": 0.25}] + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": None}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"data": 7}, + {"data": ["not an object"]}, + {"data": [{"index": None, "relevance_score": 0.5}]}, + {"data": [{"index": 0, "relevance_score": [0.5]}]}, + {"id": 7, "data": [{"index": 0, "relevance_score": 0.5}]}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"data": [{"index": {"leaked": "document text"}, "relevance_score": 0.5}]}) + + assert "document text" not in str(exc_info.value) diff --git a/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py b/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py index d9509856759..2a69926414a 100644 --- a/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py +++ b/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py @@ -1,4 +1,8 @@ +from unittest.mock import Mock + import httpx +import pytest +from pydantic import ValidationError from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.llms.gemini.image_generation.transformation import GoogleImageGenConfig @@ -419,3 +423,47 @@ def test_gemini_image_generation_response_without_grounding_has_no_web_search_re ) assert getattr(result.usage, "web_search_requests", None) is None + + +def _transform_response(model: str, payload: object) -> ImageResponse: + return GoogleImageGenConfig().transform_image_generation_response( + model=model, + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize("predictions", [[], "", {}]) +def test_imagen_generation_response_with_empty_predictions_has_no_images(predictions: object): + assert _transform_response("gemini/imagen-4.0-generate-001", {"predictions": predictions}).data == [] + + +def test_imagen_generation_response_keeps_predictions_without_image_bytes(): + result = _transform_response( + "gemini/imagen-4.0-generate-001", + {"predictions": [{"bytesBase64Encoded": "first-image"}, {"mimeType": "image/png"}]}, + ) + + assert [image.b64_json for image in result.data or []] == ["first-image", None] + + +@pytest.mark.parametrize( + ("model", "payload"), + [ + ("gemini/imagen-4.0-generate-001", {"predictions": None}), + ("gemini/imagen-4.0-generate-001", {"predictions": "not a list"}), + ("gemini/imagen-4.0-generate-001", {"predictions": [{"bytesBase64Encoded": "a"}, "not an object"]}), + ("gemini-3.1-flash-image-preview", {"usageMetadata": "not an object"}), + ("gemini-3.1-flash-image-preview", {"usageMetadata": ["not an object"]}), + ], +) +def test_image_generation_response_rejects_malformed_payloads_without_echoing_them(model: str, payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform_response(model, payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py b/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py index 01fe66ca4c7..e9c65cbb591 100644 --- a/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py +++ b/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py @@ -12,6 +12,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm import LlmProviders from litellm.llms.gigachat.embedding.transformation import ( @@ -339,4 +340,64 @@ class TestGetErrorClass: ) assert isinstance(error, GigaChatEmbeddingError) assert error.status_code == 400 - assert error.message == "embedding failed" \ No newline at end of file + assert error.message == "embedding failed" + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return GigaChatEmbeddingConfig().transform_embedding_response( + model="Embeddings", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_moves_per_item_usage_into_the_response_usage(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "model": "Embeddings", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2], "usage": {"prompt_tokens": 3}}, + {"object": "embedding", "index": 1, "embedding": [0.5], "usage": {"prompt_tokens": 4}}, + ], + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "Embeddings", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 7, + "total_tokens": 7, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7"]) +def test_transform_embedding_response_body_that_is_not_an_object_raises_type_error(body: bytes): + with pytest.raises(TypeError): + _transform(httpx.Response(200, content=body)) + + +def test_transform_embedding_response_invalid_envelope_field_is_reported_by_name(): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"data": [], "model": 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [("model",)] diff --git a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py index f43e2e4d1cb..83c0cb2fec2 100644 --- a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py +++ b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py @@ -1,6 +1,8 @@ from unittest.mock import MagicMock, patch +import httpx import pytest +from pydantic import ValidationError from litellm.exceptions import AuthenticationError @@ -8,6 +10,7 @@ from litellm.llms.github_copilot.embedding.transformation import ( GithubCopilotEmbeddingConfig, ) from litellm.llms.github_copilot.common_utils import GetAPIKeyError +from litellm.types.utils import EmbeddingResponse def test_github_copilot_embedding_config_validate_environment(): @@ -201,3 +204,57 @@ def test_github_copilot_embedding_config_transform_response(): assert len(response.data) == 1 assert response.data[0]["embedding"] == [0.1, 0.2, 0.3] assert response.model == "text-embedding-3-small" + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return GithubCopilotEmbeddingConfig().transform_embedding_response( + model="github_copilot/text-embedding-3-small", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_keeps_the_openai_envelope(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 5, "total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "text-embedding-3-small", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 5, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7", b'"leaked payload text"', b'[{"data": "leaked payload text"}]']) +def test_transform_embedding_response_rejects_a_body_that_is_not_an_object(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "leaked payload text" not in str(exc_info.value) + + +def test_transform_embedding_response_object_without_data_is_an_invalid_response_object(): + with pytest.raises(Exception, match="Invalid response object"): + _transform(httpx.Response(200, json={"model": "text-embedding-3-small"})) diff --git a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py index 38355f32da1..44e8145ff33 100644 --- a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py +++ b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py @@ -1,9 +1,12 @@ from unittest.mock import MagicMock +import httpx import pytest +from pydantic import ValidationError from litellm.llms.jina_ai.embedding.transformation import JinaAIEmbeddingConfig +from litellm.types.utils import EmbeddingResponse JINA_KEY_ENV_NAMES = ("JINA_AI_API_KEY", "JINA_API_KEY", "JINA_AI_TOKEN") @@ -131,3 +134,61 @@ class TestJinaAIEmbeddingTransform: "input": expected_input, } assert result == expected_result + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return JinaAIEmbeddingConfig().transform_embedding_response( + model="jina-embeddings-v3", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_builds_the_response_from_the_body(): + response = _transform( + httpx.Response( + 200, + json={ + "model": "jina-embeddings-v3", + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "usage": {"prompt_tokens": 3, "total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "jina-embeddings-v3", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 3, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"7", b'["leaked payload text"]', b'"leaked payload text"']) +def test_transform_embedding_response_rejects_a_body_that_is_not_an_object(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "leaked payload text" not in str(exc_info.value) + + +@pytest.mark.parametrize("field", ["model", "data", "usage"]) +def test_transform_embedding_response_invalid_field_is_reported_by_name(field: str): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={field: 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [(field,)] diff --git a/tests/unit/llms/jina_ai/rerank/__init__.py b/tests/unit/llms/jina_ai/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/jina_ai/rerank/test_transformation.py b/tests/unit/llms/jina_ai/rerank/test_transformation.py new file mode 100644 index 00000000000..9ec02a4548b --- /dev/null +++ b/tests/unit/llms/jina_ai/rerank/test_transformation.py @@ -0,0 +1,111 @@ +from unittest.mock import MagicMock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.jina_ai.rerank.transformation import JinaAIRerankConfig +from litellm.types.rerank import RerankResponse + + +def _transform(payload: object, status_code: int = 200) -> RerankResponse: + return JinaAIRerankConfig().transform_rerank_response( + model="jina-reranker-v2-base-multilingual", + raw_response=httpx.Response(status_code, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_maps_results_and_usage(): + response = _transform( + { + "id": "rerank-1", + "results": [ + {"index": 1, "relevance_score": 0.72, "document": "hello"}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "world", "extra": 1}}, + {"index": 2, "relevance_score": 0, "extra": True}, + ], + "usage": {"total_tokens": 21, "prompt_tokens": 21}, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.72, "document": {"text": "hello"}}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "world"}}, + {"index": 2, "relevance_score": 0.0}, + ] + assert response.meta == {"billed_units": {"total_tokens": 21}, "tokens": {}} + + +@pytest.mark.parametrize( + ("payload", "expected_meta"), + [ + ({"results": []}, {"billed_units": {}, "tokens": {}}), + ({"results": [], "usage": {}}, {"billed_units": {}, "tokens": {}}), + ( + {"results": [], "usage": {"total_tokens": 12, "input_tokens": 7, "output_tokens": 3, "unknown": 1}}, + {"billed_units": {"total_tokens": 12}, "tokens": {"input_tokens": 7, "output_tokens": 3}}, + ), + ], +) +def test_transform_rerank_response_keeps_only_known_usage_counters( + payload: dict[str, object], expected_meta: dict[str, object] +): + assert _transform(payload).meta == expected_meta + + +def test_transform_rerank_response_generates_an_id_when_the_provider_sends_none(): + response = _transform({"id": None, "results": [{"index": 0, "relevance_score": 0.5}]}) + + assert isinstance(response.id, str) + assert response.id != "" + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"results": [], "usage": None}, + {"results": [], "usage": {"total_tokens": 1.5}}, + {"results": [], "usage": {"input_tokens": "many"}}, + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + {"results": [{"index": None, "relevance_score": 0.5}]}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_without_results_raises_value_error_naming_the_body(): + with pytest.raises(ValueError, match="No results found") as exc_info: + _transform({"id": "rerank-1"}) + + assert str(exc_info.value) == "No results found in the response={'id': 'rerank-1'}" + + +def test_transform_rerank_response_result_without_score_raises_key_error(): + with pytest.raises(KeyError, match="relevance_score"): + _transform({"results": [{"index": 0}]}) + + +def test_transform_rerank_response_non_200_raises_with_the_response_text(): + with pytest.raises(Exception, match="quota exceeded") as exc_info: + _transform({"detail": "quota exceeded"}, status_code=429) + + assert type(exc_info.value) is Exception + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"results": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/manus/files/__init__.py b/tests/unit/llms/manus/files/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/manus/files/test_transformation.py b/tests/unit/llms/manus/files/test_transformation.py new file mode 100644 index 00000000000..a47e1c8cb9b --- /dev/null +++ b/tests/unit/llms/manus/files/test_transformation.py @@ -0,0 +1,128 @@ +import time +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.manus.files.transformation import ManusFilesConfig +from litellm.types.llms.openai import OpenAIFileObject + + +def _list_files(body: object) -> list[OpenAIFileObject]: + return ManusFilesConfig().transform_list_files_response( + raw_response=httpx.Response(200, json=body), logging_obj=Mock(), litellm_params={} + ) + + +def test_delete_file_response_is_read_into_a_file_deleted_object(): + deleted = ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, json={"id": "file-1", "deleted": True, "object": "file", "region": "eu"}), + logging_obj=Mock(), + litellm_params={}, + ) + + assert deleted.model_dump() == {"id": "file-1", "deleted": True, "object": "file", "region": "eu"} + + +@pytest.mark.parametrize("body", [b'["secret-file"]', b'"secret-file"', b"7", b"null"]) +def test_delete_file_response_rejects_a_body_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, content=body), logging_obj=Mock(), litellm_params={} + ) + + assert "secret-file" not in str(exc_info.value) + + +def test_delete_file_response_requires_the_file_deleted_fields(): + with pytest.raises(ValidationError, match="FileDeleted"): + ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, json={"id": "file-1"}), logging_obj=Mock(), litellm_params={} + ) + + +def test_list_files_response_maps_every_listed_file(): + files = _list_files( + { + "object": "list", + "data": [ + { + "id": "file-1", + "bytes": 12, + "filename": "a.pdf", + "purpose": "batch", + "status": "processed", + "status_details": "done", + "created_at": "2024-01-02T03:04:05Z", + }, + {"id": "file-2", "created_at": "2024-01-02T03:04:05.123456+00:00"}, + ], + } + ) + + created_at = int(time.mktime(time.strptime("2024-01-02T03:04:05", "%Y-%m-%dT%H:%M:%S"))) + assert [file.model_dump() for file in files] == [ + { + "id": "file-1", + "bytes": 12, + "created_at": created_at, + "filename": "a.pdf", + "object": "file", + "purpose": "batch", + "status": "processed", + "expires_at": None, + "status_details": "done", + }, + { + "id": "file-2", + "bytes": 0, + "created_at": created_at, + "filename": "", + "object": "file", + "purpose": "assistants", + "status": "uploaded", + "expires_at": None, + "status_details": None, + }, + ] + + +@pytest.mark.parametrize("created_at", ["not-a-date", "", None, 0]) +def test_list_files_response_falls_back_to_now_for_an_unusable_created_at(created_at: object): + before = int(time.time()) + + (file,) = _list_files({"data": [{"id": "file-1", "created_at": created_at}]}) + + assert before <= file.created_at <= int(time.time()) + + +@pytest.mark.parametrize("body", [{}, {"data": []}]) +def test_list_files_response_is_empty_without_listed_files(body: dict[str, object]): + assert _list_files(body) == [] + + +@pytest.mark.parametrize( + "body", + [ + ["secret-file"], + "secret-file", + {"data": None}, + {"data": 7}, + {"data": "secret-file"}, + {"data": ["secret-file"]}, + {"data": [{"id": "file-1"}, ["secret-file"]]}, + {"data": [{"id": "file-1", "created_at": ["secret-file"]}]}, + {"data": [{"id": "file-1", "created_at": 1700000000}]}, + ], +) +def test_list_files_response_rejects_malformed_listings_without_echoing_them(body: object): + with pytest.raises(ValidationError) as exc_info: + _list_files(body) + + assert "secret-file" not in str(exc_info.value) + + +def test_list_files_response_rejects_a_file_the_file_object_cannot_hold(): + with pytest.raises(ValidationError, match="OpenAIFileObject"): + _list_files({"data": [{"id": "file-1", "purpose": "not-a-purpose"}]}) diff --git a/tests/unit/llms/minimax/text_to_speech/__init__.py b/tests/unit/llms/minimax/text_to_speech/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/minimax/text_to_speech/test_transformation.py b/tests/unit/llms/minimax/text_to_speech/test_transformation.py new file mode 100644 index 00000000000..a183c314a61 --- /dev/null +++ b/tests/unit/llms/minimax/text_to_speech/test_transformation.py @@ -0,0 +1,122 @@ +import base64 +from typing import Final +from unittest.mock import Mock + +import httpx +import pytest + +from litellm.llms.minimax.text_to_speech.transformation import MinimaxException, MinimaxTextToSpeechConfig +from litellm.types.llms.openai import HttpxBinaryResponseContent + +_AUDIO: Final = b"ID3\x04minimax-audio" +_REQUEST: Final = httpx.Request("POST", "https://api.minimax.io/v1/t2a_v2") + + +def _transform(payload: object, status_code: int = 200) -> HttpxBinaryResponseContent: + return MinimaxTextToSpeechConfig().transform_text_to_speech_response( + model="speech-02-hd", + raw_response=httpx.Response( + status_code, + json=payload, + request=_REQUEST, + headers={"content-encoding": "identity", "x-trace": "abc"}, + ), + logging_obj=Mock(), + ) + + +@pytest.mark.parametrize( + "payload", + [ + {"data": {"audio": _AUDIO.hex()}, "status": 0, "extra_info": {"audio_length": 5}}, + {"data": {"audio": _AUDIO.hex(), "audio_url": ""}}, + {"data": {"audio": ""}, "audio_file": _AUDIO.hex()}, + {"data": {"audio": None}, "audio_file": base64.b64encode(_AUDIO).decode()}, + {"base_resp": {"status_code": 0}, "audio_file": base64.b64encode(_AUDIO).decode()}, + ], +) +def test_transform_text_to_speech_response_decodes_hex_or_base64_audio(payload: dict[str, object]): + response = _transform(payload).response + + assert response.status_code == 200 + assert response.content == _AUDIO + assert response.headers["content-length"] == str(len(_AUDIO)) + assert response.headers["x-trace"] == "abc" + assert "content-encoding" not in response.headers + + +@pytest.mark.parametrize( + ("payload", "detail"), + [ + ({"status": 2, "ced": "invalid api key"}, "invalid api key"), + ({"status": 2}, "Unknown error"), + ({"status": 1004, "ced": ""}, "API returned status 1004"), + ({"status": "failed", "ced": None, "data": {"audio": _AUDIO.hex()}}, "API returned status failed"), + ], +) +def test_transform_text_to_speech_response_reports_api_status_errors(payload: dict[str, object], detail: str): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload, status_code=401) + + assert exc_info.value.message == f"MiniMax TTS error: {detail}" + assert exc_info.value.status_code == 401 + + +def test_transform_text_to_speech_response_refuses_url_output(): + with pytest.raises(MinimaxException) as exc_info: + _transform({"data": {"audio_url": "https://cdn.example/a.mp3", "audio": _AUDIO.hex()}}) + + assert exc_info.value.message == ( + "URL output format is not yet supported. Use 'hex' format or fetch from URL: https://cdn.example/a.mp3" + ) + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize( + ("payload", "keys"), + [ + ({}, []), + ({"data": {}, "status": 0}, ["data", "status"]), + ({"data": {"audio": ""}, "audio_file": ""}, ["data", "audio_file"]), + ({"data": {"audio": None}, "audio_file": None}, ["data", "audio_file"]), + ({"audio_file": 0, "data": {"audio": []}}, ["audio_file", "data"]), + ], +) +def test_transform_text_to_speech_response_without_audio_lists_the_response_keys( + payload: dict[str, object], keys: list[str] +): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload) + + assert exc_info.value.message == f"No audio data in MiniMax response. Response keys: {keys}" + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "not an object", + {"data": None}, + {"data": "not an object"}, + {"data": ["not", "an", "object"]}, + {"data": {"audio": 7}}, + {"data": {"audio": ["49", "44"]}}, + {"audio_file": {"hex": "4944"}}, + ], +) +def test_transform_text_to_speech_response_wraps_malformed_payloads_without_echoing_them(payload: object): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload) + + assert exc_info.value.message.startswith("Error processing MiniMax response: ") + assert "input_value" not in exc_info.value.message + assert exc_info.value.status_code == 500 + + +def test_transform_text_to_speech_response_reports_undecodable_audio(): + with pytest.raises(MinimaxException) as exc_info: + _transform({"data": {"audio": "zzz"}}) + + assert exc_info.value.message.startswith("Failed to decode audio data: ") + assert exc_info.value.status_code == 500 diff --git a/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py b/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py index 2ffe7c3e686..5b0ca09ed13 100644 --- a/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py +++ b/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py @@ -7,9 +7,12 @@ transformation between OpenAI-compatible format and ModelScope API format. from unittest.mock import MagicMock, patch +import httpx import pytest +from pydantic import ValidationError +import litellm from litellm.llms.modelscope.image_generation.transformation import ( ModelScopeImageGenerationConfig, ) @@ -449,3 +452,83 @@ class TestModelScopeImageGenerationTransformation: ) assert isinstance(error, BadRequestError) + + +def _transform_generation_response(payload: object, status_code: int = 200) -> ImageResponse: + return ModelScopeImageGenerationConfig().transform_image_generation_response( + model="Qwen/Qwen-Image", + raw_response=httpx.Response(status_code, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("item", "expected"), + [ + ({"url": "https://a.example/i.png"}, ("https://a.example/i.png", None, None)), + ({"b64_json": "aGVsbG8="}, (None, "aGVsbG8=", None)), + ( + {"url": "https://a.example/i.png", "revised_prompt": "a calmer cat", "seed": 7}, + ("https://a.example/i.png", None, "a calmer cat"), + ), + ({}, (None, None, None)), + ], +) +def test_transform_image_generation_response_maps_one_image( + item: dict[str, object], expected: tuple[str | None, str | None, str | None] +) -> None: + response = _transform_generation_response({"created": 1, "data": [item]}) + + assert [(image.url, image.b64_json, image.revised_prompt) for image in response.data] == [expected] + + +@pytest.mark.parametrize("payload", [{}, {"data": []}, {"data": ""}, {"data": {}}, {"created": 1}]) +def test_transform_image_generation_response_without_images_is_empty(payload: dict[str, object]) -> None: + assert _transform_generation_response(payload).data == [] + + +@pytest.mark.parametrize( + ("error", "status_code", "expected_class", "expected_message"), + [ + ({"message": "Invalid prompt"}, 400, litellm.BadRequestError, "ModelScope error: Invalid prompt"), + ({"message": "Bad key"}, 401, litellm.AuthenticationError, "ModelScope error: Bad key"), + ({"code": "overloaded"}, 503, litellm.InternalServerError, "ModelScope error: {'code': 'overloaded'}"), + ({}, 200, litellm.BadRequestError, "ModelScope error: {}"), + ], +) +def test_transform_image_generation_response_reports_api_error_bodies( + error: dict[str, object], status_code: int, expected_class: type[Exception], expected_message: str +) -> None: + with pytest.raises(expected_class) as exc_info: + _transform_generation_response({"error": error, "data": [{"url": "ignored"}]}, status_code) + + assert str(exc_info.value).endswith(expected_message) + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "error", + 7, + {"error": "plain text error"}, + {"error": None}, + {"error": ["not", "an", "object"]}, + {"data": 7}, + {"data": None}, + {"data": ["not an object"]}, + {"data": [{"url": "https://a.example/i.png"}, None]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads_without_echoing_them( + payload: object, +) -> None: + with pytest.raises(ValidationError) as exc_info: + _transform_generation_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py b/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py index 2b03b2d807b..39ba396d8d2 100644 --- a/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py +++ b/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py @@ -10,7 +10,9 @@ truncate. Two defects are covered here: import json from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.nvidia_nim.rerank.ranking_transformation import ( @@ -239,3 +241,87 @@ class TestNvidiaNimRetrievalRerankRequestTransform: doc = {"title": "no supported fields here"} request_data = self._build_request([doc]) assert request_data["passages"] == [{"text": json.dumps(doc)}] + + +def _transform_retrieval_response(payload: object) -> RerankResponse: + return NvidiaNimRerankConfig().transform_rerank_response( + model="nvidia/llama-3_2-nv-rerankqa-1b-v2", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + request_data={"passages": [{"text": "first"}, {"text": "second"}]}, + ) + + +@pytest.mark.parametrize( + ("usage", "expected_total_tokens"), + [ + ({"total_tokens": 42}, 42), + ({"total_tokens": 0}, 2), + ({}, 2), + ], +) +def test_transform_rerank_response_bills_reported_tokens_or_falls_back_to_result_count( + usage: dict[str, object], expected_total_tokens: int +): + response = _transform_retrieval_response( + {"rankings": [{"index": 1, "logit": 2.5}, {"index": 0, "logit": -1}], "usage": usage} + ) + + assert response.meta == {"billed_units": {"total_tokens": expected_total_tokens}} + assert response.results == [ + {"index": 1, "relevance_score": 2.5, "document": {"text": "second"}}, + {"index": 0, "relevance_score": -1.0, "document": {"text": "first"}}, + ] + + +def test_transform_rerank_response_keeps_provider_id_and_bills_one_token_per_result_without_usage(): + response = _transform_retrieval_response({"id": "rank-1", "rankings": [{"index": 0, "logit": 0.5}]}) + + assert response.id == "rank-1" + assert response.meta == {"billed_units": {"total_tokens": 1}} + + +@pytest.mark.parametrize( + "payload", + [ + {"rankings": [], "usage": None}, + {"rankings": [], "usage": {"total_tokens": None}}, + {"rankings": [], "usage": {"total_tokens": "42"}}, + {"rankings": [], "usage": {"total_tokens": 1.5}}, + {"rankings": [], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_usage_and_id(payload: dict[str, object]): + with pytest.raises(ValidationError): + _transform_retrieval_response(payload) + + +def test_transform_rerank_response_non_object_body_raises_attribute_error(): + with pytest.raises(AttributeError): + _transform_retrieval_response(["not", "an", "object"]) + + +def test_transform_rerank_response_usage_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform_retrieval_response({"rankings": [], "usage": {"total_tokens": "leaked usage text"}}) + + assert "leaked usage text" not in str(exc_info.value) + + +def test_map_cohere_rerank_params_passes_provider_params_through_and_maps_top_n(): + params = NvidiaNimRerankConfig().map_cohere_rerank_params( + non_default_params={"truncate": "END"}, + model="nvidia/llama-3_2-nv-rerankqa-1b-v2", + drop_params=False, + query="which passage shows a cat?", + documents=["a", {"text": "b"}], + top_n=1, + ) + + assert params == { + "query": "which passage shows a cat?", + "documents": ["a", {"text": "b"}], + "top_k": 1, + "truncate": "END", + } diff --git a/tests/unit/llms/oci/embed/test_oci_embed_transformation.py b/tests/unit/llms/oci/embed/test_oci_embed_transformation.py index 4ffd79ff147..01be7c7b904 100644 --- a/tests/unit/llms/oci/embed/test_oci_embed_transformation.py +++ b/tests/unit/llms/oci/embed/test_oci_embed_transformation.py @@ -379,3 +379,78 @@ class TestOCIEmbedConfig: litellm_params={}, ) assert "eu-frankfurt-1" in url + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return OCIEmbedConfig().transform_embedding_response( + model="cohere.embed-v3.0", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key=None, + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +@pytest.mark.parametrize( + ("token_fields", "expected_prompt_tokens", "expected_total_tokens"), + [ + ({"inputTextTokenCounts": [5, 6]}, 11, 11), + ({"usage": {"promptTokens": 3, "totalTokens": 4}}, 3, 4), + ({"inputTextTokenCounts": [5, 6], "usage": {"promptTokens": 3, "totalTokens": 4}}, 11, 11), + ({}, 0, 0), + ], +) +def test_transform_embedding_response_reads_usage_from_whichever_token_field_is_present( + token_fields: dict[str, object], expected_prompt_tokens: int, expected_total_tokens: int +): + response = _transform( + httpx.Response( + 200, + json={ + "embeddings": [[0.1, 1], [0.5]], + "modelId": "cohere.embed-v3.0", + "modelVersion": "3.0", + "unknown": "ignored", + **token_fields, + }, + ) + ) + + assert response.model_dump() == { + "model": "cohere.embed-v3.0", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 1.0]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": expected_prompt_tokens, + "total_tokens": expected_total_tokens, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"7", b'["leaked payload text"]', b'"leaked payload text"']) +def test_transform_embedding_response_body_that_is_not_an_object_is_a_schema_error(body: bytes): + with pytest.raises(OCIError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert exc_info.value.status_code == 500 + assert exc_info.value.message.startswith("OCI embed response does not match expected schema: ") + assert "leaked payload text" not in exc_info.value.message + + +def test_transform_embedding_response_object_missing_required_fields_names_them(): + with pytest.raises(OCIError) as exc_info: + _transform(httpx.Response(200, json={"embeddings": [[0.1]]})) + + assert exc_info.value.status_code == 500 + assert exc_info.value.message.startswith("OCI embed response does not match expected schema: ") + assert "modelId" in exc_info.value.message + assert "modelVersion" in exc_info.value.message diff --git a/tests/unit/llms/openai/responses/count_tokens/__init__.py b/tests/unit/llms/openai/responses/count_tokens/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/openai/responses/count_tokens/test_handler.py b/tests/unit/llms/openai/responses/count_tokens/test_handler.py new file mode 100644 index 00000000000..2427b8d272f --- /dev/null +++ b/tests/unit/llms/openai/responses/count_tokens/test_handler.py @@ -0,0 +1,26 @@ +import pytest + +from litellm.llms.openai.common_utils import OpenAIError +from litellm.llms.openai.responses.count_tokens.handler import OpenAICountTokensHandler + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "input_value", "expected_message"), + [ + ("", "hello", "CountTokens processing error: model parameter is required"), + ("gpt-4o", "", "CountTokens processing error: input parameter is required"), + ("gpt-4o", [], "CountTokens processing error: input parameter is required"), + ], +) +async def test_request_without_model_or_input_is_rejected_before_calling_openai( + model: str, input_value: str | list[object], expected_message: str +) -> None: + with pytest.raises(OpenAIError) as rejected: + await OpenAICountTokensHandler().handle_count_tokens_request( + model=model, + input=input_value, + api_key="sk-test", + ) + + assert (rejected.value.status_code, rejected.value.message) == (500, expected_message) diff --git a/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py b/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py index e45270fb5e3..edd6168c6c3 100644 --- a/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py +++ b/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py @@ -580,3 +580,86 @@ class TestOpenRouterImageGenerationTransformation: assert isinstance(error, OpenRouterException) assert "Test error" in str(error) assert error.status_code == 400 + + +def _transform_generation_response(payload: object) -> ImageResponse: + return OpenRouterImageGenerationConfig().transform_image_generation_response( + model="google/gemini-2.5-flash-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + "payload", + [ + {}, + {"choices": []}, + {"choices": ""}, + {"choices": {}}, + {"choices": [{}]}, + {"choices": [{"message": {}}]}, + {"choices": [{"message": {"images": ""}}]}, + {"choices": [{"message": {"images": [{}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": None}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": ""}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": 0}}]}}]}, + ], +) +def test_transform_image_generation_response_without_usable_image_url_has_no_images( + payload: dict[str, object], +): + assert _transform_generation_response(payload).data == [] + + +@pytest.mark.parametrize( + ("url", "expected"), + [ + ("data:image/png;base64,aGVsbG8=", ("aGVsbG8=", None)), + ("data:image/png;base64,a,b", ("a,b", None)), + ("data:no-comma", (None, None)), + ("https://example.com/a.png", (None, "https://example.com/a.png")), + ], +) +def test_transform_image_generation_response_maps_one_image_url( + url: str, expected: tuple[str | None, str | None] +): + response = _transform_generation_response( + {"choices": [{"message": {"images": [{"image_url": {"url": url}}]}}]} + ) + + assert [(image.b64_json, image.url) for image in response.data] == [expected] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"choices": 7}, + {"choices": ["not an object"]}, + {"choices": [{"message": "not an object"}]}, + {"choices": [{"message": {"images": 7}}]}, + {"choices": [{"message": {"images": ["not an object"]}}]}, + {"choices": [{"message": {"images": [{"image_url": "not an object"}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": ["a"]}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": 7}}]}}]}, + ], +) +def test_transform_image_generation_response_wraps_malformed_payloads_without_echoing_them( + payload: object, +): + with pytest.raises(OpenRouterException) as exc_info: + _transform_generation_response(payload) + + message = str(exc_info.value) + assert message.startswith( + "Error transforming OpenRouter image generation response: " + ) + assert "input_value" not in message + assert exc_info.value.status_code == 500 diff --git a/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py index 87e54dfba9b..422658932e6 100644 --- a/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py +++ b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py @@ -1,4 +1,9 @@ +import httpx +import pytest +from pydantic import ValidationError +from litellm.llms.ovhcloud.audio_transcription.transformation import OVHCloudAudioTranscriptionConfig +from litellm.types.utils import TranscriptionResponse @@ -56,3 +61,40 @@ class TestOVHCloudDurationFieldMigration: mock_response.json.return_value = {"text": "silence", "seconds": 0.0} result = config.transform_audio_transcription_response(mock_response) assert result._hidden_params["duration"] == 0.0 + + +def _transform(payload: object) -> TranscriptionResponse: + return OVHCloudAudioTranscriptionConfig().transform_audio_transcription_response(httpx.Response(200, json=payload)) + + +@pytest.mark.parametrize( + ("payload", "expected_text", "expected_hidden_params"), + [ + ( + {"text": "hello", "seconds": 3.5, "duration": 9}, + "hello", + {"text": "hello", "seconds": 3.5, "duration": 3.5}, + ), + ( + {"transcript": "from transcript", "seconds": None, "duration": 4}, + "from transcript", + {"transcript": "from transcript", "seconds": None, "duration": 4}, + ), + ({"language": "en"}, "", {"language": "en"}), + ], +) +def test_transform_audio_transcription_response_normalizes_text_and_duration( + payload: dict[str, object], expected_text: str, expected_hidden_params: dict[str, object] +): + response = _transform(payload) + + assert response.text == expected_text + assert response._hidden_params == expected_hidden_params + + +@pytest.mark.parametrize("payload", [7, "spoken secret", [{"text": "spoken secret"}]]) +def test_transform_audio_transcription_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "spoken secret" not in str(exc_info.value) diff --git a/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py b/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py index 2dfe33b828c..4b677675694 100644 --- a/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py +++ b/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.recraft.image_generation.transformation import ( @@ -256,3 +257,56 @@ class TestRecraftImageGenerationTransformation: ) assert "Error transforming image generation response" in str(exc_info.value) + + +def _transform(payload: object) -> ImageResponse: + return RecraftImageGenerationConfig().transform_image_generation_response( + model="recraftv3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_each_image_object(): + response = _transform({"data": [{"url": "https://img.recraft.ai/a.png"}, {"b64_json": "QUJD"}, {}]}) + + assert [(image.url, image.b64_json) for image in response.data] == [ + ("https://img.recraft.ai/a.png", None), + (None, "QUJD"), + (None, None), + ] + + +@pytest.mark.parametrize("data", [[], "", {}]) +def test_transform_image_generation_response_with_empty_data_has_no_images(data: object): + assert _transform({"data": data}).data == [] + + +@pytest.mark.parametrize( + "payload", + [ + 7, + "https://img.recraft.ai/a.png", + [{"url": "https://img.recraft.ai/a.png"}], + {"data": None}, + {"data": "https://img.recraft.ai/a.png"}, + {"data": {"url": "https://img.recraft.ai/a.png"}}, + {"data": ["https://img.recraft.ai/a.png"]}, + {"data": [{"url": "https://img.recraft.ai/a.png"}, 7]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "img.recraft.ai" not in str(exc_info.value) + + +def test_transform_image_generation_response_requires_data(): + with pytest.raises(KeyError): + _transform({"created": 1}) diff --git a/tests/unit/llms/replicate/__init__.py b/tests/unit/llms/replicate/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/replicate/chat/__init__.py b/tests/unit/llms/replicate/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/replicate/chat/test_transformation.py b/tests/unit/llms/replicate/chat/test_transformation.py new file mode 100644 index 00000000000..4c2b1840664 --- /dev/null +++ b/tests/unit/llms/replicate/chat/test_transformation.py @@ -0,0 +1,54 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.replicate.chat.transformation import ReplicateConfig +from litellm.types.utils import ModelResponse + + +def _transform(raw_response: httpx.Response) -> ModelResponse: + return ReplicateConfig().transform_response( + model="acme/echo-model", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={"input": {"prompt": "Hello"}}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("output", "content"), + [ + (["Hello", ", ", "world"], "Hello, world"), + ("Hello", "Hello"), + ([], " "), + ("", " "), + ([""], " "), + ], +) +def test_transform_response_joins_the_prediction_output_into_the_message_content(output: object, content: str): + response = _transform(httpx.Response(200, json={"status": "succeeded", "output": output})) + + assert response.choices[0].message.content == content + assert response.model == "replicate/acme/echo-model" + assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens + + +def test_transform_response_uses_a_blank_message_when_the_prediction_has_no_output(): + response = _transform(httpx.Response(200, json={"status": "succeeded"})) + + assert response.choices[0].message.content == " " + + +@pytest.mark.parametrize("output", [None, 7, True, ["secret-output", 7], ["secret-output", None], [["secret-output"]]]) +def test_transform_response_rejects_an_output_that_is_not_made_of_strings_without_echoing_it(output: object): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"status": "succeeded", "output": output})) + + assert "secret-output" not in str(exc_info.value) diff --git a/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py b/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py index e2dd3bca74f..51a2ccb5640 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py @@ -70,6 +70,19 @@ def test_sagemaker_response_stream_shape_load_failure_returns_none(): assert shape is None +@pytest.mark.parametrize("service_model", [["shapes"], None]) +def test_sagemaker_response_stream_shape_is_none_for_a_service_model_that_is_not_a_mapping( + service_model: object, +): + pytest.importorskip("botocore") + from unittest.mock import patch + + import litellm.llms.sagemaker.common_utils as mod + + with patch("botocore.loaders.Loader.load_service_model", return_value=service_model): + assert mod._load_sagemaker_response_stream_shape() is None + + def test_sagemaker_response_stream_shape_is_structure_shape(): """ The loaded shape should be the botocore StructureShape for diff --git a/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py b/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py index 8f4551bc79b..ff2af8a8d7b 100644 --- a/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py +++ b/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py @@ -2,9 +2,15 @@ import os import json import copy -from unittest.mock import patch +from unittest.mock import MagicMock, patch + +import httpx +import pytest +from pydantic import ValidationError import litellm +from litellm.llms.snowflake.embedding.transformation import SnowflakeEmbeddingConfig +from litellm.types.utils import EmbeddingResponse model_name = "snowflake-arctic-embed" @@ -94,3 +100,59 @@ def test_snowflake_env(mock_post): os.environ.pop("SNOWFLAKE_ACCOUNT_ID", None) os.environ.pop("SNOWFLAKE_JWT", None) + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return SnowflakeEmbeddingConfig().transform_embedding_response( + model="snowflake-arctic-embed-m", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_flattens_each_vector_and_prefixes_the_model(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "model": "snowflake-arctic-embed-m", + "data": [{"object": "embedding", "index": 0, "embedding": [[0.1, 2]]}], + "usage": {"total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "snowflake/snowflake-arctic-embed-m", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 0, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + assert response._hidden_params["model"] == "snowflake-arctic-embed-m" + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7"]) +def test_transform_embedding_response_body_that_is_not_an_object_raises_type_error(body: bytes): + with pytest.raises(TypeError): + _transform(httpx.Response(200, content=body)) + + +def test_transform_embedding_response_invalid_envelope_field_is_reported_by_name(): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"data": [], "model": 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [("model",)] diff --git a/tests/unit/llms/stability/image_generation/test_stability_image_generation.py b/tests/unit/llms/stability/image_generation/test_stability_image_generation.py index c5a78603f9c..42eb1a2365b 100644 --- a/tests/unit/llms/stability/image_generation/test_stability_image_generation.py +++ b/tests/unit/llms/stability/image_generation/test_stability_image_generation.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.stability.image_generation import StabilityImageGenerationConfig from litellm.types.llms.stability import ( @@ -304,3 +305,35 @@ class TestStabilityGenerationModels: STABILITY_GENERATION_MODELS["stable-image-core"] == "/v2beta/stable-image/generate/core" ) + + +def _transform(payload: object) -> ImageResponse: + return StabilityImageGenerationConfig().transform_image_generation_response( + model="sd3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_returns_the_base64_image(): + response = _transform({"image": "QUJD", "finish_reason": "SUCCESS", "seed": 7}) + + assert [(image.b64_json, image.url) for image in response.data] == [("QUJD", None)] + + +@pytest.mark.parametrize("payload", [{}, {"image": ""}, {"image": None, "finish_reason": None}]) +def test_transform_image_generation_response_without_an_image_has_no_data(payload: dict[str, object]): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize("payload", ["base64encodedimage==", ["base64encodedimage=="]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "base64encodedimage" not in str(exc_info.value) diff --git a/tests/unit/llms/tinyfish/test_tinyfish_search.py b/tests/unit/llms/tinyfish/test_tinyfish_search.py index 69afbb416aa..e7020238ba6 100644 --- a/tests/unit/llms/tinyfish/test_tinyfish_search.py +++ b/tests/unit/llms/tinyfish/test_tinyfish_search.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.tinyfish.search.transformation import ( TinyfishSearchConfig, _append_domain_filters, @@ -821,3 +822,56 @@ class TestDefaultMissingResultFields: assert raw_json["results"][0] == "string item" assert raw_json["results"][1] == 42 assert raw_json["results"][2] == {"title": "ok", "url": "", "snippet": ""} + + +def test_transform_search_response_defaults_missing_fields_of_a_decoded_http_body(): + payload = { + "query": "q", + "results": [ + {"title": "T", "url": "https://a.example", "snippet": "S", "position": 1}, + {"title": None, "position": 2}, + ], + } + + response = TinyfishSearchConfig().transform_search_response( + raw_response=httpx.Response(200, json=payload, headers={"x-request-id": "req-1"}), + logging_obj=None, + ) + + assert [result.model_dump(exclude_none=True) for result in response.results] == [ + {"title": "T", "url": "https://a.example", "snippet": "S", "position": 1}, + {"title": "", "url": "", "snippet": "", "position": 2}, + ] + assert response.model_dump()["query"] == "q" + assert response._hidden_params["headers"]["x-request-id"] == "req-1" + + +@pytest.mark.parametrize("payload", [[], "text", 5, {}, {"results": "text"}, {"results": [5]}]) +def test_transform_search_response_wraps_a_decoded_body_of_the_wrong_shape(payload: object): + with pytest.raises(BaseLLMException, match="TinyFish Search: Response shape does not match") as exc_info: + TinyfishSearchConfig().transform_search_response( + raw_response=httpx.Response(200, json=payload), logging_obj=None + ) + + assert exc_info.value.status_code == 200 + + +@pytest.mark.parametrize( + ("body", "expected"), + [ + ('{"error": {"code": "INVALID_INPUT", "message": "query is required"}}', "query is required"), + ('{"error": {"message": ""}}', '{"error": {"message": ""}}'), + ('{"error": {"message": 5}}', '{"error": {"message": 5}}'), + ('{"error": "flat"}', '{"error": "flat"}'), + ('["not", "an", "envelope"]', '["not", "an", "envelope"]'), + ("Bad Gateway", "Bad Gateway"), + ], +) +def test_transform_search_response_unwraps_only_the_tinyfish_error_envelope(body: str, expected: str): + with pytest.raises(BaseLLMException) as exc_info: + TinyfishSearchConfig().transform_search_response(raw_response=httpx.Response(400, text=body), logging_obj=None) + + assert ( + exc_info.value.message == f"TinyFish Search: {expected}. See https://docs.tinyfish.ai/search-api for details." + ) + assert exc_info.value.status_code == 400 diff --git a/tests/unit/llms/together_ai/rerank/__init__.py b/tests/unit/llms/together_ai/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/together_ai/rerank/test_handler.py b/tests/unit/llms/together_ai/rerank/test_handler.py new file mode 100644 index 00000000000..bc2f30364d2 --- /dev/null +++ b/tests/unit/llms/together_ai/rerank/test_handler.py @@ -0,0 +1,96 @@ +import json +from collections.abc import Iterator + +import httpx +import pytest +import respx +from pydantic import ValidationError + +import litellm +from litellm.llms.together_ai.rerank.handler import TogetherAIRerank + +API_BASE = "https://api.together.example/v1" +RERANKED = { + "id": "rr-1", + "results": [ + {"index": 1, "relevance_score": 0.9, "document": {"text": "Paris"}}, + {"index": 0, "relevance_score": 0.1}, + ], + "usage": {"total_tokens": 7}, +} +NON_OBJECT_BODIES = [7, "sensitive-document", [{"id": "sensitive-document"}]] + + +@pytest.fixture +def rerank_route(monkeypatch: pytest.MonkeyPatch) -> Iterator[respx.Route]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with respx.mock as router: + yield router.post(f"{API_BASE}/rerank") + + +def _rerank(*, is_async: bool): + return TogetherAIRerank().rerank( + model="Salesforce/Llama-Rank-V1", + api_key="sk-test", + api_base=API_BASE, + query="capital of france", + documents=["Berlin", {"text": "Paris"}], + top_n=1, + _is_async=is_async, + ) + + +def _assert_is_the_reranking(response: litellm.RerankResponse, route: respx.Route) -> None: + assert response.id == "rr-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.9, "document": {"text": "Paris"}}, + {"index": 0, "relevance_score": 0.1}, + ] + assert response.meta == {"billed_units": {"total_tokens": 7}, "tokens": {}} + assert route.calls.last.request.headers["authorization"] == "Bearer sk-test" + assert json.loads(route.calls.last.request.content) == { + "model": "Salesforce/Llama-Rank-V1", + "query": "capital of france", + "top_n": 1, + "documents": ["Berlin", {"text": "Paris"}], + "return_documents": True, + } + + +def test_rerank_returns_the_upstream_ranking(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json=RERANKED)) + + _assert_is_the_reranking(_rerank(is_async=False), rerank_route) + + +async def test_async_rerank_returns_the_upstream_ranking(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json=RERANKED)) + + _assert_is_the_reranking(await _rerank(is_async=True), rerank_route) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +def test_rerank_rejects_non_object_bodies(rerank_route: respx.Route, payload: object): + rerank_route.mock(return_value=httpx.Response(200, json=payload)) + + with pytest.raises(ValidationError) as exc_info: + _rerank(is_async=False) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +async def test_async_rerank_rejects_non_object_bodies(rerank_route: respx.Route, payload: object): + rerank_route.mock(return_value=httpx.Response(200, json=payload)) + + with pytest.raises(ValidationError) as exc_info: + await _rerank(is_async=True) + + assert "sensitive-document" not in str(exc_info.value) + + +def test_rerank_without_results_is_a_value_error(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json={"id": "rr-1"})) + + with pytest.raises(ValueError, match="No results found"): + _rerank(is_async=False) diff --git a/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py b/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py new file mode 100644 index 00000000000..ac5e14c0388 --- /dev/null +++ b/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py @@ -0,0 +1,43 @@ +import pytest +from pydantic import ValidationError + +from litellm.llms.vertex_ai.image_generation.image_generation_handler import VertexImageGeneration +from litellm.types.utils import ImageResponse + + +def _process(json_response: dict[str, object]) -> ImageResponse: + return VertexImageGeneration().process_image_generation_response( + json_response, ImageResponse(), "imagegeneration@006" + ) + + +@pytest.mark.parametrize( + ("predictions", "expected"), + [ + ([{"bytesBase64Encoded": "QUJD", "mimeType": "image/png"}, {"bytesBase64Encoded": "REVG"}], ["QUJD", "REVG"]), + ([{"bytesBase64Encoded": None}], [None]), + ([], []), + ], +) +def test_process_image_generation_response_maps_each_prediction( + predictions: list[dict[str, object]], expected: list[str | None] +): + response = _process({"predictions": predictions, "deployedModelId": "1"}) + + assert [image.b64_json for image in response.data] == expected + + +@pytest.mark.parametrize( + "predictions", + [None, 7, ["QUJD"], [{"bytesBase64Encoded": "QUJD"}, None], [{"bytesBase64Encoded": 7}]], +) +def test_process_image_generation_response_rejects_malformed_predictions(predictions: object): + with pytest.raises(ValidationError) as exc_info: + _process({"predictions": predictions}) + + assert "QUJD" not in str(exc_info.value) + + +def test_process_image_generation_response_requires_the_encoded_bytes_key(): + with pytest.raises(KeyError): + _process({"predictions": [{"mimeType": "image/png"}]}) diff --git a/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py index a72a570c2a2..0461c739395 100644 --- a/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py +++ b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -1,6 +1,8 @@ from unittest.mock import MagicMock, patch import httpx +import pytest +from pydantic import ValidationError from litellm.llms.vertex_ai.image_generation import ( @@ -12,6 +14,7 @@ from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import ( VertexAIImagenImageGenerationConfig, ) +from litellm.types.utils import ImageResponse class TestVertexAIGeminiImageGenerationConfig: @@ -635,3 +638,64 @@ class TestVertexAIImageGenerationIntegration: assert "us-central1" in url assert "imagegeneration@006" in url assert "predict" in url + + +def _transform_gemini_response(payload: object) -> ImageResponse: + return VertexAIGeminiImageGenerationConfig().transform_image_generation_response( + model="gemini-2.5-flash-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_gemini_image_generation_response_maps_usage_by_modality(): + response = _transform_gemini_response( + { + "candidates": [{"content": {"parts": [{"inlineData": {"data": "aGVsbG8="}}]}}], + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 7, + "totalTokenCount": 12, + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 3}, + {"modality": "IMAGE", "tokenCount": 2}, + ], + }, + } + ) + + assert [image.b64_json for image in response.data] == ["aGVsbG8="] + assert response.usage.input_tokens == 5 + assert response.usage.output_tokens == 7 + assert response.usage.total_tokens == 12 + assert response.usage.input_tokens_details.text_tokens == 3 + assert response.usage.input_tokens_details.image_tokens == 2 + + +@pytest.mark.parametrize("usage_metadata", [None, {}, [], "", 0]) +def test_gemini_image_generation_response_with_falsy_usage_metadata_keeps_zeroed_usage(usage_metadata: object): + response = _transform_gemini_response( + { + "candidates": [{"content": {"parts": []}, "groundingMetadata": {"webSearchQueries": ["a"]}}], + "usageMetadata": usage_metadata, + } + ) + + assert response.data == [] + assert (response.usage.input_tokens, response.usage.output_tokens, response.usage.total_tokens) == (0, 0, 0) + assert response.usage.web_search_requests == 1 + + +@pytest.mark.parametrize("usage_metadata", ["not an object", ["not", "an", "object"], 7, True]) +def test_gemini_image_generation_response_rejects_non_object_usage_metadata_without_echoing_it( + usage_metadata: object, +): + with pytest.raises(ValidationError) as exc_info: + _transform_gemini_response({"candidates": [], "usageMetadata": usage_metadata}) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py index fd667e8f425..8c73b72a65a 100644 --- a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py +++ b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, Mock, patch import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.vertex_ai.text_to_speech.transformation import ( @@ -633,3 +634,32 @@ def test_litellm_speech_vertex_ai_chirp(mock_get_token, mock_ensure_token, mock_ assert "headers" in call_kwargs assert "Authorization" in call_kwargs["headers"] assert call_kwargs["headers"]["Authorization"] == "Bearer mock-token" + + +@pytest.mark.parametrize("payload", [{}, {"audioContent": ""}, {"audioContent": None}, {"audioContent": []}]) +def test_transform_text_to_speech_response_without_audio_content_reports_it_missing(payload: dict[str, object]): + with pytest.raises(ValueError, match="No audioContent in Vertex AI TTS response"): + VertexAITextToSpeechConfig().transform_text_to_speech_response( + model="vertex_ai/chirp", + raw_response=httpx.Response(200, json=payload), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"audioContent": 7}, + {"audioContent": ["UklGRiQAAABXQVZFZm10IA=="]}, + ], +) +def test_transform_text_to_speech_response_rejects_malformed_payloads_without_echoing_them(payload: object): + with pytest.raises(ValidationError) as exc_info: + VertexAITextToSpeechConfig().transform_text_to_speech_response( + model="vertex_ai/chirp", + raw_response=httpx.Response(200, json=payload), + logging_obj=MagicMock(), + ) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py index 5eb4bf31845..777e52f4e2a 100644 --- a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py +++ b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py @@ -8,6 +8,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.voyage.rerank.transformation import VoyageRerankConfig from litellm.types.rerank import RerankResponse @@ -343,3 +344,81 @@ class TestVoyageRerankTransform: assert prompt_cost == 0.0 assert completion_cost == 0.0 + + +def _transform(payload: object) -> RerankResponse: + return VoyageRerankConfig().transform_rerank_response( + model="rerank-2.5", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_keeps_provider_id_and_drops_unknown_result_fields(): + response = _transform( + { + "id": "rerank-1", + "data": [ + {"index": 2, "relevance_score": 0.25, "document": {"text": "doc", "extra": 1}, "extra": 2}, + {"index": 0, "relevance_score": 1, "document": "plain"}, + ], + "usage": {"total_tokens": 9}, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 2, "relevance_score": 0.25, "document": {"text": "doc"}}, + {"index": 0, "relevance_score": 1.0, "document": {"text": "plain"}}, + ] + + +@pytest.mark.parametrize( + ("payload", "expected_total_tokens"), + [ + ({"data": []}, 0), + ({"data": [], "usage": {}}, 0), + ({"data": [], "usage": {"total_tokens": 79}}, 79), + ({"data": [], "usage": {"total_tokens": None}}, None), + ], +) +def test_transform_rerank_response_reads_total_tokens_from_usage( + payload: dict[str, object], expected_total_tokens: int | None +): + assert _transform(payload).meta == { + "billed_units": {"total_tokens": expected_total_tokens}, + "tokens": {"input_tokens": expected_total_tokens, "output_tokens": 0}, + } + + +def test_transform_rerank_response_empty_data_list_yields_no_results(): + assert _transform({"id": "rerank-1", "data": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"data": 7}, + {"data": ["not an object"]}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": None}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"data": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_result_without_index_raises_key_error(): + with pytest.raises(KeyError, match="index"): + _transform({"data": [{"relevance_score": 0.5}]}) + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"data": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py b/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py index efe592f515e..b6e5007c456 100644 --- a/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py +++ b/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py @@ -6,6 +6,10 @@ Validates the WatsonX transcription response transformation. from unittest.mock import MagicMock +import httpx +import pytest +from pydantic import ValidationError + from litellm.llms.watsonx.audio_transcription.transformation import ( IBMWatsonXAudioTranscriptionConfig, ) @@ -83,3 +87,57 @@ class TestWatsonXAudioTranscription: # Verify duration is set via dictionary assignment assert result["duration"] == 5.5 + + +def _transform_transcription_response(payload: object) -> TranscriptionResponse: + return IBMWatsonXAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + +@pytest.mark.parametrize( + ("payload", "expected_extras"), + [ + ({"text": "hello"}, {}), + ({"text": "hello", "model": "whisper-large-v3-turbo"}, {}), + ( + {"text": "hello", "duration": 1.5, "language": "en", "task": "transcribe"}, + {"duration": 1.5, "language": "en", "task": "transcribe"}, + ), + ( + {"text": "hello", "segments": [{"id": 0, "text": "hello"}], "words": None}, + {"segments": [{"id": 0, "text": "hello"}], "words": None}, + ), + ], +) +def test_transform_audio_transcription_response_copies_every_field_except_model( + payload: dict[str, object], expected_extras: dict[str, object] +) -> None: + response = _transform_transcription_response(payload) + + assert response.text == "hello" + assert not hasattr(response, "model") + assert {key: response[key] for key in expected_extras} == expected_extras + + +@pytest.mark.parametrize("payload", [{}, {"model": "whisper"}, {"duration": 2.0}, {"text": None, "usage": None}]) +def test_transform_audio_transcription_response_without_text_or_usage_reports_the_body( + payload: dict[str, object], +) -> None: + with pytest.raises(ValueError, match="Invalid response format") as exc_info: + _transform_transcription_response(payload) + + assert exc_info.value.args == ( + "Invalid response format. Received response does not match the expected format. Got: ", + payload, + ) + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "plain text", 7, True]) +def test_transform_audio_transcription_response_rejects_non_object_bodies_without_echoing_them( + payload: object, +) -> None: + with pytest.raises(ValidationError) as exc_info: + _transform_transcription_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py b/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py index 58f6bb23498..afcebcf0990 100644 --- a/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py +++ b/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py @@ -1,8 +1,13 @@ +from unittest.mock import MagicMock + +import httpx import pytest +from pydantic import ValidationError from litellm.llms.watsonx.embed.transformation import IBMWatsonXEmbeddingConfig +from litellm.types.utils import EmbeddingResponse class TestIBMWatsonXEmbeddingConfig: @@ -42,3 +47,88 @@ class TestIBMWatsonXEmbeddingConfig: optional_params={"project_id": "test-project-id"}, headers={}, ) + + +def _transform(payload: object) -> EmbeddingResponse: + return IBMWatsonXEmbeddingConfig().transform_embedding_response( + model="ibm/slate-125m-english-rtrvr", + raw_response=httpx.Response(200, json=payload), + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_numbers_each_result_and_bills_the_input_tokens(): + response = _transform( + { + "model_id": "ibm/slate-125m-english-rtrvr", + "results": [{"embedding": [0.1, 2], "input": "ignored"}, {"embedding": [0.5]}], + "input_token_count": 7, + } + ) + + assert response.object == "list" + assert response.data == [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ] + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (7, 0, 7) + + +@pytest.mark.parametrize( + ("payload", "expected_tokens"), + [ + ({"input_token_count": None}, 0), + ({"results": []}, 0), + ({}, 0), + ], +) +def test_transform_embedding_response_defaults_missing_results_and_token_count( + payload: dict[str, object], expected_tokens: int +): + response = _transform(payload) + + assert response.data == [] + assert (response.usage.prompt_tokens, response.usage.total_tokens) == (expected_tokens, expected_tokens) + + +@pytest.mark.parametrize( + "payload", + [ + {"results": None}, + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"embedding": [0.1]}, None]}, + {"input_token_count": 1.5}, + {"input_token_count": "many"}, + ], +) +def test_transform_embedding_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "not an object"]) +def test_transform_embedding_response_non_object_body_raises_attribute_error(payload: object): + with pytest.raises(AttributeError): + _transform(payload) + + +def test_transform_embedding_response_result_without_embedding_raises_key_error(): + with pytest.raises(KeyError, match="embedding"): + _transform({"results": [{"input": "x"}]}) + + +@pytest.mark.parametrize( + "payload", + [{"results": ["leaked payload text"]}, {"results": {"leaked payload text": 1}}], +) +def test_transform_embedding_response_shape_errors_do_not_echo_the_payload(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "leaked payload text" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py b/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py index ccbd318959f..fbdbddacd5f 100644 --- a/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py +++ b/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.watsonx.common_utils import ( WatsonXAIError, @@ -260,3 +261,64 @@ class TestIBMWatsonXRerankTransform: assert "return_documents" in supported_params assert "max_tokens_per_doc" in supported_params assert len(supported_params) == 5 + + +def _transform(payload: object) -> RerankResponse: + return IBMWatsonXRerankConfig().transform_rerank_response( + model="watsonx/cross-encoder/ms-marco-minilm-l-12-v2", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_keeps_the_upstream_id_documents_and_token_count(): + response = _transform( + { + "id": "rerank-1", + "results": [ + {"index": 1, "score": 0.5, "input": "plain text"}, + {"index": 0, "score": 0.25, "input": {"text": "object text"}}, + ], + "input_token_count": 62, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.5, "document": {"text": "plain text"}}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "object text"}}, + ] + assert response.meta == {"tokens": {"input_tokens": 62}} + + +def test_transform_rerank_response_without_a_token_count_reports_zero_tokens(): + response = _transform({"results": []}) + + assert response.results == [] + assert response.meta == {"tokens": {"input_tokens": 0}} + + +@pytest.mark.parametrize( + "payload", + [ + ["sensitive-document"], + "sensitive-document", + {"results": "sensitive-document"}, + {"results": 7}, + {"results": ["sensitive-document"]}, + {"results": [{"index": 0, "score": 0.5}, None]}, + {"results": [{"index": 0, "score": 0.5}], "id": 7}, + {"results": [{"index": 0, "score": 0.5}], "input_token_count": "many"}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "sensitive-document" not in str(exc_info.value) + + +def test_transform_rerank_response_requires_index_and_score_on_each_result(): + with pytest.raises(KeyError): + _transform({"results": [{"index": 0}]}) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index cfcff73b857..2e255ddf853 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -173,6 +173,49 @@ def test_mcp_oauth_token_identity_detects_change_under_encryption(): assert mcp_oauth_token_identity(unchanged) != mcp_oauth_token_identity(changed) +def test_mcp_oauth_token_identity_reads_client_and_scopes_from_json_string_credentials(): + from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity + + credentials: Final = json.dumps( + {"client_id": "cid", "client_secret": "csec", "scopes": ["a"], "upstream_resource": "api://audience"} + ) + + assert mcp_oauth_token_identity(_identity_server(credentials=credentials)) == ( + "https://up.example.com/mcp", + None, + "oauth2", + "authorization_code", + None, + "https://idp.example.com/authorize", + "https://idp.example.com/token", + "https://idp.example.com/register", + "cid", + "csec", + ["a"], + "api://audience", + ) + + +@pytest.mark.parametrize("credentials", ["[]", '["client_id"]', '"cid"', "5", "null", "not json"]) +def test_mcp_oauth_token_identity_treats_stored_credentials_that_are_not_a_json_object_as_empty(credentials): + from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity + + assert mcp_oauth_token_identity(_identity_server(credentials=credentials)) == ( + "https://up.example.com/mcp", + None, + "oauth2", + "authorization_code", + None, + "https://idp.example.com/authorize", + "https://idp.example.com/token", + "https://idp.example.com/register", + None, + None, + None, + None, + ) + + def _oauth_row(user_id: str, server_id: str = "srv-1"): """A stored per-user OAuth token row (payload tagged type=oauth2, legacy plain-base64 encoding).""" row = _legacy_row(json.dumps({"type": "oauth2", "access_token": "tok-" + user_id})) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py index fff4221f243..6a08e8eeab6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -1016,6 +1016,34 @@ async def test_get_user_env_vars_returns_empty_for_missing_row(): assert await get_user_env_vars(prisma, "alice", "srv-1") == {} +@pytest.mark.asyncio +async def test_get_user_env_vars_stringifies_non_string_json_values(env_vars_salt_key): + from types import SimpleNamespace + + from litellm.proxy._experimental.mcp_server.db import get_user_env_vars + + row: Final = SimpleNamespace(values_b64=_encrypted_user_env_blob({"PORT": 8080, "DEBUG": True, "EMPTY": None})) + + assert await get_user_env_vars(_mock_env_vars_prisma(row=row), "alice", "srv-1") == { + "PORT": "8080", + "DEBUG": "True", + "EMPTY": "None", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stored_json", ["[]", '[{"TOKEN": "t"}]', '"TOKEN"', "5", "null", "true", "not json"]) +async def test_get_user_env_vars_treats_a_blob_that_is_not_a_json_object_as_unset(env_vars_salt_key, stored_json): + from types import SimpleNamespace + + from litellm.proxy._experimental.mcp_server.db import get_user_env_vars + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + row: Final = SimpleNamespace(values_b64=encrypt_value_helper(stored_json)) + + assert await get_user_env_vars(_mock_env_vars_prisma(row=row), "alice", "srv-1") == {} + + @pytest.mark.asyncio async def test_decode_user_env_vars_warns_when_undecryptable( env_vars_salt_key, monkeypatch diff --git a/tests/unit/proxy/client/cli/test_auth_commands.py b/tests/unit/proxy/client/cli/test_auth_commands.py index 3a7792db1fe..b4a6aaf706e 100644 --- a/tests/unit/proxy/client/cli/test_auth_commands.py +++ b/tests/unit/proxy/client/cli/test_auth_commands.py @@ -7,6 +7,7 @@ from unittest.mock import Mock, patch import pytest +import responses from click.testing import CliRunner from litellm.constants import CLI_JWT_EXPIRATION_HOURS @@ -2044,3 +2045,33 @@ class TestGetStoredApiKeyRefresh: assert captured.out == "" assert captured.err == "Could not renew the key: token request failed with 503: temporarily_unavailable\n" save.assert_not_called() + + +@pytest.mark.parametrize( + "body, expected_detail", + [ + ('{"detail": "Too many CLI login attempts.", "retry": {"after": [30, null]}}', ": Too many CLI login attempts."), + ('{"detail": 5}', ""), + ('{"detail": ""}', ""), + ('{"detail": {"message": "nested"}}', ""), + ("{}", ""), + ('["detail"]', ""), + ('"detail"', ""), + ("7", ""), + ("true", ""), + ("null", ""), + ("gateway", ""), + ("", ""), + ], +) +@responses.activate +def test_login_shows_the_error_detail_only_when_the_proxy_answers_with_a_json_object(body, expected_detail): + responses.post("https://test.example.com/sso/cli/start", body=body, status=429) + + result = CliRunner().invoke(login, obj={"base_url": "https://test.example.com"}) + + assert result.exit_code == 0 + assert result.output == ( + "Authentication failed: Starting CLI login failed: HTTP 429 from https://test.example.com/sso/cli/start" + f"{expected_detail}\n" + ) diff --git a/tests/unit/proxy/client/test_chat.py b/tests/unit/proxy/client/test_chat.py index 8fe1bfcbb2f..9ab7150c288 100644 --- a/tests/unit/proxy/client/test_chat.py +++ b/tests/unit/proxy/client/test_chat.py @@ -7,6 +7,7 @@ import sys import pytest import requests +from pydantic import ValidationError from litellm.proxy.client.chat import ChatClient from litellm.proxy.client.exceptions import UnauthorizedError @@ -256,3 +257,53 @@ def test_completions_stream_gives_up_at_the_timeout_instead_of_hanging(hanging_s next(client.completions_stream(model="gpt-5.4", messages=[{"role": "user", "content": "hi"}])) assert time.monotonic() - started < 10 + + +@pytest.mark.parametrize( + "body, expected", + [ + ( + b'data: {"choices": [{"delta": {"content": "Hel"}}]}\n\n' + b'data: {"choices": [{"delta": {"content": "lo"}}]}\n\n' + b"data: [DONE]\n\n", + [{"choices": [{"delta": {"content": "Hel"}}]}, {"choices": [{"delta": {"content": "lo"}}]}], + ), + ( + b': keep-alive\n\nevent: ping\n\ndata: not json\n\ndata: {"id": "caf\xc3\xa9"}\n\ndata: 7\n\n', + [{"id": "café"}, 7], + ), + (b'data: {"id": 1}\n\ndata: [DONE] \n\ndata: {"id": 2}\n\n', [{"id": 1}]), + (b"", []), + ], +) +@responses.activate +def test_completions_stream_yields_parsed_sse_chunks(client, base_url, sample_messages, body, expected): + responses.add(responses.POST, f"{base_url}/chat/completions", body=body, status=200) + + assert list(client.completions_stream(model="gpt-4", messages=sample_messages)) == expected + + +@responses.activate +def test_completions_stream_rejects_bytes_that_are_not_utf8(client, base_url, sample_messages): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"data: \xff\n\n", status=200) + + with pytest.raises(UnicodeDecodeError): + list(client.completions_stream(model="gpt-4", messages=sample_messages)) + + +@pytest.mark.parametrize("line", ['data: {"id": 1}', 5, memoryview(b'data: {"id": 1}')]) +@responses.activate +def test_completions_stream_rejects_lines_that_are_not_bytes(client, base_url, sample_messages, monkeypatch, line): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"", status=200) + monkeypatch.setattr(requests.Response, "iter_lines", lambda self: iter([line])) + + with pytest.raises(ValidationError): + list(client.completions_stream(model="gpt-4", messages=sample_messages)) + + +@responses.activate +def test_completions_stream_accepts_bytearray_lines(client, base_url, sample_messages, monkeypatch): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"", status=200) + monkeypatch.setattr(requests.Response, "iter_lines", lambda self: iter([bytearray(b'data: {"id": 1}')])) + + assert list(client.completions_stream(model="gpt-4", messages=sample_messages)) == [{"id": 1}] diff --git a/tests/unit/proxy/db/test_object_permission_repository.py b/tests/unit/proxy/db/test_object_permission_repository.py new file mode 100644 index 00000000000..e2d79592915 --- /dev/null +++ b/tests/unit/proxy/db/test_object_permission_repository.py @@ -0,0 +1,73 @@ +from collections.abc import Mapping +from types import SimpleNamespace +from typing import Final + +import pytest + +from litellm.repositories.object_permission_repository import ObjectPermissionRepository + + +class _RecordingPermissionTable: + def __init__(self, stored: Mapping[str, object]) -> None: + self.stored: Final = stored + self.created: Final[list[Mapping[str, object]]] = [] + self.updated: Final[list[tuple[Mapping[str, object], Mapping[str, object]]]] = [] + + async def create(self, data: Mapping[str, object]) -> Mapping[str, object]: + self.created.append(data) + return {**self.stored, **data} + + async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> Mapping[str, object]: + self.updated.append((where, data)) + return {**self.stored, **data} + + +def _repository(table: _RecordingPermissionTable) -> ObjectPermissionRepository: + return ObjectPermissionRepository(SimpleNamespace(db=SimpleNamespace(litellm_objectpermissiontable=table))) + + +@pytest.mark.asyncio +async def test_create_permission_writes_only_the_fields_it_was_given() -> None: + table: Final = _RecordingPermissionTable({"object_permission_id": "perm-1"}) + + permission: Final = await _repository(table).create_permission( + mcp_servers=["server-1"], mcp_tool_permissions={"server-1": ["search"]}, models=[] + ) + + assert table.created == [ + {"mcp_servers": ["server-1"], "mcp_tool_permissions": {"server-1": ["search"]}, "models": []} + ] + assert permission.model_dump(exclude_unset=True) == { + "object_permission_id": "perm-1", + "mcp_servers": ["server-1"], + "mcp_tool_permissions": {"server-1": ["search"]}, + "models": [], + } + + +@pytest.mark.asyncio +async def test_create_permission_without_fields_writes_an_empty_row() -> None: + table: Final = _RecordingPermissionTable({"object_permission_id": "perm-1"}) + + permission: Final = await _repository(table).create_permission() + + assert table.created == [{}] + assert permission.model_dump(exclude_unset=True) == {"object_permission_id": "perm-1"} + + +@pytest.mark.asyncio +async def test_update_permission_changes_only_the_fields_it_was_given() -> None: + table: Final = _RecordingPermissionTable( + {"object_permission_id": "perm-1", "models": ["gpt-4o"], "agents": ["agent-1"]} + ) + + permission: Final = await _repository(table).update_permission("perm-1", models=["gpt-4o-mini"], skills=[]) + + assert table.updated == [({"object_permission_id": "perm-1"}, {"models": ["gpt-4o-mini"], "skills": []})] + assert permission is not None + assert permission.model_dump(exclude_unset=True) == { + "object_permission_id": "perm-1", + "models": ["gpt-4o-mini"], + "agents": ["agent-1"], + "skills": [], + } diff --git a/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py b/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py index 7ed1a436cb6..d1b110f288a 100644 --- a/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py +++ b/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py @@ -13,9 +13,12 @@ seam stayed untouched, so a guard that raises after the provider call would stil import base64 from contextlib import ExitStack from dataclasses import dataclass +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +import respx from fastapi import Response @@ -273,3 +276,68 @@ async def test_cancel__unified_job_id_allowed_when_managed_files_required(seams) await _cancel(_unified_job_id()) assert seams.router.acancel_fine_tuning_job.call_count == 1 + + +FINE_TUNING_API_BASE: Final = "https://fine-tuning.test/v1" +PROVIDER_JOBS_PAGE: Final = { + "object": "list", + "data": [ + { + "id": RAW_JOB_ID, + "created_at": 1234567890, + "fine_tuned_model": None, + "finished_at": None, + "hyperparameters": {"n_epochs": 1}, + "model": "gpt-4o-mini", + "object": "fine_tuning.job", + "organization_id": "org-test", + "result_files": [], + "seed": 0, + "status": "running", + "trained_tokens": None, + "training_file": RAW_FILE_ID, + "validation_file": None, + } + ], + "has_more": False, +} + + +async def _list(custom_llm_provider: str | None): + return await endpoints.list_fine_tuning_jobs( + request=FakeRequest(), + fastapi_response=Response(), + custom_llm_provider=custom_llm_provider, + target_model_names=None, + after=None, + limit=None, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + +@pytest.mark.asyncio +async def test_list__rejected_when_neither_a_provider_nor_a_target_model_is_named(seams): + with pytest.raises(ProxyException) as exc: + await _list(custom_llm_provider=None) + + assert exc.value.code == "400" + assert exc.value.message == "Invalid request, No litellm managed file id or custom_llm_provider provided." + + +@pytest.mark.asyncio +@respx.mock +async def test_list__returns_the_jobs_the_configured_provider_lists(seams): + respx.get(f"{FINE_TUNING_API_BASE}/fine_tuning/jobs").mock( + return_value=httpx.Response(200, json=PROVIDER_JOBS_PAGE) + ) + provider_config: Final = { + "custom_llm_provider": "openai", + "api_key": "sk-provider", + "api_base": FINE_TUNING_API_BASE, + } + + with patch.object(endpoints, "fine_tuning_config", [provider_config]): + page: Final = await _list(custom_llm_provider="openai") + + assert [job.id for job in page.data] == [RAW_JOB_ID] + assert page.has_more is False diff --git a/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py b/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py index 86f9aeafc08..62863163b8e 100644 --- a/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py +++ b/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py @@ -6,11 +6,21 @@ without needing a running proxy. """ import pytest +from fastapi import Request +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.management_endpoints.policy_endpoints.endpoints import ( GuardrailTestResultEntry, _compute_overall_action, _test_guardrail_definitions, + list_policies, +) +from litellm.proxy.policy_engine.policy_registry import get_policy_registry +from litellm.types.proxy.policy_engine import ( + PolicyGuardrailsResponse, + PolicyListResponse, + PolicyScopeResponse, + PolicySummaryItem, ) @@ -293,3 +303,68 @@ class TestEnrichPolicyTemplateStreamKeepalive: assert b": ping\n\n" not in chunks assert chunks[0] == b'data: {"type": "competitor", "name": "Rival Air"}\n\n' assert chunks[-1].startswith(b'data: {"type": "done"') + + +def _policy_list_request() -> Request: + return Request({"type": "http", "method": "GET", "path": "/policy/list", "headers": []}) + + +def _summary_item(inherit, resolved_guardrails, inheritance_chain) -> PolicySummaryItem: + return PolicySummaryItem( + inherit=inherit, + scope=PolicyScopeResponse(), + guardrails=PolicyGuardrailsResponse(), + resolved_guardrails=resolved_guardrails, + inheritance_chain=inheritance_chain, + ) + + +@pytest.fixture +def policy_registry(): + registry = get_policy_registry() + registry.clear() + yield registry + registry.clear() + + +@pytest.mark.asyncio +async def test_list_policies_is_empty_before_any_policy_is_loaded(policy_registry): + response = await list_policies(request=_policy_list_request(), user_api_key_dict=UserAPIKeyAuth()) + + assert response == PolicyListResponse(policies={}, total_count=0) + + +@pytest.mark.parametrize( + ("policies_config", "expected_policies"), + [ + ({}, {}), + ( + {"solo": {"description": "standalone", "guardrails": {"add": ["pii"]}}}, + {"solo": _summary_item(None, ["pii"], ["solo"])}, + ), + ( + { + "base": {"guardrails": {"add": ["pii"]}}, + "child": {"inherit": "base", "guardrails": {"add": ["audit"], "remove": ["pii"]}}, + "conditional": {"guardrails": {"add": ["toxicity"]}, "condition": {"model": "gpt-4.*"}}, + "empty": {"guardrails": {"remove": ["pii"]}}, + }, + { + "base": _summary_item(None, ["pii"], ["base"]), + "child": _summary_item("base", ["audit"], ["base", "child"]), + "conditional": _summary_item(None, ["toxicity"], ["conditional"]), + "empty": _summary_item(None, [], ["empty"]), + }, + ), + ], +) +@pytest.mark.asyncio +async def test_list_policies_reports_each_loaded_policy_with_its_resolved_guardrails( + policy_registry, policies_config, expected_policies +): + policy_registry.load_policies(policies_config) + + response = await list_policies(request=_policy_list_request(), user_api_key_dict=UserAPIKeyAuth()) + + assert response == PolicyListResponse(policies=expected_policies, total_count=0) + assert list(response.policies) == list(expected_policies) diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 8159890ef16..452fdec648f 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -6167,6 +6167,53 @@ class TestStrategyRouterWriteValidation: async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=candidate_id) as table: assert hasattr(table, "create") + @pytest.mark.asyncio + async def test_slot_counts_tuned_routers_whose_stored_params_are_json_strings( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from fastapi import HTTPException + + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.model_management_endpoints import _auto_router_capability_slot + from litellm.router_utils.auto_router_tuning_baseline import snapshot_tuning_baselines + + fake = self._FakeDb([]) + fake.tx_obj.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[ + { + "model_id": "a", + "model_name": "router-a", + "litellm_params": json.dumps( + {"model": "auto_router/complexity_router", "complexity_router_config": self._TUNED_A_EDITED} + ), + "model_info": json.dumps({"id": "a"}), + } + ] + ) + monkeypatch.setattr(proxy_server._license_check, "auto_router_capability_limit", lambda: 1) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr( + proxy_server, + "heuristic_v1_tuning_baselines", + snapshot_tuning_baselines( + [self._db_router_row("a", self._TUNED_A), self._db_router_row("b", self._TUNED_B)] + ), + ) + + with pytest.raises(HTTPException) as refused: + async with _auto_router_capability_slot( + fake, + effective_params={ + "model": "auto_router/complexity_router", + "complexity_router_config": self._TUNED_B_EDITED, + }, + model_id="b", + ): + pass + + assert refused.value.status_code == 403 + assert "changed heuristic scoring rules" in str(refused.value.detail) + @pytest.mark.asyncio async def test_add_new_model_refuses_a_second_tuned_heuristic_v1_router_without_a_model_id(self) -> None: """A create request carries no model_info at all, yet the quota still judges it: Deployment mints the diff --git a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py index 440f93d1387..7d3bd049a03 100644 --- a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py @@ -1219,6 +1219,34 @@ async def test_legacy_update_without_budget_fields_skips_budget_write(monkeypatc assert prisma.db.litellm_organizationtable.update.await_args.kwargs["data"]["organization_alias"] == "renamed" +@pytest.mark.asyncio +async def test_legacy_update_writes_sent_metadata_to_the_organization_row(monkeypatch): + prisma = await _run_legacy_update_organization( + monkeypatch, + body={"organization_id": "org-1", "metadata": {"team": "search", "limits": {"rpm": 5}}}, + existing_budget_id="budget-1", + ) + + organization_write = prisma.db.litellm_organizationtable.update.await_args + assert organization_write.kwargs["where"] == {"organization_id": "org-1"} + assert json.loads(organization_write.kwargs["data"]["metadata"]) == {"team": "search", "limits": {"rpm": 5}} + prisma.db.litellm_budgettable.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [[], ["litellm_budget_table"], "organization_id", 5, 1.5, True, None]) +async def test_legacy_update_rejects_json_body_that_is_not_an_object_before_reading_the_database(monkeypatch, body): + from pydantic import ValidationError + + from litellm.proxy import proxy_server + + with pytest.raises(ValidationError): + await _run_legacy_update_organization(monkeypatch, body=body, existing_budget_id="budget-1") + + proxy_server.prisma_client.db.litellm_organizationtable.find_unique.assert_not_awaited() + proxy_server.prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + + def test_build_budget_write_data_recomputes_reset_at_on_duration(): """A sent budget_duration recomputes budget_reset_at so the reset window follows the new duration.""" from litellm.proxy.management_endpoints.organization_endpoints import build_budget_write_data diff --git a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py index 14ee9db6ffd..d5dbd3df6a6 100644 --- a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py @@ -2,6 +2,7 @@ import inspect import json from collections.abc import Mapping, Sequence from contextlib import contextmanager +from datetime import datetime from types import MappingProxyType, SimpleNamespace from typing import Final, cast from unittest.mock import AsyncMock, Mock, patch @@ -1521,3 +1522,67 @@ async def test_add_tag_to_deployment_model_not_found(): assert exc_info.value.status_code == 500 assert "not found in database" in str(exc_info.value.detail) + + +class _StoredTagTable: + def __init__(self, model_info: object) -> None: + self.model_info = model_info + + async def find_many(self, where: object = None, include: object = None) -> list[SimpleNamespace]: + return [ + SimpleNamespace( + tag_name="routed-tag", + description="Routes to one model", + models=["model-1"], + model_info=self.model_info, + budget_id=None, + created_at=datetime(2025, 1, 1), + updated_at=datetime(2025, 1, 2), + created_by="user-123", + litellm_budget_table=None, + ) + ] + + +class _NoDynamicTagSpend: + async def group_by(self, by: object, where: object, min: object, max: object) -> list[object]: + return [] + + +@pytest.mark.parametrize( + ("stored_model_info", "returned_model_info"), + [ + ('{"model-1": "gpt-4o"}', {"model-1": "gpt-4o"}), + ({"model-1": "gpt-4o"}, {"model-1": "gpt-4o"}), + (None, {}), + ], +) +def test_tag_info_and_tag_list_return_the_stored_model_info_decoded( + monkeypatch, stored_model_info, returned_model_info +): + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace( + db=SimpleNamespace( + litellm_tagtable=_StoredTagTable(stored_model_info), + litellm_dailytagspend=_NoDynamicTagSpend(), + ) + ), + ) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + info_response = client.post("/tag/info", json={"names": ["routed-tag"]}) + list_response = client.get("/tag/list") + finally: + app.dependency_overrides.clear() + + assert info_response.status_code == 200 + assert info_response.json()["routed-tag"]["model_info"] == returned_model_info + assert list_response.status_code == 200 + assert [tag["model_info"] for tag in list_response.json()] == [returned_model_info] diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 094320225cd..1c9bd470e52 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -2533,3 +2533,22 @@ class TestRecordPartialUsageForFailure: assert "combined_usage_object" not in logging_obj.model_call_details assert "response_cost" not in logging_obj.model_call_details + + +@pytest.mark.parametrize( + ("all_chunks", "interrupted"), + [ + (['data: {"type": "content_block_delta"}'], True), + (['data: {"type": "message_delta"}', 'data: {"type": "message_stop"}'], False), + ([b'data: {"type": "message_delta"}\ndata: {"type": "message_stop"}\n'], False), + (['data: {"type": "message_delta"}', "data: [1, 2]", 'data: "text"', "data: 7", "data: null"], False), + (['data: {"type": "content_block_stop"}', "data: [1, 2]", "data: not json"], True), + (['data: {"type": "message_start"}', 'data: {"type": ["message_delta"]}', "data: {}"], True), + (["data: [1, 2]", "data: null", "event: message_delta"], True), + ([], True), + ], +) +def test_stream_was_interrupted_skips_data_lines_that_are_not_json_objects( + all_chunks: list[str | bytes], interrupted: bool +): + assert AnthropicPassthroughLoggingHandler._stream_was_interrupted(all_chunks) is interrupted 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 new file mode 100644 index 00000000000..f9df5a120f5 --- /dev/null +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py @@ -0,0 +1,95 @@ +from datetime import datetime + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.proxy._types import PassThroughEndpointLoggingTypedDict +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, +) +from litellm.types.utils import EmbeddingResponse, ModelResponse + +_PREDICT_ROUTE = "/v1/projects/p/locations/us-central1/publishers/google/models/text-embedding-004:predict" +_INTERACTIONS_ROUTE = "https://aiplatform.googleapis.com/v1beta1/projects/p/locations/global/interactions" + + +def _handle(url_route: str, payload: object) -> PassThroughEndpointLoggingTypedDict: + logging_obj = Logging( + model="unknown", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="pass_through_endpoint", + start_time=datetime(2026, 1, 1), + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.optional_params = {} + response = httpx.Response(200, json=payload) + return VertexPassthroughLoggingHandler.vertex_passthrough_handler( + httpx_response=response, + logging_obj=logging_obj, + url_route=url_route, + result=response.text, + start_time=datetime(2026, 1, 1), + end_time=datetime(2026, 1, 1), + cache_hit=False, + request_body={"model": "gemini-omni-flash-preview"}, + ) + + +def test_predict_response_with_text_embeddings_is_logged_as_an_embedding_response(): + result = _handle( + _PREDICT_ROUTE, + { + "predictions": [ + {"embeddings": {"values": [0.1, 0.2], "statistics": {"token_count": 3}}}, + {"embeddings": {"values": [0.3, 0.4], "statistics": {"token_count": 4}}}, + ], + "metadata": {"billableCharacterCount": 9}, + }, + ) + + response = result["result"] + assert isinstance(response, EmbeddingResponse) + assert response.data == [ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}, + {"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}, + ] + assert response.usage.prompt_tokens == 7 + assert result["kwargs"]["model"] == "text-embedding-004" + assert result["kwargs"]["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "text", 7, 1.5, True]) +def test_predict_response_that_is_not_a_json_object_is_rejected_without_echoing_it(payload: object): + with pytest.raises(ValidationError) as exc_info: + _handle(_PREDICT_ROUTE, payload) + + assert "input_value" not in str(exc_info.value) + + +def test_interactions_usage_object_is_read_into_prompt_and_completion_tokens(): + result = _handle( + _INTERACTIONS_ROUTE, + { + "id": "interactions/abc", + "model": "gemini-omni-flash-preview", + "usage": { + "total_tokens": 41, + "total_input_tokens": 12, + "input_tokens_by_modality": [{"modality": "text", "tokens": 12}], + "total_output_tokens": 9, + "output_tokens_by_modality": [{"modality": "text", "tokens": 9}], + "total_thought_tokens": 20, + }, + }, + ) + + response = result["result"] + assert isinstance(response, ModelResponse) + assert response.usage.prompt_tokens == 12 + assert response.usage.completion_tokens == 29 + assert response.usage.completion_tokens_details.text_tokens == 9 + assert result["kwargs"]["custom_llm_provider"] == "vertex_ai" diff --git a/tests/unit/proxy/policy_engine/test_policy_registry.py b/tests/unit/proxy/policy_engine/test_policy_registry.py new file mode 100644 index 00000000000..db7775b6238 --- /dev/null +++ b/tests/unit/proxy/policy_engine/test_policy_registry.py @@ -0,0 +1,154 @@ +import json +from collections.abc import Mapping +from datetime import datetime, timezone +from types import SimpleNamespace + +import pytest + +from litellm.proxy.policy_engine.policy_registry import PolicyRegistry +from litellm.types.proxy.policy_engine import PolicyCreateRequest, PolicyUpdateRequest + +_NOW = datetime(2026, 1, 1, tzinfo=timezone.utc) + +_REQUESTED_PIPELINE = {"mode": "pre_call", "steps": [{"guardrail": "pii-guard", "on_fail": "next"}]} + +_STORED_PIPELINE = { + "mode": "pre_call", + "steps": [ + { + "guardrail": "pii-guard", + "on_fail": "next", + "on_pass": "allow", + "on_error": None, + "pass_data": False, + "modify_response_message": None, + } + ], +} + +_MALFORMED_PIPELINES = [ + pytest.param({"mode": "pre_call", "steps": []}, "steps", id="no-steps"), + pytest.param({"mode": "pre_call"}, "steps", id="steps-missing"), + pytest.param({"steps": [{"guardrail": "pii-guard"}]}, "mode", id="mode-missing"), + pytest.param({"mode": "during_call", "steps": [{"guardrail": "pii-guard"}]}, "mode", id="unknown-mode"), + pytest.param( + {"mode": "pre_call", "steps": [{"guardrail": "pii-guard"}], "order": "strict"}, "order", id="unknown-key" + ), + pytest.param({"mode": "pre_call", "steps": "pii-guard"}, "steps", id="steps-not-a-list"), +] + + +class _PolicyTable: + def __init__(self) -> None: + self.writes: list[Mapping[str, object]] = [] + + def _row(self, data: Mapping[str, object]) -> SimpleNamespace: + pipeline = data.get("pipeline") + return SimpleNamespace( + policy_id="policy-1", + policy_name=data.get("policy_name", "pii-policy"), + version_number=1, + version_status=data.get("version_status", "draft"), + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description=None, + guardrails_add=data.get("guardrails_add", []), + guardrails_remove=data.get("guardrails_remove", []), + condition=None, + pipeline=json.loads(pipeline) if isinstance(pipeline, str) else None, + created_at=_NOW, + updated_at=_NOW, + created_by=None, + updated_by=None, + ) + + async def create(self, data: Mapping[str, object]) -> SimpleNamespace: + self.writes.append(data) + return self._row(data) + + async def find_unique(self, where: Mapping[str, object]) -> SimpleNamespace: + return self._row({"policy_name": "pii-policy", "version_status": "draft"}) + + async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SimpleNamespace: + self.writes.append(data) + return self._row(data) + + +def _prisma_client(table: _PolicyTable) -> SimpleNamespace: + return SimpleNamespace(db=SimpleNamespace(litellm_policytable=table)) + + +@pytest.mark.asyncio +async def test_add_policy_to_db_stores_the_pipeline_with_step_defaults_filled_in(): + table = _PolicyTable() + registry = PolicyRegistry() + + response = await registry.add_policy_to_db( + PolicyCreateRequest(policy_name="pii-policy", guardrails_add=["pii-guard"], pipeline=_REQUESTED_PIPELINE), + _prisma_client(table), + ) + + assert [json.loads(str(write["pipeline"])) for write in table.writes] == [_STORED_PIPELINE] + assert response.pipeline == _STORED_PIPELINE + stored_policy = registry.get_policy("pii-policy") + assert stored_policy is not None + assert stored_policy.pipeline is not None + assert stored_policy.pipeline.model_dump() == _STORED_PIPELINE + + +@pytest.mark.asyncio +async def test_update_policy_in_db_stores_the_pipeline_with_step_defaults_filled_in(): + table = _PolicyTable() + + response = await PolicyRegistry().update_policy_in_db( + "policy-1", + PolicyUpdateRequest(pipeline=_REQUESTED_PIPELINE), + _prisma_client(table), + ) + + assert [json.loads(str(write["pipeline"])) for write in table.writes] == [_STORED_PIPELINE] + assert response.pipeline == _STORED_PIPELINE + + +@pytest.mark.parametrize(("pipeline", "rejected_field"), _MALFORMED_PIPELINES) +@pytest.mark.asyncio +async def test_add_policy_to_db_rejects_a_malformed_pipeline_before_writing( + pipeline: dict[str, object], rejected_field: str +): + table = _PolicyTable() + registry = PolicyRegistry() + + with pytest.raises(Exception, match="Error adding policy to DB: 1 validation error") as raised: + await registry.add_policy_to_db( + PolicyCreateRequest(policy_name="pii-policy", pipeline=pipeline), + _prisma_client(table), + ) + + assert str(raised.value).startswith( + f"Error adding policy to DB: 1 validation error for GuardrailPipeline\n{rejected_field}\n" + ) + assert table.writes == [] + assert registry.get_policy("pii-policy") is None + + +@pytest.mark.parametrize(("pipeline", "rejected_field"), _MALFORMED_PIPELINES) +@pytest.mark.asyncio +async def test_update_policy_in_db_rejects_a_malformed_pipeline_before_writing( + pipeline: dict[str, object], rejected_field: str +): + table = _PolicyTable() + + with pytest.raises(Exception, match="Error updating policy in DB: 1 validation error") as raised: + await PolicyRegistry().update_policy_in_db( + "policy-1", + PolicyUpdateRequest(pipeline=pipeline), + _prisma_client(table), + ) + + assert str(raised.value).startswith( + f"Error updating policy in DB: 1 validation error for GuardrailPipeline\n{rejected_field}\n" + ) + assert table.writes == [] diff --git a/tests/unit/proxy/spend_tracking/test_budget_reservation.py b/tests/unit/proxy/spend_tracking/test_budget_reservation.py index 9df8e6f4d67..024fe5229be 100644 --- a/tests/unit/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/unit/proxy/spend_tracking/test_budget_reservation.py @@ -11,6 +11,8 @@ from litellm.caching import DualCache from litellm.models.budget import LiteLLM_BudgetTable from litellm.proxy import proxy_server from litellm.proxy._types import ( + LiteLLM_OrganizationTable, + LiteLLM_ProjectTableCachedObj, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, @@ -18,11 +20,13 @@ from litellm.proxy._types import ( ) from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, + project_cache_key, team_membership_reservation_cache_key, ) from litellm.proxy.spend_tracking.budget_reservation import ( _get_team_member_budget_counter, estimate_request_max_cost, + get_budget_window_start, release_unbound_budget_reservation, reserve_budget_for_request, ) @@ -342,3 +346,83 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"]) assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_reservation_holds_cost_against_cached_org_and_project_budgets(spend_counter_cache: DualCache): + cache: Final = UserApiKeyCache() + await cache.async_set_cache( + key="org_id:org-budgeted:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-budgeted", + budget_id="org-budget", + spend=1.5, + models=[], + created_by="admin", + updated_by="admin", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ), + model_type=LiteLLM_OrganizationTable, + ) + await cache.async_set_cache( + key=project_cache_key("project-budgeted"), + value=LiteLLM_ProjectTableCachedObj( + project_id="project-budgeted", spend=2.5, litellm_budget_table=LiteLLM_BudgetTable(max_budget=20.0) + ), + model_type=LiteLLM_ProjectTableCachedObj, + ) + + reservation: Final = await reserve_budget_for_request( + request_body={"model": "gpt-4o", "input": "hello"}, + route="/v1/responses", + llm_router=None, + valid_token=UserAPIKeyAuth(token="hashed-org-project", org_id="org-budgeted", project_id="project-budgeted"), + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=cache, + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + assert reservation is not None + reserved_cost: Final = reservation["reserved_cost"] + assert reserved_cost > 0 + assert reservation["entries"] == [ + { + "counter_key": "spend:org:org-budgeted", + "entity_type": "Organization", + "entity_id": "org-budgeted", + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + }, + { + "counter_key": "spend:project:project-budgeted", + "entity_type": "Project", + "entity_id": "project-budgeted", + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + }, + ] + assert spend_counter_cache.in_memory_cache.get_cache(key="spend:org:org-budgeted") == pytest.approx(reserved_cost) + assert spend_counter_cache.in_memory_cache.get_cache(key="spend:project:project-budgeted") == pytest.approx( + reserved_cost + ) + + +@pytest.mark.parametrize( + ("window", "expected"), + [ + ( + '{"budget_duration": "1h", "reset_at": "2030-01-01T01:00:00Z"}', + datetime(2030, 1, 1, 0, 0, tzinfo=timezone.utc), + ), + ('{"reset_at": "2030-01-01T01:00:00Z"}', None), + ("{}", None), + ('["budget_duration", "1h"]', None), + ('"1h"', None), + ("null", None), + ("not json", None), + ], +) +def test_budget_window_start_reads_json_encoded_windows(window: str, expected: datetime | None) -> None: + assert get_budget_window_start(window) == expected diff --git a/tests/unit/rag/ingestion/test_base_ingestion.py b/tests/unit/rag/ingestion/test_base_ingestion.py new file mode 100644 index 00000000000..d20b46d356b --- /dev/null +++ b/tests/unit/rag/ingestion/test_base_ingestion.py @@ -0,0 +1,101 @@ +from collections import UserString +from typing import Final + +import httpx +import pytest +import respx + +import litellm +from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion +from litellm.types.utils import CredentialItem + +_FILE_URL: Final = "https://files.example/docs/report.pdf" +_STORED_CREDENTIAL: Final = CredentialItem( + credential_name="gemini-prod", + credential_info={}, + credential_values={"api_key": "stored-key"}, +) + + +def _vector_store_options(litellm_credential_name: object) -> dict[str, object]: + return { + "custom_llm_provider": "gemini", + "litellm_credential_name": litellm_credential_name, + "api_key": "caller-key", + "api_base": "https://caller.example", + } + + +def test_a_stored_credential_named_by_the_vector_store_replaces_the_caller_supplied_values( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + vector_store: Final = _vector_store_options("gemini-prod") + + ingestion: Final = GeminiRAGIngestion(ingest_options={"vector_store": vector_store}) + + assert ingestion.vector_store_config is vector_store + assert vector_store == { + "custom_llm_provider": "gemini", + "litellm_credential_name": "gemini-prod", + "api_key": "stored-key", + } + + +@pytest.mark.parametrize( + "litellm_credential_name", + ["gemini-staging", "", None, 7, True, ["gemini-prod"], {"credential_name": "gemini-prod"}], +) +def test_a_credential_name_that_matches_no_stored_credential_leaves_the_vector_store_config_alone( + monkeypatch: pytest.MonkeyPatch, litellm_credential_name: object +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + + ingestion: Final = GeminiRAGIngestion( + ingest_options={"vector_store": _vector_store_options(litellm_credential_name)} + ) + + assert ingestion.vector_store_config == _vector_store_options(litellm_credential_name) + + +def test_a_credential_name_that_is_not_a_string_is_not_resolved_even_when_it_equals_a_stored_name( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + litellm_credential_name: Final = UserString("gemini-prod") + + ingestion: Final = GeminiRAGIngestion( + ingest_options={"vector_store": _vector_store_options(litellm_credential_name)} + ) + + assert litellm_credential_name == "gemini-prod" + assert ingestion.vector_store_config == _vector_store_options(litellm_credential_name) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("headers", "expected_content_type"), + [ + ([("content-type", "application/pdf")], "application/pdf"), + ([("Content-Type", "text/plain; charset=utf-8")], "text/plain; charset=utf-8"), + ([("content-type", "")], ""), + ([("content-type", "text/plain"), ("content-type", "text/html")], "text/plain, text/html"), + ([], "application/octet-stream"), + ([("content-length", "8")], "application/octet-stream"), + ], +) +async def test_upload_from_a_url_takes_the_content_type_from_the_response_header( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + headers: list[tuple[str, str]], + expected_content_type: str, +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "user_url_validation", False) + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.get(_FILE_URL).mock(return_value=httpx.Response(200, content=b"%PDF-1.7", headers=headers)) + ingestion: Final = GeminiRAGIngestion(ingest_options={"vector_store": {"custom_llm_provider": "gemini"}}) + + uploaded: Final = await ingestion.upload(file_url=_FILE_URL) + + assert uploaded == ("report.pdf", b"%PDF-1.7", expected_content_type, None) diff --git a/tests/unit/rag/ingestion/test_gemini_ingestion.py b/tests/unit/rag/ingestion/test_gemini_ingestion.py new file mode 100644 index 00000000000..8122a9d334b --- /dev/null +++ b/tests/unit/rag/ingestion/test_gemini_ingestion.py @@ -0,0 +1,111 @@ +import json +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +import httpx +import pytest +import respx +from pydantic import ValidationError + +import litellm +from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion + +_API_BASE: Final = "https://gemini.example" +_STORE: Final = "fileSearchStores/docs-1" +_START_UPLOAD_URL: Final = f"{_API_BASE}/upload/v1beta/{_STORE}:uploadToFileSearchStore" +_UPLOAD_SESSION_URL: Final = "https://gemini.example/upload/session-1" +_DOCUMENT: Final = "fileSearchStores/docs-1/documents/notes-1" + + +def _ingestion(chunking_strategy: object) -> GeminiRAGIngestion: + return GeminiRAGIngestion( + ingest_options={ + "chunking_strategy": chunking_strategy, + "vector_store": { + "custom_llm_provider": "gemini", + "vector_store_id": _STORE, + "api_key": "test-key", + "api_base": _API_BASE, + }, + } + ) + + +async def _store_notes(ingestion: GeminiRAGIngestion) -> tuple[str | None, str | None]: + return await ingestion.store( + file_content=b"first note", + filename="notes.txt", + content_type="text/plain", + chunks=[], + embeddings=None, + ) + + +@pytest.fixture +def start_upload(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.put(_UPLOAD_SESSION_URL).mock(return_value=httpx.Response(200, json={"name": _DOCUMENT})) + return respx_mock.post(_START_UPLOAD_URL).mock( + return_value=httpx.Response(200, headers={"x-goog-upload-url": _UPLOAD_SESSION_URL}) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("white_space_config", "expected"), + [ + ( + {"max_tokens_per_chunk": 200, "max_overlap_tokens": 20}, + {"maxTokensPerChunk": 200, "maxOverlapTokens": 20}, + ), + ({"max_tokens_per_chunk": 200}, {"maxTokensPerChunk": 200, "maxOverlapTokens": 400}), + ({"unrelated": True}, {"maxTokensPerChunk": 800, "maxOverlapTokens": 400}), + ( + {"max_tokens_per_chunk": "200", "max_overlap_tokens": None}, + {"maxTokensPerChunk": "200", "maxOverlapTokens": None}, + ), + ({7: "not a field", "max_overlap_tokens": 1.5}, {"maxTokensPerChunk": 800, "maxOverlapTokens": 1.5}), + (MappingProxyType({"max_overlap_tokens": 0}), {"maxTokensPerChunk": 800, "maxOverlapTokens": 0}), + ], +) +async def test_white_space_config_is_sent_as_the_chunking_config_of_the_upload( + start_upload: respx.Route, white_space_config: Mapping[object, object], expected: Mapping[str, object] +): + stored: Final = await _store_notes(_ingestion({"white_space_config": white_space_config})) + + assert stored == (_STORE, _DOCUMENT) + assert json.loads(start_upload.calls.last.request.content) == { + "displayName": "notes.txt", + "chunkingConfig": {"whiteSpaceConfig": expected}, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "chunking_strategy", + [None, {"type": "auto"}, {"white_space_config": None}, {"white_space_config": {}}, {"white_space_config": 0}], +) +async def test_upload_without_a_white_space_config_sends_no_chunking_config( + start_upload: respx.Route, chunking_strategy: object +): + stored: Final = await _store_notes(_ingestion(chunking_strategy)) + + assert stored == (_STORE, _DOCUMENT) + assert json.loads(start_upload.calls.last.request.content) == {"displayName": "notes.txt"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "white_space_config", + ["private-setting", [800, 400], [("max_tokens_per_chunk", 200)], 800, True, 1.5], +) +async def test_a_white_space_config_that_is_not_a_mapping_is_rejected_before_any_upload( + start_upload: respx.Route, white_space_config: object +): + with pytest.raises(ValidationError) as raised: + await _store_notes(_ingestion({"white_space_config": white_space_config})) + + assert "private-setting" not in str(raised.value) + assert not start_upload.called diff --git a/tests/unit/rag/ingestion/test_openai_ingestion.py b/tests/unit/rag/ingestion/test_openai_ingestion.py new file mode 100644 index 00000000000..74b68c9a517 --- /dev/null +++ b/tests/unit/rag/ingestion/test_openai_ingestion.py @@ -0,0 +1,101 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm +from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion + +_API_BASE: Final = "https://openai.example/v1" +_STATIC_CHUNKING: Final = {"type": "static", "static": {"max_chunk_size_tokens": 800, "chunk_overlap_tokens": 400}} +_ATTACHED_FILE: Final = { + "id": "file-1", + "object": "vector_store.file", + "created_at": 1767323045, + "vector_store_id": "vs_1", + "status": "completed", + "usage_bytes": 10, +} +_UPLOADED_FILE: Final = { + "id": "file-1", + "object": "file", + "bytes": 10, + "created_at": 1767323045, + "filename": "notes.txt", + "purpose": "assistants", + "status": "processed", +} + + +def _ingestion(chunking_strategy: object) -> OpenAIRAGIngestion: + return OpenAIRAGIngestion( + ingest_options={ + "chunking_strategy": chunking_strategy, + "vector_store": { + "custom_llm_provider": "openai", + "vector_store_id": "vs_1", + "api_key": "sk-test", + "api_base": _API_BASE, + }, + } + ) + + +@pytest.fixture +def attach_file(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + return respx_mock.post(f"{_API_BASE}/vector_stores/vs_1/files").mock( + return_value=httpx.Response(200, json=_ATTACHED_FILE) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("chunking_strategy", "sent_chunking_strategy"), + [(None, {"type": "auto"}), ({"type": "auto"}, {"type": "auto"}), (_STATIC_CHUNKING, _STATIC_CHUNKING)], +) +async def test_an_uploaded_file_is_attached_to_the_vector_store_with_the_chunking_strategy( + attach_file: respx.Route, respx_mock: respx.MockRouter, chunking_strategy: object, sent_chunking_strategy: object +): + respx_mock.post(f"{_API_BASE}/files").mock(return_value=httpx.Response(200, json=_UPLOADED_FILE)) + + stored: Final = await _ingestion(chunking_strategy).store( + file_content=b"first note", + filename="notes.txt", + content_type="text/plain", + chunks=[], + embeddings=None, + ) + + assert stored == ("vs_1", "file-1") + assert json.loads(attach_file.calls.last.request.content) == { + "file_id": "file-1", + "chunking_strategy": sent_chunking_strategy, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("chunking_strategy", "sent_chunking_strategy"), + [(None, {"type": "auto"}), (_STATIC_CHUNKING, _STATIC_CHUNKING)], +) +async def test_an_existing_file_is_attached_to_the_vector_store_with_the_chunking_strategy( + attach_file: respx.Route, chunking_strategy: object, sent_chunking_strategy: object +): + stored: Final = await _ingestion(chunking_strategy).store( + file_content=None, + filename=None, + content_type=None, + chunks=[], + embeddings=None, + existing_file_id="file-9", + ) + + assert stored == ("vs_1", "file-9") + assert json.loads(attach_file.calls.last.request.content) == { + "file_id": "file-9", + "chunking_strategy": sent_chunking_strategy, + } diff --git a/tests/unit/rag/test_main.py b/tests/unit/rag/test_main.py index 54748efb480..81318a1e113 100644 --- a/tests/unit/rag/test_main.py +++ b/tests/unit/rag/test_main.py @@ -12,12 +12,14 @@ aquery carries the completion response with real usage and cost. import asyncio import json +from types import MappingProxyType from typing import Final from unittest.mock import patch import httpx import pytest import respx +from pydantic import ValidationError import litellm from litellm._internal_context import is_internal_call @@ -538,6 +540,91 @@ async def test_aquery_forwards_vector_store_params_to_search_but_not_completion( assert not ({"milvus_text_field", "outputFields"} & set(completion_kwargs)) +_UNDECODABLE_FILE: Final = {"filename": "notes.txt", "content": "x"} + + +@pytest.mark.parametrize( + ("ingest_options", "expected_provider"), + [ + ({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"), + ({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"), + ({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None), + ({}, None), + ], +) +def test_ingest_failure_is_attributed_to_the_vector_store_provider( + ingest_options: dict[str, object], expected_provider: str | None +) -> None: + with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised: + litellm.ingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE) + + assert raised.value.llm_provider == expected_provider + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("ingest_options", "expected_provider"), + [ + ({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"), + ({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"), + ({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None), + ({}, None), + ], +) +async def test_aingest_failure_is_attributed_to_the_vector_store_provider( + ingest_options: dict[str, object], expected_provider: str | None +) -> None: + with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised: + await litellm.aingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE) + + assert raised.value.llm_provider == expected_provider + + +@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]]) +def test_ingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error( + vector_store: object, +) -> None: + with pytest.raises(ValidationError) as raised: + litellm.ingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE) + + assert "sk-test" not in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]]) +async def test_aingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error( + vector_store: object, +) -> None: + with pytest.raises(ValidationError) as raised: + await litellm.aingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE) + + assert "sk-test" not in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", [None, 5]) +async def test_aquery_with_a_provider_that_is_not_a_string_bills_only_the_completion(provider: object) -> None: + messages: Final = [{"role": "user", "content": "hello"}] + default_provider_response: Final = await litellm.aquery( + model="gpt-4o-mini", + messages=messages, + retrieval_config={"vector_store_id": "vs_test_123"}, + mock_response="hi there", + ) + + response: Final = await litellm.aquery( + model="gpt-4o-mini", + messages=messages, + retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": provider}, + mock_response="hi there", + ) + + await _drain_logging_worker() + + assert response._hidden_params["response_cost"] > 0 + assert response._hidden_params["response_cost"] == default_provider_response._hidden_params["response_cost"] + + def test_rag_call_types_are_registered(): """ query/aquery/ingest/aingest are @client-decorated entry points, so their diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index bae6db9ee88..3aed934f40c 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -11,11 +11,13 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from prisma import models as prisma_models from prisma.builder import QueryBuilder +from pydantic import ValidationError from litellm.models.base import DomainModel from litellm.models.budget import LiteLLM_BudgetTable from litellm.models.credentials import CredentialItem from litellm.models.team import LiteLLM_TeamTable +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.repositories.base_repository import BaseRepository from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE @@ -308,6 +310,29 @@ class TestBudgetRepository: assert budget.budget_id == "budget-1" +_ROW_TIMESTAMP: Final = datetime(2026, 1, 1, 12, 0, 0) + + +def _stored_proxy_model_row(*, litellm_params: str, model_info: str) -> prisma_models.LiteLLM_ProxyModelTable: + return prisma_models.LiteLLM_ProxyModelTable( + model_id="other-model", + model_name="gpt-4o", + litellm_params=litellm_params, + model_info=model_info, + blocked=False, + created_at=_ROW_TIMESTAMP, + created_by="admin", + updated_at=_ROW_TIMESTAMP, + updated_by="admin", + ) + + +def _proxy_model_client(rows: list[prisma_models.LiteLLM_ProxyModelTable]) -> SimpleNamespace: + return SimpleNamespace( + db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=AsyncMock(return_value=rows))) + ) + + class TestModelRepository: @pytest.fixture def repo(self): @@ -331,6 +356,56 @@ class TestModelRepository: ).build_query() assert 'where: { model_id: { not: "current-model" } }' in " ".join(query.split()) + @pytest.mark.parametrize("double_encoded", [False, True]) + @pytest.mark.asyncio + async def test_find_all_except_returns_stored_rows_with_decrypted_params( + self, monkeypatch: pytest.MonkeyPatch, double_encoded: bool + ) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt") + stored_params: Final = json.dumps( + {"model": "openai/gpt-4o", "api_key": encrypt_value_helper("sk-secret"), "rpm": 5, "tags": ["prod"]} + ) + stored_info: Final = json.dumps({"id": "other-model", "team_id": "team-1"}) + row: Final = _stored_proxy_model_row( + litellm_params=json.dumps(stored_params) if double_encoded else stored_params, + model_info=json.dumps(stored_info) if double_encoded else stored_info, + ) + + models: Final = await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [model.model_dump() for model in models] == [ + { + "model_id": "other-model", + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-secret", "rpm": 5, "tags": ["prod"]}, + "model_info": {"id": "other-model", "team_id": "team-1"}, + "blocked": False, + "created_at": _ROW_TIMESTAMP, + "created_by": "admin", + "updated_at": _ROW_TIMESTAMP, + "updated_by": "admin", + } + ] + + @pytest.mark.asyncio + async def test_find_all_except_keeps_empty_params_and_missing_model_info(self) -> None: + row: Final = _stored_proxy_model_row(litellm_params="{}", model_info="null") + + models: Final = await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [(model.litellm_params, model.model_info) for model in models] == [({}, None)] + + @pytest.mark.parametrize("stored_params", ['["sk-secret"]', "1", "true", json.dumps('["sk-secret"]'), '"1"']) + @pytest.mark.asyncio + async def test_find_all_except_rejects_rows_whose_params_are_not_a_json_object(self, stored_params: str) -> None: + row: Final = _stored_proxy_model_row(litellm_params=stored_params, model_info="null") + + with pytest.raises(ValidationError) as rejected: + await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [error["type"] for error in rejected.value.errors()] == ["dict_type"] + assert "sk-secret" not in str(rejected.value) + def test_table_is_wrapped_for_config_sync(self, repo): from litellm.proxy.common_utils.config_sync_pubsub import ( _PublishOnWriteActions, diff --git a/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py b/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py index d717c4e8c89..9144bbc1788 100644 --- a/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py @@ -7,6 +7,7 @@ from litellm.router_strategy.adaptive_router import adaptive_router as ar_module import pytest from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter +from litellm.router_strategy.adaptive_router.bandit import initial_cell from litellm.router_strategy.adaptive_router.signals import Turn from litellm.types.router import ( AdaptiveRouterConfig, @@ -416,3 +417,11 @@ def test_session_state_expiry_is_refreshed_on_access(): second_exp = r._session_states_expiry[("sess-A", "fast")] assert second_exp > first_exp + + +@pytest.mark.parametrize("model", ["fast", "smart"]) +def test_cell_returns_the_cold_start_prior_for_an_available_model(model): + r = _make_router() + cell = r.cell(RequestType.CODE_GENERATION, model) + assert cell == initial_cell(r.model_to_prefs[model], RequestType.CODE_GENERATION) + assert cell.total_samples == 0 diff --git a/tests/unit/secret_managers/test_aws_secret_manager.py b/tests/unit/secret_managers/test_aws_secret_manager.py new file mode 100644 index 00000000000..d9922b43ec8 --- /dev/null +++ b/tests/unit/secret_managers/test_aws_secret_manager.py @@ -0,0 +1,36 @@ +import base64 + +import pytest +from botocore.stub import Stubber + +from litellm.secret_managers.aws_secret_manager import AWSKeyManagementService_V2 + +_CIPHERTEXT = b"ciphertext-bytes" +_KEY_ID = "arn:aws:kms:us-west-2:111122223333:key/1234abcd-12ab-34cd-56ef-1234567890ab" + + +@pytest.mark.parametrize( + ("plaintext", "expected"), + [ + pytest.param(b"True", True, id="true literal"), + pytest.param(b"False", False, id="false literal"), + pytest.param(b" True\n", True, id="padded true literal"), + pytest.param(b"1", "1", id="integer literal"), + pytest.param(b"[True]", "[True]", id="list literal"), + pytest.param(b"'True'", "'True'", id="quoted string literal"), + pytest.param(b"sk-not-a-literal", "sk-not-a-literal", id="plain secret"), + pytest.param(b"", "", id="empty secret"), + ], +) +def test_decrypt_value_turns_only_boolean_literals_into_bools(monkeypatch, plaintext, expected): + monkeypatch.setenv("AWS_REGION_NAME", "us-west-2") + monkeypatch.setenv("LITELLM_LICENSE", "license-for-test") + monkeypatch.setenv("ENCRYPTED_SETTING", "aws_kms/" + base64.b64encode(_CIPHERTEXT).decode()) + kms = AWSKeyManagementService_V2() + + with Stubber(kms.kms_client) as stubber: + stubber.add_response("decrypt", {"KeyId": _KEY_ID, "Plaintext": plaintext}, {"CiphertextBlob": _CIPHERTEXT}) + decrypted = kms.decrypt_value(secret_name="ENCRYPTED_SETTING") + + assert decrypted == expected + assert type(decrypted) is type(expected) diff --git a/tests/unit/secret_managers/test_secret_managers_main.py b/tests/unit/secret_managers/test_secret_managers_main.py index acc91b691d4..32040251795 100644 --- a/tests/unit/secret_managers/test_secret_managers_main.py +++ b/tests/unit/secret_managers/test_secret_managers_main.py @@ -431,3 +431,38 @@ def test_secret_manager_would_be_consulted_is_false_without_a_client(monkeypatch monkeypatch.setattr(litellm, "secret_manager_client", None) assert secret_manager_would_be_consulted("os.environ/ANY_NAME") is False + + +class _FixedValueSecretManager(CustomSecretManager): + def __init__(self, value): + self.value = value + + def sync_read_secret(self, secret_name, optional_params=None, timeout=None): + return self.value + + async def async_read_secret(self, secret_name, optional_params=None, timeout=None): + return self.value + + +@pytest.mark.parametrize( + ("stored", "expected"), + [ + ("True", True), + ("False", False), + ("1", "1"), + ("[True]", "[True]"), + ("'True'", "'True'"), + ("sk-not-a-literal", "sk-not-a-literal"), + ("", ""), + ], +) +def test_get_secret_turns_only_boolean_literals_from_the_secret_manager_into_bools(monkeypatch, stored, expected): + monkeypatch.setattr(litellm, "secret_manager_client", _FixedValueSecretManager(stored)) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode="read_only")) + monkeypatch.delenv("STORED_FLAG", raising=False) + + secret = get_secret("os.environ/STORED_FLAG") + + assert secret == expected + assert type(secret) is type(expected) diff --git a/tests/unit/test_dashscope_image_generation.py b/tests/unit/test_dashscope_image_generation.py index 6f91fe9a0e0..960807a033b 100644 --- a/tests/unit/test_dashscope_image_generation.py +++ b/tests/unit/test_dashscope_image_generation.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.dashscope.image_generation.transformation import ( @@ -455,3 +456,98 @@ def test_litellm_image_generation_dashscope_end_to_end(model: str): assert "input" in body assert "messages" in body["input"] assert body["parameters"]["size"] == "1024*1024" + + +def _transform_response(payload: object) -> ImageResponse: + return DashScopeImageGenerationConfig().transform_image_generation_response( + model="qwen-image-2.0", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + "payload", + [ + {}, + {"output": {}}, + {"output": {"choices": []}}, + {"output": {"choices": ""}}, + {"output": {"choices": [{}]}}, + {"output": {"choices": [{"message": {}}]}}, + {"output": {"choices": [{"message": {"content": {}}}]}}, + {"output": {"choices": [{"message": {"content": [{}]}}]}}, + {"output": {"choices": [{"message": {"content": [{"image": ""}]}}]}}, + {"output": {"choices": [{"message": {"content": [{"text": "hi"}]}}]}}, + {"code": "Partial", "output": {"choices": []}}, + ], +) +def test_transform_response_without_image_content_has_no_images( + payload: dict[str, object], +): + assert _transform_response(payload).data == [] + + +def test_transform_response_skips_content_items_without_an_image(): + response = _transform_response( + { + "output": { + "choices": [ + {"message": {"content": [{"text": "caption"}, {"image": "https://a.example/1.png"}]}}, + {"message": {"content": [{"image": None}, {"image": "https://a.example/2.png"}]}}, + ] + } + } + ) + + assert [image.url for image in response.data] == [ + "https://a.example/1.png", + "https://a.example/2.png", + ] + + +@pytest.mark.parametrize( + ("payload", "message"), + [ + ({"code": "InvalidParameter", "message": "Size not supported"}, "Size not supported"), + ({"code": "Throttled"}, "{'code': 'Throttled'}"), + ({"code": 429, "message": {"detail": "slow down"}}, "{'detail': 'slow down'}"), + ], +) +def test_transform_response_reports_api_error_bodies( + payload: dict[str, object], message: str +): + with pytest.raises(BaseLLMException) as exc_info: + _transform_response(payload) + + assert exc_info.value.message == message + assert exc_info.value.status_code == 200 + + +@pytest.mark.parametrize( + "payload", + [ + ["code"], + "code", + 7, + {"output": None}, + {"output": ["not", "an", "object"]}, + {"output": {"choices": 7}}, + {"output": {"choices": ["not an object"]}}, + {"output": {"choices": [{"message": "not an object"}]}}, + {"output": {"choices": [{"message": {"content": None}}]}}, + {"output": {"choices": [{"message": {"content": ["not an object"]}}]}}, + ], +) +def test_transform_response_rejects_malformed_payloads_without_echoing_them( + payload: object, +): + with pytest.raises(ValidationError) as exc_info: + _transform_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/vector_stores/test_vector_store_registry.py b/tests/unit/vector_stores/test_vector_store_registry.py index 762176d6a81..acfaccf8e2d 100644 --- a/tests/unit/vector_stores/test_vector_store_registry.py +++ b/tests/unit/vector_stores/test_vector_store_registry.py @@ -1,10 +1,14 @@ import json +from collections.abc import Mapping, Sequence +from types import SimpleNamespace +from typing import Final from unittest.mock import patch import httpx import pytest import respx from fastapi.testclient import TestClient +from pydantic import BaseModel, ValidationError from datetime import datetime, timezone @@ -13,7 +17,7 @@ from unittest.mock import AsyncMock, MagicMock import litellm from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.vector_stores.main import search -from litellm.vector_stores.vector_store_registry import VectorStoreRegistry +from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry, VectorStoreRegistry @pytest.fixture(autouse=True) @@ -249,3 +253,93 @@ async def test_config_owned_store_survives_db_liveness_check_while_missing_db_st prisma_client.db.litellm_managedvectorstorestable.find_unique.assert_awaited_once_with( where={"vector_store_id": "vs_from_db"} ) + + +_INDEX_PARAMS: Final = {"vector_store_index": "real-index-name", "vector_store_name": "azure-ai-search-store"} +_INDEX_CREATED_AT: Final = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) +_INDEX_ROW: Final = { + "id": "idx-1", + "index_name": "team-docs", + "litellm_params": _INDEX_PARAMS, + "index_info": {"dimensions": 1536}, + "created_at": _INDEX_CREATED_AT, + "created_by": "user-1", + "updated_at": _INDEX_CREATED_AT, + "updated_by": "user-2", +} + + +class GeneratedIndexRow(BaseModel): + id: str + index_name: str + litellm_params: dict[str, str] + index_info: dict[str, int] | None = None + created_at: datetime + created_by: str | None = None + updated_at: datetime + updated_by: str | None = None + + +def _database_listing_index_rows(rows: Sequence[object]) -> SimpleNamespace: + async def find_many(order: Mapping[str, str]) -> Sequence[object]: + return rows if order == {"created_at": "desc"} else [] + + return SimpleNamespace( + db=SimpleNamespace(litellm_managedvectorstoreindextable=SimpleNamespace(find_many=find_many)) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "row", + [ + _INDEX_ROW, + GeneratedIndexRow(**_INDEX_ROW), + list(_INDEX_ROW.items()), + {**_INDEX_ROW, "column_added_by_a_later_migration": "ignored"}, + ], +) +async def test_vector_store_index_rows_from_the_db_are_returned_as_indexes(row: object) -> None: + indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + _database_listing_index_rows([row]) + ) + + assert [index.model_dump() for index in indexes] == [_INDEX_ROW] + + +@pytest.mark.asyncio +async def test_vector_store_index_row_without_optional_columns_gets_empty_defaults() -> None: + indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + _database_listing_index_rows([{"id": "idx-1", "index_name": "team-docs", "litellm_params": _INDEX_PARAMS}]) + ) + + assert [index.model_dump() for index in indexes] == [ + { + "id": "idx-1", + "index_name": "team-docs", + "litellm_params": _INDEX_PARAMS, + "index_info": None, + "created_at": None, + "created_by": None, + "updated_at": None, + "updated_by": None, + } + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "row", + [ + {"index_name": "team-docs", "litellm_params": _INDEX_PARAMS}, + {**_INDEX_ROW, "id": 7}, + {**_INDEX_ROW, "litellm_params": None}, + {**_INDEX_ROW, "litellm_params": {"vector_store_index": "real-index-name"}}, + {**_INDEX_ROW, "index_info": ["not", "a", "mapping"]}, + {**_INDEX_ROW, 7: "column names must be strings"}, + {**_INDEX_ROW, b"index_name": "column names are not decoded"}, + ], +) +async def test_malformed_vector_store_index_row_from_the_db_raises_a_validation_error(row: object) -> None: + with pytest.raises(ValidationError): + await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(_database_listing_index_rows([row]))