From d80f8c28ca7e2fba4257b4b97457d3b313ff0d6a Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 5 Oct 2026 18:01:45 +0000 Subject: [PATCH] refactor(types): replace Any with proven types in 137 files (#44478) * refactor(types): prove runtime types at harness, search, rag and client boundaries Replace Any with adapter-validated types in the litellm.agent() harness, the search provider transformations, RAG ingestion and query, the vector store pre-call hook and registry, the galileo and opik logging integrations and the proxy client CLI. Each boundary gets unit tests for well-formed and malformed payloads. * chore(typing): prove types at more provider boundaries and restore search transformations Second pass over a2a, embedding, rerank, image, audio and small provider modules. Search transformations go back to their previous form because validating their response bodies would change the proxy status for malformed upstream bodies from 400 to 500. * refactor(types): prove types at logging, files, rerank, image, audio and management boundaries Replace Any with validated or annotation-only types in 37 more files: logging integrations, token counters, provider files/rerank/image generation/audio transcription transformations, pass-through logging handlers and management endpoints. No proxy HTTP status or error type changes. * refactor(types): prove types at repository, spend, files and router boundaries Replace Any with repository table accessors, validated mappings and annotation-only types in 33 more files: Prisma repositories, the enterprise batch and responses cost checkers, budget reservation, files endpoints, management endpoints, the policy registry, the adaptive and complexity routers and the secret managers. No proxy HTTP status or error type changes. * test(types): run the aiohttp transformation test in-process and cover repository row conversion The aiohttp chat transformation test no longer starts a server. It feeds the transformation a response whose json() returns the body under test. The proxy unit shards now exercise stored model rows whose params are JSON strings and the object permission create and update paths. --- .../send_emails/base_email.py | 6 +- .../proxy/common_utils/check_batch_cost.py | 18 +- .../common_utils/check_responses_cost.py | 4 +- .../proxy/hooks/managed_files.py | 4 +- litellm/a2a_protocol/providers/base.py | 9 +- .../providers/bedrock_agentcore/config.py | 10 +- .../providers/bedrock_agentcore/handler.py | 10 +- .../bedrock_agentcore/transformation.py | 2 +- .../a2a_protocol/providers/langflow/config.py | 9 +- .../providers/pydantic_ai_agents/config.py | 16 +- .../providers/pydantic_ai_agents/handler.py | 16 +- .../pydantic_ai_agents/transformation.py | 12 +- .../providers/watsonx_orchestrate/config.py | 16 +- .../providers/watsonx_orchestrate/handler.py | 4 +- litellm/caching/_embedding_router.py | 8 +- litellm/files/main.py | 2 +- litellm/files/streaming.py | 4 +- litellm/harness/context.py | 6 +- litellm/harness/handlers/cli_handler.py | 4 +- litellm/harness/options.py | 8 +- litellm/harness/runtime.py | 42 ++-- litellm/harness/sync.py | 44 ++-- litellm/harness/types.py | 6 +- litellm/integrations/arize/_utils.py | 15 +- .../integrations/datadog/datadog_llm_obs.py | 8 +- litellm/integrations/galileo.py | 15 +- litellm/integrations/opentelemetry.py | 8 +- .../opik/opik_payload_builder/api.py | 11 +- .../opik/opik_payload_builder/extractors.py | 6 +- .../opik/opik_payload_builder/types.py | 15 +- litellm/integrations/opik/utils.py | 9 +- litellm/integrations/otel/model/config.py | 4 +- litellm/integrations/otel/plumbing/metrics.py | 4 +- litellm/integrations/s3_v2.py | 4 +- .../vector_store_pre_call_hook.py | 20 +- litellm/llms/a2a/chat/transformation.py | 14 +- .../aiohttp_openai/chat/transformation.py | 10 +- .../anthropic/count_tokens/token_counter.py | 2 +- .../llms/anthropic/files/transformation.py | 15 +- .../audio_transcription/transformation.py | 18 +- .../anthropic/count_tokens/handler.py | 2 +- .../llms/base_llm/harness/transformation.py | 4 +- litellm/llms/bedrock/batches/handler.py | 16 +- .../image_edit/stability_transformation.py | 3 +- .../amazon_titan_transformation.py | 10 +- litellm/llms/bedrock/rerank/handler.py | 9 +- litellm/llms/bedrock/search/transformation.py | 7 +- .../image_generation/transformation.py | 8 +- litellm/llms/clarifai/chat/transformation.py | 6 +- .../claude_code/harness/transformation.py | 8 +- litellm/llms/codex/harness/transformation.py | 16 +- litellm/llms/cohere/rerank/transformation.py | 9 +- .../image_generation/transformation.py | 8 +- .../image_generation/transformation.py | 16 +- .../llms/dashscope/rerank/transformation.py | 23 +- .../llms/deepagents/harness/transformation.py | 6 +- .../audio_transcription/transformation.py | 9 +- litellm/llms/e2b/sandbox/transformation.py | 36 ++- .../audio_transcription/transformation.py | 14 +- .../flux_pro_v11_ultra_transformation.py | 19 +- .../ideogram_v3_transformation.py | 11 +- .../recraft_v3_transformation.py | 8 +- .../stable_diffusion_transformation.py | 21 +- .../fireworks_ai/rerank/transformation.py | 55 +++-- .../gemini/image_generation/transformation.py | 11 +- .../llms/gigachat/embedding/transformation.py | 6 +- .../embedding/transformation.py | 5 +- .../llms/jina_ai/embedding/transformation.py | 6 +- litellm/llms/jina_ai/rerank/transformation.py | 25 +- litellm/llms/manus/files/transformation.py | 40 +-- .../minimax/text_to_speech/transformation.py | 10 +- .../image_generation/transformation.py | 13 +- .../llms/nvidia_nim/rerank/transformation.py | 26 +- .../audio_transcription/transformation.py | 10 +- litellm/llms/oci/embed/transformation.py | 6 +- .../openai/responses/count_tokens/handler.py | 5 +- .../llms/opencode/harness/transformation.py | 8 +- .../image_generation/transformation.py | 20 +- .../audio_transcription/transformation.py | 11 +- .../image_generation/transformation.py | 8 +- litellm/llms/replicate/chat/transformation.py | 6 +- litellm/llms/sagemaker/common_utils.py | 7 +- .../snowflake/embedding/transformation.py | 6 +- .../image_generation/transformation.py | 9 +- .../llms/tinyfish/search/transformation.py | 5 +- litellm/llms/together_ai/rerank/handler.py | 14 +- .../image_generation_handler.py | 11 +- .../vertex_gemini_transformation.py | 12 +- .../text_to_speech/transformation.py | 10 +- litellm/llms/voyage/rerank/transformation.py | 19 +- .../audio_transcription/transformation.py | 13 +- litellm/llms/watsonx/embed/transformation.py | 9 +- litellm/llms/watsonx/rerank/transformation.py | 17 +- litellm/proxy/_experimental/mcp_server/db.py | 11 +- .../proxy/agent_endpoints/a2a_endpoints.py | 7 +- litellm/proxy/client/chat.py | 7 +- litellm/proxy/client/cli/commands/auth.py | 5 +- litellm/proxy/client/cli/commands/debug.py | 3 +- litellm/proxy/db/model_insights_tasks.py | 10 +- .../proxy/fine_tuning_endpoints/endpoints.py | 2 +- .../health_endpoints/_health_endpoints.py | 14 +- .../auto_router_endpoints.py | 8 +- .../callback_management_endpoints.py | 5 +- .../key_management_endpoints.py | 5 +- .../management_v1/budgets.py | 6 +- .../mcp_management_endpoints.py | 6 +- .../model_insights_endpoints.py | 16 +- .../model_management_endpoints.py | 10 +- .../organization_endpoints.py | 18 +- .../policy_endpoints/endpoints.py | 37 ++- .../router_settings_endpoints.py | 2 +- .../tag_management_endpoints.py | 7 +- .../management_endpoints/team_endpoints.py | 2 +- .../file_content_streaming_handler.py | 6 +- .../openai_files_endpoints/files_endpoints.py | 3 +- .../anthropic_passthrough_logging_handler.py | 7 +- .../vertex_passthrough_logging_handler.py | 7 +- .../proxy/policy_engine/policy_registry.py | 4 +- .../public_endpoints/public_endpoints.py | 7 +- .../spend_tracking/budget_reservation.py | 19 +- .../spend_tracking/spend_log_error_logger.py | 4 +- litellm/rag/ingestion/base_ingestion.py | 6 +- litellm/rag/ingestion/gemini_ingestion.py | 16 +- litellm/rag/ingestion/openai_ingestion.py | 6 +- litellm/rag/ingestion/s3_vectors_ingestion.py | 4 +- litellm/rag/main.py | 36 ++- .../autorouter_session_repository.py | 15 +- litellm/repositories/model_repository.py | 28 ++- .../object_permission_repository.py | 22 +- litellm/repositories/project_repository.py | 15 +- litellm/repositories/user_repository.py | 18 +- .../adaptive_router/adaptive_router.py | 3 + .../complexity_router/complexity_router.py | 8 +- .../reasoning_effort_capability.py | 2 +- litellm/secret_managers/aws_secret_manager.py | 6 +- litellm/secret_managers/main.py | 6 +- .../vector_stores/vector_store_registry.py | 8 +- .../test_pydantic_ai_agent_headers.py | 71 +++++- .../watsonx_orchestrate/test_config.py | 161 ++++++++++++ .../send_emails/test_base_email.py | 75 +++++- tests/unit/files/test_main.py | 52 ++++ .../integrations/arize/test_arize_utils.py | 41 ++++ .../datadog/test_datadog_llm_obs.py | 28 +++ .../opik/opik_payload_builder/__init__.py | 0 .../opik/opik_payload_builder/test_api.py | 230 ++++++++++++++++++ .../integrations/opik/test_opik_extractors.py | 32 +++ tests/unit/integrations/test_galileo.py | 79 ++++++ tests/unit/integrations/test_opentelemetry.py | 34 +++ tests/unit/integrations/test_opik_utils.py | 21 +- .../test_vector_store_pre_call_hook.py | 163 ++++++++++++- .../a2a/chat/test_a2a_chat_transformation.py | 36 +++ tests/unit/llms/aiohttp_openai/__init__.py | 0 .../unit/llms/aiohttp_openai/chat/__init__.py | 0 .../chat/test_transformation.py | 101 ++++++++ .../test_anthropic_files_transformation.py | 107 ++++++++ .../test_azure_speech_audio_transcription.py | 20 ++ .../test_stability_transformation.py | 14 ++ .../llms/bedrock/image_generation/__init__.py | 0 .../test_amazon_titan_transformation.py | 57 +++++ .../test_bedrock_rerank_header_forwarding.py | 59 +++++ .../test_agentcore_search_transformation.py | 43 ++++ ...est_bfl_image_generation_transformation.py | 46 ++++ tests/unit/llms/clarifai/__init__.py | 0 tests/unit/llms/clarifai/chat/__init__.py | 0 .../llms/clarifai/chat/test_transformation.py | 70 ++++++ .../harness/test_transformation.py | 23 ++ .../llms/codex/harness/test_transformation.py | 32 ++- .../llms/cohere/rerank/test_transformation.py | 72 ++++++ .../cometapi/image_generation/__init__.py | 0 .../image_generation/test_transformation.py | 54 ++++ .../test_dashscope_rerank_transformation.py | 87 +++++++ ...gram_audio_transcription_transformation.py | 80 ++++++ tests/unit/llms/e2b/__init__.py | 0 tests/unit/llms/e2b/sandbox/__init__.py | 0 .../llms/e2b/sandbox/test_transformation.py | 121 +++++++++ .../audio_transcription/__init__.py | 0 .../test_transformation.py | 100 ++++++++ .../test_flux_pro_v11_ultra_transformation.py | 56 +++++ .../test_ideogram_v3_transformation.py | 71 ++++++ .../test_recraft_v3_transformation.py | 64 +++++ .../test_stable_diffusion_transformation.py | 76 ++++++ ...test_fireworks_ai_rerank_transformation.py | 66 +++++ ..._gemini_image_generation_transformation.py | 48 ++++ .../test_gigachat_embedding_transformation.py | 63 ++++- ...github_copilot_embedding_transformation.py | 57 +++++ .../test_jina_embedding_transformation.py | 61 +++++ tests/unit/llms/jina_ai/rerank/__init__.py | 0 .../jina_ai/rerank/test_transformation.py | 111 +++++++++ tests/unit/llms/manus/files/__init__.py | 0 .../llms/manus/files/test_transformation.py | 128 ++++++++++ .../llms/minimax/text_to_speech/__init__.py | 0 .../text_to_speech/test_transformation.py | 122 ++++++++++ ...est_modelscope_image_gen_transformation.py | 83 +++++++ .../test_nvidia_nim_rerank_transformation.py | 86 +++++++ .../embed/test_oci_embed_transformation.py | 75 ++++++ .../openai/responses/count_tokens/__init__.py | 0 .../responses/count_tokens/test_handler.py | 26 ++ ...est_openrouter_image_gen_transformation.py | 83 +++++++ ...loud_audio_transcription_transformation.py | 42 ++++ .../test_recraft_image_gen_transformation.py | 54 ++++ tests/unit/llms/replicate/__init__.py | 0 tests/unit/llms/replicate/chat/__init__.py | 0 .../replicate/chat/test_transformation.py | 54 ++++ .../sagemaker/test_sagemaker_common_utils.py | 13 + .../embedding/test_snowflake_embedding.py | 64 ++++- .../test_stability_image_generation.py | 33 +++ .../llms/tinyfish/test_tinyfish_search.py | 54 ++++ .../unit/llms/together_ai/rerank/__init__.py | 0 .../llms/together_ai/rerank/test_handler.py | 96 ++++++++ .../test_image_generation_handler.py | 43 ++++ ...rtex_ai_image_generation_transformation.py | 64 +++++ .../text_to_speech/test_transformation.py | 30 +++ .../test_voyage_rerank_transformation.py | 79 ++++++ ...sonx_audio_transcription_transformation.py | 58 +++++ .../test_watsonx_embedding_transformation.py | 90 +++++++ .../watsonx/rerank/test_watsonx_rerank.py | 62 +++++ .../mcp_server/test_db_credentials.py | 43 ++++ .../mcp_server/test_mcp_env_vars.py | 28 +++ .../proxy/client/cli/test_auth_commands.py | 31 +++ tests/unit/proxy/client/test_chat.py | 51 ++++ .../db/test_object_permission_repository.py | 73 ++++++ .../fine_tuning_endpoints/test_endpoints.py | 68 ++++++ .../policy_endpoints/test_endpoints.py | 75 ++++++ .../test_model_management_endpoints.py | 47 ++++ .../test_organization_endpoints.py | 28 +++ .../test_tag_management_endpoints.py | 65 +++++ ...t_anthropic_passthrough_logging_handler.py | 19 ++ ...test_vertex_passthrough_logging_handler.py | 95 ++++++++ .../policy_engine/test_policy_registry.py | 154 ++++++++++++ .../spend_tracking/test_budget_reservation.py | 84 +++++++ .../unit/rag/ingestion/test_base_ingestion.py | 101 ++++++++ .../rag/ingestion/test_gemini_ingestion.py | 111 +++++++++ .../rag/ingestion/test_openai_ingestion.py | 101 ++++++++ tests/unit/rag/test_main.py | 87 +++++++ tests/unit/repositories/test_repositories.py | 75 ++++++ .../adaptive_router/test_adaptive_router.py | 9 + .../test_aws_secret_manager.py | 36 +++ .../test_secret_managers_main.py | 35 +++ tests/unit/test_dashscope_image_generation.py | 96 ++++++++ .../test_vector_store_registry.py | 96 +++++++- 240 files changed, 6820 insertions(+), 556 deletions(-) create mode 100644 tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py create mode 100644 tests/unit/integrations/opik/opik_payload_builder/__init__.py create mode 100644 tests/unit/integrations/opik/opik_payload_builder/test_api.py create mode 100644 tests/unit/llms/aiohttp_openai/__init__.py create mode 100644 tests/unit/llms/aiohttp_openai/chat/__init__.py create mode 100644 tests/unit/llms/aiohttp_openai/chat/test_transformation.py create mode 100644 tests/unit/llms/bedrock/image_edit/test_stability_transformation.py create mode 100644 tests/unit/llms/bedrock/image_generation/__init__.py create mode 100644 tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py create mode 100644 tests/unit/llms/clarifai/__init__.py create mode 100644 tests/unit/llms/clarifai/chat/__init__.py create mode 100644 tests/unit/llms/clarifai/chat/test_transformation.py create mode 100644 tests/unit/llms/cohere/rerank/test_transformation.py create mode 100644 tests/unit/llms/cometapi/image_generation/__init__.py create mode 100644 tests/unit/llms/cometapi/image_generation/test_transformation.py create mode 100644 tests/unit/llms/e2b/__init__.py create mode 100644 tests/unit/llms/e2b/sandbox/__init__.py create mode 100644 tests/unit/llms/e2b/sandbox/test_transformation.py create mode 100644 tests/unit/llms/elevenlabs/audio_transcription/__init__.py create mode 100644 tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py create mode 100644 tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py create mode 100644 tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py create mode 100644 tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py create mode 100644 tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py create mode 100644 tests/unit/llms/jina_ai/rerank/__init__.py create mode 100644 tests/unit/llms/jina_ai/rerank/test_transformation.py create mode 100644 tests/unit/llms/manus/files/__init__.py create mode 100644 tests/unit/llms/manus/files/test_transformation.py create mode 100644 tests/unit/llms/minimax/text_to_speech/__init__.py create mode 100644 tests/unit/llms/minimax/text_to_speech/test_transformation.py create mode 100644 tests/unit/llms/openai/responses/count_tokens/__init__.py create mode 100644 tests/unit/llms/openai/responses/count_tokens/test_handler.py create mode 100644 tests/unit/llms/replicate/__init__.py create mode 100644 tests/unit/llms/replicate/chat/__init__.py create mode 100644 tests/unit/llms/replicate/chat/test_transformation.py create mode 100644 tests/unit/llms/together_ai/rerank/__init__.py create mode 100644 tests/unit/llms/together_ai/rerank/test_handler.py create mode 100644 tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py create mode 100644 tests/unit/proxy/db/test_object_permission_repository.py create mode 100644 tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py create mode 100644 tests/unit/proxy/policy_engine/test_policy_registry.py create mode 100644 tests/unit/rag/ingestion/test_base_ingestion.py create mode 100644 tests/unit/rag/ingestion/test_gemini_ingestion.py create mode 100644 tests/unit/rag/ingestion/test_openai_ingestion.py create mode 100644 tests/unit/secret_managers/test_aws_secret_manager.py 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]))