From b604e2b20c6db2099085a2f0e59b7e99e87eed6f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 1 Aug 2026 15:43:29 -0700 Subject: [PATCH] refactor(lint): apply every safe ruff autofix and zero 28 strict-rule budgets About 35,000 fixes ruff marks safe across 32 rules (UP006/UP045/UP007 modern annotations, UP032 f-strings, SIM114/SIM118, RET501, and friends), removal of the 1,296 typing imports the rewrite orphaned, and hand fixes for what the fixers could not see: five star-import freeloaders of typing names, two F823 late-import annotations, the /get/config/list introspection crash on types.UnionType, redundant function-local RoleMappings imports in ui_sso.py that shadowed the module-level name once the annotation lost its quotes, and one FURB168 tautology. B009/B010/PIE804/RUF019 are excluded on purpose: their safe fixes rewrite getattr/setattr/**-splat/key-in-dict escape hatches into forms basedpyright then rejects (283 new errors measured), so their budgets stay at base values. ruff-strict-budget.json drops by 39,579 this commit (39,968 across the branch) with 28 rules at an actual 0 and 9 more sharply down. type-discipline-budget.json ratchets LIT002/LIT006/LIT009 down; LIT001 moves to the now-honest total: the checker matches the spelling `set` but not the alias `Set`, so the 160 typing.Set annotations rewritten to set[...] were always mutable-set annotations and only now count. --- litellm/_lazy_imports.py | 10 +- litellm/_logging.py | 12 +- litellm/_redis.py | 56 +- litellm/_redis_credential_provider.py | 14 +- litellm/_service_logger.py | 38 +- litellm/a2a_protocol/card_resolver.py | 4 +- litellm/a2a_protocol/client.py | 6 +- litellm/a2a_protocol/cost_calculator.py | 8 +- .../a2a_protocol/exception_mapping_utils.py | 10 +- litellm/a2a_protocol/exceptions.py | 36 +- .../litellm_completion_bridge/__init__.py | 2 +- .../litellm_completion_bridge/handler.py | 46 +- .../transformation.py | 46 +- litellm/a2a_protocol/main.py | 38 +- litellm/a2a_protocol/providers/__init__.py | 2 +- litellm/a2a_protocol/providers/base.py | 15 +- .../providers/bedrock_agentcore/config.py | 14 +- .../providers/bedrock_agentcore/handler.py | 18 +- .../bedrock_agentcore/transformation.py | 20 +- .../a2a_protocol/providers/config_manager.py | 8 +- .../a2a_protocol/providers/langflow/config.py | 14 +- .../providers/pydantic_ai_agents/__init__.py | 2 +- .../providers/pydantic_ai_agents/config.py | 14 +- .../providers/pydantic_ai_agents/handler.py | 18 +- .../pydantic_ai_agents/transformation.py | 30 +- .../providers/watsonx_orchestrate/config.py | 14 +- .../providers/watsonx_orchestrate/handler.py | 44 +- .../watsonx_orchestrate/transformation.py | 16 +- litellm/a2a_protocol/streaming_iterator.py | 12 +- litellm/a2a_protocol/utils.py | 16 +- litellm/anthropic_beta_headers_manager.py | 19 +- .../exceptions/__init__.py | 4 +- .../exceptions/exception_mapping_utils.py | 9 +- .../anthropic_interface/messages/__init__.py | 64 +- litellm/assistants/main.py | 214 +-- litellm/assistants/utils.py | 43 +- litellm/batch_completion/main.py | 35 +- litellm/batches/batch_utils.py | 42 +- litellm/batches/main.py | 104 +- litellm/budget_manager.py | 16 +- litellm/caching/__init__.py | 2 +- litellm/caching/_internal_lru_cache.py | 4 +- litellm/caching/azure_blob_cache.py | 2 +- litellm/caching/base_cache.py | 6 +- litellm/caching/caching.py | 159 +- litellm/caching/caching_handler.py | 116 +- litellm/caching/disk_cache.py | 4 +- litellm/caching/dual_cache.py | 64 +- litellm/caching/gcs_cache.py | 12 +- litellm/caching/in_memory_cache.py | 29 +- litellm/caching/qdrant_semantic_cache.py | 14 +- litellm/caching/redis_cache.py | 136 +- litellm/caching/redis_cluster_cache.py | 16 +- litellm/caching/redis_semantic_cache.py | 54 +- litellm/caching/s3_cache.py | 5 +- litellm/caching/valkey_semantic_cache.py | 10 +- .../handler.py | 4 +- .../transformation.py | 149 +- litellm/compression/compress.py | 92 +- litellm/compression/message_stubbing.py | 3 +- litellm/compression/retrieval_tool.py | 4 +- litellm/compression/scoring/bm25.py | 15 +- .../compression/scoring/embedding_scorer.py | 16 +- litellm/constants.py | 18 +- litellm/containers/endpoint_factory.py | 36 +- litellm/containers/main.py | 302 ++- litellm/containers/utils.py | 12 +- litellm/cost_calculator.py | 381 ++-- .../speech_to_completion_bridge/handler.py | 8 +- .../transformation.py | 4 +- litellm/evals/__init__.py | 20 +- litellm/evals/main.py | 328 ++-- litellm/exceptions.py | 244 ++- litellm/experimental_mcp_client/__init__.py | 2 +- litellm/experimental_mcp_client/client.py | 117 +- litellm/experimental_mcp_client/tools.py | 6 +- litellm/files/main.py | 129 +- litellm/files/streaming.py | 20 +- litellm/files/types.py | 6 +- litellm/files/utils.py | 6 +- litellm/fine_tuning/main.py | 94 +- litellm/google_genai/__init__.py | 4 +- litellm/google_genai/adapters/__init__.py | 2 +- litellm/google_genai/adapters/handler.py | 34 +- .../google_genai/adapters/transformation.py | 69 +- litellm/google_genai/main.py | 80 +- litellm/google_genai/streaming_iterator.py | 22 +- litellm/images/main.py | 164 +- litellm/images/utils.py | 10 +- .../SlackAlerting/batching_handler.py | 2 +- .../SlackAlerting/budget_alert_types.py | 2 - .../SlackAlerting/hanging_request_check.py | 10 +- .../SlackAlerting/slack_alerting.py | 159 +- litellm/integrations/SlackAlerting/utils.py | 14 +- .../integrations/additional_logging_utils.py | 8 +- litellm/integrations/agentops/agentops.py | 13 +- .../anthropic_cache_control_hook.py | 112 +- litellm/integrations/argilla.py | 34 +- litellm/integrations/arize/__init__.py | 6 +- litellm/integrations/arize/_utils.py | 30 +- litellm/integrations/arize/arize.py | 31 +- litellm/integrations/arize/arize_phoenix.py | 20 +- .../arize/arize_phoenix_client.py | 6 +- .../arize/arize_phoenix_prompt_manager.py | 96 +- litellm/integrations/athina.py | 1 - .../azure_sentinel/azure_sentinel.py | 36 +- .../azure_storage/azure_storage.py | 45 +- litellm/integrations/bitbucket/__init__.py | 11 +- .../bitbucket/bitbucket_client.py | 14 +- .../bitbucket/bitbucket_prompt_manager.py | 128 +- litellm/integrations/braintrust_logging.py | 15 +- litellm/integrations/cloudzero/cloudzero.py | 24 +- .../cloudzero/cz_resource_names.py | 1 - .../integrations/cloudzero/cz_stream_api.py | 10 +- litellm/integrations/cloudzero/database.py | 12 +- litellm/integrations/cloudzero/transform.py | 4 +- .../code_interpreter_interception/handler.py | 6 +- .../compression_interception/handler.py | 72 +- litellm/integrations/custom_batch_logger.py | 11 +- litellm/integrations/custom_guardrail.py | 146 +- litellm/integrations/custom_logger.py | 204 +- .../integrations/custom_prompt_management.py | 48 +- litellm/integrations/custom_secret_manager.py | 36 +- litellm/integrations/datadog/datadog.py | 95 +- .../datadog/datadog_cost_management.py | 28 +- .../integrations/datadog/datadog_handler.py | 9 +- .../integrations/datadog/datadog_llm_obs.py | 80 +- .../integrations/datadog/datadog_metrics.py | 23 +- .../datadog/datadog_team_handler.py | 12 +- litellm/integrations/deepeval/api.py | 4 +- litellm/integrations/deepeval/deepeval.py | 3 +- litellm/integrations/deepeval/types.py | 41 +- litellm/integrations/deepeval/utils.py | 1 + litellm/integrations/dotprompt/__init__.py | 11 +- .../dotprompt/dotprompt_manager.py | 84 +- .../integrations/dotprompt/prompt_manager.py | 48 +- litellm/integrations/dynamodb.py | 5 +- litellm/integrations/email_alerting.py | 5 +- litellm/integrations/focus/database.py | 10 +- .../focus/destinations/__init__.py | 6 +- .../focus/destinations/factory.py | 10 +- .../focus/destinations/gcs_destination.py | 4 +- .../focus/destinations/mavvrik_destination.py | 8 +- .../focus/destinations/s3_destination.py | 4 +- .../focus/destinations/vantage_destination.py | 8 +- litellm/integrations/focus/export_engine.py | 10 +- litellm/integrations/focus/focus_logger.py | 39 +- .../focus/serializers/__init__.py | 2 +- litellm/integrations/galileo.py | 82 +- litellm/integrations/gcs_bucket/gcs_bucket.py | 44 +- .../gcs_bucket/gcs_bucket_base.py | 52 +- .../gcs_bucket/gcs_bucket_mock_client.py | 2 +- litellm/integrations/gcs_pubsub/pub_sub.py | 21 +- .../generic_api/generic_api_callback.py | 42 +- .../generic_prompt_management/__init__.py | 11 +- .../generic_prompt_manager.py | 106 +- litellm/integrations/gitlab/__init__.py | 23 +- litellm/integrations/gitlab/gitlab_client.py | 26 +- .../gitlab/gitlab_prompt_manager.py | 183 +- litellm/integrations/greenscale.py | 1 - litellm/integrations/helicone.py | 6 +- litellm/integrations/humanloop.py | 52 +- litellm/integrations/lago.py | 24 +- litellm/integrations/langfuse/langfuse.py | 108 +- .../integrations/langfuse/langfuse_handler.py | 14 +- .../langfuse/langfuse_mock_client.py | 1 + .../integrations/langfuse/langfuse_otel.py | 17 +- .../langfuse/langfuse_otel_attributes.py | 30 +- .../langfuse/langfuse_prompt_management.py | 68 +- litellm/integrations/langsmith.py | 62 +- litellm/integrations/levo/levo.py | 4 +- .../litellm_agent_model_resolver.py | 40 +- litellm/integrations/literal_ai.py | 7 +- litellm/integrations/logfire_logger.py | 11 +- litellm/integrations/lunary.py | 1 - .../mavvrik_focus/mavvrik_focus_logger.py | 14 +- litellm/integrations/mlflow.py | 12 +- litellm/integrations/mock_client_factory.py | 24 +- litellm/integrations/newrelic/newrelic.py | 76 +- litellm/integrations/openmeter.py | 6 +- litellm/integrations/opentelemetry.py | 197 +- .../base_otel_llm_obs_attributes.py | 6 +- .../opentelemetry_utils/gen_ai_semconv.py | 14 +- litellm/integrations/opik/opik.py | 38 +- .../opik/opik_payload_builder/api.py | 10 +- .../opik/opik_payload_builder/extractors.py | 36 +- .../opik_payload_builder/payload_builders.py | 24 +- .../opik/opik_payload_builder/types.py | 22 +- litellm/integrations/opik/utils.py | 16 +- litellm/integrations/otel/__init__.py | 10 +- litellm/integrations/otel/emitter.py | 2 +- litellm/integrations/otel/logger.py | 6 +- litellm/integrations/otel/mappers/__init__.py | 2 +- litellm/integrations/otel/model/config.py | 10 +- litellm/integrations/otel/model/metadata.py | 14 +- litellm/integrations/otel/model/payloads.py | 24 +- litellm/integrations/otel/model/utils.py | 4 +- litellm/integrations/otel/plumbing/metrics.py | 8 +- litellm/integrations/otel/plumbing/routing.py | 2 +- litellm/integrations/otel/presets/__init__.py | 4 +- litellm/integrations/otel/runtime.py | 4 +- litellm/integrations/posthog.py | 49 +- litellm/integrations/prometheus.py | 342 ++-- .../prometheus_helpers/__init__.py | 18 +- .../bounded_prometheus_series_tracker.py | 22 +- .../prometheus_helpers/prometheus_api.py | 7 +- litellm/integrations/prometheus_services.py | 28 +- litellm/integrations/prompt_layer.py | 1 - .../integrations/prompt_management_base.py | 110 +- litellm/integrations/rubrik.py | 11 +- litellm/integrations/s3.py | 15 +- litellm/integrations/s3_v2.py | 99 +- litellm/integrations/sqs.py | 80 +- litellm/integrations/supabase.py | 2 - .../integrations/vantage/vantage_logger.py | 18 +- .../vector_store_pre_call_hook.py | 58 +- litellm/integrations/weave/weave_otel.py | 9 +- .../websearch_interception/handler.py | 174 +- .../websearch_interception/tools.py | 12 +- .../websearch_interception/transformation.py | 38 +- litellm/integrations/weights_biases.py | 23 +- litellm/interactions/agents/__init__.py | 14 +- litellm/interactions/agents/http_handler.py | 76 +- litellm/interactions/agents/main.py | 94 +- litellm/interactions/agents/utils.py | 13 +- litellm/interactions/http_handler.py | 101 +- .../__init__.py | 2 +- .../handler.py | 32 +- .../streaming_iterator.py | 28 +- .../transformation.py | 24 +- litellm/interactions/main.py | 140 +- litellm/interactions/streaming_iterator.py | 27 +- litellm/interactions/utils.py | 8 +- .../api_route_to_call_types.py | 4 +- litellm/litellm_core_utils/app_crypto.py | 5 +- litellm/litellm_core_utils/asyncify.py | 3 +- .../litellm_core_utils/audio_utils/utils.py | 13 +- litellm/litellm_core_utils/cached_imports.py | 8 +- litellm/litellm_core_utils/cli_token_utils.py | 9 +- .../cloud_storage_security.py | 14 +- .../litellm_core_utils/completion_timeout.py | 11 +- litellm/litellm_core_utils/core_helpers.py | 38 +- .../litellm_core_utils/coroutine_checker.py | 1 + .../litellm_core_utils/credential_accessor.py | 4 +- .../custom_logger_registry.py | 8 +- litellm/litellm_core_utils/dd_tracing.py | 6 +- .../litellm_core_utils/default_encoding.py | 5 +- .../dot_notation_indexing.py | 10 +- litellm/litellm_core_utils/duration_parser.py | 11 +- .../exception_mapping_utils.py | 55 +- .../fallback_generalizations.py | 14 +- litellm/litellm_core_utils/fallback_utils.py | 10 +- litellm/litellm_core_utils/get_blog_posts.py | 17 +- .../litellm_core_utils/get_litellm_params.py | 50 +- .../get_llm_provider_logic.py | 66 +- .../litellm_core_utils/get_model_cost_map.py | 13 +- .../get_provider_specific_headers.py | 8 +- .../get_supported_openai_params.py | 12 +- .../health_check_helpers.py | 8 +- .../initialize_dynamic_callback_params.py | 4 +- .../json_validation_rule.py | 10 +- litellm/litellm_core_utils/litellm_logging.py | 547 +++--- .../llm_cost_calc/tiered_pricing.py | 14 +- .../llm_cost_calc/tool_call_cost_tracking.py | 110 +- .../usage_object_transformation.py | 6 +- .../litellm_core_utils/llm_cost_calc/utils.py | 136 +- .../litellm_core_utils/llm_request_utils.py | 6 +- .../convert_dict_to_response.py | 97 +- .../llm_response_utils/get_api_base.py | 24 +- .../get_formatted_prompt.py | 4 +- .../llm_response_utils/get_headers.py | 5 +- .../llm_response_utils/response_metadata.py | 10 +- .../logging_callback_manager.py | 55 +- litellm/litellm_core_utils/logging_utils.py | 30 +- litellm/litellm_core_utils/logging_worker.py | 11 +- litellm/litellm_core_utils/mock_functions.py | 4 +- .../litellm_core_utils/model_param_helper.py | 35 +- .../prompt_templates/common_utils.py | 148 +- .../prompt_templates/factory.py | 500 +++-- .../huggingface_template_handler.py | 12 +- .../prompt_templates/image_handling.py | 1 - .../litellm_core_utils/realtime_streaming.py | 106 +- litellm/litellm_core_utils/redact_messages.py | 8 +- .../request_timeout_resolver.py | 4 +- litellm/litellm_core_utils/rules.py | 4 +- litellm/litellm_core_utils/safe_json_dumps.py | 4 +- litellm/litellm_core_utils/safe_json_loads.py | 2 +- .../litellm_core_utils/secret_redaction.py | 3 +- .../sensitive_data_masker.py | 28 +- .../specialty_caches/dynamic_logging_cache.py | 5 +- .../streaming_chunk_builder_utils.py | 136 +- .../litellm_core_utils/streaming_handler.py | 96 +- litellm/litellm_core_utils/token_counter.py | 73 +- litellm/litellm_core_utils/url_utils.py | 20 +- litellm/llms/__init__.py | 14 +- .../a2a/chat/guardrail_translation/handler.py | 56 +- litellm/llms/a2a/chat/streaming_iterator.py | 8 +- litellm/llms/a2a/chat/transformation.py | 46 +- litellm/llms/a2a/common_utils.py | 12 +- litellm/llms/ai21/chat/transformation.py | 44 +- litellm/llms/aiml/chat/transformation.py | 10 +- .../aiml/image_generation/transformation.py | 29 +- .../aiohttp_openai/chat/transformation.py | 20 +- .../llms/amazon_nova/chat/transformation.py | 42 +- litellm/llms/amazon_nova/cost_calculation.py | 4 +- litellm/llms/anthropic/__init__.py | 4 +- litellm/llms/anthropic/batches/__init__.py | 2 +- litellm/llms/anthropic/batches/handler.py | 24 +- .../llms/anthropic/batches/transformation.py | 52 +- .../chat/guardrail_translation/handler.py | 130 +- litellm/llms/anthropic/chat/handler.py | 61 +- litellm/llms/anthropic/chat/transformation.py | 315 ++-- litellm/llms/anthropic/common_utils.py | 103 +- .../anthropic/completion/transformation.py | 53 +- litellm/llms/anthropic/cost_calculation.py | 4 +- .../llms/anthropic/count_tokens/__init__.py | 2 +- .../llms/anthropic/count_tokens/handler.py | 20 +- .../anthropic/count_tokens/token_counter.py | 16 +- .../anthropic/count_tokens/transformation.py | 18 +- .../adapters/handler.py | 168 +- .../adapters/streaming_iterator.py | 43 +- .../adapters/transformation.py | 175 +- .../context_management/__init__.py | 4 +- .../context_management/dispatcher.py | 22 +- .../editors/clear_tool_uses.py | 38 +- .../context_management/editors/compact.py | 116 +- .../context_management/placeholders.py | 4 +- .../context_management/result.py | 16 +- .../messages/agentic_streaming_iterator.py | 38 +- .../messages/fake_stream_iterator.py | 12 +- .../messages/handler.py | 108 +- .../messages/interceptors/__init__.py | 6 +- .../messages/interceptors/advisor.py | 56 +- .../messages/interceptors/base.py | 15 +- .../messages/mcp_handler.py | 8 +- .../messages/streaming_iterator.py | 8 +- .../messages/transformation.py | 32 +- .../messages/utils.py | 8 +- .../responses_adapters/handler.py | 104 +- .../responses_adapters/streaming_iterator.py | 12 +- .../responses_adapters/transformation.py | 66 +- .../experimental_pass_through/utils.py | 5 +- litellm/llms/anthropic/files/__init__.py | 2 +- litellm/llms/anthropic/files/handler.py | 20 +- .../llms/anthropic/files/transformation.py | 26 +- .../llms/anthropic/skills/transformation.py | 18 +- litellm/llms/apiserpent/search/defaults.py | 16 +- .../llms/apiserpent/search/transformation.py | 26 +- .../text_to_speech/transformation.py | 50 +- litellm/llms/azure/assistants.py | 576 +++--- .../audio_transcription/transformation.py | 26 +- litellm/llms/azure/audio_transcriptions.py | 22 +- litellm/llms/azure/azure.py | 168 +- .../llms/azure/chat/gpt_5_transformation.py | 6 +- litellm/llms/azure/chat/gpt_transformation.py | 44 +- litellm/llms/azure/chat/o_series_handler.py | 28 +- .../azure/chat/o_series_transformation.py | 10 +- litellm/llms/azure/common_utils.py | 102 +- litellm/llms/azure/completion/handler.py | 22 +- .../llms/azure/completion/transformation.py | 18 +- .../llms/azure/containers/transformation.py | 9 +- litellm/llms/azure/cost_calculation.py | 8 +- litellm/llms/azure/exception_mapping.py | 8 +- litellm/llms/azure/files/handler.py | 176 +- litellm/llms/azure/fine_tuning/handler.py | 89 +- .../llms/azure/image_edit/transformation.py | 12 +- .../dall_e_2_transformation.py | 2 - .../dall_e_3_transformation.py | 2 - .../image_generation/gpt_transformation.py | 2 - .../llms/azure/passthrough/transformation.py | 28 +- litellm/llms/azure/realtime/handler.py | 29 +- .../azure/realtime/http_transformation.py | 16 +- .../responses/o_series_transformation.py | 4 +- .../llms/azure/responses/transformation.py | 46 +- .../azure/text_to_speech/transformation.py | 52 +- .../azure/vector_stores/transformation.py | 6 +- litellm/llms/azure/videos/transformation.py | 16 +- litellm/llms/azure_ai/agents/handler.py | 62 +- .../llms/azure_ai/agents/transformation.py | 54 +- .../anthropic/count_tokens/__init__.py | 2 +- .../anthropic/count_tokens/handler.py | 20 +- .../anthropic/count_tokens/token_counter.py | 16 +- litellm/llms/azure_ai/anthropic/handler.py | 6 +- .../anthropic/messages_transformation.py | 34 +- .../llms/azure_ai/anthropic/transformation.py | 18 +- .../azure_model_router/transformation.py | 12 +- litellm/llms/azure_ai/chat/transformation.py | 46 +- litellm/llms/azure_ai/common_utils.py | 20 +- litellm/llms/azure_ai/cost_calculator.py | 10 +- .../azure_ai/embed/cohere_transformation.py | 20 +- litellm/llms/azure_ai/embed/handler.py | 52 +- litellm/llms/azure_ai/image_edit/__init__.py | 2 +- .../image_edit/flux2_transformation.py | 24 +- .../azure_ai/image_edit/mai_transformation.py | 24 +- .../azure_ai/image_edit/transformation.py | 10 +- .../azure_ai/image_generation/__init__.py | 4 +- .../dall_e_2_transformation.py | 2 - .../dall_e_3_transformation.py | 2 - .../image_generation/flux_transformation.py | 6 +- .../image_generation/gpt_transformation.py | 2 - .../image_generation/mai_transformation.py | 18 +- .../document_intelligence/transformation.py | 17 +- litellm/llms/azure_ai/ocr/transformation.py | 6 +- .../llms/azure_ai/rerank/transformation.py | 14 +- .../azure_ai/vector_stores/transformation.py | 18 +- litellm/llms/base.py | 12 +- litellm/llms/base_llm/__init__.py | 8 +- .../llms/base_llm/agents/transformation.py | 40 +- .../anthropic_messages/transformation.py | 34 +- .../audio_transcription/transformation.py | 26 +- litellm/llms/base_llm/base_model_iterator.py | 20 +- litellm/llms/base_llm/base_utils.py | 57 +- .../llms/base_llm/batches/transformation.py | 36 +- .../bridges/completion_transformation.py | 14 +- litellm/llms/base_llm/chat/transformation.py | 88 +- .../base_llm/completion/transformation.py | 18 +- .../llms/base_llm/embedding/transformation.py | 18 +- litellm/llms/base_llm/evals/transformation.py | 52 +- .../files/azure_blob_storage_backend.py | 13 +- .../llms/base_llm/files/storage_backend.py | 6 +- litellm/llms/base_llm/files/transformation.py | 53 +- .../base_llm/google_genai/transformation.py | 30 +- .../guardrail_translation/base_translation.py | 12 +- .../base_llm/guardrail_translation/utils.py | 14 +- .../base_llm/image_edit/transformation.py | 20 +- .../image_generation/transformation.py | 24 +- .../image_variations/transformation.py | 34 +- .../base_llm/interactions/transformation.py | 55 +- .../base_llm/managed_resources/__init__.py | 14 +- .../base_managed_resource.py | 77 +- .../base_llm/managed_resources/isolation.py | 12 +- .../llms/base_llm/managed_resources/utils.py | 24 +- litellm/llms/base_llm/ocr/__init__.py | 4 +- litellm/llms/base_llm/ocr/transformation.py | 18 +- .../base_llm/passthrough/transformation.py | 24 +- .../base_llm/realtime/http_transformation.py | 20 +- .../llms/base_llm/realtime/transformation.py | 23 +- .../llms/base_llm/rerank/transformation.py | 16 +- .../llms/base_llm/responses/transformation.py | 67 +- .../llms/base_llm/sandbox/transformation.py | 7 +- .../llms/base_llm/search/transformation.py | 24 +- .../llms/base_llm/skills/transformation.py | 24 +- .../base_llm/text_to_speech/transformation.py | 32 +- .../base_llm/vector_store/transformation.py | 36 +- .../vector_store_files/transformation.py | 52 +- .../llms/base_llm/videos/transformation.py | 79 +- litellm/llms/baseten/chat.py | 49 +- .../bedrock/audio_transcription/__init__.py | 5 +- litellm/llms/bedrock/base_aws_llm.py | 178 +- litellm/llms/bedrock/batches/handler.py | 12 +- .../llms/bedrock/batches/transformation.py | 44 +- litellm/llms/bedrock/chat/__init__.py | 4 +- .../bedrock/chat/agentcore/transformation.py | 80 +- litellm/llms/bedrock/chat/converse_handler.py | 38 +- .../bedrock/chat/converse_transformation.py | 274 ++- .../chat/invoke_agent/transformation.py | 70 +- litellm/llms/bedrock/chat/invoke_handler.py | 146 +- .../amazon_ai21_transformation.py | 31 +- .../amazon_cohere_transformation.py | 13 +- .../amazon_deepseek_transformation.py | 12 +- .../amazon_llama_transformation.py | 15 +- .../amazon_mistral_transformation.py | 24 +- .../amazon_moonshot_transformation.py | 30 +- .../amazon_nova_transformation.py | 11 +- .../amazon_openai_transformation.py | 38 +- .../amazon_qwen2_transformation.py | 14 +- .../amazon_qwen3_transformation.py | 40 +- .../amazon_titan_transformation.py | 21 +- ...mazon_twelvelabs_pegasus_transformation.py | 24 +- .../anthropic_claude2_transformation.py | 25 +- .../anthropic_claude3_transformation.py | 26 +- .../base_invoke_transformation.py | 88 +- .../bedrock/chat/mantle/transformation.py | 20 +- .../llms/bedrock/claude_platform/__init__.py | 6 +- .../bedrock/claude_platform/common_utils.py | 20 +- .../messages_transformation.py | 16 +- .../bedrock/claude_platform/transformation.py | 14 +- litellm/llms/bedrock/common_utils.py | 97 +- litellm/llms/bedrock/cost_calculation.py | 4 +- .../count_tokens/bedrock_token_counter.py | 18 +- litellm/llms/bedrock/count_tokens/handler.py | 14 +- .../bedrock/count_tokens/transformation.py | 31 +- .../embed/amazon_nova_transformation.py | 18 +- .../embed/amazon_titan_g1_transformation.py | 7 +- .../amazon_titan_multimodal_transformation.py | 12 +- .../embed/amazon_titan_v2_transformation.py | 15 +- .../bedrock/embed/cohere_transformation.py | 6 +- litellm/llms/bedrock/embed/embedding.py | 70 +- .../twelvelabs_marengo_transformation.py | 16 +- litellm/llms/bedrock/files/handler.py | 18 +- litellm/llms/bedrock/files/transformation.py | 74 +- ...n_nova_canvas_image_edit_transformation.py | 62 +- litellm/llms/bedrock/image_edit/handler.py | 32 +- .../image_edit/stability_transformation.py | 33 +- .../amazon_nova_canvas_transformation.py | 20 +- .../amazon_stability1_transformation.py | 29 +- .../amazon_stability3_transformation.py | 14 +- .../amazon_titan_transformation.py | 41 +- .../image_generation/cost_calculator.py | 6 +- .../bedrock/image_generation/image_handler.py | 22 +- .../anthropic_claude3_transformation.py | 81 +- .../bedrock/messages/mantle_transformation.py | 22 +- .../guardrail_translation/handler.py | 42 +- .../bedrock/passthrough/transformation.py | 23 +- litellm/llms/bedrock/realtime/handler.py | 32 +- .../llms/bedrock/realtime/transformation.py | 116 +- litellm/llms/bedrock/rerank/handler.py | 30 +- litellm/llms/bedrock/rerank/transformation.py | 6 +- .../bedrock/vector_stores/transformation.py | 44 +- .../bedrock_mantle/chat/transformation.py | 22 +- litellm/llms/bedrock_mantle/common_utils.py | 3 +- .../responses/transformation.py | 16 +- litellm/llms/black_forest_labs/__init__.py | 6 +- .../llms/black_forest_labs/common_utils.py | 7 +- .../black_forest_labs/image_edit/__init__.py | 2 +- .../black_forest_labs/image_edit/handler.py | 40 +- .../image_edit/transformation.py | 30 +- .../image_generation/__init__.py | 4 +- .../image_generation/handler.py | 32 +- .../image_generation/transformation.py | 26 +- litellm/llms/brave/search/transformation.py | 45 +- litellm/llms/bytez/chat/transformation.py | 58 +- litellm/llms/bytez/common_utils.py | 4 +- litellm/llms/cerebras/chat.py | 42 +- litellm/llms/chatgpt/authenticator.py | 36 +- litellm/llms/chatgpt/chat/streaming_utils.py | 6 +- litellm/llms/chatgpt/chat/transformation.py | 18 +- litellm/llms/chatgpt/common_utils.py | 20 +- .../llms/chatgpt/responses/transformation.py | 22 +- litellm/llms/clarifai/chat/transformation.py | 32 +- .../llms/cloudflare/chat/transformation.py | 20 +- litellm/llms/codestral/completion/handler.py | 25 +- .../codestral/completion/transformation.py | 29 +- litellm/llms/cohere/chat/transformation.py | 106 +- litellm/llms/cohere/chat/v2_transformation.py | 112 +- litellm/llms/cohere/common_utils.py | 49 +- litellm/llms/cohere/embed/handler.py | 24 +- litellm/llms/cohere/embed/transformation.py | 38 +- .../llms/cohere/embed/v1_transformation.py | 18 +- .../rerank/guardrail_translation/__init__.py | 2 +- .../rerank/guardrail_translation/handler.py | 10 +- litellm/llms/cohere/rerank/transformation.py | 14 +- .../llms/cohere/rerank_v2/transformation.py | 10 +- litellm/llms/cometapi/chat/transformation.py | 24 +- litellm/llms/cometapi/common_utils.py | 2 - litellm/llms/cometapi/embed/transformation.py | 22 +- .../image_generation/transformation.py | 26 +- .../llms/compactifai/chat/transformation.py | 22 +- litellm/llms/custom_httpx/aiohttp_handler.py | 70 +- .../llms/custom_httpx/aiohttp_transport.py | 24 +- .../llms/custom_httpx/container_handler.py | 45 +- litellm/llms/custom_httpx/http_handler.py | 213 +-- litellm/llms/custom_httpx/httpx_handler.py | 9 +- litellm/llms/custom_httpx/llm_http_handler.py | 1602 ++++++++-------- litellm/llms/custom_httpx/mock_transport.py | 3 +- litellm/llms/custom_llm.py | 72 +- litellm/llms/dashscope/chat/transformation.py | 30 +- litellm/llms/dashscope/common_utils.py | 4 +- litellm/llms/dashscope/cost_calculator.py | 7 +- .../llms/dashscope/embed/transformation.py | 22 +- .../image_generation/transformation.py | 20 +- .../llms/dashscope/rerank/transformation.py | 18 +- .../llms/databricks/chat/transformation.py | 126 +- litellm/llms/databricks/common_utils.py | 34 +- litellm/llms/databricks/cost_calculator.py | 13 +- litellm/llms/databricks/embed/handler.py | 11 +- .../llms/databricks/embed/transformation.py | 5 +- .../databricks/responses/transformation.py | 15 +- litellm/llms/databricks/streaming_utils.py | 7 +- .../llms/dataforseo/search/transformation.py | 22 +- litellm/llms/datarobot/chat/transformation.py | 19 +- .../audio_transcription/transformation.py | 19 +- litellm/llms/deepinfra/chat/transformation.py | 78 +- .../llms/deepinfra/rerank/transformation.py | 14 +- litellm/llms/deepseek/chat/transformation.py | 32 +- litellm/llms/deepseek/cost_calculator.py | 4 +- .../llms/deepseek/messages/transformation.py | 37 +- .../llms/deprecated_providers/aleph_alpha.py | 125 +- litellm/llms/deprecated_providers/palm.py | 29 +- .../chat/transformation.py | 24 +- .../llms/duckduckgo/search/transformation.py | 18 +- litellm/llms/e2b/sandbox/transformation.py | 10 +- .../audio_transcription/transformation.py | 20 +- .../text_to_speech/transformation.py | 48 +- litellm/llms/exa_ai/search/transformation.py | 28 +- litellm/llms/fal_ai/__init__.py | 8 +- .../llms/fal_ai/image_generation/__init__.py | 24 +- .../image_generation/bria_transformation.py | 12 +- .../bytedance_transformation.py | 4 +- .../flux_pro_v11_transformation.py | 4 +- .../flux_pro_v11_ultra_transformation.py | 12 +- .../flux_schnell_transformation.py | 4 +- .../ideogram_v3_transformation.py | 10 +- .../imagen4_transformation.py | 12 +- .../nano_banana_transformation.py | 12 +- .../recraft_v3_transformation.py | 12 +- .../stable_diffusion_transformation.py | 18 +- .../fal_ai/image_generation/transformation.py | 26 +- litellm/llms/fastcrw/search/transformation.py | 12 +- .../featherless_ai/chat/transformation.py | 62 +- .../llms/firecrawl/search/transformation.py | 24 +- .../llms/fireworks_ai/chat/transformation.py | 126 +- litellm/llms/fireworks_ai/common_utils.py | 14 +- .../fireworks_ai/completion/transformation.py | 4 +- litellm/llms/fireworks_ai/cost_calculator.py | 4 +- .../fireworks_ai/rerank/transformation.py | 22 +- litellm/llms/gdc/chat/transformation.py | 2 +- litellm/llms/gemini/agents/transformation.py | 50 +- litellm/llms/gemini/chat/transformation.py | 53 +- litellm/llms/gemini/common_utils.py | 84 +- litellm/llms/gemini/cost_calculator.py | 4 +- litellm/llms/gemini/count_tokens/handler.py | 26 +- litellm/llms/gemini/files/transformation.py | 38 +- .../gemini/google_genai/transformation.py | 46 +- litellm/llms/gemini/image_edit/__init__.py | 4 +- .../llms/gemini/image_edit/transformation.py | 38 +- .../gemini/image_generation/transformation.py | 22 +- .../gemini/interactions/transformation.py | 43 +- .../llms/gemini/realtime/transformation.py | 182 +- .../gemini/vector_stores/transformation.py | 32 +- litellm/llms/gemini/videos/transformation.py | 72 +- litellm/llms/gigachat/authenticator.py | 27 +- litellm/llms/gigachat/chat/__init__.py | 2 +- litellm/llms/gigachat/chat/streaming.py | 14 +- litellm/llms/gigachat/chat/transformation.py | 64 +- .../llms/gigachat/embedding/transformation.py | 34 +- litellm/llms/gigachat/file_handler.py | 21 +- litellm/llms/github_copilot/authenticator.py | 60 +- .../github_copilot/chat/transformation.py | 21 +- litellm/llms/github_copilot/common_utils.py | 9 +- .../embedding/transformation.py | 17 +- .../github_copilot/messages/transformation.py | 16 +- .../responses/transformation.py | 23 +- .../llms/google_pse/search/transformation.py | 22 +- .../llms/gradient_ai/chat/transformation.py | 70 +- litellm/llms/groq/chat/handler.py | 16 +- litellm/llms/groq/chat/transformation.py | 88 +- litellm/llms/groq/stt/transformation.py | 59 +- litellm/llms/heroku/chat/transformation.py | 24 +- .../llms/hosted_vllm/chat/transformation.py | 31 +- .../hosted_vllm/embedding/transformation.py | 22 +- .../llms/hosted_vllm/rerank/transformation.py | 20 +- .../hosted_vllm/responses/transformation.py | 6 +- .../transcriptions/transformation.py | 10 +- .../llms/huggingface/chat/transformation.py | 24 +- litellm/llms/huggingface/common_utils.py | 10 +- litellm/llms/huggingface/embedding/handler.py | 56 +- .../huggingface/embedding/transformation.py | 114 +- .../llms/huggingface/rerank/transformation.py | 19 +- .../llms/hyperbolic/chat/transformation.py | 8 +- litellm/llms/inception/chat/transformation.py | 10 +- .../inception/completion/transformation.py | 4 +- litellm/llms/infinity/common_utils.py | 3 +- .../llms/infinity/embedding/transformation.py | 20 +- .../llms/infinity/rerank/transformation.py | 16 +- .../llms/jina_ai/embedding/transformation.py | 38 +- litellm/llms/jina_ai/rerank/transformation.py | 28 +- litellm/llms/lambda_ai/chat/transformation.py | 8 +- litellm/llms/langflow/a2a.py | 14 +- litellm/llms/langflow/chat/transformation.py | 58 +- litellm/llms/langgraph/chat/sse_iterator.py | 20 +- litellm/llms/langgraph/chat/transformation.py | 66 +- litellm/llms/lemonade/chat/transformation.py | 78 +- litellm/llms/lemonade/cost_calculator.py | 4 +- litellm/llms/linkup/search/transformation.py | 22 +- .../llms/litellm_proxy/chat/transformation.py | 22 +- .../image_edit/transformation.py | 10 +- .../image_generation/transformation.py | 12 +- .../litellm_proxy/responses/transformation.py | 4 +- litellm/llms/litellm_proxy/skills/__init__.py | 16 +- .../litellm_proxy/skills/code_execution.py | 32 +- litellm/llms/litellm_proxy/skills/handler.py | 26 +- .../litellm_proxy/skills/prompt_injection.py | 20 +- .../litellm_proxy/skills/sandbox_executor.py | 20 +- .../litellm_proxy/skills/transformation.py | 52 +- litellm/llms/llamafile/chat/transformation.py | 10 +- litellm/llms/lm_studio/chat/transformation.py | 6 +- .../llms/lm_studio/embed/transformation.py | 3 +- litellm/llms/manus/files/transformation.py | 30 +- .../llms/manus/responses/transformation.py | 24 +- litellm/llms/maritalk.py | 32 +- .../milvus/vector_stores/transformation.py | 22 +- litellm/llms/minimax/__init__.py | 2 +- litellm/llms/minimax/chat/transformation.py | 18 +- .../llms/minimax/messages/transformation.py | 14 +- .../llms/minimax/text_to_speech/__init__.py | 2 +- .../minimax/text_to_speech/transformation.py | 42 +- .../audio_transcription/transformation.py | 20 +- litellm/llms/mistral/chat/transformation.py | 105 +- .../ocr/guardrail_translation/__init__.py | 2 +- .../ocr/guardrail_translation/handler.py | 16 +- litellm/llms/mistral/ocr/transformation.py | 6 +- .../llms/modelscope/chat/transformation.py | 14 +- .../image_generation/transformation.py | 20 +- litellm/llms/moonshot/chat/transformation.py | 41 +- litellm/llms/morph/chat/transformation.py | 8 +- litellm/llms/nlp_cloud/chat/handler.py | 3 +- litellm/llms/nlp_cloud/chat/transformation.py | 70 +- litellm/llms/nlp_cloud/common_utils.py | 4 +- litellm/llms/novita/chat/transformation.py | 8 +- litellm/llms/nscale/chat/transformation.py | 12 +- litellm/llms/nvidia_nim/embed.py | 19 +- .../rerank/ranking_transformation.py | 19 +- .../llms/nvidia_nim/rerank/transformation.py | 29 +- .../audio_transcription/audio_utils.py | 4 +- .../audio_transcription/handler.py | 36 +- .../audio_transcription/transformation.py | 28 +- litellm/llms/nvidia_riva/common_utils.py | 8 +- litellm/llms/oci/chat/cohere.py | 36 +- litellm/llms/oci/chat/generic.py | 39 +- litellm/llms/oci/chat/transformation.py | 73 +- litellm/llms/oci/common_utils.py | 36 +- litellm/llms/oci/embed/transformation.py | 30 +- litellm/llms/ollama/chat/transformation.py | 101 +- litellm/llms/ollama/common_utils.py | 24 +- litellm/llms/ollama/completion/handler.py | 16 +- .../llms/ollama/completion/transformation.py | 132 +- litellm/llms/oobabooga/chat/oobabooga.py | 8 +- litellm/llms/oobabooga/chat/transformation.py | 16 +- litellm/llms/oobabooga/common_utils.py | 4 +- .../llms/openai/chat/gpt_5_transformation.py | 22 +- .../llms/openai/chat/gpt_transformation.py | 132 +- .../chat/guardrail_translation/handler.py | 116 +- .../openai/chat/o_series_transformation.py | 22 +- litellm/llms/openai/common_utils.py | 48 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 10 +- litellm/llms/openai/completion/handler.py | 17 +- .../llms/openai/completion/transformation.py | 56 +- litellm/llms/openai/completion/utils.py | 8 +- .../llms/openai/containers/transformation.py | 48 +- litellm/llms/openai/cost_calculation.py | 22 +- litellm/llms/openai/data_residency.py | 5 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 16 +- litellm/llms/openai/fine_tuning/handler.py | 117 +- litellm/llms/openai/image_edit/__init__.py | 2 +- .../image_edit/dalle2_transformation.py | 14 +- .../llms/openai/image_edit/transformation.py | 26 +- .../image_generation/cost_calculator.py | 4 +- .../dall_e_2_transformation.py | 12 +- .../dall_e_3_transformation.py | 12 +- .../image_generation/gpt_transformation.py | 12 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 10 +- .../llms/openai/image_variations/handler.py | 23 +- .../openai/image_variations/transformation.py | 16 +- litellm/llms/openai/openai.py | 906 +++++---- litellm/llms/openai/realtime/handler.py | 20 +- .../openai/realtime/http_transformation.py | 25 +- .../openai/responses/count_tokens/__init__.py | 2 +- .../openai/responses/count_tokens/handler.py | 20 +- .../responses/count_tokens/token_counter.py | 16 +- .../responses/count_tokens/transformation.py | 30 +- .../guardrail_translation/handler.py | 126 +- .../llms/openai/responses/transformation.py | 72 +- .../speech/guardrail_translation/__init__.py | 2 +- .../speech/guardrail_translation/handler.py | 10 +- .../transcriptions/gpt_transformation.py | 4 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 10 +- litellm/llms/openai/transcriptions/handler.py | 16 +- .../transcriptions/whisper_transformation.py | 17 +- .../vector_store_files/transformation.py | 34 +- .../openai/vector_stores/transformation.py | 14 +- litellm/llms/openai/videos/transformation.py | 103 +- litellm/llms/openai_like/chat/handler.py | 39 +- .../llms/openai_like/chat/transformation.py | 26 +- litellm/llms/openai_like/common_utils.py | 18 +- litellm/llms/openai_like/dynamic_config.py | 32 +- litellm/llms/openai_like/embedding/handler.py | 11 +- litellm/llms/openai_like/json_loader.py | 5 +- .../openai_like/messages/transformation.py | 35 +- .../openai_like/responses/transformation.py | 8 +- .../llms/openrouter/chat/transformation.py | 28 +- .../openrouter/embedding/transformation.py | 14 +- .../openrouter/image_edit/transformation.py | 36 +- .../image_generation/transformation.py | 34 +- .../openrouter/responses/transformation.py | 6 +- .../opensandbox/sandbox/transformation.py | 12 +- .../audio_transcription/transformation.py | 20 +- litellm/llms/ovhcloud/chat/transformation.py | 24 +- .../llms/ovhcloud/embedding/transformation.py | 20 +- litellm/llms/ovhcloud/utils.py | 2 - .../llms/parallel_ai/search/transformation.py | 28 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 28 +- .../llms/perplexity/chat/transformation.py | 24 +- litellm/llms/perplexity/cost_calculator.py | 6 +- .../perplexity/embedding/transformation.py | 24 +- .../perplexity/responses/transformation.py | 18 +- .../llms/perplexity/search/transformation.py | 22 +- litellm/llms/petals/common_utils.py | 4 +- litellm/llms/petals/completion/handler.py | 7 +- .../llms/petals/completion/transformation.py | 48 +- .../pg_vector/vector_stores/transformation.py | 12 +- litellm/llms/predibase/chat/handler.py | 17 +- litellm/llms/predibase/chat/transformation.py | 82 +- litellm/llms/predibase/common_utils.py | 8 +- litellm/llms/ragflow/chat/transformation.py | 24 +- .../ragflow/vector_stores/transformation.py | 16 +- .../llms/recraft/image_edit/transformation.py | 32 +- .../image_generation/transformation.py | 26 +- litellm/llms/reducto/common.py | 24 +- litellm/llms/reducto/ocr/transformation.py | 32 +- litellm/llms/replicate/chat/handler.py | 7 +- litellm/llms/replicate/chat/transformation.py | 68 +- litellm/llms/replicate/common_utils.py | 4 +- .../image_generation/transformation.py | 38 +- .../runwayml/text_to_speech/transformation.py | 46 +- .../llms/runwayml/videos/transformation.py | 72 +- .../vector_stores/transformation.py | 36 +- litellm/llms/sagemaker/chat/handler.py | 7 +- litellm/llms/sagemaker/chat/transformation.py | 40 +- litellm/llms/sagemaker/common_utils.py | 13 +- litellm/llms/sagemaker/completion/handler.py | 22 +- .../sagemaker/completion/transformation.py | 51 +- .../embedding/cohere_transformation.py | 18 +- .../sagemaker/embedding/transformation.py | 16 +- litellm/llms/sagemaker/nova/transformation.py | 6 +- litellm/llms/sambanova/chat.py | 54 +- .../sambanova/embedding/transformation.py | 20 +- litellm/llms/sap/chat/handler.py | 9 +- litellm/llms/sap/chat/models.py | 110 +- litellm/llms/sap/chat/transformation.py | 100 +- litellm/llms/sap/credentials.py | 74 +- litellm/llms/sap/embed/transformation.py | 30 +- .../audio_transcription/transformation.py | 20 +- .../llms/searchapi/search/transformation.py | 28 +- litellm/llms/searxng/search/transformation.py | 18 +- litellm/llms/serper/search/transformation.py | 18 +- litellm/llms/snowflake/chat/transformation.py | 54 +- litellm/llms/snowflake/common_utils.py | 5 +- .../snowflake/embedding/transformation.py | 16 +- litellm/llms/snowflake/utils.py | 18 +- .../soniox/audio_transcription/handler.py | 119 +- .../audio_transcription/transformation.py | 44 +- litellm/llms/soniox/common_utils.py | 57 +- .../stability/image_edit/transformations.py | 40 +- .../image_generation/transformation.py | 25 +- litellm/llms/tavily/search/transformation.py | 26 +- litellm/llms/tencent/chat/transformation.py | 12 +- .../llms/tencent/messages/transformation.py | 20 +- litellm/llms/together_ai/chat.py | 7 +- .../together_ai/completion/transformation.py | 6 +- litellm/llms/together_ai/rerank/handler.py | 16 +- .../llms/together_ai/rerank/transformation.py | 6 +- litellm/llms/topaz/common_utils.py | 14 +- .../topaz/image_variations/transformation.py | 26 +- litellm/llms/triton/common_utils.py | 4 +- .../llms/triton/completion/transformation.py | 84 +- .../llms/triton/embedding/transformation.py | 18 +- litellm/llms/v0/chat/transformation.py | 8 +- .../vercel_ai_gateway/chat/transformation.py | 22 +- .../embedding/transformation.py | 14 +- .../vertex_ai/agent_engine/sse_iterator.py | 4 +- .../vertex_ai/agent_engine/transformation.py | 63 +- litellm/llms/vertex_ai/batches/handler.py | 78 +- .../llms/vertex_ai/batches/transformation.py | 10 +- litellm/llms/vertex_ai/common_utils.py | 130 +- .../context_caching/transformation.py | 24 +- .../vertex_ai_context_caching.py | 92 +- litellm/llms/vertex_ai/cost_calculator.py | 65 +- .../llms/vertex_ai/count_tokens/handler.py | 12 +- litellm/llms/vertex_ai/files/handler.py | 34 +- .../llms/vertex_ai/files/transformation.py | 69 +- litellm/llms/vertex_ai/fine_tuning/handler.py | 35 +- .../llms/vertex_ai/gemini/transformation.py | 170 +- .../vertex_and_google_ai_studio_gemini.py | 450 +++-- .../batch_embed_content_handler.py | 34 +- .../batch_embed_content_transformation.py | 21 +- .../vertex_ai/google_genai/transformation.py | 16 +- litellm/llms/vertex_ai/image_edit/__init__.py | 2 +- .../vertex_gemini_transformation.py | 50 +- .../vertex_imagen_transformation.py | 52 +- .../image_generation_handler.py | 48 +- .../vertex_gemini_transformation.py | 27 +- .../vertex_imagen_transformation.py | 28 +- .../embedding_handler.py | 16 +- .../multimodal_embeddings/transformation.py | 28 +- .../vertex_ai/ocr/deepseek_transformation.py | 6 +- litellm/llms/vertex_ai/ocr/transformation.py | 6 +- .../llms/vertex_ai/rag_engine/ingestion.py | 26 +- .../vertex_ai/rag_engine/transformation.py | 16 +- .../llms/vertex_ai/realtime/transformation.py | 11 +- .../llms/vertex_ai/rerank/transformation.py | 14 +- .../text_to_speech/text_to_speech_handler.py | 26 +- .../text_to_speech/transformation.py | 52 +- .../llms/vertex_ai/vector_stores/__init__.py | 2 +- .../vector_stores/rag_api/transformation.py | 18 +- .../search_api/transformation.py | 22 +- litellm/llms/vertex_ai/vertex_ai_aws_wif.py | 4 +- .../llms/vertex_ai/vertex_ai_non_gemini.py | 4 +- .../ai21/transformation.py | 3 +- .../transformation.py | 24 +- .../anthropic/transformation.py | 16 +- .../count_tokens/handler.py | 10 +- .../llama3/transformation.py | 20 +- .../vertex_ai_partner_models/main.py | 9 +- .../llms/vertex_ai/vertex_embeddings/bge.py | 12 +- .../vertex_embeddings/embedding_handler.py | 46 +- .../vertex_embeddings/transformation.py | 52 +- .../llms/vertex_ai/vertex_embeddings/types.py | 31 +- .../vertex_ai/vertex_gemma_models/main.py | 7 +- .../vertex_gemma_models/transformation.py | 26 +- litellm/llms/vertex_ai/vertex_llm_base.py | 190 +- .../vertex_ai/vertex_model_garden/main.py | 11 +- .../llms/vertex_ai/videos/transformation.py | 94 +- litellm/llms/vllm/common_utils.py | 26 +- .../llms/vllm/completion/transformation.py | 2 - .../llms/vllm/passthrough/transformation.py | 10 +- litellm/llms/volcengine/__init__.py | 2 +- .../llms/volcengine/chat/transformation.py | 46 +- litellm/llms/volcengine/common_utils.py | 8 +- .../volcengine/embedding/transformation.py | 41 +- .../llms/voyage/embedding/transformation.py | 22 +- .../embedding/transformation_contextual.py | 24 +- .../embedding/transformation_multimodal.py | 26 +- litellm/llms/voyage/rerank/transformation.py | 34 +- .../audio_transcription/transformation.py | 29 +- litellm/llms/watsonx/chat/handler.py | 15 +- litellm/llms/watsonx/chat/transformation.py | 22 +- litellm/llms/watsonx/common_utils.py | 58 +- .../llms/watsonx/completion/transformation.py | 118 +- litellm/llms/watsonx/embed/transformation.py | 10 +- .../watsonx/passthrough/transformation.py | 22 +- litellm/llms/watsonx/rerank/transformation.py | 18 +- litellm/llms/xai/chat/transformation.py | 44 +- litellm/llms/xai/common_utils.py | 20 +- litellm/llms/xai/cost_calculator.py | 6 +- litellm/llms/xai/oauth.py | 40 +- litellm/llms/xai/realtime/transformation.py | 4 +- litellm/llms/xai/responses/transformation.py | 18 +- .../image_generation/transformation.py | 8 +- litellm/llms/you_com/search/transformation.py | 24 +- litellm/llms/zai/chat/transformation.py | 14 +- litellm/main.py | 580 +++--- litellm/models/__init__.py | 18 +- litellm/models/access_group.py | 21 +- litellm/models/base.py | 8 +- litellm/models/budget.py | 29 +- litellm/models/config.py | 4 +- litellm/models/credentials.py | 6 +- litellm/models/end_user.py | 16 +- litellm/models/managed_files.py | 58 +- litellm/models/mcp_server.py | 100 +- litellm/models/model.py | 15 +- litellm/models/object_permission.py | 24 +- litellm/models/organization.py | 20 +- litellm/models/organization_membership.py | 12 +- litellm/models/project.py | 33 +- litellm/models/skills.py | 26 +- litellm/models/spend_logs.py | 55 +- litellm/models/tag.py | 17 +- litellm/models/team.py | 82 +- litellm/models/team_membership.py | 14 +- litellm/models/user.py | 49 +- litellm/models/verification_token.py | 95 +- litellm/ocr/__init__.py | 2 +- litellm/passthrough/__init__.py | 2 +- litellm/passthrough/main.py | 74 +- litellm/passthrough/timeout_utils.py | 9 +- litellm/passthrough/utils.py | 9 +- .../mcp_server/auth/litellm_auth_handler.py | 18 +- .../mcp_server/auth/user_api_key_auth_mcp.py | 191 +- .../mcp_server/bridge_token_flow.py | 4 +- .../mcp_server/byok_oauth_endpoints.py | 20 +- .../mcp_server/cost_calculator.py | 4 +- litellm/proxy/_experimental/mcp_server/db.py | 6 +- .../mcp_server/discoverable_endpoints.py | 153 +- .../mcp_server/elicitation_handler.py | 11 +- .../_experimental/mcp_server/exceptions.py | 10 +- .../guardrail_translation/__init__.py | 2 +- .../guardrail_translation/handler.py | 14 +- .../_experimental/mcp_server/mcp_context.py | 7 +- .../_experimental/mcp_server/mcp_debug.py | 48 +- .../mcp_server/mcp_server_manager.py | 451 +++-- .../mcp_server/oauth2_flow_backfill.py | 14 +- .../mcp_server/oauth2_token_cache.py | 26 +- .../_experimental/mcp_server/oauth_utils.py | 32 +- .../mcp_server/openapi_to_mcp_generator.py | 50 +- .../outbound_credentials/__init__.py | 54 +- .../outbound_credentials/adapter.py | 14 +- .../outbound_credentials/resolver.py | 1 - .../outbound_credentials/token_exchanger.py | 6 +- .../mcp_server/rest_endpoints.py | 16 +- .../mcp_server/sampling_handler.py | 46 +- .../mcp_server/semantic_tool_filter.py | 33 +- .../proxy/_experimental/mcp_server/server.py | 38 +- .../_experimental/mcp_server/sse_transport.py | 2 +- .../_experimental/mcp_server/tool_registry.py | 16 +- .../_experimental/mcp_server/tool_search.py | 28 +- .../_experimental/mcp_server/toolset_db.py | 17 +- .../mcp_server/ui_session_utils.py | 8 +- .../proxy/_experimental/mcp_server/utils.py | 56 +- litellm/proxy/_lazy_features.py | 14 +- litellm/proxy/_lazy_openapi_snapshot.py | 11 +- litellm/proxy/_logging.py | 2 +- litellm/proxy/_types.py | 1666 ++++++++--------- litellm/proxy/a2a/__init__.py | 2 +- litellm/proxy/a2a/agent_card.py | 18 +- litellm/proxy/a2a/discovery.py | 18 +- litellm/proxy/a2a/endpoints.py | 6 +- .../proxy/agent_endpoints/a2a_endpoints.py | 36 +- litellm/proxy/agent_endpoints/a2a_routing.py | 8 +- .../proxy/agent_endpoints/agent_registry.py | 16 +- .../auth/agent_permission_handler.py | 72 +- .../proxy/agent_endpoints/databricks_oauth.py | 19 +- litellm/proxy/agent_endpoints/endpoints.py | 18 +- .../agent_endpoints/model_list_helpers.py | 10 +- .../claude_code_marketplace.py | 12 +- .../proxy/anthropic_endpoints/endpoints.py | 14 +- .../anthropic_endpoints/skills_endpoints.py | 16 +- litellm/proxy/auth/auth_checks.py | 552 +++--- .../proxy/auth/auth_checks_organization.py | 23 +- litellm/proxy/auth/auth_exception_handler.py | 15 +- litellm/proxy/auth/auth_utils.py | 130 +- litellm/proxy/auth/budget_throttle.py | 5 +- litellm/proxy/auth/handle_jwt.py | 254 +-- litellm/proxy/auth/ip_address_utils.py | 30 +- litellm/proxy/auth/litellm_license.py | 40 +- litellm/proxy/auth/login_utils.py | 23 +- litellm/proxy/auth/model_checks.py | 88 +- litellm/proxy/auth/oauth2_check.py | 18 +- litellm/proxy/auth/oauth2_proxy_hook.py | 7 +- litellm/proxy/auth/rds_iam_token.py | 20 +- litellm/proxy/auth/resolvers/exceptions.py | 4 +- litellm/proxy/auth/route_checks.py | 54 +- litellm/proxy/auth/trusted_proxy_utils.py | 8 +- litellm/proxy/auth/user_api_key_auth.py | 196 +- litellm/proxy/batches_endpoints/endpoints.py | 48 +- litellm/proxy/caching_routes.py | 20 +- litellm/proxy/client/__init__.py | 18 +- litellm/proxy/client/chat.py | 46 +- litellm/proxy/client/cli/commands/agents.py | 39 +- litellm/proxy/client/cli/commands/auth.py | 46 +- .../client/cli/commands/autoroute/config.py | 2 +- litellm/proxy/client/cli/commands/chat.py | 36 +- .../proxy/client/cli/commands/credentials.py | 5 +- .../proxy/client/cli/commands/encryption.py | 1 - litellm/proxy/client/cli/commands/http.py | 8 +- litellm/proxy/client/cli/commands/keys.py | 68 +- litellm/proxy/client/cli/commands/models.py | 21 +- litellm/proxy/client/cli/commands/teams.py | 13 +- litellm/proxy/client/cli/commands/users.py | 4 +- litellm/proxy/client/cli/interface.py | 4 - litellm/proxy/client/cli/main.py | 5 +- litellm/proxy/client/client.py | 4 +- litellm/proxy/client/credentials.py | 19 +- litellm/proxy/client/exceptions.py | 10 +- litellm/proxy/client/health.py | 9 +- litellm/proxy/client/http_client.py | 11 +- litellm/proxy/client/keys.py | 79 +- litellm/proxy/client/model_groups.py | 10 +- litellm/proxy/client/models.py | 37 +- litellm/proxy/client/teams.py | 29 +- litellm/proxy/client/users.py | 20 +- litellm/proxy/common_request_processing.py | 175 +- .../proxy/common_utils/cache_coordinator.py | 16 +- .../common_utils/cache_pydantic_utils.py | 6 +- litellm/proxy/common_utils/callback_utils.py | 38 +- .../proxy/common_utils/custom_openapi_spec.py | 44 +- litellm/proxy/common_utils/debug_utils.py | 32 +- .../common_utils/encrypt_decrypt_utils.py | 6 +- .../expired_ui_session_key_cleanup_manager.py | 8 +- litellm/proxy/common_utils/get_routes.py | 12 +- .../proxy/common_utils/http_parsing_utils.py | 54 +- .../common_utils/key_rotation_manager.py | 3 +- .../proxy/common_utils/load_config_utils.py | 13 +- .../proxy/common_utils/model_listing_utils.py | 8 +- .../common_utils/openai_endpoint_utils.py | 10 +- .../common_utils/openapi_schema_compat.py | 4 +- .../proxy/common_utils/performance_utils.py | 8 +- .../common_utils/proxy_rate_limit_error.py | 18 +- litellm/proxy/common_utils/realtime_utils.py | 3 +- .../proxy/common_utils/reset_budget_job.py | 46 +- .../proxy/common_utils/resource_ownership.py | 20 +- .../proxy/common_utils/static_asset_utils.py | 5 +- litellm/proxy/common_utils/swagger_utils.py | 4 +- .../proxy/common_utils/user_api_key_cache.py | 28 +- litellm/proxy/compliance_checks.py | 12 +- .../proxy/container_endpoints/endpoints.py | 8 +- .../container_endpoints/handler_factory.py | 14 +- .../proxy/container_endpoints/ownership.py | 30 +- .../proxy/credential_endpoints/endpoints.py | 14 +- litellm/proxy/custom_auth_auto.py | 4 +- litellm/proxy/custom_prompt_management.py | 20 +- litellm/proxy/custom_sso.py | 2 +- litellm/proxy/db/base_client.py | 9 +- litellm/proxy/db/check_migration.py | 11 +- litellm/proxy/db/create_views.py | 4 +- litellm/proxy/db/db_spend_update_writer.py | 205 +- .../db_transaction_queue/base_update_queue.py | 4 +- .../daily_spend_update_queue.py | 17 +- .../db_transaction_queue/pod_lock_manager.py | 12 +- .../redis_update_buffer.py | 94 +- .../db_transaction_queue/spend_log_cleanup.py | 9 +- .../spend_logs_partition_manager.py | 19 +- .../spend_update_queue.py | 16 +- .../tool_discovery_queue.py | 8 +- litellm/proxy/db/exception_handler.py | 12 +- litellm/proxy/db/log_db_metrics.py | 7 +- litellm/proxy/db/prisma_client.py | 4 +- litellm/proxy/db/query_engine_reaper.py | 5 +- litellm/proxy/db/routing_prisma_wrapper.py | 2 +- litellm/proxy/db/spend_counter_reseed.py | 14 +- litellm/proxy/db/tool_registry_writer.py | 46 +- litellm/proxy/dd_span_tagger.py | 6 +- .../ui_discovery_endpoints.py | 3 +- .../enterprise_billing/billing_metrics.py | 20 +- .../proxy/fine_tuning_endpoints/endpoints.py | 30 +- litellm/proxy/guardrails/_content_utils.py | 28 +- .../proxy/guardrails/guardrail_endpoints.py | 104 +- litellm/proxy/guardrails/guardrail_helpers.py | 3 +- .../guardrails/guardrail_hooks/aim/aim.py | 26 +- .../guardrails/guardrail_hooks/akto/akto.py | 51 +- .../guardrail_hooks/aporia_ai/aporia_ai.py | 23 +- .../guardrail_hooks/azure/__init__.py | 9 +- .../guardrails/guardrail_hooks/azure/base.py | 12 +- .../guardrail_hooks/azure/prompt_shield.py | 17 +- .../guardrail_hooks/azure/text_moderation.py | 23 +- .../guardrail_hooks/bedrock_guardrails.py | 150 +- .../block_code_execution/__init__.py | 10 +- .../block_code_execution.py | 44 +- .../cato_networks/cato_networks.py | 38 +- .../cisco_ai_defense/__init__.py | 4 +- .../cisco_ai_defense/cisco_ai_defense.py | 246 ++- .../cisco_ai_defense/cisco_ai_defense_mcp.py | 53 +- .../crowdstrike_aidr/crowdstrike_aidr.py | 22 +- .../guardrail_hooks/custom_code/__init__.py | 2 +- .../custom_code/custom_code_guardrail.py | 16 +- .../guardrail_hooks/custom_code/primitives.py | 82 +- .../guardrail_hooks/custom_code/sandbox.py | 10 +- .../guardrail_hooks/custom_guardrail.py | 4 +- .../guardrail_hooks/deepkeep/deepkeep.py | 6 +- .../guardrail_hooks/dynamoai/dynamoai.py | 24 +- .../guardrail_hooks/enkryptai/enkryptai.py | 19 +- .../generic_guardrail_api/__init__.py | 4 +- .../generic_guardrail_api.py | 22 +- .../guardrail_hooks/grayswan/grayswan.py | 52 +- .../guardrails_ai/guardrails_ai.py | 25 +- .../guardrail_hooks/headroom/headroom.py | 12 +- .../hiddenlayer/hiddenlayer.py | 32 +- .../ibm_guardrails/ibm_detector.py | 44 +- .../guardrail_hooks/javelin/javelin.py | 26 +- .../guardrails/guardrail_hooks/lakera_ai.py | 26 +- .../guardrail_hooks/lakera_ai_v2.py | 65 +- .../guardrails/guardrail_hooks/lasso/lasso.py | 115 +- .../competitor_intent/__init__.py | 2 +- .../competitor_intent/airline.py | 14 +- .../competitor_intent/base.py | 48 +- .../litellm_content_filter/content_filter.py | 190 +- .../guardrail_benchmarks/test_eval.py | 8 +- .../litellm_content_filter/patterns.py | 24 +- .../llm_as_a_judge/__init__.py | 30 +- .../mcp_end_user_permission/__init__.py | 4 +- .../mcp_end_user_permission.py | 22 +- .../mcp_jwt_signer/__init__.py | 2 +- .../mcp_jwt_signer/mcp_jwt_signer.py | 136 +- .../mcp_security/mcp_security_guardrail.py | 12 +- .../guardrail_hooks/microsoft_purview/base.py | 50 +- .../microsoft_purview/purview_dlp.py | 55 +- .../model_armor/model_armor.py | 47 +- .../guardrails/guardrail_hooks/noma/noma.py | 78 +- .../guardrail_hooks/noma/noma_v2.py | 22 +- .../guardrails/guardrail_hooks/onyx/onyx.py | 16 +- .../guardrails/guardrail_hooks/openai/base.py | 4 +- .../guardrail_hooks/openai/moderations.py | 48 +- .../guardrail_hooks/ovalix/ovalix.py | 36 +- .../guardrail_hooks/pangea/pangea.py | 18 +- .../panw_prisma_airs/panw_prisma_airs.py | 112 +- .../guardrail_hooks/pillar/pillar.py | 70 +- .../guardrails/guardrail_hooks/presidio.py | 141 +- .../prompt_security/prompt_security.py | 52 +- .../promptguard/promptguard.py | 19 +- .../guardrail_hooks/qohash/qohash.py | 6 +- .../guardrail_hooks/qualifire/qualifire.py | 62 +- .../guardrail_hooks/repelloai/__init__.py | 4 +- .../guardrail_hooks/repelloai/repelloai.py | 4 +- .../semantic_guard/route_loader.py | 20 +- .../semantic_guard/semantic_guard.py | 26 +- .../guardrail_hooks/tool_permission.py | 76 +- .../tool_policy/tool_policy_guardrail.py | 12 +- .../unified_guardrail/unified_guardrail.py | 10 +- .../vigil_guard/vigil_guard.py | 52 +- .../guardrail_hooks/xecguard/xecguard.py | 65 +- .../zscaler_ai_guard/zscaler_ai_guard.py | 22 +- .../guardrails/guardrail_initializers.py | 4 +- .../proxy/guardrails/guardrail_registry.py | 70 +- litellm/proxy/guardrails/init_guardrails.py | 22 +- .../proxy/guardrails/tool_name_extraction.py | 14 +- litellm/proxy/guardrails/usage_tracking.py | 12 +- litellm/proxy/health_check.py | 34 +- .../shared_health_check_manager.py | 22 +- .../health_endpoints/_health_endpoints.py | 105 +- litellm/proxy/hooks/__init__.py | 12 +- litellm/proxy/hooks/azure_content_safety.py | 7 +- litellm/proxy/hooks/batch_rate_limiter.py | 65 +- litellm/proxy/hooks/batch_redis_get.py | 12 +- litellm/proxy/hooks/cache_control_check.py | 4 +- litellm/proxy/hooks/dynamic_rate_limiter.py | 69 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 54 +- .../proxy/hooks/key_management_event_hooks.py | 29 +- .../proxy/hooks/litellm_skills/__init__.py | 10 +- litellm/proxy/hooks/litellm_skills/main.py | 66 +- litellm/proxy/hooks/max_budget_limiter.py | 6 +- .../hooks/max_budget_per_session_limiter.py | 10 +- litellm/proxy/hooks/max_iterations_limiter.py | 10 +- .../proxy/hooks/mcp_semantic_filter/hook.py | 20 +- .../proxy/hooks/model_max_budget_limiter.py | 27 +- .../proxy/hooks/parallel_request_limiter.py | 66 +- .../hooks/parallel_request_limiter_v3.py | 299 ++- .../proxy/hooks/prompt_injection_detection.py | 16 +- .../proxy/hooks/proxy_track_cost_callback.py | 1138 +++++------ litellm/proxy/hooks/rate_limiter_utils.py | 12 +- litellm/proxy/hooks/responses_id_security.py | 14 +- litellm/proxy/hooks/sensitive_data_routing.py | 12 +- .../hooks/user_management_event_hooks.py | 12 +- litellm/proxy/image_endpoints/endpoints.py | 23 +- litellm/proxy/lambda.py | 1 + litellm/proxy/litellm_pre_call_utils.py | 151 +- .../access_group_endpoints.py | 52 +- .../budget_management_endpoints.py | 21 +- .../cache_settings_endpoints.py | 50 +- .../common_daily_activity.py | 8 +- .../management_endpoints/common_utils.py | 12 +- .../config_override_endpoints.py | 36 +- .../coordination_redis_endpoints.py | 23 +- .../cost_tracking_settings.py | 28 +- .../credential_migration.py | 76 +- .../customer_endpoints.py | 59 +- .../fallback_management_endpoints.py | 18 +- .../internal_user_endpoints.py | 68 +- .../key_management_endpoints.py | 559 +++--- .../management_v1/budgets.py | 6 +- .../management_v1/list_framework.py | 5 +- .../management_v1/spend_logs.py | 2 +- .../mcp_management_endpoints.py | 36 +- ...model_access_group_management_endpoints.py | 42 +- .../model_management_endpoints.py | 116 +- .../organization_endpoints.py | 4 +- .../policy_endpoints/ai_policy_suggester.py | 7 +- .../policy_endpoints/endpoints.py | 56 +- .../router_settings_endpoints.py | 20 +- .../scim/scim_transformations.py | 16 +- .../management_endpoints/scim/scim_v2.py | 100 +- .../sso/custom_microsoft_sso.py | 7 +- .../management_endpoints/sso_helper_utils.py | 4 +- .../tag_management_endpoints.py | 8 +- .../team_callback_endpoints.py | 22 +- .../management_endpoints/team_endpoints.py | 373 ++-- .../tool_management_endpoints.py | 26 +- litellm/proxy/management_endpoints/types.py | 12 +- litellm/proxy/management_endpoints/ui_sso.py | 400 ++-- .../usage_endpoints/ai_usage_chat.py | 64 +- .../usage_endpoints/endpoints.py | 6 +- .../user_agent_analytics_endpoints.py | 56 +- .../workflow_management_endpoints.py | 36 +- .../proxy/management_helpers/audit_logs.py | 28 +- .../object_permission_utils.py | 112 +- .../team_member_permission_checks.py | 22 +- litellm/proxy/management_helpers/utils.py | 51 +- litellm/proxy/mcp_tools.py | 6 +- litellm/proxy/memory/memory_endpoints.py | 16 +- .../billable_request_metrics_middleware.py | 22 +- .../in_flight_requests_middleware.py | 6 +- .../middleware/prometheus_auth_middleware.py | 4 +- .../request_size_limit_middleware.py | 9 +- litellm/proxy/ocr_endpoints/endpoints.py | 16 +- .../proxy/openai_evals_endpoints/endpoints.py | 42 +- .../openai_files_endpoints/common_utils.py | 63 +- .../file_content_streaming_handler.py | 14 +- .../openai_files_endpoints/files_endpoints.py | 94 +- .../storage_backend_service.py | 8 +- .../jsonpath_extractor.py | 8 +- .../llm_passthrough_endpoints.py | 136 +- .../anthropic_passthrough_logging_handler.py | 66 +- .../assembly_passthrough_logging_handler.py | 34 +- .../base_passthrough_logging_handler.py | 15 +- .../cohere_passthrough_logging_handler.py | 7 +- .../cursor_passthrough_logging_handler.py | 3 +- .../gemini_passthrough_logging_handler.py | 14 +- .../openai_passthrough_logging_handler.py | 40 +- ...tex_ai_live_passthrough_logging_handler.py | 16 +- .../vertex_passthrough_logging_handler.py | 42 +- .../managed_id_codec.py | 3 +- .../managed_id_rewriter.py | 78 +- .../pass_through_endpoints.py | 276 ++- .../passthrough_endpoint_router.py | 36 +- .../passthrough_guardrails.py | 34 +- .../streaming_handler.py | 29 +- .../pass_through_endpoints/success_handler.py | 30 +- litellm/proxy/plugin_routes.py | 2 +- .../policy_engine/attachment_registry.py | 42 +- .../policy_engine/condition_evaluator.py | 7 +- litellm/proxy/policy_engine/init_policies.py | 22 +- .../proxy/policy_engine/pipeline_executor.py | 14 +- .../proxy/policy_engine/policy_endpoints.py | 4 +- litellm/proxy/policy_engine/policy_matcher.py | 14 +- .../proxy/policy_engine/policy_registry.py | 30 +- .../proxy/policy_engine/policy_resolver.py | 44 +- .../proxy/policy_engine/policy_validator.py | 40 +- litellm/proxy/prompts/init_prompts.py | 8 +- litellm/proxy/prompts/prompt_endpoints.py | 52 +- litellm/proxy/prompts/prompt_registry.py | 17 +- litellm/proxy/proxy_cli.py | 58 +- litellm/proxy/proxy_server.py | 949 +++++----- .../public_endpoints/public_endpoints.py | 40 +- litellm/proxy/rag_endpoints/endpoints.py | 22 +- litellm/proxy/realtime_endpoints/endpoints.py | 28 +- litellm/proxy/rerank_endpoints/endpoints.py | 4 +- .../proxy/response_api_endpoints/endpoints.py | 20 +- .../response_polling/background_streaming.py | 7 +- .../proxy/response_polling/polling_handler.py | 48 +- litellm/proxy/route_llm_request.py | 40 +- litellm/proxy/search_endpoints/endpoints.py | 4 +- .../search_tool_management.py | 16 +- .../search_endpoints/search_tool_registry.py | 33 +- .../shutdown/graceful_shutdown_manager.py | 7 +- .../spend_tracking/budget_reservation.py | 158 +- .../spend_tracking/cloudzero_endpoints.py | 28 +- .../spend_tracking/spend_log_error_logger.py | 4 +- .../spend_management_endpoints.py | 79 +- .../spend_tracking/spend_tracking_utils.py | 87 +- .../proxy/spend_tracking/vantage_endpoints.py | 28 +- litellm/proxy/types_utils/utils.py | 14 +- .../proxy_setting_endpoints.py | 56 +- litellm/proxy/utils.py | 405 ++-- .../proxy/vector_store_endpoints/endpoints.py | 18 +- .../management_endpoints.py | 52 +- litellm/proxy/vector_store_endpoints/utils.py | 26 +- .../vector_store_files_endpoints/endpoints.py | 34 +- .../vertex_ai_endpoints/langfuse_endpoints.py | 21 +- litellm/proxy/video_endpoints/endpoints.py | 12 +- litellm/proxy/video_endpoints/utils.py | 8 +- litellm/proxy_auth/__init__.py | 4 +- litellm/proxy_auth/credentials.py | 8 +- litellm/rag/__init__.py | 2 +- litellm/rag/ingestion/base_ingestion.py | 59 +- litellm/rag/ingestion/bedrock_ingestion.py | 36 +- .../rag/ingestion/file_parsers/pdf_parser.py | 4 +- litellm/rag/ingestion/gemini_ingestion.py | 34 +- litellm/rag/ingestion/openai_ingestion.py | 22 +- litellm/rag/ingestion/s3_vectors_ingestion.py | 36 +- litellm/rag/ingestion/vertex_ai_ingestion.py | 34 +- litellm/rag/main.py | 70 +- litellm/rag/rag_query.py | 18 +- .../recursive_character_text_splitter.py | 24 +- litellm/rag/utils.py | 4 +- litellm/realtime_api/main.py | 60 +- litellm/repositories/__init__.py | 92 +- litellm/repositories/base_repository.py | 36 +- litellm/repositories/budget_repository.py | 46 +- litellm/repositories/config_repository.py | 18 +- .../repositories/credentials_repository.py | 10 +- litellm/repositories/model_repository.py | 50 +- .../object_permission_repository.py | 54 +- .../repositories/organization_repository.py | 36 +- litellm/repositories/project_repository.py | 60 +- litellm/repositories/team_repository.py | 108 +- litellm/repositories/user_repository.py | 102 +- litellm/rerank_api/main.py | 18 +- litellm/rerank_api/rerank_utils.py | 8 +- .../responses/file_search/emulated_handler.py | 90 +- .../handler.py | 30 +- .../session_handler.py | 54 +- .../streaming_iterator.py | 12 +- litellm/responses/main.py | 500 +++-- .../responses/mcp/chat_completions_handler.py | 25 +- .../mcp/litellm_proxy_mcp_handler.py | 146 +- .../responses/mcp/mcp_streaming_iterator.py | 51 +- litellm/responses/mcp/request_context.py | 18 +- litellm/responses/sse_output_recovery.py | 14 +- litellm/responses/streaming_iterator.py | 209 +-- litellm/responses/utils.py | 119 +- litellm/router.py | 1007 +++++----- .../adaptive_router/adaptive_router.py | 6 +- .../router_strategy/adaptive_router/bandit.py | 15 +- .../adaptive_router/classifier.py | 3 +- .../router_strategy/adaptive_router/config.py | 4 +- .../router_strategy/adaptive_router/hooks.py | 24 +- .../adaptive_router/update_queue.py | 20 +- .../auto_router/auto_router.py | 34 +- .../auto_router/litellm_encoder.py | 4 +- .../router_strategy/base_routing_strategy.py | 31 +- litellm/router_strategy/budget_limiter.py | 116 +- .../complexity_router/__init__.py | 6 +- .../complexity_router/complexity_router.py | 6 +- .../evals/eval_complexity_router.py | 9 +- litellm/router_strategy/lar1_routing.py | 22 +- litellm/router_strategy/least_busy.py | 9 +- litellm/router_strategy/lowest_cost.py | 15 +- litellm/router_strategy/lowest_latency.py | 48 +- litellm/router_strategy/lowest_tpm_rpm.py | 23 +- litellm/router_strategy/lowest_tpm_rpm_v2.py | 86 +- .../quality_router/__init__.py | 2 +- .../router_strategy/quality_router/config.py | 14 +- .../quality_router/quality_router.py | 52 +- litellm/router_strategy/simple_shuffle.py | 6 +- litellm/router_strategy/tag_based_routing.py | 22 +- litellm/router_utils/batch_utils.py | 5 +- .../clientside_credential_handler.py | 4 +- litellm/router_utils/common_utils.py | 28 +- litellm/router_utils/cooldown_cache.py | 28 +- litellm/router_utils/cooldown_callbacks.py | 10 +- litellm/router_utils/cooldown_handlers.py | 37 +- .../router_utils/fallback_event_handlers.py | 22 +- litellm/router_utils/get_retry_from_policy.py | 8 +- litellm/router_utils/handle_error.py | 4 +- litellm/router_utils/health_state_cache.py | 10 +- .../router_utils/pattern_match_deployments.py | 23 +- .../deployment_affinity_check.py | 58 +- .../encrypted_content_affinity_check.py | 36 +- .../io_token_rate_limit_check.py | 56 +- .../pre_call_checks/model_rate_limit_check.py | 26 +- .../prompt_caching_deployment_check.py | 18 +- .../responses_api_deployment_check.py | 11 +- litellm/router_utils/prompt_caching_cache.py | 44 +- litellm/router_utils/search_api_router.py | 14 +- litellm/rust_bridge/messages.py | 6 +- litellm/rust_bridge/ocr.py | 6 +- litellm/rust_bridge/timeouts.py | 4 +- litellm/rust_bridge/transcription.py | 6 +- litellm/sandbox/main.py | 12 +- litellm/scheduler.py | 16 +- litellm/search/__init__.py | 2 +- litellm/search/cost_calculator.py | 8 +- litellm/search/main.py | 56 +- litellm/secret_managers/aws_secret_manager.py | 17 +- .../secret_managers/aws_secret_manager_v2.py | 72 +- .../secret_managers/base_secret_manager.py | 38 +- .../custom_secret_manager_loader.py | 3 +- .../cyberark_secret_manager.py | 30 +- .../get_azure_ad_token_provider.py | 36 +- litellm/secret_managers/google_kms.py | 3 +- .../secret_managers/google_secret_manager.py | 7 +- .../hashicorp_secret_manager.py | 50 +- litellm/secret_managers/main.py | 21 +- .../secret_managers/secret_manager_handler.py | 14 +- litellm/setup_wizard.py | 25 +- litellm/skills/__init__.py | 12 +- litellm/skills/main.py | 138 +- litellm/utils.py | 764 ++++---- litellm/vector_store_files/__init__.py | 20 +- litellm/vector_store_files/main.py | 148 +- litellm/vector_store_files/utils.py | 10 +- litellm/vector_stores/__init__.py | 2 +- litellm/vector_stores/main.py | 214 +-- litellm/vector_stores/utils.py | 6 +- .../vector_stores/vector_store_registry.py | 94 +- litellm/videos/__init__.py | 26 +- litellm/videos/main.py | 345 ++-- litellm/videos/utils.py | 10 +- ruff-strict-budget.json | 74 +- .../test_customer_endpoints.py | 3 +- type-discipline-budget.json | 8 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 +- 1457 files changed, 29765 insertions(+), 32318 deletions(-) diff --git a/litellm/_lazy_imports.py b/litellm/_lazy_imports.py index 9f7b26da7cd..8f9cd74f171 100644 --- a/litellm/_lazy_imports.py +++ b/litellm/_lazy_imports.py @@ -18,7 +18,7 @@ until they're actually needed. import importlib import sys from collections.abc import Callable -from typing import Any, Optional, cast +from typing import Any, cast # Import all the data structures that define what can be lazy-loaded # These are just lists of names and maps of where to find them @@ -78,7 +78,7 @@ def _get_utils_globals() -> dict: # They're separate from the main lazy import system because they have specific use cases # Lazy loader for default encoding - avoids importing heavy tiktoken library at startup -_default_encoding: Optional[Any] = None +_default_encoding: Any | None = None def _get_default_encoding() -> Any: @@ -100,7 +100,7 @@ def _get_default_encoding() -> Any: # Lazy loader for get_modified_max_tokens to avoid importing token_counter at module import time -_get_modified_max_tokens_func: Optional[Any] = None +_get_modified_max_tokens_func: Any | None = None def _get_modified_max_tokens() -> Any: @@ -124,7 +124,7 @@ def _get_modified_max_tokens() -> Any: # Lazy loader for token_counter to avoid importing token_counter module at module import time -_token_counter_new_func: Optional[Any] = None +_token_counter_new_func: Any | None = None def _get_token_counter_new() -> Any: @@ -154,7 +154,7 @@ def _get_token_counter_new() -> Any: # This registry maps attribute names (like "ModelResponse") to handler functions # It's built once the first time someone accesses a lazy-loaded attribute # Example: {"ModelResponse": _lazy_import_utils, "Cache": _lazy_import_caching, ...} -_LAZY_IMPORT_REGISTRY: Optional[dict[str, Callable[[str], Any]]] = None +_LAZY_IMPORT_REGISTRY: dict[str, Callable[[str], Any]] | None = None def _get_lazy_import_registry() -> dict[str, Callable[[str], Any]]: diff --git a/litellm/_logging.py b/litellm/_logging.py index 5f3c483869d..a41784e9170 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -4,11 +4,11 @@ import os import sys from datetime import datetime from logging import Formatter -from typing import Any, Dict, Optional +from typing import Any -from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads +from litellm.litellm_core_utils.secret_redaction import redact_string set_verbose = False @@ -86,7 +86,7 @@ handler.setLevel(numeric_level) handler.addFilter(_secret_filter) -def _try_parse_json_message(message: str) -> Optional[Dict[str, Any]]: +def _try_parse_json_message(message: str) -> dict[str, Any] | None: """ Try to parse a log message as JSON. Returns parsed dict if valid, else None. Handles messages that are entirely valid JSON (e.g. json.dumps output). @@ -103,7 +103,7 @@ def _try_parse_json_message(message: str) -> Optional[Dict[str, Any]]: return parsed -def _try_parse_embedded_python_dict(message: str) -> Optional[Dict[str, Any]]: +def _try_parse_embedded_python_dict(message: str) -> dict[str, Any] | None: """ Try to find and parse a Python dict repr (e.g. str(d) or repr(d)) embedded in the message. Handles patterns like: @@ -149,7 +149,7 @@ _STANDARD_RECORD_ATTRS = _get_standard_record_attrs() class JsonFormatter(Formatter): def __init__(self): - super(JsonFormatter, self).__init__() + super().__init__() def formatTime(self, record, datefmt=None): # Use datetime to format the timestamp in ISO 8601 format @@ -158,7 +158,7 @@ class JsonFormatter(Formatter): def format(self, record): message_str = record.getMessage() - json_record: Dict[str, Any] = { + json_record: dict[str, Any] = { "message": message_str, "level": record.levelname, "timestamp": self.formatTime(record), diff --git a/litellm/_redis.py b/litellm/_redis.py index 2728d45e0f3..693f9582705 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -13,7 +13,6 @@ import json # s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation import os from collections.abc import Callable -from typing import List, Optional, Union import redis # type: ignore import redis.asyncio as async_redis # type: ignore @@ -77,7 +76,7 @@ def _init_arg_names(cls: type) -> frozenset[str]: ) -def _get_redis_url_kwargs(client: Optional[type] = None) -> tuple[str, ...]: +def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]: """Connection kwargs that redis-py forwards from ``from_url`` down to the connection. ``from_url`` is declared as ``(cls, url, **kwargs)``, so introspecting it yields no @@ -161,7 +160,7 @@ def _redis_kwargs_from_environment(): def create_gcp_iam_redis_connect_func( service_account: str, - ssl_ca_certs: Optional[str] = None, + ssl_ca_certs: str | None = None, ) -> Callable: """ Creates a custom Redis connection function for GCP IAM authentication. @@ -204,9 +203,9 @@ def create_gcp_iam_redis_connect_func( def _build_azure_credential( - azure_client_id: Optional[str] = None, - azure_tenant_id: Optional[str] = None, - azure_client_secret: Optional[str] = None, + azure_client_id: str | None = None, + azure_tenant_id: str | None = None, + azure_client_secret: str | None = None, ): """ Build a long-lived Azure credential object. @@ -242,9 +241,9 @@ def _build_azure_credential( def _generate_azure_ad_redis_token( - azure_client_id: Optional[str] = None, - azure_tenant_id: Optional[str] = None, - azure_client_secret: Optional[str] = None, + azure_client_id: str | None = None, + azure_tenant_id: str | None = None, + azure_client_secret: str | None = None, ) -> str: """ One-shot helper that builds a credential and fetches a single Azure AD @@ -264,9 +263,9 @@ def _generate_azure_ad_redis_token( def create_azure_ad_redis_connect_func( - azure_client_id: Optional[str] = None, - azure_tenant_id: Optional[str] = None, - azure_client_secret: Optional[str] = None, + azure_client_id: str | None = None, + azure_tenant_id: str | None = None, + azure_client_secret: str | None = None, ) -> Callable: """ Creates a custom Redis connection function for Azure AD authentication. @@ -370,7 +369,7 @@ def _get_redis_client_logic(**env_overrides): **env_overrides, } - _startup_nodes: Optional[Union[str, list]] = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore + _startup_nodes: str | list | None = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore "REDIS_CLUSTER_NODES" ) @@ -381,21 +380,21 @@ def _get_redis_client_logic(**env_overrides): elif _startup_nodes is None: redis_kwargs.pop("startup_nodes", None) - _sentinel_nodes: Optional[Union[str, list]] = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore + _sentinel_nodes: str | list | None = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore "REDIS_SENTINEL_NODES" ) if _sentinel_nodes is not None and isinstance(_sentinel_nodes, str): redis_kwargs["sentinel_nodes"] = json.loads(_sentinel_nodes) - _sentinel_password: Optional[str] = redis_kwargs.get("sentinel_password", None) or get_secret_str( + _sentinel_password: str | None = redis_kwargs.get("sentinel_password", None) or get_secret_str( "REDIS_SENTINEL_PASSWORD" ) if _sentinel_password is not None: redis_kwargs["sentinel_password"] = _sentinel_password - _service_name: Optional[str] = redis_kwargs.get("service_name", None) or get_secret( # type: ignore + _service_name: str | None = redis_kwargs.get("service_name", None) or get_secret( # type: ignore "REDIS_SERVICE_NAME" ) @@ -466,9 +465,12 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("port", None) redis_kwargs.pop("db", None) redis_kwargs.pop("password", None) - elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None: - pass - elif "sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None: + elif ( + "startup_nodes" in redis_kwargs + and redis_kwargs["startup_nodes"] is not None + or "sentinel_nodes" in redis_kwargs + and redis_kwargs["sentinel_nodes"] is not None + ): pass elif "host" not in redis_kwargs or redis_kwargs["host"] is None: raise ValueError("Either 'host' or 'url' must be specified for redis.") @@ -478,7 +480,7 @@ def _get_redis_client_logic(**env_overrides): def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: - _redis_cluster_nodes_in_env: Optional[str] = get_secret("REDIS_CLUSTER_NODES") # type: ignore + _redis_cluster_nodes_in_env: str | None = get_secret("REDIS_CLUSTER_NODES") # type: ignore if _redis_cluster_nodes_in_env is not None: try: redis_kwargs["startup_nodes"] = json.loads(_redis_cluster_nodes_in_env) @@ -496,7 +498,7 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: if arg in args: cluster_kwargs[arg] = redis_kwargs[arg] - new_startup_nodes: List[ClusterNode] = [] + new_startup_nodes: list[ClusterNode] = [] for item in redis_kwargs["startup_nodes"]: new_startup_nodes.append(ClusterNode(**item)) @@ -588,9 +590,9 @@ def get_redis_client(**env_overrides): def get_redis_async_client( - connection_pool: Optional[async_redis.BlockingConnectionPool] = None, + connection_pool: async_redis.BlockingConnectionPool | None = None, **env_overrides, -) -> Union[async_redis.Redis, async_redis.RedisCluster]: +) -> async_redis.Redis | async_redis.RedisCluster: redis_kwargs = _get_redis_client_logic(**env_overrides) if "startup_nodes" in redis_kwargs: @@ -619,7 +621,7 @@ def get_redis_async_client( username=os.environ.get("REDIS_USERNAME") or None, ) - new_startup_nodes: List[ClusterNode] = [] + new_startup_nodes: list[ClusterNode] = [] for item in redis_kwargs["startup_nodes"]: new_startup_nodes.append(ClusterNode(**item)) @@ -649,9 +651,7 @@ def get_redis_async_client( if arg in args: url_kwargs[arg] = redis_kwargs[arg] else: - verbose_logger.debug( - "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(arg) - ) + verbose_logger.debug(f"REDIS: ignoring argument: {arg}. Not an allowed async_redis.Redis.from_url arg.") return async_redis.Redis.from_url(**url_kwargs) # Check for Redis Sentinel @@ -683,7 +683,7 @@ def get_redis_async_client( def get_redis_connection_pool( **env_overrides, -) -> Optional[async_redis.BlockingConnectionPool]: +) -> async_redis.BlockingConnectionPool | None: redis_kwargs = _get_redis_client_logic(**env_overrides) verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index b973e292a17..762e8bcd928 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,7 +1,7 @@ import asyncio import threading import time -from typing import Any, Dict, Optional, Tuple, Union +from typing import Any from redis.credentials import CredentialProvider # type: ignore[attr-defined] @@ -14,7 +14,7 @@ _GCP_IAM_TOKEN_TTL_SECONDS = 3300 # Module-level cache shared across all GCPIAMCredentialProvider instances for the # same service account, so multiple Redis connections on the same pod share one token. # Keyed by service_account → (token, expiry_monotonic_timestamp). -_token_cache: Dict[str, Tuple[str, float]] = {} +_token_cache: dict[str, tuple[str, float]] = {} _token_cache_lock = threading.Lock() @@ -95,11 +95,11 @@ class GCPIAMCredentialProvider(CredentialProvider): def __init__(self, gcp_service_account: str) -> None: self._gcp_service_account = gcp_service_account - def get_credentials(self) -> Tuple[str]: + def get_credentials(self) -> tuple[str]: token = _get_cached_gcp_iam_token(self._gcp_service_account) return (token,) - async def get_credentials_async(self) -> Tuple[str]: + async def get_credentials_async(self) -> tuple[str]: token = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account) return (token,) @@ -115,17 +115,17 @@ class AzureADCredentialProvider(CredentialProvider): fail authentication after the initial token expired (~1 hour TTL). """ - def __init__(self, credential: Any, username: Optional[str] = None) -> None: + def __init__(self, credential: Any, username: str | None = None) -> None: self._credential = credential self._username = username - def get_credentials(self) -> Union[Tuple[str], Tuple[str, str]]: + def get_credentials(self) -> tuple[str] | tuple[str, str]: token = self._credential.get_token(AZURE_REDIS_SCOPE).token if self._username: return (self._username, token) return (token,) - async def get_credentials_async(self) -> Union[Tuple[str], Tuple[str, str]]: + async def get_credentials_async(self) -> tuple[str] | tuple[str, str]: token_obj = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE) if self._username: return (self._username, token_obj.token) diff --git a/litellm/_service_logger.py b/litellm/_service_logger.py index b1bd0a3bba2..7eea82b4e74 100644 --- a/litellm/_service_logger.py +++ b/litellm/_service_logger.py @@ -1,6 +1,6 @@ import asyncio from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union import litellm from litellm._logging import verbose_logger @@ -24,7 +24,7 @@ else: UserAPIKeyAuth = Any -def _get_otel_v2_class() -> Optional[type]: +def _get_otel_v2_class() -> type | None: """Return the ``OpenTelemetryV2`` class, or ``None`` if the OTel SDK is absent. Imported lazily: ``litellm.integrations.otel.logger`` imports the OpenTelemetry @@ -54,7 +54,7 @@ class ServiceLogging(CustomLogger): if "prometheus_system" in litellm.service_callback: self.prometheusServicesLogger = PrometheusServicesLogger() - def _resolve_otel_service_logger(self, callback: Any) -> Optional[Any]: + def _resolve_otel_service_logger(self, callback: Any) -> Any | None: """Resolve the OTel logger (legacy or V2) to emit a service span on. Returns the logger instance whose ``async_service_*_hook`` should fire for @@ -88,9 +88,9 @@ class ServiceLogging(CustomLogger): service: ServiceTypes, duration: float, call_type: str, - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[float, datetime]] = None, + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: float | datetime | None = None, ): """ Handles both sync and async monitoring by checking for existing event loop. @@ -152,10 +152,10 @@ class ServiceLogging(CustomLogger): service: ServiceTypes, call_type: str, duration: float, - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[datetime, float]] = None, - event_metadata: Optional[dict] = None, + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: datetime | float | None = None, + event_metadata: dict | None = None, ): """ - For counting if the redis, postgres call is successful @@ -218,7 +218,6 @@ class ServiceLogging(CustomLogger): self.prometheusServicesLogger = PrometheusServicesLogger() elif self.prometheusServicesLogger is None: self.prometheusServicesLogger = self.prometheusServicesLogger() - return async def init_datadog_logger_if_none(self): """ @@ -230,8 +229,6 @@ class ServiceLogging(CustomLogger): if not hasattr(self, "dd_logger"): self.dd_logger: DataDogLogger = DataDogLogger() - return - async def init_otel_logger_if_none(self): """ initializes otel_logger if it is None or no attribute exists on ServiceLogging Object @@ -246,18 +243,17 @@ class ServiceLogging(CustomLogger): verbose_logger.warning( "ServiceLogger: open_telemetry_logger is None or not an instance of OpenTelemetry" ) - return async def async_service_failure_hook( self, service: ServiceTypes, duration: float, - error: Union[str, Exception], + error: str | Exception, call_type: str, - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[float, datetime]] = None, - event_metadata: Optional[dict] = None, + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: float | datetime | None = None, + event_metadata: dict | None = None, ): """ - For counting if the redis, postgres call is unsuccessful @@ -324,7 +320,7 @@ class ServiceLogging(CustomLogger): request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, + traceback_str: str | None = None, ): """ Hook to track failed litellm-service calls @@ -347,7 +343,7 @@ class ServiceLogging(CustomLogger): pass else: raise Exception( - "Duration={} is not a float or timedelta object. type={}".format(_duration, type(_duration)) + f"Duration={_duration} is not a float or timedelta object. type={type(_duration)}" ) # invalid _duration value # Batch polling callbacks (check_batch_cost) don't include call_type in kwargs. # Use .get() to avoid KeyError. diff --git a/litellm/a2a_protocol/card_resolver.py b/litellm/a2a_protocol/card_resolver.py index 412c7a0897d..e4cce56d0e4 100644 --- a/litellm/a2a_protocol/card_resolver.py +++ b/litellm/a2a_protocol/card_resolver.py @@ -4,7 +4,7 @@ Custom A2A Card Resolver for LiteLLM. Extends the A2A SDK's card resolver to support multiple well-known paths. """ -from typing import TYPE_CHECKING, Any, Dict +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.constants import LOCALHOST_URL_PATTERNS @@ -114,7 +114,7 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc] async def get_agent_card( self, relative_card_path: str | None = None, - http_kwargs: Dict[str, Any] | None = None, + http_kwargs: dict[str, Any] | None = None, ) -> "AgentCard": """ Fetch the agent card, trying multiple well-known paths. diff --git a/litellm/a2a_protocol/client.py b/litellm/a2a_protocol/client.py index a8ee24f0c09..8fbd4b8b81c 100644 --- a/litellm/a2a_protocol/client.py +++ b/litellm/a2a_protocol/client.py @@ -5,7 +5,7 @@ Provides a class-based interface for A2A agent invocation. """ from collections.abc import AsyncIterator -from typing import TYPE_CHECKING, Dict, Optional +from typing import TYPE_CHECKING from litellm.types.agents import LiteLLMSendMessageResponse @@ -51,7 +51,7 @@ class A2AClient: self, base_url: str, timeout: float = 60.0, - extra_headers: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, ): """ Initialize the A2A client wrapper. @@ -64,7 +64,7 @@ class A2AClient: self.base_url = base_url self.timeout = timeout self.extra_headers = extra_headers - self._a2a_client: Optional["A2AClientType"] = None + self._a2a_client: A2AClientType | None = None async def _get_client(self) -> "A2AClientType": """Get or create the underlying A2A client.""" diff --git a/litellm/a2a_protocol/cost_calculator.py b/litellm/a2a_protocol/cost_calculator.py index f3e84c5b84d..7e6c20a31e5 100644 --- a/litellm/a2a_protocol/cost_calculator.py +++ b/litellm/a2a_protocol/cost_calculator.py @@ -5,7 +5,7 @@ Supports dynamic cost parameters that allow platform owners to define custom costs per agent query or per token. """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( @@ -18,7 +18,7 @@ else: class A2ACostCalculator: @staticmethod def calculate_a2a_cost( - litellm_logging_obj: Optional[LitellmLoggingObject], + litellm_logging_obj: LitellmLoggingObject | None, ) -> float: """ Calculate the cost of an A2A send_message call. @@ -73,8 +73,8 @@ class A2ACostCalculator: @staticmethod def _calculate_token_based_cost( model_call_details: dict, - input_cost_per_token: Optional[float], - output_cost_per_token: Optional[float], + input_cost_per_token: float | None, + output_cost_per_token: float | None, ) -> float: """ Calculate cost based on token usage and per-token pricing. diff --git a/litellm/a2a_protocol/exception_mapping_utils.py b/litellm/a2a_protocol/exception_mapping_utils.py index 89b831351ab..4d24dd4f1d8 100644 --- a/litellm/a2a_protocol/exception_mapping_utils.py +++ b/litellm/a2a_protocol/exception_mapping_utils.py @@ -4,7 +4,7 @@ A2A Protocol Exception Mapping Utils. Maps A2A SDK exceptions to LiteLLM A2A exception types. """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.a2a_protocol.card_resolver import ( @@ -57,7 +57,7 @@ class A2AExceptionCheckers: return any(pattern in error_str_lower for pattern in CONNECTION_ERROR_PATTERNS) @staticmethod - def is_localhost_url(url: Optional[str]) -> bool: + def is_localhost_url(url: str | None) -> bool: """ Check if a URL is a localhost/internal URL. @@ -96,9 +96,9 @@ class A2AExceptionCheckers: def map_a2a_exception( original_exception: Exception, - card_url: Optional[str] = None, - api_base: Optional[str] = None, - model: Optional[str] = None, + card_url: str | None = None, + api_base: str | None = None, + model: str | None = None, ) -> Exception: """ Map an A2A SDK exception to a LiteLLM A2A exception type. diff --git a/litellm/a2a_protocol/exceptions.py b/litellm/a2a_protocol/exceptions.py index b672971e727..2542cbc67b0 100644 --- a/litellm/a2a_protocol/exceptions.py +++ b/litellm/a2a_protocol/exceptions.py @@ -4,8 +4,6 @@ A2A Protocol Exceptions. Custom exception types for A2A protocol operations, following LiteLLM's exception pattern. """ -from typing import Optional - import httpx @@ -21,11 +19,11 @@ class A2AError(Exception): message: str, status_code: int = 500, llm_provider: str = "a2a_agent", - model: Optional[str] = None, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + model: str | None = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = status_code self.message = f"litellm.A2AError: {message}" @@ -65,12 +63,12 @@ class A2AConnectionError(A2AError): def __init__( self, message: str, - url: Optional[str] = None, - model: Optional[str] = None, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + url: str | None = None, + model: str | None = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.url = url super().__init__( @@ -98,10 +96,10 @@ class A2AAgentCardError(A2AError): def __init__( self, message: str, - url: Optional[str] = None, - model: Optional[str] = None, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, + url: str | None = None, + model: str | None = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, ): self.url = url super().__init__( @@ -132,8 +130,8 @@ class A2ALocalhostURLError(A2AConnectionError): self, localhost_url: str, base_url: str, - original_error: Optional[Exception] = None, - model: Optional[str] = None, + original_error: Exception | None = None, + model: str | None = None, ): self.localhost_url = localhost_url self.base_url = base_url diff --git a/litellm/a2a_protocol/litellm_completion_bridge/__init__.py b/litellm/a2a_protocol/litellm_completion_bridge/__init__.py index 6c9df0ee285..a81f5304f7c 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/__init__.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/__init__.py @@ -16,8 +16,8 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import ( ) __all__ = [ - "A2ACompletionBridgeTransformation", "A2ACompletionBridgeHandler", + "A2ACompletionBridgeTransformation", "handle_a2a_completion", "handle_a2a_completion_streaming", ] diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 94d4bb5d899..1d46e5c700f 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -11,7 +11,7 @@ A2A Streaming Events (in order): """ from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any import litellm from litellm._logging import verbose_logger @@ -47,13 +47,13 @@ class A2ACompletionBridgeHandler: @staticmethod async def handle_non_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - api_base: Optional[str] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, + params: dict[str, Any], + litellm_params: dict[str, Any], + api_base: str | None = None, + agent_extra_headers: dict[str, str] | None = None, *, _skip_a2a_provider_routing: bool = False, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Handle non-streaming A2A request via litellm.acompletion. @@ -106,7 +106,7 @@ class A2ACompletionBridgeHandler: verbose_logger.info(f"A2A completion bridge: model={full_model}, api_base={api_base}") # Build completion params dict - completion_params: Dict[str, Any] = { + completion_params: dict[str, Any] = { "model": full_model, "messages": openai_messages, "api_base": api_base, @@ -150,13 +150,13 @@ class A2ACompletionBridgeHandler: @staticmethod async def handle_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - api_base: Optional[str] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, + params: dict[str, Any], + litellm_params: dict[str, Any], + api_base: str | None = None, + agent_extra_headers: dict[str, str] | None = None, *, _skip_a2a_provider_routing: bool = False, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """ Handle streaming A2A request via litellm.acompletion with stream=True. @@ -224,7 +224,7 @@ class A2ACompletionBridgeHandler: verbose_logger.info(f"A2A completion bridge streaming: model={full_model}, api_base={api_base}") # Build completion params dict - completion_params: Dict[str, Any] = { + completion_params: dict[str, Any] = { "model": full_model, "messages": openai_messages, "api_base": api_base, @@ -306,11 +306,11 @@ class A2ACompletionBridgeHandler: # Convenience functions that delegate to the class methods async def handle_a2a_completion( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - api_base: Optional[str] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, -) -> Dict[str, Any]: + params: dict[str, Any], + litellm_params: dict[str, Any], + api_base: str | None = None, + agent_extra_headers: dict[str, str] | None = None, +) -> dict[str, Any]: """Convenience function for non-streaming A2A completion.""" return await A2ACompletionBridgeHandler.handle_non_streaming( request_id=request_id, @@ -323,11 +323,11 @@ async def handle_a2a_completion( async def handle_a2a_completion_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - api_base: Optional[str] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, -) -> AsyncIterator[Dict[str, Any]]: + params: dict[str, Any], + litellm_params: dict[str, Any], + api_base: str | None = None, + agent_extra_headers: dict[str, str] | None = None, +) -> AsyncIterator[dict[str, Any]]: """Convenience function for streaming A2A completion.""" async for chunk in A2ACompletionBridgeHandler.handle_streaming( request_id=request_id, diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index b32963dd6fb..a63221b4a77 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -18,7 +18,7 @@ A2A Streaming Events: """ from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from uuid import uuid4 from litellm._logging import verbose_logger @@ -30,7 +30,7 @@ class A2AStreamingContext: Tracks task_id, context_id, and message accumulation. """ - def __init__(self, request_id: str, input_message: Dict[str, Any]): + def __init__(self, request_id: str, input_message: dict[str, Any]): self.request_id = request_id self.task_id = str(uuid4()) self.context_id = str(uuid4()) @@ -46,9 +46,9 @@ class A2ACompletionBridgeTransformation: """ @staticmethod - def _extract_text_from_a2a_parts(parts: List[Dict[str, Any]]) -> str: + def _extract_text_from_a2a_parts(parts: list[dict[str, Any]]) -> str: """Extract text from A2A parts (with or without explicit ``kind``).""" - content_parts: List[str] = [] + content_parts: list[str] = [] for part in parts: if not isinstance(part, dict): continue @@ -62,16 +62,16 @@ class A2ACompletionBridgeTransformation: @staticmethod def get_forward_metadata( - a2a_message: Dict[str, Any], - params: Optional[Dict[str, Any]] = None, - ) -> Optional[Dict[str, Any]]: + a2a_message: dict[str, Any], + params: dict[str, Any] | None = None, + ) -> dict[str, Any] | None: """ Merge A2A metadata from MessageSendParams and the message for downstream providers. Forwarded once on the LangGraph run payload (``metadata``), not duplicated on each input message — see ``apply_forward_metadata_to_completion_params``. """ - merged: Dict[str, Any] = {} + merged: dict[str, Any] = {} if params and isinstance(params.get("metadata"), dict): merged.update(params["metadata"]) message_metadata = a2a_message.get("metadata") @@ -81,9 +81,9 @@ class A2ACompletionBridgeTransformation: @staticmethod def apply_forward_metadata_to_completion_params( - completion_params: Dict[str, Any], - a2a_message: Dict[str, Any], - params: Optional[Dict[str, Any]] = None, + completion_params: dict[str, Any], + a2a_message: dict[str, Any], + params: dict[str, Any] | None = None, ) -> None: """ Attach A2A metadata to completion kwargs for provider bridges (e.g. LangGraph). @@ -104,8 +104,8 @@ class A2ACompletionBridgeTransformation: # ``extra_body.metadata`` so the configured keys remain authoritative # and an A2A caller cannot overwrite server-set run metadata. existing_metadata = extra_body.get("metadata") - existing_dict: Dict[str, Any] = existing_metadata if isinstance(existing_metadata, dict) else {} - merged_metadata: Dict[str, Any] = {**forward_metadata, **existing_dict} + existing_dict: dict[str, Any] = existing_metadata if isinstance(existing_metadata, dict) else {} + merged_metadata: dict[str, Any] = {**forward_metadata, **existing_dict} extra_body = {**extra_body, "metadata": merged_metadata} completion_params["extra_body"] = extra_body @@ -113,8 +113,8 @@ class A2ACompletionBridgeTransformation: @staticmethod def a2a_message_to_openai_messages( - a2a_message: Dict[str, Any], - ) -> List[Dict[str, Any]]: + a2a_message: dict[str, Any], + ) -> list[dict[str, Any]]: """ Transform an A2A message to OpenAI message format. @@ -143,7 +143,7 @@ class A2ACompletionBridgeTransformation: # Do not attach A2A message.metadata here — the completion bridge forwards it # once at run level via extra_body.metadata (LangGraph POST /runs/wait shape). - openai_message: Dict[str, Any] = {"role": openai_role, "content": content} + openai_message: dict[str, Any] = {"role": openai_role, "content": content} verbose_logger.debug(f"A2A -> OpenAI transform: role={role} -> {openai_role}, content_length={len(content)}") @@ -152,8 +152,8 @@ class A2ACompletionBridgeTransformation: @staticmethod def openai_response_to_a2a_response( response: Any, - request_id: Optional[str] = None, - ) -> Dict[str, Any]: + request_id: str | None = None, + ) -> dict[str, Any]: """ Transform a LiteLLM ModelResponse to A2A SendMessageResponse format. @@ -198,7 +198,7 @@ class A2ACompletionBridgeTransformation: @staticmethod def create_task_event( ctx: A2AStreamingContext, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Create the initial task event with status 'submitted'. @@ -232,8 +232,8 @@ class A2ACompletionBridgeTransformation: ctx: A2AStreamingContext, state: str, final: bool = False, - message_text: Optional[str] = None, - ) -> Dict[str, Any]: + message_text: str | None = None, + ) -> dict[str, Any]: """ Create a status update event. @@ -243,7 +243,7 @@ class A2ACompletionBridgeTransformation: final: Whether this is the final event message_text: Optional message text for 'working' status """ - status: Dict[str, Any] = { + status: dict[str, Any] = { "state": state, "timestamp": A2ACompletionBridgeTransformation._get_timestamp(), } @@ -275,7 +275,7 @@ class A2ACompletionBridgeTransformation: def create_artifact_update_event( ctx: A2AStreamingContext, text: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Create an artifact update event with content. diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index a2985f2c325..52d35a988c6 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -16,9 +16,7 @@ from collections.abc import AsyncIterator, Coroutine from typing import ( TYPE_CHECKING, Any, - Dict, Optional, - Union, cast, ) @@ -86,7 +84,7 @@ A2ACardResolver = LiteLLMA2ACardResolver def _set_usage_on_logging_obj( - kwargs: Dict[str, Any], + kwargs: dict[str, Any], prompt_tokens: int, completion_tokens: int, ) -> None: @@ -109,7 +107,7 @@ def _set_usage_on_logging_obj( def _set_agent_id_on_logging_obj( - kwargs: Dict[str, Any], + kwargs: dict[str, Any], agent_id: str | None, ) -> None: """ @@ -155,7 +153,7 @@ def _set_litellm_params_on_logging_obj( logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params} -def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str: +def _get_a2a_model_info(a2a_client: Any, kwargs: dict[str, Any]) -> str: """ Extract agent info and set model/custom_llm_provider for cost tracking. @@ -198,8 +196,8 @@ async def _send_message_via_completion_bridge( request: "SendMessageRequest", custom_llm_provider: str, api_base: str | None, - litellm_params: Dict[str, Any], - agent_extra_headers: Dict[str, str] | None = None, + litellm_params: dict[str, Any], + agent_extra_headers: dict[str, str] | None = None, ) -> LiteLLMSendMessageResponse: """ Route a send_message through the LiteLLM completion bridge (e.g. LangGraph, Bedrock AgentCore). @@ -369,9 +367,9 @@ async def asend_message( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendMessageRequest"] = None, api_base: str | None = None, - litellm_params: Dict[str, Any] | None = None, + litellm_params: dict[str, Any] | None = None, agent_id: str | None = None, - agent_extra_headers: Dict[str, str] | None = None, + agent_extra_headers: dict[str, str] | None = None, **kwargs: Any, ) -> LiteLLMSendMessageResponse: """ @@ -452,7 +450,7 @@ async def asend_message( if api_base is None: raise ValueError("Either a2a_client or api_base is required for standard A2A flow") trace_id = trace_id or str(uuid.uuid4()) - extra_headers: Dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id} + extra_headers: dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id} if agent_id: extra_headers["X-LiteLLM-Agent-Id"] = agent_id # Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones) @@ -517,7 +515,7 @@ def send_message( a2a_client: "A2AClientType", request: "SendMessageRequest", **kwargs: Any, -) -> Union[LiteLLMSendMessageResponse, Coroutine[Any, Any, LiteLLMSendMessageResponse]]: +) -> LiteLLMSendMessageResponse | Coroutine[Any, Any, LiteLLMSendMessageResponse]: """ Sync: Send a message to an A2A agent. @@ -546,9 +544,9 @@ def _build_streaming_logging_obj( request: "SendStreamingMessageRequest", agent_name: str, agent_id: str | None, - litellm_params: Dict[str, Any] | None, - metadata: Dict[str, Any] | None, - proxy_server_request: Dict[str, Any] | None, + litellm_params: dict[str, Any] | None, + metadata: dict[str, Any] | None, + proxy_server_request: dict[str, Any] | None, ) -> Logging: """Build logging object for streaming A2A requests.""" start_time = datetime.datetime.now() @@ -589,11 +587,11 @@ async def asend_message_streaming( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendStreamingMessageRequest"] = None, api_base: str | None = None, - litellm_params: Dict[str, Any] | None = None, + litellm_params: dict[str, Any] | None = None, agent_id: str | None = None, - metadata: Dict[str, Any] | None = None, - proxy_server_request: Dict[str, Any] | None = None, - agent_extra_headers: Dict[str, str] | None = None, + metadata: dict[str, Any] | None = None, + proxy_server_request: dict[str, Any] | None = None, + agent_extra_headers: dict[str, str] | None = None, **kwargs: object, ) -> AsyncIterator[Any]: """ @@ -727,7 +725,7 @@ async def asend_message_streaming( async def create_a2a_client( base_url: str, timeout: float = DEFAULT_A2A_AGENT_TIMEOUT, - extra_headers: Dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, streaming: bool = False, ) -> "A2AClientType": """ @@ -808,7 +806,7 @@ async def create_a2a_client( async def aget_agent_card( base_url: str, timeout: float = DEFAULT_A2A_AGENT_TIMEOUT, - extra_headers: Dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, ) -> "AgentCard": """ Fetch the agent card from an A2A agent. diff --git a/litellm/a2a_protocol/providers/__init__.py b/litellm/a2a_protocol/providers/__init__.py index a21fa5f8f5e..8f16fcf15c8 100644 --- a/litellm/a2a_protocol/providers/__init__.py +++ b/litellm/a2a_protocol/providers/__init__.py @@ -7,4 +7,4 @@ This module contains provider-specific implementations for the A2A protocol. from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager -__all__ = ["BaseA2AProviderConfig", "A2AProviderConfigManager"] +__all__ = ["A2AProviderConfigManager", "BaseA2AProviderConfig"] diff --git a/litellm/a2a_protocol/providers/base.py b/litellm/a2a_protocol/providers/base.py index f546bb0c501..5a5eff8cf35 100644 --- a/litellm/a2a_protocol/providers/base.py +++ b/litellm/a2a_protocol/providers/base.py @@ -4,7 +4,7 @@ Base configuration for A2A protocol providers. from abc import ABC, abstractmethod from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any class BaseA2AProviderConfig(ABC): @@ -19,10 +19,10 @@ class BaseA2AProviderConfig(ABC): async def handle_non_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Handle non-streaming A2A request. @@ -35,16 +35,15 @@ class BaseA2AProviderConfig(ABC): Returns: A2A SendMessageResponse dict """ - pass @abstractmethod async def handle_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """ 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 d22379916e9..9390b0a94e2 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, Dict, Optional +from typing import Any 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], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Handle non-streaming request to AgentCore A2A agent.""" litellm_params = 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], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """Handle streaming request to AgentCore A2A agent.""" litellm_params = 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 5cbbdc62f4d..56f5f806e7b 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 -from typing import Any, Dict, Optional, cast +from typing import Any, cast from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( @@ -28,10 +28,10 @@ class BedrockAgentCoreA2AHandler: @staticmethod async def handle_non_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + params: dict[str, Any], + litellm_params: dict[str, Any], + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ Handle non-streaming A2A request to AgentCore. @@ -74,10 +74,10 @@ class BedrockAgentCoreA2AHandler: @staticmethod async def handle_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> AsyncIterator[Dict[str, Any]]: + params: dict[str, Any], + litellm_params: dict[str, Any], + agent_extra_headers: dict[str, str] | None = None, + ) -> AsyncIterator[dict[str, Any]]: """ 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 3aef69ca775..f9343d2d3b4 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py @@ -7,7 +7,7 @@ and signs requests via AmazonAgentCoreConfig (SigV4 or JWT). import json from collections.abc import AsyncIterator, Mapping -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm._logging import verbose_logger from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig @@ -29,15 +29,15 @@ _RESERVED_EXACT_HEADERS = frozenset( "host", } ) -_RESERVED_PREFIX_HEADERS: Tuple[str, ...] = ( +_RESERVED_PREFIX_HEADERS: tuple[str, ...] = ( "x-amzn-bedrock-agentcore-runtime-", "x-amz-", ) def _filter_reserved_headers( - agent_extra_headers: Optional[Mapping[str, str]], -) -> Optional[Dict[str, str]]: + agent_extra_headers: Mapping[str, str] | None, +) -> dict[str, str] | None: """ Strip reserved AWS / AgentCore headers from caller-supplied ``agent_extra_headers`` before they are merged into the signed request. @@ -47,7 +47,7 @@ def _filter_reserved_headers( if not agent_extra_headers: return None - filtered: Dict[str, str] = {} + filtered: dict[str, str] = {} dropped: list = [] for k, v in agent_extra_headers.items(): k_lower = k.lower() @@ -77,12 +77,12 @@ class BedrockAgentCoreA2ATransformation: @staticmethod def get_url_and_signed_request( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], + params: dict[str, Any], + litellm_params: dict[str, Any], method: str = "message/send", stream: bool = False, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Tuple[str, dict, bytes]: + agent_extra_headers: dict[str, str] | None = None, + ) -> tuple[str, dict, bytes]: """ Build the AgentCore URL, construct a JSON-RPC envelope, and sign the request. @@ -170,7 +170,7 @@ class BedrockAgentCoreA2ATransformation: return url, signed_headers, signed_body @staticmethod - async def parse_sse_events(response: Any) -> AsyncIterator[Dict[str, Any]]: + async def parse_sse_events(response: Any) -> AsyncIterator[dict[str, Any]]: """ Parse SSE events from an httpx streaming response. diff --git a/litellm/a2a_protocol/providers/config_manager.py b/litellm/a2a_protocol/providers/config_manager.py index a421afec184..2eab2adb1ba 100644 --- a/litellm/a2a_protocol/providers/config_manager.py +++ b/litellm/a2a_protocol/providers/config_manager.py @@ -4,8 +4,6 @@ A2A Provider Config Manager. Manages provider-specific configurations for A2A protocol. """ -from typing import Optional - from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig @@ -18,9 +16,9 @@ class A2AProviderConfigManager: @staticmethod def get_provider_config( - custom_llm_provider: Optional[str], - model: Optional[str] = None, - ) -> Optional[BaseA2AProviderConfig]: + custom_llm_provider: str | None, + model: str | None = None, + ) -> BaseA2AProviderConfig | None: """ Get the provider configuration for a given custom_llm_provider. diff --git a/litellm/a2a_protocol/providers/langflow/config.py b/litellm/a2a_protocol/providers/langflow/config.py index 179b201cb44..54d403f88c0 100644 --- a/litellm/a2a_protocol/providers/langflow/config.py +++ b/litellm/a2a_protocol/providers/langflow/config.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any from litellm.a2a_protocol.litellm_completion_bridge.handler import ( A2A_USER_API_KEY_HASH_PARAM, @@ -16,10 +16,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( @@ -39,10 +39,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py index 8e9cd6fc87e..078e0633e04 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py @@ -13,4 +13,4 @@ from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( PydanticAITransformation, ) -__all__ = ["PydanticAIHandler", "PydanticAITransformation", "PydanticAIProviderConfig"] +__all__ = ["PydanticAIHandler", "PydanticAIProviderConfig", "PydanticAITransformation"] diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py index dabf1e387a3..b7546e1a2a1 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py @@ -3,7 +3,7 @@ Pydantic AI provider configuration. """ from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.pydantic_ai_agents.handler import PydanticAIHandler @@ -20,10 +20,10 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs: Any, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Handle non-streaming request to Pydantic AI agent.""" if api_base is None: raise ValueError("api_base is required for PydanticAIProviderConfig") @@ -38,10 +38,10 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """Handle streaming request with fake streaming.""" if not api_base: raise ValueError("api_base is required for Pydantic AI agents") diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py index d9b060c2640..86cb2d47ad3 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py @@ -6,7 +6,7 @@ This handler provides fake streaming by converting non-streaming responses into """ from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any 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], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, timeout: float = 60.0, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ 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], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, timeout: float = 60.0, chunk_size: int = 50, delay_ms: int = 10, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> AsyncIterator[Dict[str, Any]]: + agent_extra_headers: dict[str, str] | None = None, + ) -> AsyncIterator[dict[str, Any]]: """ 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 632428d5a9c..37127f2fcab 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 -from typing import Any, Dict, Optional, cast +from typing import Any, cast from uuid import uuid4 from litellm._logging import verbose_logger @@ -49,7 +49,7 @@ class PydanticAITransformation: return obj @staticmethod - def _params_to_dict(params: Any) -> Dict[str, Any]: + def _params_to_dict(params: Any) -> dict[str, Any]: """ Convert params to a dict, handling Pydantic models. @@ -79,8 +79,8 @@ class PydanticAITransformation: request_id: str, max_attempts: int = 30, poll_interval: float = 0.5, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ Poll for task completion using tasks/get method. @@ -135,8 +135,8 @@ class PydanticAITransformation: request_id: str, params: Any, timeout: float = 60.0, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ Send a request to Pydantic AI agent and return the raw task response. @@ -219,8 +219,8 @@ class PydanticAITransformation: request_id: str, params: Any, timeout: float = 60.0, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ Send a non-streaming A2A request to Pydantic AI agent and wait for completion. @@ -255,8 +255,8 @@ class PydanticAITransformation: request_id: str, params: Any, timeout: float = 60.0, - agent_extra_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, Any]: + agent_extra_headers: dict[str, str] | None = None, + ) -> dict[str, Any]: """ Send a request to Pydantic AI agent and return the raw task response. @@ -282,9 +282,9 @@ class PydanticAITransformation: @staticmethod def _transform_to_a2a_response( - response_data: Dict[str, Any], + response_data: dict[str, Any], request_id: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform Pydantic AI task response to standard A2A non-streaming format. @@ -328,7 +328,7 @@ class PydanticAITransformation: } @staticmethod - def _extract_response_text(response_data: Dict[str, Any]) -> tuple[str, str, list]: + def _extract_response_text(response_data: dict[str, Any]) -> tuple[str, str, list]: """ Extract response text from completed task response. @@ -383,11 +383,11 @@ class PydanticAITransformation: @staticmethod async def fake_streaming_from_response( - response_data: Dict[str, Any], + response_data: dict[str, Any], request_id: str, chunk_size: int = 50, delay_ms: int = 10, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """ Convert a non-streaming A2A response into fake streaming chunks. diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index 093b1f0c7d4..7c526b89c35 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -3,7 +3,7 @@ A2A provider configuration for IBM watsonx Orchestrate (WXO). """ from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import ( @@ -17,10 +17,10 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs: Any, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Handle a non-streaming A2A request via WXO runs API.""" litellm_params = kwargs.get("litellm_params") if not litellm_params: @@ -37,10 +37,10 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: Dict[str, Any], - api_base: Optional[str] = None, + params: dict[str, Any], + api_base: str | None = None, **kwargs: Any, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """Handle a streaming A2A request via WXO streaming runs API.""" litellm_params = kwargs.get("litellm_params") if not litellm_params: diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index 31a042907f2..efb2b38b912 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, Dict, NamedTuple, Optional, Tuple, cast +from typing import Any, NamedTuple, cast import httpx @@ -25,7 +25,7 @@ _IBM_CLOUD_IAM_URL = "https://iam.cloud.ibm.com/identity/token" _POLL_INTERVAL_S = 2.0 _MAX_POLL_ATTEMPTS = 90 _TOKEN_CACHE_TTL_BUFFER_S = 60 -_token_cache: Dict[str, Tuple[str, float]] = {} +_token_cache: dict[str, tuple[str, float]] = {} class WXORequestParams(NamedTuple): @@ -33,9 +33,9 @@ class WXORequestParams(NamedTuple): instance_id: str wxo_agent_id: str api_key: str - username: Optional[str] + username: str | None auth_mode: str - thread_id: Optional[str] + thread_id: str | None class WatsonxOrchestrateHandler: @@ -51,13 +51,13 @@ class WatsonxOrchestrateHandler: auth_mode: str, cp4d_host: str, api_key: str, - username: Optional[str], + username: str | None, ) -> str: material = f"{auth_mode}:{cp4d_host}:{username or ''}:{api_key}" return hashlib.sha256(material.encode()).hexdigest() @staticmethod - def _cp4d_token_ttl_seconds(expiration: Any, now_wall: Optional[float] = None) -> int: + def _cp4d_token_ttl_seconds(expiration: Any, now_wall: float | None = None) -> int: # CP4D returns expiration as absolute Unix epoch seconds, not a duration. expires_at = int(expiration) wall = now_wall if now_wall is not None else time.time() @@ -68,8 +68,8 @@ class WatsonxOrchestrateHandler: cp4d_host: str, auth_mode: str, api_key: str, - username: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, + username: str | None = None, + client: AsyncHTTPHandler | None = None, ) -> str: cache_key = WatsonxOrchestrateHandler._token_cache_key(auth_mode, cp4d_host, api_key, username) now = time.monotonic() @@ -122,18 +122,18 @@ class WatsonxOrchestrateHandler: async def _poll_run( base_url: str, run_id: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], client: AsyncHTTPHandler, max_attempts: int = _MAX_POLL_ATTEMPTS, interval_s: float = _POLL_INTERVAL_S, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: url = f"{base_url}/v1/orchestrate/runs/{run_id}" for attempt in range(max_attempts): await asyncio.sleep(interval_s) response = await client.get(url, headers=auth_headers) response.raise_for_status() - result: Dict[str, Any] = response.json() + result: dict[str, Any] = response.json() status = result.get("status", "") verbose_logger.debug(f"WXO: Poll {attempt + 1}/{max_attempts} run='{run_id}' status='{status}'") if status in WatsonxOrchestrateTransformation.TERMINAL_STATES: @@ -145,11 +145,11 @@ class WatsonxOrchestrateHandler: @staticmethod async def _get_successful_run_data( - run_data: Dict[str, Any], + run_data: dict[str, Any], base_url: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], client: AsyncHTTPHandler, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: status = run_data.get("status", "") if status not in WatsonxOrchestrateTransformation.TERMINAL_STATES: run_id = run_data.get("run_id") or run_data.get("id") or "" @@ -187,7 +187,7 @@ class WatsonxOrchestrateHandler: return accumulated_text @staticmethod - def _extract_litellm_params(litellm_params: Dict[str, Any]) -> WXORequestParams: + def _extract_litellm_params(litellm_params: dict[str, Any]) -> WXORequestParams: cp4d_host = litellm_params.get("cp4d_host") or "" instance_id = litellm_params.get("instance_id") or "" wxo_agent_id = litellm_params.get("wxo_agent_id") or "" @@ -215,9 +215,9 @@ class WatsonxOrchestrateHandler: @staticmethod async def handle_non_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], - ) -> Dict[str, Any]: + params: dict[str, Any], + litellm_params: dict[str, Any], + ) -> dict[str, Any]: wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) client = WatsonxOrchestrateHandler._http_client(timeout=90.0) @@ -246,7 +246,7 @@ class WatsonxOrchestrateHandler: headers=auth_headers, ) run_response.raise_for_status() - run_data: Dict[str, Any] = run_response.json() + run_data: dict[str, Any] = run_response.json() run_data = await WatsonxOrchestrateHandler._get_successful_run_data( run_data=run_data, @@ -261,11 +261,11 @@ class WatsonxOrchestrateHandler: @staticmethod async def handle_streaming( request_id: str, - params: Dict[str, Any], - litellm_params: Dict[str, Any], + params: dict[str, Any], + litellm_params: dict[str, Any], chunk_size: int = 50, delay_ms: int = 10, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) client = WatsonxOrchestrateHandler._http_client(timeout=120.0) diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py index e945c721f52..ab7b8abb3ba 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py @@ -9,7 +9,7 @@ WXO uses a REST API (not A2A/JSON-RPC) with an async-poll execution model: import asyncio from collections.abc import AsyncIterator -from typing import Any, Dict, Optional +from typing import Any from uuid import uuid4 from litellm._logging import verbose_logger @@ -29,7 +29,7 @@ class WatsonxOrchestrateTransformation: return f"{cp4d_host.rstrip('/')}/orchestrate/cpd/instances/{instance_id}" @staticmethod - def extract_text_from_a2a_params(params: Dict[str, Any]) -> str: + def extract_text_from_a2a_params(params: dict[str, Any]) -> str: """ Extract user message text from A2A MessageSendParams. @@ -50,10 +50,10 @@ class WatsonxOrchestrateTransformation: def build_wxo_run_body( wxo_agent_id: str, text: str, - thread_id: Optional[str] = None, - ) -> Dict[str, Any]: + thread_id: str | None = None, + ) -> dict[str, Any]: """Build the WXO POST /v1/orchestrate/runs request body.""" - body: Dict[str, Any] = { + body: dict[str, Any] = { "agent_id": wxo_agent_id, "message": { "role": "user", @@ -103,7 +103,7 @@ class WatsonxOrchestrateTransformation: return "" @staticmethod - def extract_text_from_a2a_message_response(a2a_response: Dict[str, Any]) -> str: + def extract_text_from_a2a_message_response(a2a_response: dict[str, Any]) -> str: result = a2a_response.get("result") if not isinstance(result, dict): verbose_logger.warning("WXO: A2A response missing result object") @@ -119,7 +119,7 @@ class WatsonxOrchestrateTransformation: return "" @staticmethod - def build_a2a_message_response(request_id: str, text: str) -> Dict[str, Any]: + def build_a2a_message_response(request_id: str, text: str) -> dict[str, Any]: """ Build a standard A2A non-streaming SendMessageResponse (kind=message). """ @@ -140,7 +140,7 @@ class WatsonxOrchestrateTransformation: request_id: str, chunk_size: int = 50, delay_ms: int = 10, - ) -> AsyncIterator[Dict[str, Any]]: + ) -> AsyncIterator[dict[str, Any]]: """ Emit standard A2A streaming events from a completed text response. diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index 2885d9b2fb0..79056ca336f 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support. import asyncio from collections.abc import AsyncIterator from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -38,9 +38,9 @@ class A2AStreamingIterator: self.start_time = datetime.now() # Collect chunks for token counting - self.chunks: List[Any] = [] - self.collected_text_parts: List[str] = [] - self.final_chunk: Optional[Any] = None + self.chunks: list[Any] = [] + self.collected_text_parts: list[str] = [] + self.final_chunk: Any | None = None def __aiter__(self): return self @@ -146,9 +146,9 @@ class A2AStreamingIterator: except Exception as e: verbose_logger.debug(f"Error in A2A streaming completion handler: {e}") - def _build_logging_result(self, usage: litellm.Usage) -> Dict[str, Any]: + def _build_logging_result(self, usage: litellm.Usage) -> dict[str, Any]: """Build a result dict for logging.""" - result: Dict[str, Any] = { + result: dict[str, Any] = { "id": getattr(self.request, "id", "unknown"), "jsonrpc": "2.0", "usage": (usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)), diff --git a/litellm/a2a_protocol/utils.py b/litellm/a2a_protocol/utils.py index ce5a168c3ac..d6a45e39a02 100644 --- a/litellm/a2a_protocol/utils.py +++ b/litellm/a2a_protocol/utils.py @@ -2,7 +2,7 @@ Utility functions for A2A protocol. """ -from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -34,7 +34,7 @@ class A2ARequestUtils: else: parts = getattr(message, "parts", []) or [] - text_parts: List[str] = [] + text_parts: list[str] = [] for part in parts: if isinstance(part, dict): if part.get("kind") == "text": @@ -46,7 +46,7 @@ class A2ARequestUtils: return " ".join(text_parts) @staticmethod - def extract_text_from_response(response_dict: Dict[str, Any]) -> str: + def extract_text_from_response(response_dict: dict[str, Any]) -> str: """ Extract text content from A2A response result. @@ -71,7 +71,7 @@ class A2ARequestUtils: @staticmethod def get_input_message_from_request( - request: "Union[SendMessageRequest, SendStreamingMessageRequest]", + request: "SendMessageRequest | SendStreamingMessageRequest", ) -> Any: """ Extract the input message from an A2A request. @@ -108,9 +108,9 @@ class A2ARequestUtils: @staticmethod def calculate_usage_from_request_response( - request: "Union[SendMessageRequest, SendStreamingMessageRequest]", - response_dict: Dict[str, Any], - ) -> Tuple[int, int, int]: + request: "SendMessageRequest | SendStreamingMessageRequest", + response_dict: dict[str, Any], + ) -> tuple[int, int, int]: """ Calculate token usage from A2A request and response. @@ -145,5 +145,5 @@ def extract_text_from_a2a_message(message: Any) -> str: return A2ARequestUtils.extract_text_from_message(message) -def extract_text_from_a2a_response(response_dict: Dict[str, Any]) -> str: +def extract_text_from_a2a_response(response_dict: dict[str, Any]) -> str: return A2ARequestUtils.extract_text_from_response(response_dict) diff --git a/litellm/anthropic_beta_headers_manager.py b/litellm/anthropic_beta_headers_manager.py index d0082498b09..542885b5130 100644 --- a/litellm/anthropic_beta_headers_manager.py +++ b/litellm/anthropic_beta_headers_manager.py @@ -25,14 +25,13 @@ Environment Variables: import json import os from importlib.resources import files -from typing import Dict, List, Optional, Set import httpx from litellm.litellm_core_utils.litellm_logging import verbose_logger # Cache for the loaded configuration -_BETA_HEADERS_CONFIG: Optional[Dict] = None +_BETA_HEADERS_CONFIG: dict | None = None class GetAnthropicBetaHeadersConfig: @@ -44,7 +43,7 @@ class GetAnthropicBetaHeadersConfig: """ @staticmethod - def load_local_beta_headers_config() -> Dict: + def load_local_beta_headers_config() -> dict: """Load the local backup beta headers config bundled with the package.""" try: content = json.loads( @@ -159,7 +158,7 @@ def get_beta_headers_config(url: str) -> dict: return content -def _load_beta_headers_config() -> Dict: +def _load_beta_headers_config() -> dict: """ Load the beta headers configuration. Uses caching to avoid repeated fetches/file reads. @@ -183,7 +182,7 @@ def _load_beta_headers_config() -> Dict: return _BETA_HEADERS_CONFIG -def reload_beta_headers_config() -> Dict: +def reload_beta_headers_config() -> dict: """ Force reload the beta headers configuration from source (remote or local). Clears the cache and fetches fresh configuration. @@ -213,9 +212,9 @@ def get_provider_name(provider: str) -> str: def filter_and_transform_beta_headers( - beta_headers: List[str], + beta_headers: list[str], provider: str, -) -> List[str]: +) -> list[str]: """ Filter and transform beta headers based on provider's mapping configuration. @@ -240,7 +239,7 @@ def filter_and_transform_beta_headers( # Get the header mapping for this provider provider_mapping = config.get(provider, {}) - filtered_headers: Set[str] = set() + filtered_headers: set[str] = set() for header in beta_headers: header = header.strip() @@ -289,7 +288,7 @@ def is_beta_header_supported( def get_provider_beta_header( anthropic_beta_header: str, provider: str, -) -> Optional[str]: +) -> str | None: """ Get the provider-specific beta header name for a given Anthropic beta header. @@ -390,7 +389,7 @@ def update_request_with_filtered_beta( return headers, request_data -def get_unsupported_headers(provider: str) -> List[str]: +def get_unsupported_headers(provider: str) -> list[str]: """ Get all beta headers that are unsupported by a provider (have null values in mapping). diff --git a/litellm/anthropic_interface/exceptions/__init__.py b/litellm/anthropic_interface/exceptions/__init__.py index 875b09e3da3..7f2de0e60dc 100644 --- a/litellm/anthropic_interface/exceptions/__init__.py +++ b/litellm/anthropic_interface/exceptions/__init__.py @@ -11,9 +11,9 @@ from .exceptions import ( ) __all__ = [ - "AnthropicErrorType", + "ANTHROPIC_ERROR_TYPE_MAP", "AnthropicErrorDetail", "AnthropicErrorResponse", - "ANTHROPIC_ERROR_TYPE_MAP", + "AnthropicErrorType", "AnthropicExceptionMapping", ] diff --git a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py index b4ec83517ee..c0038dd7d83 100644 --- a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py +++ b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py @@ -5,13 +5,12 @@ Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthrop """ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from typing import Dict, Optional from .exceptions import AnthropicErrorResponse, AnthropicErrorType # HTTP status code -> Anthropic error type # Source: https://docs.anthropic.com/en/api/errors -ANTHROPIC_ERROR_TYPE_MAP: Dict[int, AnthropicErrorType] = { +ANTHROPIC_ERROR_TYPE_MAP: dict[int, AnthropicErrorType] = { 400: "invalid_request_error", 401: "authentication_error", 403: "permission_error", @@ -39,7 +38,7 @@ class AnthropicExceptionMapping: def create_error_response( status_code: int, message: str, - request_id: Optional[str] = None, + request_id: str | None = None, ) -> AnthropicErrorResponse: """ Create an Anthropic-formatted error response dict. @@ -124,7 +123,7 @@ class AnthropicExceptionMapping: def transform_to_anthropic_error( status_code: int, raw_message: str, - request_id: Optional[str] = None, + request_id: str | None = None, ) -> AnthropicErrorResponse: """ Transform an error message to Anthropic format. @@ -143,7 +142,7 @@ class AnthropicExceptionMapping: AnthropicErrorResponse dict """ # Try to parse as JSON once - parsed: Optional[dict] = safe_json_loads(raw_message) + parsed: dict | None = safe_json_loads(raw_message) if not isinstance(parsed, dict): parsed = None diff --git a/litellm/anthropic_interface/messages/__init__.py b/litellm/anthropic_interface/messages/__init__.py index 8fb9a9d1d04..2698cff5980 100644 --- a/litellm/anthropic_interface/messages/__init__.py +++ b/litellm/anthropic_interface/messages/__init__.py @@ -11,7 +11,7 @@ This is an __init__.py file to allow the following interface """ from collections.abc import AsyncIterator, Coroutine, Iterator -from typing import Any, Dict, List, Optional, Union +from typing import Any from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( anthropic_messages as _async_anthropic_messages, @@ -26,21 +26,21 @@ from litellm.types.llms.anthropic_messages.anthropic_response import ( async def acreate( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - container: Optional[Dict] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + container: dict | None = None, **kwargs, -) -> Union[AnthropicMessagesResponse, AsyncIterator]: +) -> AnthropicMessagesResponse | AsyncIterator: """ Async wrapper for Anthropic's messages API @@ -85,26 +85,26 @@ async def acreate( def create( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - container: Optional[Dict] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + container: dict | None = None, **kwargs, -) -> Union[ - AnthropicMessagesResponse, - Iterator[bytes], - AsyncIterator[Any], - Coroutine[Any, Any, Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]]], -]: +) -> ( + AnthropicMessagesResponse + | Iterator[bytes] + | AsyncIterator[Any] + | Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]] +): """ Async wrapper for Anthropic's messages API diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index 7219f4385d1..b476b3993d4 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -5,7 +5,7 @@ import contextvars import os from collections.abc import Coroutine, Iterable from functools import partial -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal import httpx from openai import AsyncOpenAI, OpenAI @@ -37,7 +37,7 @@ azure_assistants_api = AzureAssistantsAPI() async def aget_assistants( custom_llm_provider: Literal["openai", "azure"], - client: Optional[AsyncOpenAI] = None, + client: AsyncOpenAI | None = None, **kwargs, ) -> AsyncCursorPage[Assistant]: loop = asyncio.get_event_loop() @@ -74,13 +74,13 @@ async def aget_assistants( def get_assistants( custom_llm_provider: Literal["openai", "azure"], - client: Optional[Any] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + client: Any | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, **kwargs, ) -> SyncCursorPage[Assistant]: - aget_assistants: Optional[bool] = kwargs.pop("aget_assistants", None) + aget_assistants: bool | None = kwargs.pop("aget_assistants", None) if aget_assistants is not None and not isinstance(aget_assistants, bool): raise Exception("Invalid value passed in for aget_assistants. Only bool or None allowed") optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) @@ -102,7 +102,7 @@ def get_assistants( elif timeout is None: timeout = 600.0 - response: Optional[SyncCursorPage[Assistant]] = None + response: SyncCursorPage[Assistant] | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -148,7 +148,7 @@ def get_assistants( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -167,9 +167,7 @@ def get_assistants( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'get_assistants'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'get_assistants'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -181,9 +179,7 @@ def get_assistants( if response is None: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'get_assistants'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'get_assistants'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -198,7 +194,7 @@ def get_assistants( async def acreate_assistants( custom_llm_provider: Literal["openai", "azure"], - client: Optional[AsyncOpenAI] = None, + client: AsyncOpenAI | None = None, **kwargs, ) -> Assistant: loop = asyncio.get_event_loop() @@ -238,22 +234,22 @@ async def acreate_assistants( def create_assistants( custom_llm_provider: Literal["openai", "azure"], model: str, - name: Optional[str] = None, - description: Optional[str] = None, - instructions: Optional[str] = None, - tools: Optional[List[Dict[str, Any]]] = None, - tool_resources: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, str]] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - response_format: Optional[Union[str, Dict[str, str]]] = None, - client: Optional[Any] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + name: str | None = None, + description: str | None = None, + instructions: str | None = None, + tools: list[dict[str, Any]] | None = None, + tool_resources: dict[str, Any] | None = None, + metadata: dict[str, str] | None = None, + temperature: float | None = None, + top_p: float | None = None, + response_format: str | dict[str, str] | None = None, + client: Any | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, **kwargs, -) -> Union[Assistant, Coroutine[Any, Any, Assistant]]: - async_create_assistants: Optional[bool] = kwargs.pop("async_create_assistants", None) +) -> Assistant | Coroutine[Any, Any, Assistant]: + async_create_assistants: bool | None = kwargs.pop("async_create_assistants", None) if async_create_assistants is not None and not isinstance(async_create_assistants, bool): raise ValueError("Invalid value passed in for async_create_assistants. Only bool or None allowed") optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) @@ -291,7 +287,7 @@ def create_assistants( # only send params that are not None create_assistant_data = {k: v for k, v in create_assistant_data.items() if v is not None} - response: Optional[Union[Coroutine[Any, Any, Assistant], Assistant]] = None + response: Coroutine[Any, Any, Assistant] | Assistant | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -338,7 +334,7 @@ def create_assistants( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -361,9 +357,7 @@ def create_assistants( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_assistants'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_assistants'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -383,7 +377,7 @@ def create_assistants( async def adelete_assistant( custom_llm_provider: Literal["openai", "azure"], - client: Optional[AsyncOpenAI] = None, + client: AsyncOpenAI | None = None, **kwargs, ) -> AssistantDeleted: loop = asyncio.get_event_loop() @@ -422,17 +416,17 @@ async def adelete_assistant( def delete_assistant( custom_llm_provider: Literal["openai", "azure"], assistant_id: str, - client: Optional[Any] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + client: Any | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, **kwargs, -) -> Union[AssistantDeleted, Coroutine[Any, Any, AssistantDeleted]]: +) -> AssistantDeleted | Coroutine[Any, Any, AssistantDeleted]: optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) litellm_params_dict = get_litellm_params(**kwargs) - async_delete_assistants: Optional[bool] = kwargs.pop("async_delete_assistants", None) + async_delete_assistants: bool | None = kwargs.pop("async_delete_assistants", None) if async_delete_assistants is not None and not isinstance(async_delete_assistants, bool): raise ValueError("Invalid value passed in for async_delete_assistants. Only bool or None allowed") @@ -452,7 +446,7 @@ def delete_assistant( elif timeout is None: timeout = 600.0 - response: Optional[Union[AssistantDeleted, Coroutine[Any, Any, AssistantDeleted]]] = None + response: AssistantDeleted | Coroutine[Any, Any, AssistantDeleted] | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base @@ -491,7 +485,7 @@ def delete_assistant( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -514,9 +508,7 @@ def delete_assistant( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'delete_assistant'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'delete_assistant'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -572,10 +564,10 @@ async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwar def create_thread( custom_llm_provider: Literal["openai", "azure"], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]] = None, - metadata: Optional[dict] = None, - tool_resources: Optional[OpenAICreateThreadParamsToolResources] = None, - client: Optional[OpenAI] = None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None = None, + metadata: dict | None = None, + tool_resources: OpenAICreateThreadParamsToolResources | None = None, + client: OpenAI | None = None, **kwargs, ) -> Thread: """ @@ -620,10 +612,10 @@ def create_thread( elif timeout is None: timeout = 600.0 - api_base: Optional[str] = None - api_key: Optional[str] = None + api_base: str | None = None + api_key: str | None = None - response: Optional[Thread] = None + response: Thread | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -667,12 +659,10 @@ def create_thread( or get_secret("AZURE_API_KEY") ) # type: ignore - api_version: Optional[str] = ( - optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") - ) # type: ignore + api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -696,9 +686,7 @@ def create_thread( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_thread'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_thread'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -713,7 +701,7 @@ def create_thread( async def aget_thread( custom_llm_provider: Literal["openai", "azure"], thread_id: str, - client: Optional[AsyncOpenAI] = None, + client: AsyncOpenAI | None = None, **kwargs, ) -> Thread: loop = asyncio.get_event_loop() @@ -773,9 +761,9 @@ def get_thread( timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - api_base: Optional[str] = None - api_key: Optional[str] = None - response: Optional[Thread] = None + api_base: str | None = None + api_key: str | None = None + response: Thread | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -811,9 +799,7 @@ def get_thread( elif custom_llm_provider == "azure": api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore - api_version: Optional[str] = ( - optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") - ) # type: ignore + api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore api_key = ( optional_params.api_key @@ -824,7 +810,7 @@ def get_thread( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -847,9 +833,7 @@ def get_thread( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'get_thread'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'get_thread'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -869,8 +853,8 @@ async def a_add_message( thread_id: str, role: Literal["user", "assistant"], content: str, - attachments: Optional[List[Attachment]] = None, - metadata: Optional[dict] = None, + attachments: list[Attachment] | None = None, + metadata: dict | None = None, client=None, **kwargs, ) -> OpenAIMessage: @@ -922,8 +906,8 @@ def add_message( thread_id: str, role: Literal["user", "assistant"], content: str, - attachments: Optional[List[Attachment]] = None, - metadata: Optional[dict] = None, + attachments: list[Attachment] | None = None, + metadata: dict | None = None, client=None, **kwargs, ) -> OpenAIMessage: @@ -956,9 +940,9 @@ def add_message( timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - api_key: Optional[str] = None - api_base: Optional[str] = None - response: Optional[OpenAIMessage] = None + api_key: str | None = None + api_base: str | None = None + response: OpenAIMessage | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -994,9 +978,7 @@ def add_message( elif custom_llm_provider == "azure": api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore - api_version: Optional[str] = ( - optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") - ) # type: ignore + api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore api_key = ( optional_params.api_key @@ -1007,7 +989,7 @@ def add_message( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -1028,9 +1010,7 @@ def add_message( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_thread'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_thread'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -1046,7 +1026,7 @@ def add_message( async def aget_messages( custom_llm_provider: Literal["openai", "azure"], thread_id: str, - client: Optional[AsyncOpenAI] = None, + client: AsyncOpenAI | None = None, **kwargs, ) -> AsyncCursorPage[OpenAIMessage]: loop = asyncio.get_event_loop() @@ -1091,7 +1071,7 @@ async def aget_messages( def get_messages( custom_llm_provider: Literal["openai", "azure"], thread_id: str, - client: Optional[Any] = None, + client: Any | None = None, **kwargs, ) -> SyncCursorPage[OpenAIMessage]: aget_messages = kwargs.pop("aget_messages", None) @@ -1114,9 +1094,9 @@ def get_messages( elif timeout is None: timeout = 600.0 - response: Optional[SyncCursorPage[OpenAIMessage]] = None - api_key: Optional[str] = None - api_base: Optional[str] = None + response: SyncCursorPage[OpenAIMessage] | None = None + api_key: str | None = None + api_base: str | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -1151,9 +1131,7 @@ def get_messages( elif custom_llm_provider == "azure": api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore - api_version: Optional[str] = ( - optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") - ) # type: ignore + api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore api_key = ( optional_params.api_key @@ -1164,7 +1142,7 @@ def get_messages( ) # type: ignore extra_body = optional_params.get("extra_body", {}) - azure_ad_token: Optional[str] = None + azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) else: @@ -1184,9 +1162,7 @@ def get_messages( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'get_messages'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'get_messages'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -1202,7 +1178,7 @@ def get_messages( ### RUNS ### def arun_thread_stream( *, - event_handler: Optional[AssistantEventHandler] = None, + event_handler: AssistantEventHandler | None = None, **kwargs, ) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: kwargs["arun_thread"] = True @@ -1213,13 +1189,13 @@ async def arun_thread( custom_llm_provider: Literal["openai", "azure"], thread_id: str, assistant_id: str, - additional_instructions: Optional[str] = None, - instructions: Optional[str] = None, - metadata: Optional[dict] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - tools: Optional[Iterable[AssistantToolParam]] = None, - client: Optional[Any] = None, + additional_instructions: str | None = None, + instructions: str | None = None, + metadata: dict | None = None, + model: str | None = None, + stream: bool | None = None, + tools: Iterable[AssistantToolParam] | None = None, + client: Any | None = None, **kwargs, ) -> Run: loop = asyncio.get_event_loop() @@ -1270,7 +1246,7 @@ async def arun_thread( def run_thread_stream( *, - event_handler: Optional[AssistantEventHandler] = None, + event_handler: AssistantEventHandler | None = None, **kwargs, ) -> AssistantStreamManager[AssistantEventHandler]: return run_thread(stream=True, event_handler=event_handler, **kwargs) # type: ignore @@ -1280,14 +1256,14 @@ def run_thread( custom_llm_provider: Literal["openai", "azure"], thread_id: str, assistant_id: str, - additional_instructions: Optional[str] = None, - instructions: Optional[str] = None, - metadata: Optional[dict] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - tools: Optional[Iterable[AssistantToolParam]] = None, - client: Optional[Any] = None, - event_handler: Optional[AssistantEventHandler] = None, # for stream=True calls + additional_instructions: str | None = None, + instructions: str | None = None, + metadata: dict | None = None, + model: str | None = None, + stream: bool | None = None, + tools: Iterable[AssistantToolParam] | None = None, + client: Any | None = None, + event_handler: AssistantEventHandler | None = None, # for stream=True calls **kwargs, ) -> Run: """Run a given thread + assistant.""" @@ -1311,7 +1287,7 @@ def run_thread( elif timeout is None: timeout = 600.0 - response: Optional[Run] = None + response: Run | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -1393,9 +1369,7 @@ def run_thread( ) # type: ignore else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'run_thread'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'run_thread'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( diff --git a/litellm/assistants/utils.py b/litellm/assistants/utils.py index f775c1b6508..d23dcd973b4 100644 --- a/litellm/assistants/utils.py +++ b/litellm/assistants/utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import litellm from ..exceptions import UnsupportedParamsError @@ -7,21 +5,10 @@ from ..types.llms.openai import * def get_optional_params_add_message( - role: Optional[str], - content: Optional[ - Union[ - str, - List[ - Union[ - MessageContentTextObject, - MessageContentImageFileObject, - MessageContentImageURLObject, - ] - ], - ] - ], - attachments: Optional[List[Attachment]], - metadata: Optional[dict], + role: str | None, + content: str | List[MessageContentTextObject | MessageContentImageFileObject | MessageContentImageURLObject] | None, + attachments: List[Attachment] | None, + metadata: dict | None, custom_llm_provider: str, **kwargs, ): @@ -56,9 +43,7 @@ def get_optional_params_add_message( elif k not in supported_params: raise litellm.utils.UnsupportedParamsError( status_code=500, - message="k={}, not supported by {}. Supported params={}. To drop it from the call, set `litellm.drop_params = True`.".format( - k, custom_llm_provider, supported_params - ), + message=f"k={k}, not supported by {custom_llm_provider}. Supported params={supported_params}. To drop it from the call, set `litellm.drop_params = True`.", ) return non_default_params @@ -71,19 +56,19 @@ def get_optional_params_add_message( non_default_params=non_default_params, optional_params=optional_params ) for k in passed_params.keys(): - if k not in default_params.keys(): + if k not in default_params: optional_params[k] = passed_params[k] return optional_params def get_optional_params_image_gen( - n: Optional[int] = None, - quality: Optional[str] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + n: int | None = None, + quality: str | None = None, + response_format: str | None = None, + size: str | None = None, + style: str | None = None, + user: str | None = None, + custom_llm_provider: str | None = None, **kwargs, ): # retrieve all parameters passed to the function @@ -142,6 +127,6 @@ def get_optional_params_image_gen( optional_params["sampleCount"] = int(n) for k in passed_params.keys(): - if k not in default_params.keys(): + if k not in default_params: optional_params[k] = passed_params[k] return optional_params diff --git a/litellm/batch_completion/main.py b/litellm/batch_completion/main.py index 664977dc8d6..792be3ff7ad 100644 --- a/litellm/batch_completion/main.py +++ b/litellm/batch_completion/main.py @@ -1,5 +1,4 @@ from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait -from typing import List, Optional import litellm from litellm._logging import print_verbose @@ -11,23 +10,23 @@ from ..llms.vllm.completion import handler as vllm_handler def batch_completion( model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create - messages: List = [], - functions: Optional[List] = None, - function_call: Optional[str] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - n: Optional[int] = None, - stream: Optional[bool] = None, + messages: list = [], + functions: list | None = None, + function_call: str | None = None, + temperature: float | None = None, + top_p: float | None = None, + n: int | None = None, + stream: bool | None = None, stop=None, - max_tokens: Optional[int] = None, - presence_penalty: Optional[float] = None, - frequency_penalty: Optional[float] = None, - logit_bias: Optional[dict] = None, - user: Optional[str] = None, + max_tokens: int | None = None, + presence_penalty: float | None = None, + frequency_penalty: float | None = None, + logit_bias: dict | None = None, + user: str | None = None, deployment_id=None, - request_timeout: Optional[int] = None, - timeout: Optional[int] = 600, - max_workers: Optional[int] = 100, + request_timeout: int | None = None, + timeout: int | None = 600, + max_workers: int | None = 100, # Optional liteLLM function params **kwargs, ): @@ -164,7 +163,7 @@ def batch_completion_models(*args, **kwargs): futures = {} with ThreadPoolExecutor(max_workers=len(deployments)) as executor: for deployment in deployments: - for key in kwargs.keys(): + for key in kwargs: if key not in deployment: # don't override deployment values e.g. model name, api base, etc. deployment[key] = kwargs[key] kwargs = {**deployment, **nested_kwargs} @@ -250,7 +249,7 @@ def batch_completion_models_all_responses(*args, **kwargs): if result is not None: responses.append(result) except Exception as e: - print_verbose(f"batch_completion_models_all_responses: model request failed: {str(e)}") + print_verbose(f"batch_completion_models_all_responses: model request failed: {e!s}") continue return responses diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 8303ec75393..2e28aaa14df 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -1,7 +1,7 @@ import json from collections.abc import Iterable, Iterator from dataclasses import dataclass -from typing import Any, List, Literal, Optional, Tuple +from typing import Any, Literal import litellm from litellm._logging import verbose_logger @@ -12,11 +12,11 @@ from litellm.utils import token_counter async def calculate_batch_cost_and_usage( - file_content_dictionary: List[dict], + file_content_dictionary: list[dict], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"], - model_name: Optional[str] = None, - model_info: Optional[ModelInfo] = None, -) -> Tuple[float, Usage, List[str]]: + model_name: str | None = None, + model_info: ModelInfo | None = None, +) -> tuple[float, Usage, list[str]]: """ Calculate the cost and usage of a batch. @@ -45,9 +45,9 @@ async def calculate_batch_cost_and_usage( async def _handle_completed_batch( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"], - model_name: Optional[str] = None, - litellm_params: Optional[dict] = None, -) -> Tuple[float, Usage, List[str]]: + model_name: str | None = None, + litellm_params: dict | None = None, +) -> tuple[float, Usage, list[str]]: """Fetch a completed batch's output file and aggregate its cost, usage, and models in a single pass over the JSONL lines, so the parsed file content is never materialized in memory. @@ -85,14 +85,14 @@ class _BatchOutputLineStats: total_tokens: int cache_read_tokens: int cache_creation_tokens: int - model: Optional[str] + model: str | None def _iter_successful_output_line_stats( entries: Iterable[dict], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"], - model_name: Optional[str], - model_info: Optional[ModelInfo], + model_name: str | None, + model_info: ModelInfo | None, ) -> Iterator[_BatchOutputLineStats]: from litellm.cost_calculator import batch_cost_calculator @@ -136,9 +136,9 @@ def _iter_successful_output_line_stats( def _aggregate_batch_cost_usage_models( entries: Iterable[dict], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"], - model_name: Optional[str] = None, - model_info: Optional[ModelInfo] = None, -) -> Tuple[float, Usage, List[str]]: + model_name: str | None = None, + model_info: ModelInfo | None = None, +) -> tuple[float, Usage, list[str]]: """Aggregate cost, usage, and models from batch output entries in a single pass, holding one small stats record per line instead of the parsed file.""" line_stats = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info)) @@ -164,9 +164,9 @@ def _aggregate_batch_cost_usage_models( def calculate_vertex_ai_batch_cost_and_usage( - vertex_ai_batch_responses: List[dict], - model_name: Optional[str] = None, -) -> Tuple[float, Usage]: + vertex_ai_batch_responses: list[dict], + model_name: str | None = None, +) -> tuple[float, Usage]: """ Calculate both cost and usage from raw Vertex AI batch responses. @@ -234,7 +234,7 @@ def calculate_vertex_ai_batch_cost_and_usage( async def _fetch_batch_output_file_content( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai", - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> bytes: """ Fetch the batch output file and return its raw JSONL bytes @@ -278,7 +278,7 @@ async def _fetch_batch_output_file_content( return _file_content.content -def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict: +def _extract_file_access_credentials(litellm_params: dict | None) -> dict: """ Extract credentials from litellm_params for file access operations. @@ -317,7 +317,7 @@ def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict: return credentials -def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]: +def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]: """ Get the file content as a list of dictionaries from JSON Lines format """ @@ -367,7 +367,7 @@ def _estimate_batch_entry_tokens(raw_line: bytes) -> int: def _count_entry_tokens( entry: dict, - model_name: Optional[str] = None, + model_name: str | None = None, ) -> int: """Token-count a single batch input entry's body (chat / text / embedding).""" body = entry.get("body", {}) or {} diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 9db2b0c2b5f..b27939be8bf 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -15,7 +15,7 @@ import contextvars import os from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, Literal, Optional, Union, cast +from typing import Any, Literal, cast import httpx from openai.types.batch import BatchRequestCounts @@ -64,7 +64,7 @@ base_llm_http_handler = BaseLLMHTTPHandler() def _resolve_timeout( optional_params: GenericLiteLLMParams, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], custom_llm_provider: str, default_timeout: float = 600.0, ) -> float: @@ -107,10 +107,10 @@ async def acreate_batch( endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"], input_file_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, - output_expires_after: Optional[Dict[str, Any]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, + output_expires_after: dict[str, Any] | None = None, **kwargs, ) -> LiteLLMBatch: """ @@ -157,12 +157,12 @@ def create_batch( endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"], input_file_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, - output_expires_after: Optional[Dict[str, Any]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, + output_expires_after: dict[str, Any] | None = None, **kwargs, -) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: +) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: """ Creates and executes a batch from an uploaded file of request @@ -173,7 +173,7 @@ def create_batch( litellm_call_id = kwargs.get("litellm_call_id", None) proxy_server_request = kwargs.get("proxy_server_request", None) model_info = kwargs.get("model_info", None) - model: Optional[str] = kwargs.get("model", None) + model: str | None = kwargs.get("model", None) try: if model is not None: model, _, _, _ = get_llm_provider( @@ -182,7 +182,7 @@ def create_batch( ) except Exception as e: verbose_logger.exception( - f"litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - {str(e)}" + f"litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - {e!s}" ) _is_async = kwargs.pop("acreate_batch", False) is True @@ -238,7 +238,7 @@ def create_batch( model=model, ) return response - api_base: Optional[str] = None + api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there api_base = ( @@ -321,7 +321,7 @@ def create_batch( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support custom_llm_provider={} for 'create_batch'".format(custom_llm_provider), + message=f"LiteLLM doesn't support custom_llm_provider={custom_llm_provider} for 'create_batch'", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -339,9 +339,9 @@ def create_batch( async def aretrieve_batch( batch_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> LiteLLMBatch: """ @@ -380,14 +380,14 @@ async def aretrieve_batch( def _handle_retrieve_batch_providers_without_provider_config( batch_id: str, optional_params: GenericLiteLLMParams, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, _retrieve_batch_request: RetrieveBatchRequest, _is_async: bool, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai", - logging_obj: Optional[Any] = None, + logging_obj: Any | None = None, ): - api_base: Optional[str] = None + api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there api_base = ( @@ -489,10 +489,10 @@ def _handle_retrieve_batch_providers_without_provider_config( else: raise litellm.exceptions.BadRequestError( message=( - "LiteLLM doesn't support custom_llm_provider={} for 'retrieve_batch' without a `model` kwarg. " + f"LiteLLM doesn't support custom_llm_provider={custom_llm_provider} for 'retrieve_batch' without a `model` kwarg. " "Supported via this path: 'openai', 'azure', 'vertex_ai', 'anthropic'. " "'bedrock' is supported but requires `model` to be passed so the provider config can be loaded." - ).format(custom_llm_provider), + ), model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -508,11 +508,11 @@ def _handle_retrieve_batch_providers_without_provider_config( def retrieve_batch( batch_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: +) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: """ Retrieves a batch. @@ -520,7 +520,7 @@ def retrieve_batch( """ try: optional_params = GenericLiteLLMParams(**kwargs) - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) + litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj", None) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 litellm_params = get_litellm_params( @@ -589,7 +589,7 @@ def retrieve_batch( ) # Try to use provider config first (for providers like bedrock) - model: Optional[str] = kwargs.get("model", None) + model: str | None = kwargs.get("model", None) if model is not None: provider_config = ProviderConfigManager.get_provider_batches_config( model=model, @@ -643,12 +643,12 @@ def retrieve_batch( @client async def alist_batches( - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, custom_llm_provider: ListBatchesSupportedProvider = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -686,11 +686,11 @@ async def alist_batches( @client def list_batches( - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, custom_llm_provider: ListBatchesSupportedProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -823,11 +823,11 @@ def list_batches( async def acancel_batch( batch_id: str, - model: Optional[str] = None, + model: str | None = None, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> LiteLLMBatch: """ @@ -869,13 +869,13 @@ async def acancel_batch( def cancel_batch( batch_id: str, - model: Optional[str] = None, - custom_llm_provider: Union[Literal["openai", "azure", "vertex_ai"], str] = "openai", - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + model: str | None = None, + custom_llm_provider: Literal["openai", "azure", "vertex_ai"] | str = "openai", + metadata: dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: +) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: """ Cancels a batch. @@ -890,7 +890,7 @@ def cancel_batch( ) except Exception as e: verbose_logger.exception( - f"litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - {str(e)}" + f"litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - {e!s}" ) optional_params = GenericLiteLLMParams(**kwargs) litellm_params = get_litellm_params( @@ -920,7 +920,7 @@ def cancel_batch( ) _is_async = kwargs.pop("acancel_batch", False) is True - api_base: Optional[str] = None + api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: api_base = ( optional_params.api_base @@ -993,9 +993,7 @@ def cancel_batch( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'cancel_batch'. Only 'openai', 'azure', and 'vertex_ai' are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', and 'vertex_ai' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( diff --git a/litellm/budget_manager.py b/litellm/budget_manager.py index 26f888c8077..cfe1775d7e8 100644 --- a/litellm/budget_manager.py +++ b/litellm/budget_manager.py @@ -11,7 +11,7 @@ import json import os import threading import time -from typing import Literal, Optional +from typing import Literal import litellm from litellm.constants import ( @@ -28,8 +28,8 @@ class BudgetManager: self, project_name: str, client_type: str = "local", - api_base: Optional[str] = None, - headers: Optional[dict] = None, + api_base: str | None = None, + headers: dict | None = None, ): self.client_type = client_type self.project_name = project_name @@ -73,7 +73,7 @@ class BudgetManager: self, total_budget: float, user: str, - duration: Optional[Literal["daily", "weekly", "monthly", "yearly"]] = None, + duration: Literal["daily", "weekly", "monthly", "yearly"] | None = None, created_at: float = time.time(), ): self.user_dict[user] = {"total_budget": total_budget} @@ -113,10 +113,10 @@ class BudgetManager: def update_cost( self, user: str, - completion_obj: Optional[ModelResponse] = None, - model: Optional[str] = None, - input_text: Optional[str] = None, - output_text: Optional[str] = None, + completion_obj: ModelResponse | None = None, + model: str | None = None, + input_text: str | None = None, + output_text: str | None = None, ): if model and input_text and output_text: prompt_tokens = litellm.token_counter(model=model, messages=[{"role": "user", "content": input_text}]) diff --git a/litellm/caching/__init__.py b/litellm/caching/__init__.py index bbe90b04121..87f4f7a7c63 100644 --- a/litellm/caching/__init__.py +++ b/litellm/caching/__init__.py @@ -2,10 +2,10 @@ from .azure_blob_cache import AzureBlobCache from .caching import Cache, LiteLLMCacheType from .disk_cache import DiskCache from .dual_cache import DualCache +from .gcs_cache import GCSCache from .in_memory_cache import InMemoryCache from .qdrant_semantic_cache import QdrantSemanticCache from .redis_cache import RedisCache from .redis_cluster_cache import RedisClusterCache from .redis_semantic_cache import RedisSemanticCache from .s3_cache import S3Cache -from .gcs_cache import GCSCache diff --git a/litellm/caching/_internal_lru_cache.py b/litellm/caching/_internal_lru_cache.py index c5285b1dfba..df6e1fc0941 100644 --- a/litellm/caching/_internal_lru_cache.py +++ b/litellm/caching/_internal_lru_cache.py @@ -1,12 +1,12 @@ from collections.abc import Callable from functools import lru_cache -from typing import Optional, TypeVar +from typing import TypeVar T = TypeVar("T") def lru_cache_wrapper( - maxsize: Optional[int] = None, + maxsize: int | None = None, ) -> Callable[[Callable[..., T]], Callable[..., T]]: """ Wrapper for lru_cache that caches success and exceptions diff --git a/litellm/caching/azure_blob_cache.py b/litellm/caching/azure_blob_cache.py index fca7cf20313..80ad645ec7b 100644 --- a/litellm/caching/azure_blob_cache.py +++ b/litellm/caching/azure_blob_cache.py @@ -19,12 +19,12 @@ from .base_cache import BaseCache class AzureBlobCache(BaseCache): def __init__(self, account_url, container) -> None: - from azure.storage.blob import BlobServiceClient from azure.core.exceptions import ResourceExistsError from azure.identity import DefaultAzureCredential from azure.identity.aio import ( DefaultAzureCredential as AsyncDefaultAzureCredential, ) + from azure.storage.blob import BlobServiceClient from azure.storage.blob.aio import BlobServiceClient as AsyncBlobServiceClient self.container_client = BlobServiceClient( diff --git a/litellm/caching/base_cache.py b/litellm/caching/base_cache.py index 81f1d61bd0d..d1965772157 100644 --- a/litellm/caching/base_cache.py +++ b/litellm/caching/base_cache.py @@ -9,7 +9,7 @@ Has 4 methods: """ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -23,8 +23,8 @@ class BaseCache(ABC): def __init__(self, default_ttl: int = 60): self.default_ttl = default_ttl - def get_ttl(self, **kwargs) -> Optional[int]: - kwargs_ttl: Optional[int] = kwargs.get("ttl") + def get_ttl(self, **kwargs) -> int | None: + kwargs_ttl: int | None = kwargs.get("ttl") if kwargs_ttl is not None: try: return int(kwargs_ttl) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 34badaa3e8a..f69c2fa3b58 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -13,7 +13,7 @@ import json import time import traceback from enum import Enum -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any from pydantic import BaseModel @@ -55,19 +55,18 @@ class CacheMode(str, Enum): class Cache: def __init__( self, - type: Optional[LiteLLMCacheType] = LiteLLMCacheType.LOCAL, - mode: Optional[ - CacheMode - ] = CacheMode.default_on, # when default_on cache is always on, when default_off cache is opt in - host: Optional[str] = None, - port: Optional[str] = None, - password: Optional[str] = None, - namespace: Optional[str] = None, - ttl: Optional[float] = None, - default_in_memory_ttl: Optional[float] = None, - default_in_redis_ttl: Optional[float] = None, - similarity_threshold: Optional[float] = None, - supported_call_types: Optional[List[CachingSupportedCallTypes]] = [ + type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL, + mode: CacheMode + | None = CacheMode.default_on, # when default_on cache is always on, when default_off cache is opt in + host: str | None = None, + port: str | None = None, + password: str | None = None, + namespace: str | None = None, + ttl: float | None = None, + default_in_memory_ttl: float | None = None, + default_in_redis_ttl: float | None = None, + similarity_threshold: float | None = None, + supported_call_types: list[CachingSupportedCallTypes] | None = [ "completion", "acompletion", "embedding", @@ -82,38 +81,38 @@ class Cache: "aresponses", ], # s3 Bucket, boto3 configuration - azure_account_url: Optional[str] = None, - azure_blob_container: Optional[str] = None, - s3_bucket_name: Optional[str] = None, - s3_region_name: Optional[str] = None, - s3_api_version: Optional[str] = None, - s3_use_ssl: Optional[bool] = True, - s3_verify: Optional[Union[bool, str]] = None, - s3_endpoint_url: Optional[str] = None, - s3_aws_access_key_id: Optional[str] = None, - s3_aws_secret_access_key: Optional[str] = None, - s3_aws_session_token: Optional[str] = None, - s3_config: Optional[Any] = None, - s3_path: Optional[str] = None, - gcs_bucket_name: Optional[str] = None, - gcs_path_service_account: Optional[str] = None, - gcs_path: Optional[str] = None, + azure_account_url: str | None = None, + azure_blob_container: str | None = None, + s3_bucket_name: str | None = None, + s3_region_name: str | None = None, + s3_api_version: str | None = None, + s3_use_ssl: bool | None = True, + s3_verify: bool | str | None = None, + s3_endpoint_url: str | None = None, + s3_aws_access_key_id: str | None = None, + s3_aws_secret_access_key: str | None = None, + s3_aws_session_token: str | None = None, + s3_config: Any | None = None, + s3_path: str | None = None, + gcs_bucket_name: str | None = None, + gcs_path_service_account: str | None = None, + gcs_path: str | None = None, redis_semantic_cache_embedding_model: str = "text-embedding-ada-002", - redis_semantic_cache_index_name: Optional[str] = None, + redis_semantic_cache_index_name: str | None = None, valkey_semantic_cache_embedding_model: str = "text-embedding-ada-002", valkey_semantic_cache_index_name: str | None = None, - redis_flush_size: Optional[int] = None, - redis_startup_nodes: Optional[List] = None, - disk_cache_dir: Optional[str] = None, - qdrant_api_base: Optional[str] = None, - qdrant_api_key: Optional[str] = None, - qdrant_collection_name: Optional[str] = None, - qdrant_quantization_config: Optional[str] = None, + redis_flush_size: int | None = None, + redis_startup_nodes: list | None = None, + disk_cache_dir: str | None = None, + qdrant_api_base: str | None = None, + qdrant_api_key: str | None = None, + qdrant_collection_name: str | None = None, + qdrant_quantization_config: str | None = None, qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002", - qdrant_semantic_cache_vector_size: Optional[int] = None, + qdrant_semantic_cache_vector_size: int | None = None, # GCP IAM authentication parameters - gcp_service_account: Optional[str] = None, - gcp_ssl_ca_certs: Optional[str] = None, + gcp_service_account: str | None = None, + gcp_ssl_ca_certs: str | None = None, **kwargs, ): """ @@ -352,15 +351,15 @@ class Cache: if param in scope_excluded_params: continue if param in combined_kwargs: - param_value: Optional[str] = self._get_param_value(param, kwargs) + param_value: str | None = self._get_param_value(param, kwargs) if param_value is not None: - cache_key += f"{str(param)}: {str(param_value)}" + cache_key += f"{param!s}: {param_value!s}" elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now if kwargs[param] is None: continue # ignore None params param_value = kwargs[param] - cache_key += f"{str(param)}: {str(param_value)}" + cache_key += f"{param!s}: {param_value!s}" if is_semantic_cache: cache_key += self._get_semantic_cache_tenant_scope(kwargs) @@ -382,7 +381,7 @@ class Cache: self, param: str, kwargs: dict, - ) -> Optional[str]: + ) -> str | None: """ Get the value for the given param from kwargs """ @@ -400,15 +399,15 @@ class Cache: 2. Else if a model_group is set, then return the model_group as the model. This is used for all requests sent through the litellm.Router() 3. Else use the `model` passed in kwargs """ - metadata: Dict = kwargs.get("metadata", {}) or {} - litellm_params: Dict = kwargs.get("litellm_params", {}) or {} - metadata_in_litellm_params: Dict = litellm_params.get("metadata", {}) or {} - model_group: Optional[str] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group") + metadata: dict = kwargs.get("metadata", {}) or {} + litellm_params: dict = kwargs.get("litellm_params", {}) or {} + metadata_in_litellm_params: dict = litellm_params.get("metadata", {}) or {} + model_group: str | None = metadata.get("model_group") or metadata_in_litellm_params.get("model_group") caching_group = self._get_caching_group(metadata, model_group) return caching_group or model_group or kwargs["model"] - def _get_caching_group(self, metadata: dict, model_group: Optional[str]) -> Optional[str]: - caching_groups: Optional[List] = metadata.get("caching_groups", []) + def _get_caching_group(self, metadata: dict, model_group: str | None) -> str | None: + caching_groups: list | None = metadata.get("caching_groups", []) if caching_groups: for group in caching_groups: if model_group in group: @@ -429,7 +428,7 @@ class Cache: or litellm_params.get("file_name") ) - def _get_preset_cache_key_from_kwargs(self, **kwargs) -> Optional[str]: + def _get_preset_cache_key_from_kwargs(self, **kwargs) -> str | None: """ Get the preset cache key from kwargs["litellm_params"] @@ -510,8 +509,8 @@ class Cache: def _get_cache_logic( self, - cached_result: Optional[Any], - max_age: Optional[float], + cached_result: Any | None, + max_age: float | None, ): """ Common get cache logic across sync + async implementations @@ -544,8 +543,8 @@ class Cache: return cached_result @staticmethod - def _get_safe_cache_lookup_kwargs(kwargs: Dict[str, Any]) -> Dict[str, Any]: - cache_lookup_kwargs: Dict[str, Any] = {} + def _get_safe_cache_lookup_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]: + cache_lookup_kwargs: dict[str, Any] = {} for prompt_kwarg in ("messages", "input"): if prompt_kwarg in kwargs: cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg] @@ -558,7 +557,7 @@ class Cache: @staticmethod def _update_metadata_from_cache_lookup_kwargs( - original_kwargs: Dict[str, Any], cache_lookup_kwargs: Dict[str, Any] + original_kwargs: dict[str, Any], cache_lookup_kwargs: dict[str, Any] ) -> None: original_metadata = original_kwargs.get("metadata") cache_lookup_metadata = cache_lookup_kwargs.get("metadata") @@ -568,7 +567,7 @@ class Cache: if "semantic-similarity" in cache_lookup_metadata: original_metadata["semantic-similarity"] = cache_lookup_metadata["semantic-similarity"] - def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs): + def get_cache(self, dynamic_cache_object: BaseCache | None = None, **kwargs): """ Retrieves the cached result for the given arguments. @@ -603,7 +602,7 @@ class Cache: print_verbose(f"An exception occurred: {traceback.format_exc()}") return None - async def async_get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs): + async def async_get_cache(self, dynamic_cache_object: BaseCache | None = None, **kwargs): """ Async get cache implementation. @@ -677,9 +676,9 @@ class Cache: cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs) self.cache.set_cache(cache_key, cached_data, **kwargs) except Exception as e: - verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}") + verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {e!s}") - async def async_add_cache(self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs): + async def async_add_cache(self, result, dynamic_cache_object: BaseCache | None = None, **kwargs): """ Async implementation of add_cache """ @@ -696,14 +695,14 @@ class Cache: else: await self.cache.async_set_cache(cache_key, cached_data, **kwargs) except Exception as e: - verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}") + verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {e!s}") def _convert_to_cached_embedding( self, embedding_response: Any, - model: Optional[str], - prompt_tokens: Optional[int] = None, - prompt_tokens_details: Optional[dict] = None, + model: str | None, + prompt_tokens: int | None = None, + prompt_tokens_details: dict | None = None, ) -> CachedEmbedding: """ Convert any embedding response into the standardized CachedEmbedding TypedDict format. @@ -745,7 +744,7 @@ class Cache: self, result: EmbeddingResponse, idx_in_result_data: int, - ) -> Optional[dict]: + ) -> dict | None: """ Extract per-item prompt_tokens_details from a response for caching. @@ -788,7 +787,7 @@ class Cache: self, result: EmbeddingResponse, idx_in_result_data: int, - ) -> Optional[int]: + ) -> int | None: """ Extract the per-item prompt_tokens from a response for caching. @@ -813,7 +812,7 @@ class Cache: input: str, kwargs: dict, idx_in_result_data: int = 0, - ) -> Tuple[str, dict, dict]: + ) -> tuple[str, dict, dict]: preset_cache_key = self.get_cache_key(**{**kwargs, "input": input}) kwargs["cache_key"] = preset_cache_key embedding_response = result.data[idx_in_result_data] @@ -843,7 +842,7 @@ class Cache: ) return cache_key, cached_data, kwargs - async def async_add_cache_pipeline(self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs): + async def async_add_cache_pipeline(self, result, dynamic_cache_object: BaseCache | None = None, **kwargs): """ Async implementation of add_cache for Embedding calls @@ -875,7 +874,7 @@ class Cache: else: await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs) except Exception as e: - verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}") + verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {e!s}") def should_use_cache(self, **kwargs): """ @@ -926,11 +925,11 @@ class Cache: def enable_cache( - type: Optional[LiteLLMCacheType] = LiteLLMCacheType.LOCAL, - host: Optional[str] = None, - port: Optional[str] = None, - password: Optional[str] = None, - supported_call_types: Optional[List[CachingSupportedCallTypes]] = [ + type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL, + host: str | None = None, + port: str | None = None, + password: str | None = None, + supported_call_types: list[CachingSupportedCallTypes] | None = [ "completion", "acompletion", "embedding", @@ -986,11 +985,11 @@ def enable_cache( def update_cache( - type: Optional[LiteLLMCacheType] = LiteLLMCacheType.LOCAL, - host: Optional[str] = None, - port: Optional[str] = None, - password: Optional[str] = None, - supported_call_types: Optional[List[CachingSupportedCallTypes]] = [ + type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL, + host: str | None = None, + port: str | None = None, + password: str | None = None, + supported_call_types: list[CachingSupportedCallTypes] | None = [ "completion", "acompletion", "embedding", diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index c547b7aff61..aed38d6ef65 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -22,11 +22,7 @@ from collections.abc import AsyncGenerator, Callable, Generator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Optional, - Tuple, - Union, ) from pydantic import BaseModel @@ -75,8 +71,8 @@ class CachingHandlerResponse(BaseModel): For embeddings there can be a cache hit for some of the inputs in the list and a cache miss for others """ - cached_result: Optional[Any] = None - final_embedding_cached_response: Optional[EmbeddingResponse] = None + cached_result: Any | None = None + final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call @@ -109,7 +105,7 @@ def _is_chat_completion_cached_dict(cached_result: dict) -> bool: return "choices" in cached_result -def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bool: +def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, Any]) -> bool: """ When stream=True, do not run success callbacks at cache-hit time. @@ -125,25 +121,24 @@ class LLMCachingHandler: def __init__( self, original_function: Callable, - request_kwargs: Dict[str, Any], + request_kwargs: dict[str, Any], start_time: datetime.datetime, ): from litellm.caching import DualCache, RedisCache - self.async_streaming_chunks: List[ModelResponse] = [] - self.sync_streaming_chunks: List[ModelResponse] = [] + self.async_streaming_chunks: list[ModelResponse] = [] + self.sync_streaming_chunks: list[ModelResponse] = [] self.request_kwargs = _drop_logging_obj_from_kwargs(request_kwargs) - self.preset_cache_key: Optional[str] = None + self.preset_cache_key: str | None = None self.original_function = original_function self.start_time = start_time if litellm.cache is not None and isinstance(litellm.cache.cache, RedisCache): - self.dual_cache: Optional[DualCache] = DualCache( + self.dual_cache: DualCache | None = DualCache( redis_cache=litellm.cache.cache, in_memory_cache=in_memory_cache_obj, ) else: self.dual_cache = None - pass async def _async_get_cache( self, @@ -152,9 +147,9 @@ class LLMCachingHandler: logging_obj: LiteLLMLoggingObj, start_time: datetime.datetime, call_type: str, - kwargs: Dict[str, Any], - args: Optional[Tuple[Any, ...]] = None, - ) -> Optional[CachingHandlerResponse]: + kwargs: dict[str, Any], + args: tuple[Any, ...] | None = None, + ) -> CachingHandlerResponse | None: """ Internal method to get from the cache. Handles different call types (embeddings, chat/completions, text_completion, transcription) @@ -182,15 +177,15 @@ class LLMCachingHandler: kwargs.get("cache", {}).get("no-cache", False) is not True ): # allow users to control returning cached responses from the completion function args = args or () - final_embedding_cached_response: Optional[EmbeddingResponse] = None + final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False - cached_result: Optional[Any] = None + cached_result: Any | None = None kwargs = kwargs.copy() ######################################################### # Init cache timing metrics ######################################################### cache_check_start_time = time.perf_counter() - cache_check_end_time: Optional[float] = None + cache_check_end_time: float | None = None ######################################################### parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) kwargs["parent_otel_span"] = parent_otel_span @@ -291,10 +286,10 @@ class LLMCachingHandler: logging_obj: LiteLLMLoggingObj, start_time: datetime.datetime, call_type: str, - kwargs: Dict[str, Any], - args: Optional[Tuple[Any, ...]] = None, + kwargs: dict[str, Any], + args: tuple[Any, ...] | None = None, ) -> CachingHandlerResponse: - cached_result: Optional[Any] = None + cached_result: Any | None = None # Check if caching should be performed BEFORE doing expensive kwargs copy if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function): @@ -369,7 +364,7 @@ class LLMCachingHandler: return CachingHandlerResponse(cached_result=cached_result) return CachingHandlerResponse(cached_result=cached_result) - def handle_kwargs_input_list_or_str(self, kwargs: Dict[str, Any]) -> List[str]: + def handle_kwargs_input_list_or_str(self, kwargs: dict[str, Any]) -> list[str]: """ Handles the input of kwargs['input'] being a list or a string """ @@ -380,7 +375,7 @@ class LLMCachingHandler: else: raise ValueError("input must be a string or a list") - def _extract_model_from_cached_results(self, non_null_list: List[Tuple[int, CachedEmbedding]]) -> Optional[str]: + def _extract_model_from_cached_results(self, non_null_list: list[tuple[int, CachedEmbedding]]) -> str | None: """ Helper method to extract the model name from cached results. @@ -397,13 +392,13 @@ class LLMCachingHandler: def _process_async_embedding_cached_response( self, - final_embedding_cached_response: Optional[EmbeddingResponse], - cached_result: List[Optional[CachedEmbedding]], - kwargs: Dict[str, Any], + final_embedding_cached_response: EmbeddingResponse | None, + cached_result: list[CachedEmbedding | None], + kwargs: dict[str, Any], logging_obj: LiteLLMLoggingObj, start_time: datetime.datetime, model: str, - ) -> Tuple[Optional[EmbeddingResponse], bool]: + ) -> tuple[EmbeddingResponse | None, bool]: """ Returns the final embedding cached response and a boolean indicating if all elements in the list have a cache hit @@ -446,7 +441,7 @@ class LLMCachingHandler: final_embedding_cached_response._hidden_params["cache_hit"] = True prompt_tokens = 0 - aggregated_details: Optional[dict] = None + aggregated_details: dict | None = None for val in non_null_list: idx, cr = val # (idx, cr) tuple if cr is not None: @@ -476,10 +471,10 @@ class LLMCachingHandler: aggregated_details[key] = value ## USAGE - prompt_tokens_details: Optional["PromptTokensDetailsWrapper"] = None - if aggregated_details: - from litellm.types.utils import PromptTokensDetailsWrapper + from litellm.types.utils import PromptTokensDetailsWrapper + prompt_tokens_details: PromptTokensDetailsWrapper | None = None + if aggregated_details: try: prompt_tokens_details = PromptTokensDetailsWrapper(**aggregated_details) except Exception: @@ -674,9 +669,7 @@ class LLMCachingHandler: cache_hit=cache_hit, ) - async def _retrieve_from_cache( - self, call_type: str, kwargs: Dict[str, Any], args: Tuple[Any, ...] - ) -> Optional[Any]: + async def _retrieve_from_cache(self, call_type: str, kwargs: dict[str, Any], args: tuple[Any, ...]) -> Any | None: """ Internal method to - get cache key @@ -709,7 +702,7 @@ class LLMCachingHandler: if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs) - cached_result: Optional[Any] = None + cached_result: Any | None = None if call_type == CallTypes.aembedding.value: if isinstance(new_kwargs["input"], str): new_kwargs["input"] = [new_kwargs["input"]] @@ -754,21 +747,20 @@ class LLMCachingHandler: self, cached_result: Any, call_type: str, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], logging_obj: LiteLLMLoggingObj, model: str, - args: Tuple[Any, ...], - custom_llm_provider: Optional[str] = None, - ) -> Optional[ - Union[ - ModelResponse, - TextCompletionResponse, - EmbeddingResponse, - RerankResponse, - TranscriptionResponse, - CustomStreamWrapper, - ] - ]: + args: tuple[Any, ...], + custom_llm_provider: str | None = None, + ) -> ( + ModelResponse + | TextCompletionResponse + | EmbeddingResponse + | RerankResponse + | TranscriptionResponse + | CustomStreamWrapper + | None + ): """ Internal method to process the cached result @@ -921,7 +913,7 @@ class LLMCachingHandler: convert_to_streaming_response_async, ) - _stream_cached_result: Union[AsyncGenerator, Generator] + _stream_cached_result: AsyncGenerator | Generator if call_type == CallTypes.acompletion.value or call_type == CallTypes.atext_completion.value: _stream_cached_result = convert_to_streaming_response_async( response_object=cached_result, @@ -941,8 +933,8 @@ class LLMCachingHandler: self, result: Any, original_function: Callable, - kwargs: Dict[str, Any], - args: Optional[Tuple[Any, ...]] = None, + kwargs: dict[str, Any], + args: tuple[Any, ...] | None = None, ): """ Internal method to check the type of the result & cache used and adds the result to the cache accordingly @@ -1007,8 +999,8 @@ class LLMCachingHandler: def sync_set_cache( self, result: Any, - kwargs: Dict[str, Any], - args: Optional[Tuple[Any, ...]] = None, + kwargs: dict[str, Any], + args: tuple[Any, ...] | None = None, ): """ Sync internal method to add the result to the cache @@ -1029,7 +1021,7 @@ class LLMCachingHandler: return - def _should_store_result_in_cache(self, original_function: Callable, kwargs: Dict[str, Any]) -> bool: + def _should_store_result_in_cache(self, original_function: Callable, kwargs: dict[str, Any]) -> bool: """ Helper function to determine if the result should be stored in the cache. @@ -1075,7 +1067,7 @@ class LLMCachingHandler: """ - complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse]] = ( + complete_streaming_response: ModelResponse | TextCompletionResponse | None = ( _assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, @@ -1097,7 +1089,7 @@ class LLMCachingHandler: """ Sync internal method to add the streaming response to the cache """ - complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse]] = ( + complete_streaming_response: ModelResponse | TextCompletionResponse | None = ( _assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, @@ -1119,12 +1111,12 @@ class LLMCachingHandler: self, logging_obj: LiteLLMLoggingObj, model: str, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], cached_result: Any, is_async: bool, is_embedding: bool = False, - custom_llm_provider: Optional[str] = None, - cache_duration_ms: Optional[float] = None, + custom_llm_provider: str | None = None, + cache_duration_ms: float | None = None, ): """ Helper function to update the LiteLLMLoggingObj environment variables. @@ -1178,8 +1170,8 @@ class LLMCachingHandler: def convert_args_to_kwargs( original_function: Callable, - args: Optional[Tuple[Any, ...]] = None, -) -> Dict[str, Any]: + args: tuple[Any, ...] | None = None, +) -> dict[str, Any]: # Get the signature of the original function signature = inspect.signature(original_function) diff --git a/litellm/caching/disk_cache.py b/litellm/caching/disk_cache.py index af8eb92849f..aec18d836b0 100644 --- a/litellm/caching/disk_cache.py +++ b/litellm/caching/disk_cache.py @@ -1,5 +1,5 @@ import json -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from .base_cache import BaseCache @@ -12,7 +12,7 @@ else: class DiskCache(BaseCache): - def __init__(self, disk_cache_dir: Optional[str] = None): + def __init__(self, disk_cache_dir: str | None = None): try: import diskcache as dc except ModuleNotFoundError as e: diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 0e3c93946fd..5b56789e8db 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -13,7 +13,7 @@ import time import traceback from concurrent.futures import ThreadPoolExecutor from threading import Lock -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union if TYPE_CHECKING: from litellm.types.caching import RedisPipelineIncrementOperation @@ -57,11 +57,11 @@ class DualCache(BaseCache): def __init__( self, - in_memory_cache: Optional[InMemoryCache] = None, - redis_cache: Optional[RedisCache] = None, - default_in_memory_ttl: Optional[float] = None, - default_redis_ttl: Optional[float] = None, - default_redis_batch_cache_expiry: Optional[float] = None, + in_memory_cache: InMemoryCache | None = None, + redis_cache: RedisCache | None = None, + default_in_memory_ttl: float | None = None, + default_redis_ttl: float | None = None, + default_redis_batch_cache_expiry: float | None = None, default_max_redis_batch_cache_size: int = DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE, ) -> None: super().__init__() @@ -77,7 +77,7 @@ class DualCache(BaseCache): self.default_in_memory_ttl = default_in_memory_ttl or litellm.default_in_memory_ttl self.default_redis_ttl = default_redis_ttl or litellm.default_redis_ttl - def update_cache_ttl(self, default_in_memory_ttl: Optional[float], default_redis_ttl: Optional[float]): + def update_cache_ttl(self, default_in_memory_ttl: float | None, default_redis_ttl: float | None): if default_in_memory_ttl is not None: self.default_in_memory_ttl = default_in_memory_ttl @@ -86,9 +86,9 @@ class DualCache(BaseCache): def attach_redis_cache( self, - redis_cache: Optional[RedisCache] = None, + redis_cache: RedisCache | None = None, *, - default_redis_ttl: Optional[float] = None, + default_redis_ttl: float | None = None, ) -> None: """ Attach a Redis backend if this DualCache does not already have one. @@ -147,13 +147,13 @@ class DualCache(BaseCache): return result except Exception as e: - verbose_logger.error(f"LiteLLM Cache: Excepton async add_cache: {str(e)}") + verbose_logger.error(f"LiteLLM Cache: Excepton async add_cache: {e!s}") raise e def get_cache( self, key, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, **kwargs, ): @@ -184,7 +184,7 @@ class DualCache(BaseCache): def batch_get_cache( self, keys: list, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, **kwargs, ): @@ -217,7 +217,7 @@ class DualCache(BaseCache): async def async_get_cache( self, key, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, **kwargs, ): @@ -250,15 +250,15 @@ class DualCache(BaseCache): def _reserve_redis_batch_keys( self, current_time: float, - keys: List[str], - result: List[Any], - ) -> Tuple[List[str], Dict[str, Optional[float]]]: + keys: list[str], + result: list[Any], + ) -> tuple[list[str], dict[str, float | None]]: """ Atomically choose keys to fetch from Redis and reserve their access time. This prevents check-then-act races under concurrent async callers. """ - sublist_keys: List[str] = [] - previous_access_times: Dict[str, Optional[float]] = {} + sublist_keys: list[str] = [] + previous_access_times: dict[str, float | None] = {} with self._last_redis_batch_access_time_lock: for key, value in zip(keys, result): @@ -275,7 +275,7 @@ class DualCache(BaseCache): return sublist_keys, previous_access_times - def _rollback_redis_batch_key_reservations(self, previous_access_times: Dict[str, Optional[float]]) -> None: + def _rollback_redis_batch_key_reservations(self, previous_access_times: dict[str, float | None]) -> None: with self._last_redis_batch_access_time_lock: for key, previous_time in previous_access_times.items(): if previous_time is None: @@ -286,7 +286,7 @@ class DualCache(BaseCache): async def async_batch_get_cache( self, keys: list, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, **kwargs, ): @@ -347,7 +347,7 @@ class DualCache(BaseCache): if self.redis_cache is not None and local_only is False: await self.redis_cache.async_set_cache(key, value, **kwargs) except Exception as e: - verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {str(e)}") + verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {e!s}") # async_batch_set_cache async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs): @@ -366,17 +366,17 @@ class DualCache(BaseCache): cache_list=cache_list, ttl=kwargs.pop("ttl", None), **kwargs ) except Exception as e: - verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {str(e)}") + verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {e!s}") async def async_increment_cache( self, key, value: float, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, refresh_ttl: bool = False, **kwargs, - ) -> Optional[float]: + ) -> float | None: """ Key - the key in cache @@ -388,7 +388,7 @@ class DualCache(BaseCache): Returns - the incremented value, or None if no cache backend is available (in_memory_cache is None and Redis failed/is absent). """ - result: Optional[float] = None + result: float | None = None try: if self.in_memory_cache is not None: result = await self.in_memory_cache.async_increment(key, value, **kwargs) @@ -412,12 +412,12 @@ class DualCache(BaseCache): async def async_increment_cache_pipeline( self, - increment_list: List["RedisPipelineIncrementOperation"], + increment_list: list["RedisPipelineIncrementOperation"], local_only: bool = False, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, **kwargs, - ) -> Optional[List[float]]: - result: Optional[List[float]] = None + ) -> list[float] | None: + result: list[float] | None = None try: if self.in_memory_cache is not None: result = await self.in_memory_cache.async_increment_pipeline( @@ -439,7 +439,7 @@ class DualCache(BaseCache): ) return result - async def async_set_cache_sadd(self, key, value: List, local_only: bool = False, **kwargs) -> None: + async def async_set_cache_sadd(self, key, value: list, local_only: bool = False, **kwargs) -> None: """ Add value to a set @@ -456,7 +456,7 @@ class DualCache(BaseCache): if self.redis_cache is not None and local_only is False: _ = await self.redis_cache.async_set_cache_sadd(key, value, ttl=kwargs.get("ttl", None)) - return None + return except Exception as e: raise e # don't log, if exception is raised @@ -484,7 +484,7 @@ class DualCache(BaseCache): if self.redis_cache is not None: await self.redis_cache.async_delete_cache(key) - async def async_get_ttl(self, key: str) -> Optional[int]: + async def async_get_ttl(self, key: str) -> int | None: """ Get the remaining TTL of a key in in-memory cache or redis """ diff --git a/litellm/caching/gcs_cache.py b/litellm/caching/gcs_cache.py index 3345f8fc5eb..d74c68de770 100644 --- a/litellm/caching/gcs_cache.py +++ b/litellm/caching/gcs_cache.py @@ -2,27 +2,27 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests. """ -import json import asyncio -from typing import Optional +import json from urllib.parse import quote from litellm._logging import print_verbose, verbose_logger from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase from litellm.llms.custom_httpx.http_handler import ( - get_async_httpx_client, _get_httpx_client, + get_async_httpx_client, httpxSpecialProvider, ) + from .base_cache import BaseCache class GCSCache(BaseCache): def __init__( self, - bucket_name: Optional[str] = None, - path_service_account: Optional[str] = None, - gcs_path: Optional[str] = None, + bucket_name: str | None = None, + path_service_account: str | None = None, + gcs_path: str | None = None, ) -> None: super().__init__() self.bucket_name = bucket_name or GCSBucketBase(bucket_name=None).BUCKET_NAME diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 36b477f7a8b..e8b071bd492 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -8,12 +8,12 @@ Has 4 methods: - async_get_cache """ +import heapq import json import sys -import time -import heapq import threading -from typing import TYPE_CHECKING, Any, List, Optional +import time +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from litellm.types.caching import RedisPipelineIncrementOperation @@ -28,11 +28,10 @@ from .base_cache import BaseCache class InMemoryCache(BaseCache): def __init__( self, - max_size_in_memory: Optional[int] = 200, - default_ttl: Optional[ - int - ] = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute - max_size_per_item: Optional[int] = 1024, # 1MB = 1024KB + max_size_in_memory: int | None = 200, + default_ttl: int + | None = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute + max_size_per_item: int | None = 1024, # 1MB = 1024KB ): """ max_size_in_memory [int]: Maximum number of items in cache. done to prevent memory leaks. Use 200 items as a default @@ -146,9 +145,7 @@ class InMemoryCache(BaseCache): Check if ttl is set for a key """ ttl_time = self.ttl_dict.get(key) - if ttl_time is None: # if ttl is not set, allow override - return True - elif float(ttl_time) < time.time(): # if ttl is expired, allow override + if ttl_time is None or float(ttl_time) < time.time(): # if ttl is not set, allow override return True else: return False @@ -184,7 +181,7 @@ class InMemoryCache(BaseCache): else: self.set_cache(key=cache_key, value=cache_value) - async def async_set_cache_sadd(self, key, value: List, ttl: Optional[float]): + async def async_set_cache_sadd(self, key, value: list, ttl: float | None): """ Add value to set """ @@ -247,8 +244,8 @@ class InMemoryCache(BaseCache): return self.increment_cache(key=key, value=value, **kwargs) async def async_increment_pipeline( - self, increment_list: List["RedisPipelineIncrementOperation"], **kwargs - ) -> Optional[List[float]]: + self, increment_list: list["RedisPipelineIncrementOperation"], **kwargs + ) -> list[float] | None: results = [] for increment in increment_list: result = await self.async_increment(increment["key"], increment["increment_value"], **kwargs) @@ -266,13 +263,13 @@ class InMemoryCache(BaseCache): def delete_cache(self, key): self._remove_key(key) - async def async_get_ttl(self, key: str) -> Optional[int]: + async def async_get_ttl(self, key: str) -> int | None: """ Get the remaining TTL of a key in in-memory cache """ return self.ttl_dict.get(key, None) - async def async_get_oldest_n_keys(self, n: int) -> List[str]: + async def async_get_oldest_n_keys(self, n: int) -> list[str]: """ Get the oldest n keys in the cache """ diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 5ed1bb47eba..6e36dfbc096 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -12,7 +12,7 @@ import ast import asyncio import json import os -from typing import Any, Dict, cast +from typing import Any, cast import litellm from litellm._logging import print_verbose @@ -104,7 +104,7 @@ class QdrantSemanticCache(BaseCache): print_verbose(f"Collection already exists.\nCollection details:{self.collection_info}") self._ensure_cache_key_payload_index() else: - quantization_params: Dict[str, Any] + quantization_params: dict[str, Any] if quantization_config is None or quantization_config == "binary": quantization_params = { "binary": { @@ -178,7 +178,7 @@ class QdrantSemanticCache(BaseCache): if response.status_code not in (200, 201): print_verbose(f"Qdrant semantic-cache could not create cache-key payload index: {response.text}") except Exception as exc: - print_verbose(f"Qdrant semantic-cache could not create cache-key payload index: {str(exc)}") + print_verbose(f"Qdrant semantic-cache could not create cache-key payload index: {exc!s}") def _payload_matches_cache_key(self, payload: dict, key: str) -> bool: # Pre-isolation points stored only prompt + response with no cache-key @@ -188,7 +188,7 @@ class QdrantSemanticCache(BaseCache): cached_key = payload.get(self.CACHE_KEY_FIELD_NAME) return cached_key is not None and str(cached_key) == str(key) - def _get_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse: + def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse: """Embed via the proxy Router when it serves the model, else direct.""" try: from litellm.proxy.proxy_server import llm_model_list, llm_router @@ -210,7 +210,7 @@ class QdrantSemanticCache(BaseCache): cache={"no-store": True, "no-cache": True}, ) - async def _get_async_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse: + async def _get_async_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse: try: from litellm.proxy.proxy_server import llm_model_list, llm_router except ImportError: @@ -270,7 +270,6 @@ class QdrantSemanticCache(BaseCache): headers=self.headers, json=data, ) - return def get_cache(self, key, **kwargs): print_verbose(f"sync qdrant semantic-cache get_cache, kwargs: {kwargs}") @@ -344,7 +343,6 @@ class QdrantSemanticCache(BaseCache): else: # cache miss ! return None - pass async def async_set_cache(self, key, value, **kwargs): from litellm._uuid import uuid @@ -381,7 +379,6 @@ class QdrantSemanticCache(BaseCache): headers=self.headers, json=data, ) - return async def async_get_cache(self, key, **kwargs): print_verbose(f"async qdrant semantic-cache get_cache, kwargs: {kwargs}") @@ -452,7 +449,6 @@ class QdrantSemanticCache(BaseCache): else: # cache miss ! return None - pass async def _collection_info(self): return self.collection_info diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 9e0f022262b..1b0aa778f4e 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -16,9 +16,9 @@ import inspect import json import time from collections.abc import Awaitable, Callable, Sequence -from datetime import timedelta from contextvars import ContextVar -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, TypeVar, Union, cast +from datetime import timedelta +from typing import TYPE_CHECKING, Any, TypeVar, Union, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -127,7 +127,7 @@ class RedisCircuitBreaker: self.recovery_timeout = recovery_timeout self.enabled = enabled self._failure_count = 0 - self._opened_at: Optional[float] = None + self._opened_at: float | None = None self._state = self.CLOSED def is_open(self) -> bool: @@ -272,10 +272,10 @@ class RedisCache(BaseCache): host=None, port=None, password=None, - redis_flush_size: Optional[int] = 100, - namespace: Optional[str] = None, - startup_nodes: Optional[List] = None, # for redis-cluster - socket_timeout: Optional[float] = 5.0, # default 5 second timeout + redis_flush_size: int | None = 100, + namespace: str | None = None, + startup_nodes: list | None = None, # for redis-cluster + socket_timeout: float | None = 5.0, # default 5 second timeout **kwargs, ): from litellm._service_logger import ServiceLogging @@ -304,7 +304,7 @@ class RedisCache(BaseCache): redis_kwargs.update(kwargs) self.redis_client = get_redis_client(**redis_kwargs) - self.redis_async_client: Optional[Union[async_redis_client, async_redis_cluster_client]] = None + self.redis_async_client: async_redis_client | async_redis_cluster_client | None = None self.redis_kwargs = redis_kwargs self.async_redis_conn_pool = get_redis_connection_pool(**redis_kwargs) @@ -346,7 +346,7 @@ class RedisCache(BaseCache): verbose_logger.debug("Ignoring async redis ping. No running event loop.") else: verbose_logger.error( - "Error connecting to Async Redis client - {}".format(str(e)), + f"Error connecting to Async Redis client - {e!s}", extra={"error": str(e)}, ) self._handle_async_ping_error(e) @@ -407,7 +407,7 @@ class RedisCache(BaseCache): def init_async_client( self, - ) -> Union[async_redis_client, async_redis_cluster_client]: + ) -> async_redis_client | async_redis_cluster_client: from litellm import in_memory_llm_clients_cache from .._redis import get_redis_async_client, get_redis_connection_pool @@ -415,7 +415,7 @@ class RedisCache(BaseCache): cache_key = self._get_async_client_cache_key() cached_client = in_memory_llm_clients_cache.get_cache(key=cache_key) if cached_client is not None: - redis_async_client = cast(Union[async_redis_client, async_redis_cluster_client], cached_client) + redis_async_client = cast(async_redis_client | async_redis_cluster_client, cached_client) else: # Create new connection pool and client for current event loop self.async_redis_conn_pool = get_redis_connection_pool(**self.redis_kwargs) @@ -483,9 +483,9 @@ class RedisCache(BaseCache): ) except Exception as e: # NON blocking - notify users Redis is throwing an exception - print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {str(e)}") + print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {e!s}") - def increment_cache(self, key, value: int, ttl: Optional[float] = None, **kwargs) -> int: + def increment_cache(self, key, value: int, ttl: float | None = None, **kwargs) -> int: _redis_client = self.redis_client start_time = time.time() set_ttl = self.get_ttl(ttl=ttl) @@ -626,7 +626,7 @@ class RedisCache(BaseCache): async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any: async def execute() -> object: - executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache( + executor: Callable[..., Awaitable[Any]] | None = litellm.in_memory_llm_clients_cache.get_cache( key=script_cache_key ) if executor is None: @@ -755,10 +755,10 @@ class RedisCache(BaseCache): async def _pipeline_helper( self, - pipe: Union[pipeline, cluster_pipeline], - cache_list: List[Tuple[Any, Any]], - ttl: Optional[float], - ) -> List: + pipe: pipeline | cluster_pipeline, + cache_list: list[tuple[Any, Any]], + ttl: float | None, + ) -> list: """ Helper function for executing a pipeline of set operations on Redis """ @@ -769,7 +769,7 @@ class RedisCache(BaseCache): print_verbose(f"Set ASYNC Redis Cache PIPELINE: key: {cache_key}\nValue {cache_value}\nttl={ttl}") json_cache_value = json.dumps(cache_value) # Set the value with a TTL if it's provided. - _td: Optional[timedelta] = None + _td: timedelta | None = None if ttl is not None: _td = timedelta(seconds=ttl) pipe.set( # type: ignore @@ -782,7 +782,7 @@ class RedisCache(BaseCache): return results @_redis_circuit_breaker_guard - async def async_set_cache_pipeline(self, cache_list: List[Tuple[Any, Any]], ttl: Optional[float] = None, **kwargs): + async def async_set_cache_pipeline(self, cache_list: list[tuple[Any, Any]], ttl: float | None = None, **kwargs): """ Use Redis Pipelines for bulk write operations """ @@ -814,7 +814,7 @@ class RedisCache(BaseCache): parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), ) ) - return None + return except Exception as e: ## LOGGING ## end_time = time.time() @@ -842,8 +842,8 @@ class RedisCache(BaseCache): self, redis_client: async_redis_client, key: str, - value: List, - ttl: Optional[float], + value: list, + ttl: float | None, ) -> None: """Helper function for async_set_cache_sadd. Separated for testing.""" ttl = self.get_ttl(ttl=ttl) @@ -856,7 +856,7 @@ class RedisCache(BaseCache): raise @_redis_circuit_breaker_guard - async def async_set_cache_sadd(self, key, value: List, ttl: Optional[float], **kwargs): + async def async_set_cache_sadd(self, key, value: list, ttl: float | None, **kwargs): from redis.asyncio import Redis start_time = time.time() @@ -938,8 +938,8 @@ class RedisCache(BaseCache): self, key, value: float, - ttl: Optional[int] = None, - parent_otel_span: Optional[Span] = None, + ttl: int | None = None, + parent_otel_span: Span | None = None, refresh_ttl: bool = False, ) -> float: from redis.asyncio import Redis @@ -1051,7 +1051,7 @@ class RedisCache(BaseCache): cached_response = ast.literal_eval(cached_response) return cached_response - def get_cache(self, key, parent_otel_span: Optional[Span] = None, **kwargs): + def get_cache(self, key, parent_otel_span: Span | None = None, **kwargs): try: key = self.check_and_fix_namespace(key=key) print_verbose(f"Get Redis Cache: key: {key}") @@ -1073,7 +1073,7 @@ class RedisCache(BaseCache): # NON blocking - notify users Redis is throwing an exception verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: ", e) - def _run_redis_mget_operation(self, keys: List[str]) -> List[Any]: + def _run_redis_mget_operation(self, keys: list[str]) -> list[Any]: """ Wrapper to call `mget` on the redis client @@ -1081,7 +1081,7 @@ class RedisCache(BaseCache): """ return self.redis_client.mget(keys=keys) # type: ignore - async def _async_run_redis_mget_operation(self, keys: List[str]) -> List[Any]: + async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]: """ Wrapper to call `mget` on the redis client @@ -1092,8 +1092,8 @@ class RedisCache(BaseCache): def batch_get_cache( self, - key_list: Union[List[str], List[Optional[str]]], - parent_otel_span: Optional[Span] = None, + key_list: list[str] | list[str | None], + parent_otel_span: Span | None = None, ) -> dict: """ Use Redis for bulk read operations @@ -1114,7 +1114,7 @@ class RedisCache(BaseCache): cache_key = self.check_and_fix_namespace(key=cache_key or "") _keys.append(cache_key) start_time = time.time() - results: List = self._run_redis_mget_operation(keys=_keys) + results: list = self._run_redis_mget_operation(keys=_keys) end_time = time.time() _duration = end_time - start_time self.service_logger_obj.service_success_hook( @@ -1139,11 +1139,11 @@ class RedisCache(BaseCache): return decoded_results except Exception as e: - verbose_logger.error(f"Error occurred in batch get cache - {str(e)}") + verbose_logger.error(f"Error occurred in batch get cache - {e!s}") return key_value_dict @_redis_circuit_breaker_guard - async def async_get_cache(self, key, parent_otel_span: Optional[Span] = None, **kwargs): + async def async_get_cache(self, key, parent_otel_span: Span | None = None, **kwargs): from redis.asyncio import Redis _redis_client: Redis = self.init_async_client() # type: ignore @@ -1185,14 +1185,14 @@ class RedisCache(BaseCache): event_metadata={"key": key}, ) ) - print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}") + print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {e!s}") _record_swallowed_redis_failure(self._circuit_breaker, e) @_redis_circuit_breaker_guard async def async_batch_get_cache( self, - key_list: Union[List[str], List[Optional[str]]], - parent_otel_span: Optional[Span] = None, + key_list: list[str] | list[str | None], + parent_otel_span: Span | None = None, ) -> dict: """ Use Redis for bulk read operations @@ -1257,7 +1257,7 @@ class RedisCache(BaseCache): parent_otel_span=parent_otel_span, ) ) - verbose_logger.error(f"Error occurred in async batch get cache - {str(e)}") + verbose_logger.error(f"Error occurred in async batch get cache - {e!s}") _record_swallowed_redis_failure(self._circuit_breaker, e) return key_value_dict @@ -1292,7 +1292,7 @@ class RedisCache(BaseCache): error=e, call_type=f"sync_ping <- {_get_call_stack_info()}", ) - verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}") + verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {e!s}") raise e async def ping(self) -> bool: @@ -1326,7 +1326,7 @@ class RedisCache(BaseCache): call_type=f"async_ping <- {_get_call_stack_info()}", ) ) - verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}") + verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {e!s}") raise e @_redis_circuit_breaker_guard @@ -1337,8 +1337,8 @@ class RedisCache(BaseCache): # keys is a list, unpack it so it gets passed as individual elements to delete await _redis_client.delete(*keys) - def client_list(self) -> List: - client_list: List = self.redis_client.client_list() # type: ignore + def client_list(self) -> list: + client_list: list = self.redis_client.client_list() # type: ignore return client_list def info(self): @@ -1388,10 +1388,10 @@ class RedisCache(BaseCache): else: return {"status": "failed", "message": "Redis ping returned False"} except Exception as e: - verbose_logger.error(f"Redis connection test failed: {str(e)}") + verbose_logger.error(f"Redis connection test failed: {e!s}") return { "status": "failed", - "message": f"Redis connection failed: {str(e)}", + "message": f"Redis connection failed: {e!s}", "error": str(e), } @@ -1410,8 +1410,8 @@ class RedisCache(BaseCache): async def _pipeline_increment_helper( self, pipe: pipeline, - increment_list: List[RedisPipelineIncrementOperation], - ) -> Optional[List[float]]: + increment_list: list[RedisPipelineIncrementOperation], + ) -> list[float] | None: """Helper function for pipeline increment operations""" # Iterate through each increment operation and add commands to pipeline for increment_op in increment_list: @@ -1431,8 +1431,8 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_increment_pipeline( - self, increment_list: List[RedisPipelineIncrementOperation], **kwargs - ) -> Optional[List[float]]: + self, increment_list: list[RedisPipelineIncrementOperation], **kwargs + ) -> list[float] | None: """ Use Redis Pipelines for bulk increment operations Args: @@ -1492,7 +1492,7 @@ class RedisCache(BaseCache): raise e @_redis_circuit_breaker_guard - async def async_get_ttl(self, key: str) -> Optional[int]: + async def async_get_ttl(self, key: str) -> int | None: """ Get the remaining TTL of a key in Redis @@ -1521,8 +1521,8 @@ class RedisCache(BaseCache): async def async_rpush( self, key: str, - values: List[Any], - parent_otel_span: Optional[Span] = None, + values: list[Any], + parent_otel_span: Span | None = None, **kwargs, ) -> int: """ @@ -1565,14 +1565,14 @@ class RedisCache(BaseCache): call_type=f"async_rpush <- {_get_call_stack_info()}", ) ) - verbose_logger.error(f"LiteLLM Redis Cache RPUSH: - Got exception from REDIS : {str(e)}") + verbose_logger.error(f"LiteLLM Redis Cache RPUSH: - Got exception from REDIS : {e!s}") raise e async def _pipeline_rpush_helper( self, pipe: pipeline, - rpush_list: List[RedisPipelineRpushOperation], - ) -> List[int]: + rpush_list: list[RedisPipelineRpushOperation], + ) -> list[int]: """Helper function for pipeline rpush operations""" for rpush_op in rpush_list: key = self.check_and_fix_namespace(key=rpush_op["key"]) @@ -1587,8 +1587,8 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_rpush_pipeline( self, - rpush_list: List[RedisPipelineRpushOperation], - ) -> List[int]: + rpush_list: list[RedisPipelineRpushOperation], + ) -> list[int]: """ Use Redis Pipelines for bulk RPUSH operations @@ -1639,8 +1639,8 @@ class RedisCache(BaseCache): ) raise e - async def handle_lpop_count_for_older_redis_versions(self, pipe: pipeline, key: str, count: int) -> List[bytes]: - result: List[bytes] = [] + async def handle_lpop_count_for_older_redis_versions(self, pipe: pipeline, key: str, count: int) -> list[bytes]: + result: list[bytes] = [] for _ in range(count): pipe.lpop(key) results = await pipe.execute() @@ -1656,10 +1656,10 @@ class RedisCache(BaseCache): async def async_lpop( self, key: str, - count: Optional[int] = None, - parent_otel_span: Optional[Span] = None, + count: int | None = None, + parent_otel_span: Span | None = None, **kwargs, - ) -> Union[Any, List[Any]]: + ) -> Any | list[Any]: _redis_client: Any = self.init_async_client() key = self.check_and_fix_namespace(key=key) start_time = time.time() @@ -1711,14 +1711,14 @@ class RedisCache(BaseCache): call_type=f"async_lpop <- {_get_call_stack_info()}", ) ) - verbose_logger.error(f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}") + verbose_logger.error(f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {e!s}") raise e async def _pipeline_lpop_helper( self, pipe: pipeline, - lpop_list: List[RedisPipelineLpopOperation], - ) -> List[Optional[List[str]]]: + lpop_list: list[RedisPipelineLpopOperation], + ) -> list[list[str] | None]: """Helper function for pipeline lpop operations. For Redis >= 7, queues one LPOP(key, count) per operation. @@ -1734,7 +1734,7 @@ class RedisCache(BaseCache): else: # For Redis < 7, LPOP doesn't support count param. # Issue `count` individual LPOP commands per key, all in one pipeline. - counts: List[int] = [] + counts: list[int] = [] for lpop_op in lpop_list: key = self.check_and_fix_namespace(key=lpop_op["key"]) count = lpop_op["count"] or 1 @@ -1757,7 +1757,7 @@ class RedisCache(BaseCache): raise r # Decode bytes -> str for each result set - decoded_results: List[Optional[List[str]]] = [] + decoded_results: list[list[str] | None] = [] for r in raw_results: if r is None: decoded_results.append(None) @@ -1776,8 +1776,8 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_lpop_pipeline( self, - lpop_list: List[RedisPipelineLpopOperation], - ) -> List[Optional[List[str]]]: + lpop_list: list[RedisPipelineLpopOperation], + ) -> list[list[str] | None]: """ Use Redis Pipelines for bulk LPOP operations diff --git a/litellm/caching/redis_cluster_cache.py b/litellm/caching/redis_cluster_cache.py index 0698ebdcf2a..1e4c4684f48 100644 --- a/litellm/caching/redis_cluster_cache.py +++ b/litellm/caching/redis_cluster_cache.py @@ -5,7 +5,7 @@ Key differences: - RedisClient NEEDs to be re-used across requests, adds 3000ms latency if it's re-created """ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union from litellm.caching.redis_cache import RedisCache @@ -26,8 +26,8 @@ else: class RedisClusterCache(RedisCache): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.redis_async_redis_cluster_client: Optional[RedisCluster] = None - self.redis_sync_redis_cluster_client: Optional[RedisCluster] = None + self.redis_async_redis_cluster_client: RedisCluster | None = None + self.redis_sync_redis_cluster_client: RedisCluster | None = None def init_async_client(self): from redis.asyncio import RedisCluster @@ -43,13 +43,13 @@ class RedisClusterCache(RedisCache): return _redis_client - def _run_redis_mget_operation(self, keys: List[str]) -> List[Any]: + def _run_redis_mget_operation(self, keys: list[str]) -> list[Any]: """ Overrides `_run_redis_mget_operation` in redis_cache.py """ return self.redis_client.mget_nonatomic(keys=keys) # type: ignore - async def _async_run_redis_mget_operation(self, keys: List[str]) -> List[Any]: + async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]: """ Overrides `_async_run_redis_mget_operation` in redis_cache.py """ @@ -71,7 +71,7 @@ class RedisClusterCache(RedisCache): cluster_kwargs = self.redis_kwargs.copy() startup_nodes = cluster_kwargs.pop("startup_nodes", []) - new_startup_nodes: List[ClusterNode] = [] + new_startup_nodes: list[ClusterNode] = [] for item in startup_nodes: new_startup_nodes.append(ClusterNode(**item)) @@ -100,9 +100,9 @@ class RedisClusterCache(RedisCache): except Exception as e: from litellm._logging import verbose_logger - verbose_logger.error(f"Redis Cluster connection test failed: {str(e)}") + verbose_logger.error(f"Redis Cluster connection test failed: {e!s}") return { "status": "failed", - "message": f"Redis Cluster connection failed: {str(e)}", + "message": f"Redis Cluster connection failed: {e!s}", "error": str(e), } diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index d4288cc777c..b2d8efa1dba 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -13,7 +13,7 @@ import ast import asyncio import json import os -from typing import Any, Dict, List, Optional, Tuple, cast +from typing import Any, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -40,13 +40,13 @@ class RedisSemanticCache(BaseCache): def __init__( self, - host: Optional[str] = None, - port: Optional[str] = None, - password: Optional[str] = None, - redis_url: Optional[str] = None, - similarity_threshold: Optional[float] = None, + host: str | None = None, + port: str | None = None, + password: str | None = None, + redis_url: str | None = None, + similarity_threshold: float | None = None, embedding_model: str = "text-embedding-ada-002", - index_name: Optional[str] = None, + index_name: str | None = None, **kwargs, ): """ @@ -142,7 +142,7 @@ class RedisSemanticCache(BaseCache): raise @classmethod - def _cache_key_filterable_field(cls) -> Dict[str, str]: + def _cache_key_filterable_field(cls) -> dict[str, str]: return { "name": cls.CACHE_KEY_FIELD_NAME, "type": "tag", @@ -203,7 +203,7 @@ class RedisSemanticCache(BaseCache): overwrite=True, ) - def _get_cache_filters(self, key: str) -> Dict[str, str]: + def _get_cache_filters(self, key: str) -> dict[str, str]: return {self.CACHE_KEY_FIELD_NAME: str(key)} def _get_cache_key_filter_expression(self, key: str) -> Any: @@ -211,7 +211,7 @@ class RedisSemanticCache(BaseCache): return Tag(self.CACHE_KEY_FIELD_NAME) == str(key) - def _cache_hit_matches_key(self, cache_hit: Dict[str, Any], key: str) -> bool: + def _cache_hit_matches_key(self, cache_hit: dict[str, Any], key: str) -> bool: # Pre-isolation entries with no ``litellm_cache_key`` field cannot be # safely reassigned to a caller's scope and are treated as misses. cached_key = cache_hit.get(self.CACHE_KEY_FIELD_NAME) @@ -219,7 +219,7 @@ class RedisSemanticCache(BaseCache): cached_key = cached_key.decode("utf-8") return cached_key is not None and str(cached_key) == str(key) - def _get_ttl(self, **kwargs) -> Optional[int]: + def _get_ttl(self, **kwargs) -> int | None: """ Get the TTL (time-to-live) value for cache entries. @@ -235,7 +235,7 @@ class RedisSemanticCache(BaseCache): return ttl @classmethod - def _get_prompt_from_kwargs(cls, **kwargs) -> Optional[str]: + def _get_prompt_from_kwargs(cls, **kwargs) -> str | None: """ Extract a semantic-cache prompt from chat or Responses API request kwargs. """ @@ -246,13 +246,13 @@ class RedisSemanticCache(BaseCache): if "input" not in kwargs: return None - prompt_parts: List[str] = [] + prompt_parts: list[str] = [] cls._collect_responses_input_text(kwargs.get("input"), prompt_parts) prompt = "\n".join(prompt_parts).strip() return prompt or None @classmethod - def _collect_responses_input_text(cls, value: Any, prompt_parts: List[str]) -> None: + def _collect_responses_input_text(cls, value: Any, prompt_parts: list[str]) -> None: value = cls._coerce_response_input_value(value) if value is None: return @@ -306,7 +306,7 @@ class RedisSemanticCache(BaseCache): return dict_method() return value - def _get_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> List[float]: + def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]: """ Routes through the proxy Router when the embedding model is a Router deployment so per-deployment auth (e.g. Bedrock aws_role_name) applies, @@ -364,7 +364,7 @@ class RedisSemanticCache(BaseCache): try: cached_response = ast.literal_eval(cached_response) except (ValueError, SyntaxError) as e: - print_verbose(f"Error parsing cached response: {str(e)}") + print_verbose(f"Error parsing cached response: {e!s}") return None return cached_response @@ -381,7 +381,7 @@ class RedisSemanticCache(BaseCache): """ print_verbose(f"Redis semantic-cache set_cache, kwargs: {kwargs}") - value_str: Optional[str] = None + value_str: str | None = None try: prompt = self._get_prompt_from_kwargs(**kwargs) if prompt is None: @@ -403,7 +403,7 @@ class RedisSemanticCache(BaseCache): store_kwargs["ttl"] = int(ttl) self.llmcache.store(prompt, value_str, **store_kwargs) except Exception as e: - print_verbose(f"Error setting {value_str or value} in the Redis semantic cache: {str(e)}") + print_verbose(f"Error setting {value_str or value} in the Redis semantic cache: {e!s}") def get_cache(self, key: str, **kwargs) -> Any: """ @@ -468,10 +468,10 @@ class RedisSemanticCache(BaseCache): return self._get_cache_logic(cached_response=cached_response) except Exception as e: - print_verbose(f"Error retrieving from Redis semantic cache: {str(e)}") + print_verbose(f"Error retrieving from Redis semantic cache: {e!s}") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 - async def _get_async_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> List[float]: + async def _get_async_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]: """ Asynchronously generate an embedding for the given prompt. @@ -505,8 +505,8 @@ class RedisSemanticCache(BaseCache): ) return embedding_response["data"][0]["embedding"] except Exception as e: - print_verbose(f"Error generating async embedding: {str(e)}") - raise ValueError(f"Failed to generate embedding: {str(e)}") from e + print_verbose(f"Error generating async embedding: {e!s}") + raise ValueError(f"Failed to generate embedding: {e!s}") from e async def async_set_cache(self, key: str, value: Any, **kwargs) -> None: """ @@ -546,7 +546,7 @@ class RedisSemanticCache(BaseCache): **store_kwargs, ) except Exception as e: - print_verbose(f"Error in async_set_cache: {str(e)}") + print_verbose(f"Error in async_set_cache: {e!s}") async def async_get_cache(self, key: str, **kwargs) -> Any: """ @@ -612,10 +612,10 @@ class RedisSemanticCache(BaseCache): return self._get_cache_logic(cached_response=cached_response) except Exception as e: - print_verbose(f"Error in async_get_cache: {str(e)}") + print_verbose(f"Error in async_get_cache: {e!s}") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 - async def _index_info(self) -> Dict[str, Any]: + async def _index_info(self) -> dict[str, Any]: """ Get information about the Redis index. @@ -625,7 +625,7 @@ class RedisSemanticCache(BaseCache): aindex = await self.llmcache._get_async_index() return await aindex.info() - async def async_set_cache_pipeline(self, cache_list: List[Tuple[str, Any]], **kwargs) -> None: + async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs) -> None: """ Asynchronously store multiple values in the semantic cache. @@ -639,4 +639,4 @@ class RedisSemanticCache(BaseCache): tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) except Exception as e: - print_verbose(f"Error in async_set_cache_pipeline: {str(e)}") + print_verbose(f"Error in async_set_cache_pipeline: {e!s}") diff --git a/litellm/caching/s3_cache.py b/litellm/caching/s3_cache.py index 1ada940a9c9..5e185de7526 100644 --- a/litellm/caching/s3_cache.py +++ b/litellm/caching/s3_cache.py @@ -11,9 +11,8 @@ Has 4 methods: import ast import asyncio import json +from datetime import datetime, timedelta, timezone from functools import partial -from typing import Optional -from datetime import datetime, timezone, timedelta from litellm._logging import print_verbose, verbose_logger @@ -26,7 +25,7 @@ class S3Cache(BaseCache): s3_bucket_name, s3_region_name=None, s3_api_version=None, - s3_use_ssl: Optional[bool] = True, + s3_use_ssl: bool | None = True, s3_verify=None, s3_endpoint_url=None, s3_aws_access_key_id=None, diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index 76b7f7d5b87..86e687c0009 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -249,7 +249,7 @@ class ValkeySemanticCache(RedisSemanticCache): if ttl is not None: self.sync_client.expire(doc_key, ttl) except Exception as e: - print_verbose(f"Error in Valkey semantic-cache set_cache: {str(e)}") + print_verbose(f"Error in Valkey semantic-cache set_cache: {e!s}") def get_cache(self, key: str, **kwargs: Any) -> Any: print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}") @@ -268,7 +268,7 @@ class ValkeySemanticCache(RedisSemanticCache): ) return self._resolve_hit(self._first_hit(search_result), key, **kwargs) except Exception as e: - print_verbose(f"Error in Valkey semantic-cache get_cache: {str(e)}") + print_verbose(f"Error in Valkey semantic-cache get_cache: {e!s}") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None: @@ -288,7 +288,7 @@ class ValkeySemanticCache(RedisSemanticCache): if ttl is not None: await self.async_client.expire(doc_key, ttl) except Exception as e: - print_verbose(f"Error in async Valkey semantic-cache set_cache: {str(e)}") + print_verbose(f"Error in async Valkey semantic-cache set_cache: {e!s}") async def async_get_cache(self, key: str, **kwargs: Any) -> Any: print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}") @@ -307,14 +307,14 @@ class ValkeySemanticCache(RedisSemanticCache): ) return self._resolve_hit(self._first_hit(search_result), key, **kwargs) except Exception as e: - print_verbose(f"Error in async Valkey semantic-cache get_cache: {str(e)}") + print_verbose(f"Error in async Valkey semantic-cache get_cache: {e!s}") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: Any) -> None: try: await asyncio.gather(*[self.async_set_cache(key, value, **kwargs) for key, value in cache_list]) except Exception as e: - print_verbose(f"Error in Valkey semantic-cache async_set_cache_pipeline: {str(e)}") + print_verbose(f"Error in Valkey semantic-cache async_set_cache_pipeline: {e!s}") async def _index_info(self) -> dict: return await self.async_client.ft(self.index_name).info() diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 1b69154bf3d..6083bc8e26b 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -3,7 +3,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req """ from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from typing_extensions import TypedDict @@ -47,7 +47,7 @@ class ResponsesToCompletionBridgeHandler: @staticmethod def _coerce_response_object( response_obj: Any, - hidden_params: Optional[dict], + hidden_params: dict | None, ) -> "ResponsesAPIResponse": if isinstance(response_obj, ResponsesAPIResponse): response = response_obj diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 3e8c78c4741..3825854852d 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -8,11 +8,7 @@ from collections.abc import AsyncIterator, Callable, Iterable, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, Union, cast, ) @@ -61,7 +57,7 @@ if TYPE_CHECKING: def _get_reasoning_items( msg: "AllMessageValues", -) -> List[ChatCompletionReasoningItem]: +) -> list[ChatCompletionReasoningItem]: """Extract reasoning_items from a message dict with proper typing.""" items = msg.get("reasoning_items") # type: ignore[union-attr] if items: @@ -71,14 +67,14 @@ def _get_reasoning_items( def _build_reasoning_item( item_id: str, - encrypted_content: Optional[str], + encrypted_content: str | None, summary_raw: Any, -) -> Dict[str, Any]: +) -> dict[str, Any]: """Build a ChatCompletionReasoningItem-shaped dict from raw response data. Handles both pydantic objects (attribute access) and plain dicts. """ - summary: List[Dict[str, Any]] = [] + summary: list[dict[str, Any]] = [] for s in summary_raw or []: if isinstance(s, dict): summary.append({"type": s.get("type", "summary_text"), "text": s.get("text", "")}) @@ -98,10 +94,10 @@ def _build_reasoning_item( def _reasoning_item_to_response_input( - r_item: Union[ChatCompletionReasoningItem, Dict[str, Any]], -) -> Dict[str, Any]: + r_item: ChatCompletionReasoningItem | dict[str, Any], +) -> dict[str, Any]: """Convert a stored ChatCompletionReasoningItem back to a Responses API input item.""" - r_input: Dict[str, Any] = { + r_input: dict[str, Any] = { "type": "reasoning", "id": r_item.get("id") or f"rs_{id(r_item)}", # summary is always required by the Responses API, even when empty @@ -134,7 +130,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return {"type": "function", "name": fn_name} return tool_choice - def _handle_raw_dict_response_item(self, item: Dict[str, Any], index: int) -> Tuple[Optional[Any], int]: + def _handle_raw_dict_response_item(self, item: dict[str, Any], index: int) -> tuple[Any | None, int]: """ Handle raw dict response items from Responses API (e.g., GPT-5 Codex format). @@ -208,10 +204,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return None, index def convert_chat_completion_messages_to_responses_api( - self, messages: List["AllMessageValues"] - ) -> Tuple[List[Any], Optional[str]]: - input_items: List[Any] = [] - instructions: Optional[str] = None + self, messages: list["AllMessageValues"] + ) -> tuple[list[Any], str | None]: + input_items: list[Any] = [] + instructions: str | None = None for msg in messages: role = msg.get("role") @@ -242,7 +238,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Convert tool message to function call output format # The Responses API expects 'output' to be a list with input_text/input_image types # Using list format for consistency across text and multimodal content - tool_output: List[Dict[str, Any]] + tool_output: list[dict[str, Any]] if content is None: tool_output = [] elif isinstance(content, str): @@ -270,7 +266,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): for tool_call in tool_calls: function = tool_call.get("function") if function: - input_tool_call: Dict[str, Any] = { + input_tool_call: dict[str, Any] = { "type": "function_call", "call_id": tool_call["id"], } @@ -308,7 +304,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): responses_api_request["max_output_tokens"] = value elif key == "tools" and value is not None: responses_api_request["tools"] = self._convert_tools_to_responses_format( - cast(List[Dict[str, Any]], value) + cast(list[dict[str, Any]], value) ) elif key == "response_format": text_format = self._transform_response_format_to_text_format(value) @@ -331,15 +327,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif key == "web_search_options": self._add_web_search_tool(responses_api_request, value) - def _build_sanitized_litellm_params(self, litellm_params: dict) -> Dict[str, Any]: + def _build_sanitized_litellm_params(self, litellm_params: dict) -> dict[str, Any]: """Build sanitized litellm_params with merged metadata.""" responses_optional_param_keys = set(ResponsesAPIOptionalRequestParams.__annotations__.keys()) - sanitized: Dict[str, Any] = { + sanitized: dict[str, Any] = { key: value for key, value in litellm_params.items() if key not in responses_optional_param_keys } legacy_metadata = litellm_params.get("metadata") existing_litellm_metadata = litellm_params.get("litellm_metadata") - merged_litellm_metadata: Dict[str, Any] = {} + merged_litellm_metadata: dict[str, Any] = {} if isinstance(legacy_metadata, dict): merged_litellm_metadata.update(legacy_metadata) if isinstance(existing_litellm_metadata, dict): @@ -352,9 +348,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def _merge_responses_api_request_into_request_data( self, - request_data: Dict[str, Any], + request_data: dict[str, Any], responses_api_request: "ResponsesAPIOptionalRequestParams", - instructions: Optional[str], + instructions: str | None, ) -> None: """Add non-None values from responses_api_request into request_data.""" for key, value in responses_api_request.items(): @@ -374,12 +370,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def transform_request( self, model: str, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, headers: dict, litellm_logging_obj: "LiteLLMLoggingObj", - client: Optional[Any] = None, + client: Any | None = None, ) -> dict: ( input_items, @@ -453,9 +449,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): @staticmethod def _convert_response_output_to_choices( - output_items: List[Any], - handle_raw_dict_callback: Optional[Callable] = None, - ) -> List[Any]: + output_items: list[Any], + handle_raw_dict_callback: Callable | None = None, + ) -> list[Any]: """ Convert Responses API output items to chat completion choices. @@ -481,14 +477,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): from litellm.types.utils import Choices, Message - choices: List[Choices] = [] + choices: list[Choices] = [] index = 0 - reasoning_content: Optional[str] = None - pending_reasoning_item: Optional[Dict[str, Any]] = None + reasoning_content: str | None = None + pending_reasoning_item: dict[str, Any] | None = None # Collect all tool calls to put them in a single choice # (Chat Completions API expects all tool calls in one message) - accumulated_tool_calls: List[Dict[str, Any]] = [] + accumulated_tool_calls: list[dict[str, Any]] = [] tool_call_index = 0 for item in output_items: @@ -514,7 +510,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): reasoning_content=reasoning_content, annotations=annotations, reasoning_items=cast( - Optional[List[ChatCompletionReasoningItem]], + list[ChatCompletionReasoningItem] | None, ([pending_reasoning_item] if pending_reasoning_item is not None else None), ), ) @@ -574,7 +570,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tool_calls=accumulated_tool_calls, reasoning_content=reasoning_content, reasoning_items=cast( - Optional[List[ChatCompletionReasoningItem]], + list[ChatCompletionReasoningItem] | None, ([pending_reasoning_item] if pending_reasoning_item is not None else None), ), ) @@ -585,22 +581,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return choices @classmethod - def _extract_output_from_completed_event(cls, parsed_chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]: + def _extract_output_from_completed_event(cls, parsed_chunk: dict[str, Any]) -> list[dict[str, Any]] | None: response_payload = parsed_chunk.get("response") if not isinstance(response_payload, dict): return None response_output = response_payload.get("output") if not isinstance(response_output, list) or len(response_output) == 0: return None - return cast(List[Dict[str, Any]], response_output) + return cast(list[dict[str, Any]], response_output) @classmethod - def _recover_output_items_from_raw_sse(cls, raw_sse: Optional[str]) -> List[Dict[str, Any]]: + def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, Any]]: if not raw_sse or not isinstance(raw_sse, str): return [] - recovered_output_items: Dict[int, Dict[str, Any]] = {} - recovered_text_only_items: Dict[int, Dict[str, Any]] = {} + recovered_output_items: dict[int, dict[str, Any]] = {} + recovered_text_only_items: dict[int, dict[str, Any]] = {} for chunk in raw_sse.splitlines(): parsed_chunk = parse_sse_json_chunk(chunk) @@ -635,7 +631,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # but text-only items at indices without a matching OUTPUT_ITEM_DONE # must still be preserved (e.g. multi-output responses where some # indices only emitted OUTPUT_TEXT_DONE). - merged_items: Dict[int, Dict[str, Any]] = {**recovered_text_only_items} + merged_items: dict[int, dict[str, Any]] = {**recovered_text_only_items} merged_items.update(recovered_output_items) if merged_items: @@ -644,7 +640,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return [] @classmethod - def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> List[Dict[str, Any]]: + def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, Any]]: model_call_details = getattr(logging_obj, "model_call_details", {}) or {} original_response = model_call_details.get("original_response") return cls._recover_output_items_from_raw_sse(original_response) @@ -656,12 +652,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): model_response: "ModelResponse", logging_obj: "LiteLLMLoggingObj", request_data: dict, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> "ModelResponse": """Transform Responses API response to chat completion response""" from litellm.responses.utils import ResponseAPILoggingUtils @@ -729,11 +725,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): self, streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"], sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> BaseModelResponseIterator: return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode) - def _convert_content_str_to_input_text(self, content: str, role: str) -> Dict[str, Any]: + def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, Any]: if role == "user" or role == "system" or role == "tool": return {"type": "input_text", "text": content} else: @@ -745,15 +741,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): from openai.types.responses import ResponseInputImageParam content_image_url = content.get("image_url") - actual_image_url: Optional[str] = None - detail: Optional[Literal["low", "high", "auto"]] = None + actual_image_url: str | None = None + detail: Literal["low", "high", "auto"] | None = None if isinstance(content_image_url, str): actual_image_url = content_image_url elif isinstance(content_image_url, dict): actual_image_url = content_image_url.get("url") detail = cast( - Optional[Literal["low", "high", "auto"]], + Literal["low", "high", "auto"] | None, content_image_url.get("detail"), ) @@ -769,21 +765,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def _convert_content_to_responses_format( self, - content: Optional[ - Union[ - str, - List[Any], - Iterable[ - Union[ - "OpenAIMessageContentListBlock", - "ChatCompletionThinkingBlock", - "ChatCompletionRedactedThinkingBlock", - ] - ], - ] - ], + content: str + | list[Any] + | Iterable[ + Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"] + ] + | None, role: str, - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """Convert chat completion content to responses API format""" from litellm.types.llms.openai import ChatCompletionImageObject @@ -863,9 +852,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): verbose_logger.debug(f"Chat provider: Other content type -> {result}") return result - def _convert_tools_to_responses_format(self, tools: List[Dict[str, Any]]) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]: + def _convert_tools_to_responses_format(self, tools: list[dict[str, Any]]) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]: """Convert chat completion tools to responses API tools format""" - responses_tools: List["ALL_RESPONSES_API_TOOL_PARAMS"] = [] + responses_tools: list[ALL_RESPONSES_API_TOOL_PARAMS] = [] for tool in tools: # convert function tool from chat completion to responses API format if tool.get("type") == "function": @@ -882,7 +871,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): else: responses_tools.append(tool) # type: ignore - return cast(List["ALL_RESPONSES_API_TOOL_PARAMS"], responses_tools) + return cast(list["ALL_RESPONSES_API_TOOL_PARAMS"], responses_tools) def _extract_extra_body_params(self, optional_params: dict): """ @@ -913,7 +902,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return optional_params - def _map_reasoning_effort(self, reasoning_effort: Union[str, Dict[str, Any]]) -> Optional[Reasoning]: + def _map_reasoning_effort(self, reasoning_effort: str | dict[str, Any]) -> Reasoning | None: # If dict is passed, convert it directly to Reasoning object if isinstance(reasoning_effort, dict): return Reasoning(**reasoning_effort) # type: ignore[typeddict-item] @@ -964,16 +953,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tools = [] responses_api_request["tools"] = tools - web_search_tool: Dict[str, Any] = {"type": "web_search"} + web_search_tool: dict[str, Any] = {"type": "web_search"} if isinstance(web_search_options, dict): web_search_tool.update(web_search_options) # Cast to Any to match the expected union type for tools list items tools.append(cast(Any, web_search_tool)) - def _transform_response_format_to_text_format( - self, response_format: Union[Dict[str, Any], Any] - ) -> Optional[Dict[str, Any]]: + def _transform_response_format_to_text_format(self, response_format: dict[str, Any] | Any) -> dict[str, Any] | None: """ Transform Chat Completion response_format parameter to Responses API text.format parameter. @@ -1022,8 +1009,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): @staticmethod def _convert_annotations_to_chat_format( - annotations: Optional[List[Any]], - ) -> Optional[List[ChatCompletionAnnotation]]: + annotations: list[Any] | None, + ) -> list[ChatCompletionAnnotation] | None: """ Convert annotations from Responses API to Chat Completions format. @@ -1033,7 +1020,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if not annotations: return None - result: List[ChatCompletionAnnotation] = [] + result: list[ChatCompletionAnnotation] = [] for annotation in annotations: try: # Convert Pydantic models to dicts (handles both v1 and v2) @@ -1056,7 +1043,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return result if result else None - def _map_responses_status_to_finish_reason(self, status: Optional[str]) -> str: + def _map_responses_status_to_finish_reason(self, status: str | None) -> str: """Map responses API status to chat completion finish_reason""" if not status: return "stop" @@ -1072,7 +1059,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): - def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): + def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): super().__init__(streaming_response, sync_stream, json_mode) self._chat_completion_id: str | None = None @@ -1095,7 +1082,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): @staticmethod def translate_responses_chunk_to_openai_stream( - parsed_chunk: Union[dict, BaseModel], + parsed_chunk: dict | BaseModel, ) -> "ModelResponseStream": """ Translate a Responses API streaming chunk to OpenAI chat completion streaming format. @@ -1196,7 +1183,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ] ) elif event_type == "response.function_call_arguments.delta": - content_part: Optional[str] = parsed_chunk.get("delta", None) + content_part: str | None = parsed_chunk.get("delta", None) if content_part: tool_call_index = parsed_chunk.get("output_index", 0) return ModelResponseStream( @@ -1319,7 +1306,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): finish_reason = "tool_calls" if has_function_calls else "stop" # Extract reasoning items with encrypted_content for round-tripping - completed_reasoning_items: Optional[List[Dict[str, Any]]] = None + completed_reasoning_items: list[dict[str, Any]] | None = None for item in output_items: if not isinstance(item, dict) or item.get("type") != "reasoning": continue @@ -1333,7 +1320,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ) ) completed_reasoning_items_typed = cast( - Optional[List[ChatCompletionReasoningItem]], + list[ChatCompletionReasoningItem] | None, completed_reasoning_items, ) diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index 9c5f57bc98f..042db79a71a 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -4,7 +4,7 @@ scoring, message stubbing, and retrieval tool injection. """ from collections.abc import Mapping, Sequence -from typing import Any, Dict, List, Optional, Set, Tuple, Union, cast +from typing import Any, cast from litellm.caching.dual_cache import DualCache from litellm.compression.message_stubbing import ( @@ -33,7 +33,7 @@ _SUPPORTED_CALL_TYPES = frozenset( ) -def _normalize_call_type(call_type: Union[CallTypes, str]) -> str: +def _normalize_call_type(call_type: CallTypes | str) -> str: """Return the string value for a ``CallTypes`` enum or a raw string.""" if isinstance(call_type, CallTypes): return call_type.value @@ -44,7 +44,7 @@ def _is_anthropic_call_type(call_type: str) -> bool: return call_type in _ANTHROPIC_CALL_TYPES -def _build_retrieval_tools(keys: List[str], call_type: str) -> List[dict]: +def _build_retrieval_tools(keys: list[str], call_type: str) -> list[dict]: """ Build retrieval tool definitions in the target request schema. @@ -63,7 +63,7 @@ def _build_retrieval_tools(keys: List[str], call_type: str) -> List[dict]: from litellm.llms.anthropic.chat.transformation import AnthropicConfig anthropic_tools, _mcp_servers = AnthropicConfig()._map_tools(openai_tools) - return cast(List[dict], anthropic_tools) + return cast(list[dict], anthropic_tools) def _content_to_text(content: Any) -> str: @@ -77,8 +77,8 @@ def _content_to_text(content: Any) -> str: Implemented iteratively (stack-based) to avoid unbounded recursion. """ - parts: List[str] = [] - stack: List[Any] = [content] + parts: list[str] = [] + stack: list[Any] = [content] while stack: item = stack.pop() if isinstance(item, str): @@ -97,9 +97,9 @@ def _content_to_text(content: Any) -> str: def _normalize_messages_for_compression( - messages: List[dict], + messages: list[dict], call_type: str, -) -> Tuple[List[dict], List[dict]]: +) -> tuple[list[dict], list[dict]]: """ Normalize each original message to a text-surrogate content for scoring. @@ -111,9 +111,9 @@ def _normalize_messages_for_compression( f"Unsupported call_type={call_type!r} for compression. Expected one of: {sorted(_SUPPORTED_CALL_TYPES)}." ) - original_messages: List[Dict[str, Any]] = [dict(m) for m in messages] + original_messages: list[dict[str, Any]] = [dict(m) for m in messages] - normalized_messages: List[dict] = [] + normalized_messages: list[dict] = [] for msg in original_messages: normalized_messages.append( { @@ -124,7 +124,7 @@ def _normalize_messages_for_compression( return normalized_messages, original_messages -def _extract_last_user_message(messages: List[dict]) -> str: +def _extract_last_user_message(messages: list[dict]) -> str: """Return the text content of the last user message.""" for msg in reversed(messages): if msg.get("role") == "user": @@ -132,10 +132,10 @@ def _extract_last_user_message(messages: List[dict]) -> str: return "" -def _extract_tool_use_ids(content: Any) -> List[str]: +def _extract_tool_use_ids(content: Any) -> list[str]: if not isinstance(content, list): return [] - tool_use_ids: List[str] = [] + tool_use_ids: list[str] = [] for part in content: if not isinstance(part, dict): continue @@ -147,10 +147,10 @@ def _extract_tool_use_ids(content: Any) -> List[str]: return tool_use_ids -def _extract_tool_result_ids(content: Any) -> Set[str]: +def _extract_tool_result_ids(content: Any) -> set[str]: if not isinstance(content, list): return set() - tool_result_ids: Set[str] = set() + tool_result_ids: set[str] = set() for part in content: if not isinstance(part, dict): continue @@ -163,15 +163,15 @@ def _extract_tool_result_ids(content: Any) -> Set[str]: def _extract_anthropic_tool_exchange_spans( - messages: List[dict], -) -> Tuple[List[Set[int]], Optional[str]]: + messages: list[dict], +) -> tuple[list[set[int]], str | None]: """ Return atomic 2-message spans for Anthropic tool exchanges. Each assistant message containing `tool_use` must be immediately followed by a user message containing matching `tool_result` blocks for all tool_use ids. """ - spans: List[Set[int]] = [] + spans: list[set[int]] = [] i = 0 while i < len(messages): current = messages[i] @@ -223,13 +223,13 @@ def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int def _combine_scores( - bm25_scores: List[float], - emb_scores: List[float], + bm25_scores: list[float], + emb_scores: list[float], bm25_weight: float = 0.4, -) -> List[float]: +) -> list[float]: """Weighted average of BM25 and embedding scores, with min-max normalization.""" - def _normalize(scores: List[float]) -> List[float]: + def _normalize(scores: list[float]) -> list[float]: min_s = min(scores) if scores else 0.0 max_s = max(scores) if scores else 0.0 rng = max_s - min_s @@ -245,14 +245,14 @@ def _combine_scores( def _select_kept_indices_for_budget( - normalized_messages: List[dict], - original_messages: List[dict], - combined_scores: List[float], + normalized_messages: list[dict], + original_messages: list[dict], + combined_scores: list[float], compression_target: int, model: str, - initial_kept_indices: Set[int], - tool_exchange_spans: List[Set[int]], -) -> Tuple[Set[int], Dict[int, dict]]: + initial_kept_indices: set[int], + tool_exchange_spans: list[set[int]], +) -> tuple[set[int], dict[int, dict]]: kept_indices = set(initial_kept_indices) current_tokens = 0 for i in kept_indices: @@ -265,14 +265,14 @@ def _select_kept_indices_for_budget( # A unit is either: # 1) a single message index, or # 2) an Anthropic tool-exchange span that must be kept/dropped atomically. - truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict - span_id_by_index: Dict[int, int] = {} + truncated_overrides: dict[int, dict] = {} # idx -> truncated message dict + span_id_by_index: dict[int, int] = {} for span_id, span in enumerate(tool_exchange_spans): for idx in span: span_id_by_index[idx] = span_id # Build single-message candidate units (non-span messages). - candidate_units: List[Tuple[float, Tuple[int, ...], bool]] = [] + candidate_units: list[tuple[float, tuple[int, ...], bool]] = [] for idx in range(len(normalized_messages)): if idx in span_id_by_index or idx in kept_indices: continue @@ -322,8 +322,8 @@ def _select_kept_indices_for_budget( return kept_indices, truncated_overrides -def _get_dropped_tool_span_indices(kept_indices: Set[int], tool_exchange_spans: List[Set[int]]) -> Set[int]: - dropped_tool_span_indices: Set[int] = set() +def _get_dropped_tool_span_indices(kept_indices: set[int], tool_exchange_spans: list[set[int]]) -> set[int]: + dropped_tool_span_indices: set[int] = set() for span in tool_exchange_spans: if not any(idx in kept_indices for idx in span): dropped_tool_span_indices.update(span) @@ -331,14 +331,14 @@ def _get_dropped_tool_span_indices(kept_indices: Set[int], tool_exchange_spans: def compress( - messages: List[dict], + messages: list[dict], model: str, - call_type: Union[CallTypes, str] = CallTypes.completion, + call_type: CallTypes | str = CallTypes.completion, compression_trigger: int = 200_000, - compression_target: Optional[int] = None, - embedding_model: Optional[str] = None, - embedding_model_params: Optional[Dict[str, Any]] = None, - compression_cache: Optional[DualCache] = None, + compression_target: int | None = None, + embedding_model: str | None = None, + embedding_model_params: dict[str, Any] | None = None, + compression_cache: DualCache | None = None, ) -> CompressedResult: """ Compress a list of messages by replacing low-relevance content with stubs. @@ -383,7 +383,7 @@ def compress( original_tokens = token_counter( model=model, - messages=cast(List[Any], original_messages), + messages=cast(list[Any], original_messages), ) # Pass through if below trigger @@ -422,9 +422,9 @@ def compress( # Protected messages are never compressed protected_indices = get_protected_indices(normalized_messages) - kept_indices: Set[int] = set(protected_indices) + kept_indices: set[int] = set(protected_indices) - tool_exchange_spans: List[Set[int]] = [] + tool_exchange_spans: list[set[int]] = [] if _is_anthropic_call_type(call_type_str): tool_exchange_spans, tool_sequence_error = _extract_anthropic_tool_exchange_spans(original_messages) if tool_sequence_error is not None: @@ -454,9 +454,9 @@ def compress( ) # Build compressed messages and cache - compressed_messages: List[dict] = [] - cache: Dict[str, str] = {} - used_keys: Set[str] = set() + compressed_messages: list[dict] = [] + cache: dict[str, str] = {} + used_keys: set[str] = set() dropped_tool_span_indices = _get_dropped_tool_span_indices( kept_indices=kept_indices, tool_exchange_spans=tool_exchange_spans ) @@ -478,7 +478,7 @@ def compress( compressed_tokens = token_counter( model=model, - messages=cast(List[Any], compressed_messages), + messages=cast(list[Any], compressed_messages), ) return CompressedResult( diff --git a/litellm/compression/message_stubbing.py b/litellm/compression/message_stubbing.py index 8d4e65752c1..ebb2d19997a 100644 --- a/litellm/compression/message_stubbing.py +++ b/litellm/compression/message_stubbing.py @@ -3,7 +3,6 @@ Replace messages with compact stubs and extract human-readable keys. """ import re -from typing import Set from litellm.compression.content_detection import detect_content_type @@ -17,7 +16,7 @@ _FILE_PATH_PATTERNS = [ ] -def extract_key(message: dict, fallback_index: int, used_keys: Set[str]) -> str: +def extract_key(message: dict, fallback_index: int, used_keys: set[str]) -> str: """ Extract a human-readable key for the message. diff --git a/litellm/compression/retrieval_tool.py b/litellm/compression/retrieval_tool.py index 99431a2a15d..ed8c486ad10 100644 --- a/litellm/compression/retrieval_tool.py +++ b/litellm/compression/retrieval_tool.py @@ -2,10 +2,8 @@ Build the litellm_content_retrieve tool definition for the LLM. """ -from typing import List - -def build_retrieval_tool(available_keys: List[str]) -> dict: +def build_retrieval_tool(available_keys: list[str]) -> dict: """ Return an OpenAI-format tool definition that lets the model retrieve the full content of a compressed message. diff --git a/litellm/compression/scoring/bm25.py b/litellm/compression/scoring/bm25.py index 7f919ef16fb..1ff5962835c 100644 --- a/litellm/compression/scoring/bm25.py +++ b/litellm/compression/scoring/bm25.py @@ -7,10 +7,9 @@ No external dependencies — uses only stdlib. import math import re from collections import Counter -from typing import Dict, List -def _tokenize(text: str) -> List[str]: +def _tokenize(text: str) -> list[str]: """Split text into lowercase tokens on word boundaries.""" return re.findall(r"[a-z0-9_]+", text.lower()) @@ -33,10 +32,10 @@ def _extract_content(message: dict) -> str: def bm25_score_messages( query: str, - messages: List[dict], + messages: list[dict], k1: float = 1.5, b: float = 0.75, -) -> List[float]: +) -> list[float]: """ Score each message's relevance to the query using BM25 (Okapi BM25). @@ -54,7 +53,7 @@ def bm25_score_messages( return [0.0] * len(messages) # Tokenize all documents - doc_tokens: List[List[str]] = [] + doc_tokens: list[list[str]] = [] for msg in messages: doc_tokens.append(_tokenize(_extract_content(msg))) @@ -67,14 +66,14 @@ def bm25_score_messages( avgdl = sum(doc_lengths) / n if n > 0 else 1.0 # Document frequency for each term - df: Dict[str, int] = {} + df: dict[str, int] = {} for dt in doc_tokens: seen = set(dt) for term in seen: df[term] = df.get(term, 0) + 1 # IDF for query terms - idf: Dict[str, float] = {} + idf: dict[str, float] = {} for term in set(query_terms): term_df = df.get(term, 0) # Standard BM25 IDF: log((N - df + 0.5) / (df + 0.5) + 1) @@ -94,7 +93,7 @@ def bm25_score_messages( return sum(count for token, count in tf_counts.items() if token != query_term and token.startswith(query_term)) # Score each document - scores: List[float] = [] + scores: list[float] = [] for i, dt in enumerate(doc_tokens): if not dt: scores.append(0.0) diff --git a/litellm/compression/scoring/embedding_scorer.py b/litellm/compression/scoring/embedding_scorer.py index f3558ae8f5c..c65856a0414 100644 --- a/litellm/compression/scoring/embedding_scorer.py +++ b/litellm/compression/scoring/embedding_scorer.py @@ -5,7 +5,7 @@ Computes cosine similarity between the query embedding and each message embeddin """ import math -from typing import Any, Dict, List, Optional +from typing import Any from litellm.caching.dual_cache import DualCache @@ -34,7 +34,7 @@ def _truncate_text(text: str, max_chars: int = 30000) -> str: return text[:half] + "\n...\n" + text[-half:] -def _cosine_similarity(a: List[float], b: List[float]) -> float: +def _cosine_similarity(a: list[float], b: list[float]) -> float: """Compute cosine similarity between two vectors.""" dot = sum(x * y for x, y in zip(a, b)) norm_a = math.sqrt(sum(x * x for x in a)) @@ -46,11 +46,11 @@ def _cosine_similarity(a: List[float], b: List[float]) -> float: def embedding_score_messages( query: str, - messages: List[dict], + messages: list[dict], model: str, - cache: Optional[DualCache] = None, - embedding_model_params: Optional[Dict[str, Any]] = None, -) -> List[float]: + cache: DualCache | None = None, + embedding_model_params: dict[str, Any] | None = None, +) -> list[float]: """ Score each message's semantic similarity to the query using embeddings. @@ -74,7 +74,7 @@ def embedding_score_messages( # Filter out empty texts — replace with a placeholder to maintain indexing processed_texts = [t if t.strip() else "empty" for t in texts] - kwargs: Dict[str, Any] = { + kwargs: dict[str, Any] = { "model": model, "input": processed_texts, "caching": cache is not None, @@ -88,7 +88,7 @@ def embedding_score_messages( embeddings = [item["embedding"] for item in response.data] query_embedding = embeddings[0] - scores: List[float] = [] + scores: list[float] = [] for i in range(1, len(embeddings)): scores.append(_cosine_similarity(query_embedding, embeddings[i])) diff --git a/litellm/constants.py b/litellm/constants.py index 78bfc6501e8..b2eaa5722d0 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1,6 +1,6 @@ import os import sys -from typing import List, Literal, Optional +from typing import Literal from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none @@ -396,14 +396,14 @@ DEFAULT_A2A_AGENT_TIMEOUT: float = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", # Patterns that indicate a localhost/internal URL in A2A agent cards that should be # replaced with the original base_url. This is a common misconfiguration where # developers deploy agents with development URLs in their agent cards. -LOCALHOST_URL_PATTERNS: List[str] = [ +LOCALHOST_URL_PATTERNS: list[str] = [ "localhost", "127.0.0.1", "0.0.0.0", "[::1]", # IPv6 localhost ] # Patterns in error messages that indicate a connection failure -CONNECTION_ERROR_PATTERNS: List[str] = [ +CONNECTION_ERROR_PATTERNS: list[str] = [ "connect", "connection", "network", @@ -685,7 +685,7 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = { "context_management": None, } -openai_compatible_endpoints: List = [ +openai_compatible_endpoints: list = [ "api.perplexity.ai", "api.endpoints.anyscale.com/v1", "api.deepinfra.com/v1/openai", @@ -731,7 +731,7 @@ openai_compatible_endpoints: List = [ ] -openai_compatible_providers: List = [ +openai_compatible_providers: list = [ "anyscale", "groq", "nvidia_nim", @@ -796,7 +796,7 @@ openai_compatible_providers: List = [ "darkbloom", "meta", # Meta Model API (Muse Spark) - JSON-configured provider ] -openai_text_completion_compatible_providers: List = [ # providers that support `/v1/completions` +openai_text_completion_compatible_providers: list = [ # providers that support `/v1/completions` "together_ai", "fireworks_ai", "hosted_vllm", @@ -819,7 +819,7 @@ openai_text_completion_compatible_providers: List = [ # providers that support "hyperbolic", "wandb", ] -_openai_like_providers: List = [ +_openai_like_providers: list = [ "predibase", "databricks", "lemonade", @@ -1362,7 +1362,7 @@ try: _raw_background_health_check_max_tokens = ( _background_health_check_max_tokens_env.strip() if _background_health_check_max_tokens_env is not None else "" ) - BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = ( + BACKGROUND_HEALTH_CHECK_MAX_TOKENS: int | None = ( int(_raw_background_health_check_max_tokens) if _raw_background_health_check_max_tokens else None ) except (ValueError, TypeError): @@ -1376,7 +1376,7 @@ try: if _background_health_check_max_tokens_reasoning_env is not None else "" ) - BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING: Optional[int] = ( + BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING: int | None = ( int(_raw_background_health_check_max_tokens_reasoning) if _raw_background_health_check_max_tokens_reasoning else None diff --git a/litellm/containers/endpoint_factory.py b/litellm/containers/endpoint_factory.py index 046beaac791..5a1b8da351f 100644 --- a/litellm/containers/endpoint_factory.py +++ b/litellm/containers/endpoint_factory.py @@ -11,7 +11,7 @@ import json from collections.abc import Callable from functools import partial from pathlib import Path -from typing import Any, Dict, List, Literal, Optional, Type +from typing import Any, Literal import litellm from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT @@ -28,21 +28,21 @@ from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client # Response type mapping -RESPONSE_TYPES: Dict[str, Type] = { +RESPONSE_TYPES: dict[str, type] = { "ContainerFileListResponse": ContainerFileListResponse, "ContainerFileObject": ContainerFileObject, "DeleteContainerFileResponse": DeleteContainerFileResponse, } -def _load_endpoints_config() -> Dict: +def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" config_path = Path(__file__).parent / "endpoints.json" with open(config_path) as f: return json.load(f) -def create_sync_endpoint_function(endpoint_config: Dict) -> Callable: +def create_sync_endpoint_function(endpoint_config: dict) -> Callable: """ Create a sync SDK function from endpoint config. @@ -56,16 +56,16 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable: def endpoint_func( timeout: int = 600, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ): local_vars = locals() try: resolved_custom_llm_provider: str = custom_llm_provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response @@ -91,10 +91,8 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable: custom_llm_provider=resolved_custom_llm_provider, litellm_params=litellm_params, ) - container_provider_config: Optional[BaseContainerConfig] = ( - ProviderConfigManager.get_provider_container_config( - provider=litellm.LlmProviders(resolved_custom_llm_provider), - ) + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( + provider=litellm.LlmProviders(resolved_custom_llm_provider), ) if container_provider_config is None: @@ -139,7 +137,7 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable: def create_async_endpoint_function( sync_func: Callable, - endpoint_config: Dict, + endpoint_config: dict, ) -> Callable: """Create an async SDK function that wraps the sync function.""" @@ -147,9 +145,9 @@ def create_async_endpoint_function( async def async_endpoint_func( timeout: int = 600, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ): local_vars = locals() @@ -189,7 +187,7 @@ def create_async_endpoint_function( return async_endpoint_func -def generate_container_endpoints() -> Dict[str, Callable]: +def generate_container_endpoints() -> dict[str, Callable]: """ Generate all container endpoint functions from the JSON config. @@ -210,7 +208,7 @@ def generate_container_endpoints() -> Dict[str, Callable]: return endpoints -def get_all_endpoint_names() -> List[str]: +def get_all_endpoint_names() -> list[str]: """Get all endpoint names (sync and async) from config.""" config = _load_endpoints_config() names = [] @@ -220,7 +218,7 @@ def get_all_endpoint_names() -> List[str]: return names -def get_async_endpoint_names() -> List[str]: +def get_async_endpoint_names() -> list[str]: """Get all async endpoint names for router registration.""" config = _load_endpoints_config() return [endpoint["async_name"] for endpoint in config["endpoints"]] diff --git a/litellm/containers/main.py b/litellm/containers/main.py index b1114e080d3..466237ebce3 100644 --- a/litellm/containers/main.py +++ b/litellm/containers/main.py @@ -3,7 +3,7 @@ import contextvars import json from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Literal, Optional, Union, overload +from typing import Any, Literal, overload import litellm from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT @@ -48,16 +48,16 @@ __all__ = [ @client async def acreate_container( name: str, - expires_after: Optional[Dict[str, Any]] = None, - file_ids: Optional[List[str]] = None, + expires_after: dict[str, Any] | None = None, + file_ids: list[str] | None = None, timeout=600, # default to 10 minutes # LiteLLM specific params, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> ContainerObject: """Asynchronously calls the `create_container` function with the given arguments and keyword arguments. @@ -120,12 +120,12 @@ async def acreate_container( @overload def create_container( name: str, - expires_after: Optional[Dict[str, Any]] = None, - file_ids: Optional[List[str]] = None, + expires_after: dict[str, Any] | None = None, + file_ids: list[str] | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, acreate_container: Literal[True], @@ -137,12 +137,12 @@ def create_container( @overload def create_container( name: str, - expires_after: Optional[Dict[str, Any]] = None, - file_ids: Optional[List[str]] = None, + expires_after: dict[str, Any] | None = None, + file_ids: list[str] | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, acreate_container: Literal[False] = False, @@ -156,23 +156,20 @@ def create_container( @client def create_container( name: str, - expires_after: Optional[Dict[str, Any]] = None, - file_ids: Optional[List[str]] = None, + expires_after: dict[str, Any] | None = None, + file_ids: list[str] | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - ContainerObject, - Coroutine[Any, Any, ContainerObject], -]: +) -> ContainerObject | Coroutine[Any, Any, ContainerObject]: """Create a container using the OpenAI Container API. Currently supports OpenAI @@ -191,7 +188,7 @@ def create_container( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -212,7 +209,7 @@ def create_container( **kwargs, ) # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(custom_llm_provider), ) @@ -226,7 +223,7 @@ def create_container( ) # Get optional parameters for the container API - container_create_request_params: Dict = ContainerRequestUtils.get_optional_params_container_create( + container_create_request_params: dict = ContainerRequestUtils.get_optional_params_container_create( container_provider_config=container_provider_config, container_create_optional_params=container_create_optional_params, ) @@ -281,16 +278,16 @@ def create_container( ##### Container List ####################### @client async def alist_containers( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> ContainerListResponse: """Asynchronously list containers. @@ -351,13 +348,13 @@ async def alist_containers( @overload def list_containers( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, alist_containers: Literal[True], @@ -368,13 +365,13 @@ def list_containers( @overload def list_containers( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, alist_containers: Literal[False] = False, @@ -387,24 +384,21 @@ def list_containers( @client def list_containers( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - ContainerListResponse, - Coroutine[Any, Any, ContainerListResponse], -]: +) -> ContainerListResponse | Coroutine[Any, Any, ContainerListResponse]: """List containers using the OpenAI Container API. Currently supports OpenAI @@ -412,7 +406,7 @@ def list_containers( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -433,7 +427,7 @@ def list_containers( **kwargs, ) # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(custom_llm_provider), ) @@ -491,9 +485,9 @@ async def aretrieve_container( custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> ContainerObject: """Asynchronously retrieve a container. @@ -552,9 +546,9 @@ async def aretrieve_container( def retrieve_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, aretrieve_container: Literal[True], @@ -567,9 +561,9 @@ def retrieve_container( def retrieve_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, aretrieve_container: Literal[False] = False, @@ -584,20 +578,17 @@ def retrieve_container( def retrieve_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - ContainerObject, - Coroutine[Any, Any, ContainerObject], -]: +) -> ContainerObject | Coroutine[Any, Any, ContainerObject]: """Retrieve a container using the OpenAI Container API. Currently supports OpenAI @@ -606,7 +597,7 @@ def retrieve_container( try: resolved_custom_llm_provider: str = custom_llm_provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -637,7 +628,7 @@ def retrieve_container( was_encoded = original_container_id != container_id # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(resolved_custom_llm_provider), ) @@ -709,9 +700,9 @@ async def adelete_container( custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> DeleteContainerResult: """Asynchronously delete a container. @@ -770,9 +761,9 @@ async def adelete_container( def delete_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, adelete_container: Literal[True], @@ -785,9 +776,9 @@ def delete_container( def delete_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, adelete_container: Literal[False] = False, @@ -802,20 +793,17 @@ def delete_container( def delete_container( container_id: str, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - DeleteContainerResult, - Coroutine[Any, Any, DeleteContainerResult], -]: +) -> DeleteContainerResult | Coroutine[Any, Any, DeleteContainerResult]: """Delete a container using the OpenAI Container API. Currently supports OpenAI @@ -824,7 +812,7 @@ def delete_container( try: resolved_custom_llm_provider: str = custom_llm_provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -855,7 +843,7 @@ def delete_container( was_encoded = original_container_id != container_id # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(resolved_custom_llm_provider), ) @@ -923,14 +911,14 @@ def delete_container( @client async def alist_container_files( container_id: str, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> ContainerFileListResponse: """Asynchronously list files in a container. @@ -994,13 +982,13 @@ async def alist_container_files( @overload def list_container_files( container_id: str, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, alist_container_files: Literal[True], @@ -1012,13 +1000,13 @@ def list_container_files( @overload def list_container_files( container_id: str, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, alist_container_files: Literal[False] = False, @@ -1032,22 +1020,19 @@ def list_container_files( @client def list_container_files( container_id: str, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - ContainerFileListResponse, - Coroutine[Any, Any, ContainerFileListResponse], -]: +) -> ContainerFileListResponse | Coroutine[Any, Any, ContainerFileListResponse]: """List files in a container using the OpenAI Container API. Currently supports OpenAI @@ -1056,7 +1041,7 @@ def list_container_files( try: resolved_custom_llm_provider: str = custom_llm_provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -1085,7 +1070,7 @@ def list_container_files( ) # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(resolved_custom_llm_provider), ) @@ -1142,9 +1127,9 @@ async def aupload_container_file( file: FileTypes, timeout=600, # default to 10 minutes custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, ) -> ContainerFileObject: """Asynchronously upload a file to a container. @@ -1227,9 +1212,9 @@ def upload_container_file( container_id: str, file: FileTypes, timeout=600, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, aupload_container_file: Literal[True], @@ -1243,9 +1228,9 @@ def upload_container_file( container_id: str, file: FileTypes, timeout=600, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", *, aupload_container_file: Literal[False] = False, @@ -1261,18 +1246,15 @@ def upload_container_file( container_id: str, file: FileTypes, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, **kwargs, -) -> Union[ - ContainerFileObject, - Coroutine[Any, Any, ContainerFileObject], -]: +) -> ContainerFileObject | Coroutine[Any, Any, ContainerFileObject]: """Upload a file to a container using the OpenAI Container API. This endpoint allows uploading files directly to a container session, @@ -1310,7 +1292,7 @@ def upload_container_file( try: resolved_custom_llm_provider: str = custom_llm_provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -1339,7 +1321,7 @@ def upload_container_file( ) # get provider config - container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config( + container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( provider=litellm.LlmProviders(resolved_custom_llm_provider), ) diff --git a/litellm/containers/utils.py b/litellm/containers/utils.py index 2b115c6b3c4..9b48a4f0f98 100644 --- a/litellm/containers/utils.py +++ b/litellm/containers/utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional, TypeVar +from typing import Any, TypeVar from litellm.llms.base_llm.containers.transformation import BaseContainerConfig from litellm.responses.utils import ResponsesAPIRequestUtils @@ -61,7 +61,7 @@ class ContainerRequestUtils: def get_optional_params_container_create( container_provider_config: BaseContainerConfig, container_create_optional_params: ContainerCreateOptionalRequestParams, - ) -> Dict: + ) -> dict: """Get the optional parameters for container creation.""" supported_params = container_provider_config.get_supported_openai_params() @@ -97,9 +97,9 @@ class ContainerRequestUtils: @staticmethod def encode_container_id_in_response( response_obj: T, - custom_llm_provider: Optional[str], - litellm_metadata: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, + custom_llm_provider: str | None, + litellm_metadata: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, ) -> T: """ Encode container_id in response object with provider/model metadata for routing. @@ -124,7 +124,7 @@ class ContainerRequestUtils: """ # Extract model_id from litellm_metadata litellm_metadata = litellm_metadata or {} - model_info: Dict[str, Any] = litellm_metadata.get("model_info", {}) or {} + model_info: dict[str, Any] = litellm_metadata.get("model_info", {}) or {} model_id = model_info.get("id") # Check if we should encode based on routing metadata diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 96aed20529f..f10a9e327d6 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -3,7 +3,7 @@ import logging import time from functools import lru_cache -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, cast from httpx import Response from pydantic import BaseModel @@ -29,8 +29,8 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import ( _parse_prompt_tokens_details, calculate_cost_component, generic_cost_per_token, - get_token_type_cost_breakdown, get_billable_input_tokens, + get_token_type_cost_breakdown, select_cost_metric_for_model, ) from litellm.llms.anthropic.cost_calculation import ( @@ -52,9 +52,6 @@ from litellm.llms.databricks.cost_calculator import ( from litellm.llms.deepseek.cost_calculator import ( cost_per_token as deepseek_cost_per_token, ) -from litellm.llms.tencent.cost_calculator import ( - cost_per_token as tencent_cost_per_token, -) from litellm.llms.fireworks_ai.cost_calculator import ( cost_per_token as fireworks_ai_cost_per_token, ) @@ -64,12 +61,19 @@ from litellm.llms.lemonade.cost_calculator import ( ) from litellm.llms.openai.cost_calculation import ( _video_output_cost_per_second, +) +from litellm.llms.openai.cost_calculation import ( cost_per_second as openai_cost_per_second, +) +from litellm.llms.openai.cost_calculation import ( cost_per_token as openai_cost_per_token, ) from litellm.llms.perplexity.cost_calculator import ( cost_per_token as perplexity_cost_per_token, ) +from litellm.llms.tencent.cost_calculator import ( + cost_per_token as tencent_cost_per_token, +) from litellm.llms.together_ai.cost_calculator import get_model_params_and_category from litellm.llms.vertex_ai.cost_calculator import ( cost_per_character as google_cost_per_character, @@ -180,13 +184,13 @@ _MCP_CALL_TYPE = CallTypes.call_mcp_tool.value def _cost_per_token_custom_pricing_helper( prompt_tokens: float = 0, completion_tokens: float = 0, - response_time_ms: Optional[float] = 0.0, + response_time_ms: float | None = 0.0, cached_tokens: float = 0, cache_creation_tokens: float = 0, ### CUSTOM PRICING ### - custom_cost_per_token: Optional[CostPerToken] = None, - custom_cost_per_second: Optional[float] = None, -) -> Optional[Tuple[float, float]]: + custom_cost_per_token: CostPerToken | None = None, + custom_cost_per_second: float | None = None, +) -> tuple[float, float] | None: """Internal helper function for calculating cost, if custom pricing given. prompt_tokens is assumed to include both cached_tokens and cache_creation_tokens @@ -230,10 +234,10 @@ def _cost_per_token_custom_pricing_helper( def _get_additional_costs( model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, prompt_tokens: int, completion_tokens: int, -) -> Optional[dict]: +) -> dict | None: """ Calculate additional costs beyond standard token costs. @@ -275,7 +279,7 @@ def _get_additional_costs( def _transcription_usage_has_token_details( - usage_block: Optional[Usage], + usage_block: Usage | None, ) -> bool: if usage_block is None: return False @@ -297,35 +301,35 @@ def cost_per_token( model: str = "", prompt_tokens: int = 0, completion_tokens: int = 0, - response_time_ms: Optional[float] = 0.0, - custom_llm_provider: Optional[str] = None, + response_time_ms: float | None = 0.0, + custom_llm_provider: str | None = None, region_name=None, ### CHARACTER PRICING ### - prompt_characters: Optional[int] = None, - completion_characters: Optional[int] = None, + prompt_characters: int | None = None, + completion_characters: int | None = None, ### PROMPT CACHING PRICING ### - used for anthropic - cache_creation_input_tokens: Optional[int] = 0, - cache_read_input_tokens: Optional[int] = 0, + cache_creation_input_tokens: int | None = 0, + cache_read_input_tokens: int | None = 0, ### CUSTOM PRICING ### - custom_cost_per_token: Optional[CostPerToken] = None, - custom_cost_per_second: Optional[float] = None, + custom_cost_per_token: CostPerToken | None = None, + custom_cost_per_second: float | None = None, ### NUMBER OF QUERIES ### - number_of_queries: Optional[int] = None, + number_of_queries: int | None = None, ### USAGE OBJECT ### - usage_object: Optional[Usage] = None, # just read the usage object if provided + usage_object: Usage | None = None, # just read the usage object if provided ### BILLED UNITS ### - rerank_billed_units: Optional[RerankBilledUnits] = None, + rerank_billed_units: RerankBilledUnits | None = None, ### CALL TYPE ### call_type: CallTypesLiteral = "completion", audio_transcription_file_duration: float = 0.0, # for audio transcription calls - the file time in seconds ### SERVICE TIER ### - service_tier: Optional[str] = None, # for OpenAI service tier pricing + service_tier: str | None = None, # for OpenAI service tier pricing ### DATA RESIDENCY ### - data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") - response: Optional[Any] = None, + data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") + response: Any | None = None, ### REQUEST MODEL ### - request_model: Optional[str] = None, # original request model for router detection -) -> Tuple[float, float]: # type: ignore + request_model: str | None = None, # original request model for router detection +) -> tuple[float, float]: # type: ignore """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -489,12 +493,7 @@ def cost_per_token( if cost_metric == "cost_per_character": if prompt_characters is None: raise ValueError( - "prompt_characters must be provided for tts calls. prompt_characters={}, model={}, custom_llm_provider={}, call_type={}".format( - prompt_characters, - model, - custom_llm_provider, - call_type, - ) + f"prompt_characters must be provided for tts calls. prompt_characters={prompt_characters}, model={model}, custom_llm_provider={custom_llm_provider}, call_type={call_type}" ) _prompt_cost, _completion_cost = _generic_cost_per_character( model=model_without_prefix, @@ -506,14 +505,7 @@ def cost_per_token( ) if _prompt_cost is None or _completion_cost is None: raise ValueError( - "cost for tts call is None. prompt_cost={}, completion_cost={}, model={}, custom_llm_provider={}, prompt_characters={}, completion_characters={}".format( - _prompt_cost, - _completion_cost, - model_without_prefix, - custom_llm_provider, - prompt_characters, - completion_characters, - ) + f"cost for tts call is None. prompt_cost={_prompt_cost}, completion_cost={_completion_cost}, model={model_without_prefix}, custom_llm_provider={custom_llm_provider}, prompt_characters={prompt_characters}, completion_characters={completion_characters}" ) prompt_cost = _prompt_cost completion_cost = _completion_cost @@ -712,9 +704,9 @@ def has_hidden_params(obj: Any) -> bool: def _get_provider_for_cost_calc( - model: Optional[str], - custom_llm_provider: Optional[str] = None, -) -> Optional[str]: + model: str | None, + custom_llm_provider: str | None = None, +) -> str | None: if custom_llm_provider is not None: return custom_llm_provider if model is None: @@ -723,7 +715,7 @@ def _get_provider_for_cost_calc( _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model) except Exception as e: verbose_logger.debug( - f"litellm.cost_calculator.py::_get_provider_for_cost_calc() - Error inferring custom_llm_provider - {str(e)}" + f"litellm.cost_calculator.py::_get_provider_for_cost_calc() - Error inferring custom_llm_provider - {e!s}" ) return None @@ -731,13 +723,13 @@ def _get_provider_for_cost_calc( def _select_model_name_for_cost_calc( - model: Optional[str], - completion_response: Optional[Any], - base_model: Optional[str] = None, - custom_pricing: Optional[bool] = None, - custom_llm_provider: Optional[str] = None, - router_model_id: Optional[str] = None, -) -> Optional[str]: + model: str | None, + completion_response: Any | None, + base_model: str | None = None, + custom_pricing: bool | None = None, + custom_llm_provider: str | None = None, + router_model_id: str | None = None, +) -> str | None: """ 1. If custom pricing is true, return received model name 2. If base_model is set (e.g. for azure models), return that @@ -745,17 +737,17 @@ def _select_model_name_for_cost_calc( 4. Check if model is passed in return that """ - return_model: Optional[str] = None - region_name: Optional[str] = None + return_model: str | None = None + region_name: str | None = None custom_llm_provider = _get_provider_for_cost_calc(model=model, custom_llm_provider=custom_llm_provider) - completion_response_model: Optional[str] = None + completion_response_model: str | None = None if completion_response is not None: if isinstance(completion_response, BaseModel): completion_response_model = getattr(completion_response, "model", None) elif isinstance(completion_response, dict): completion_response_model = completion_response.get("model", None) - hidden_params: Optional[dict] = getattr(completion_response, "_hidden_params", None) + hidden_params: dict | None = getattr(completion_response, "_hidden_params", None) if custom_pricing is True: if router_model_id is not None and router_model_id in litellm.model_cost: @@ -808,7 +800,7 @@ def _model_contains_known_llm_provider(model: str) -> bool: return _provider_prefix in LlmProvidersSet -def _get_response_model(completion_response: Any) -> Optional[str]: +def _get_response_model(completion_response: Any) -> str | None: """ Extract the model name from a completion response object. @@ -837,7 +829,7 @@ _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: dict = { } -def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[str]: +def _map_traffic_type_to_service_tier(traffic_type: str | None) -> str | None: """ Map a Gemini usageMetadata.trafficType value to a LiteLLM service_tier string. @@ -872,9 +864,9 @@ def _normalize_service_tier(service_tier: object) -> str | None: def _get_usage_object( completion_response: Any, -) -> Optional[Usage]: +) -> Usage | None: usage_obj = cast( - Union[Usage, ResponseAPIUsage, dict, BaseModel], + Usage | ResponseAPIUsage | dict | BaseModel, ( completion_response.get("usage") if isinstance(completion_response, dict) @@ -895,7 +887,7 @@ def _get_usage_object( elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj): return TranscriptionUsageObjectTransformation.transform_transcription_usage_object( cast( - Union[TranscriptionUsageDurationObject, TranscriptionUsageTokensObject], + TranscriptionUsageDurationObject | TranscriptionUsageTokensObject, usage_obj, ) ) @@ -917,7 +909,7 @@ def _is_known_usage_objects(usage_obj): ) -def _infer_call_type(call_type: Optional[CallTypesLiteral], completion_response: Any) -> Optional[CallTypesLiteral]: +def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: Any) -> CallTypesLiteral | None: if call_type is not None: return call_type @@ -946,8 +938,8 @@ def _infer_call_type(call_type: Optional[CallTypesLiteral], completion_response: def _apply_cost_discount( base_cost: float, - custom_llm_provider: Optional[str], -) -> Tuple[float, float, float]: + custom_llm_provider: str | None, +) -> tuple[float, float, float]: """ Apply provider-specific cost discount from module-level config. @@ -980,8 +972,8 @@ def _apply_cost_discount( def _apply_cost_margin( base_cost: float, - custom_llm_provider: Optional[str], -) -> Tuple[float, float, float, float]: + custom_llm_provider: str | None, +) -> tuple[float, float, float, float]: """ Apply provider-specific or global cost margin from module-level config. @@ -1044,21 +1036,21 @@ def _apply_cost_margin( def _store_cost_breakdown_in_logging_obj( - litellm_logging_obj: Optional[LitellmLoggingObject], + litellm_logging_obj: LitellmLoggingObject | None, prompt_tokens_cost_usd_dollar: float, completion_tokens_cost_usd_dollar: float, cost_for_built_in_tools_cost_usd_dollar: float, total_cost_usd_dollar: float, - additional_costs: Optional[dict] = None, - original_cost: Optional[float] = None, - discount_percent: Optional[float] = None, - discount_amount: Optional[float] = None, - margin_percent: Optional[float] = None, - margin_fixed_amount: Optional[float] = None, - margin_total_amount: Optional[float] = None, - cache_read_cost: Optional[float] = None, - cache_creation_cost: Optional[float] = None, - reasoning_cost: Optional[float] = None, + additional_costs: dict | None = None, + original_cost: float | None = None, + discount_percent: float | None = None, + discount_amount: float | None = None, + margin_percent: float | None = None, + margin_fixed_amount: float | None = None, + margin_total_amount: float | None = None, + cache_read_cost: float | None = None, + cache_creation_cost: float | None = None, + reasoning_cost: float | None = None, ) -> None: """ Helper function to store cost breakdown in the logging object. @@ -1100,40 +1092,39 @@ def _store_cost_breakdown_in_logging_obj( ) except Exception as breakdown_error: - verbose_logger.debug(f"Error storing cost breakdown: {str(breakdown_error)}") + verbose_logger.debug(f"Error storing cost breakdown: {breakdown_error!s}") # Don't fail the main cost calculation if breakdown storage fails - pass def completion_cost( completion_response=None, - model: Optional[str] = None, + model: str | None = None, prompt="", - messages: List = [], + messages: list = [], completion="", - total_time: Optional[float] = 0.0, # used for replicate, sagemaker - call_type: Optional[CallTypesLiteral] = None, + total_time: float | None = 0.0, # used for replicate, sagemaker + call_type: CallTypesLiteral | None = None, ### REGION ### custom_llm_provider=None, region_name=None, # used for bedrock pricing ### IMAGE GEN ### - size: Optional[str] = None, - quality: Optional[str] = None, - n: Optional[int] = None, # number of images + size: str | None = None, + quality: str | None = None, + n: int | None = None, # number of images ### CUSTOM PRICING ### - custom_cost_per_token: Optional[CostPerToken] = None, - custom_cost_per_second: Optional[float] = None, - optional_params: Optional[dict] = None, - custom_pricing: Optional[bool] = None, - base_model: Optional[str] = None, - standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None, - litellm_model_name: Optional[str] = None, - router_model_id: Optional[str] = None, - litellm_logging_obj: Optional[LitellmLoggingObject] = None, + custom_cost_per_token: CostPerToken | None = None, + custom_cost_per_second: float | None = None, + optional_params: dict | None = None, + custom_pricing: bool | None = None, + base_model: str | None = None, + standard_built_in_tools_params: StandardBuiltInToolsParams | None = None, + litellm_model_name: str | None = None, + router_model_id: str | None = None, + litellm_logging_obj: LitellmLoggingObject | None = None, ### SERVICE TIER ### - service_tier: Optional[str] = None, # for OpenAI service tier pricing + service_tier: str | None = None, # for OpenAI service tier pricing ### DATA RESIDENCY ### - data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") + data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") ) -> float: """ Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm. @@ -1176,14 +1167,14 @@ def completion_cost( model = "dall-e-2" # for dall-e-2, azure expects an empty model name # Handle Inputs to completion_cost prompt_tokens = 0 - prompt_characters: Optional[int] = None + prompt_characters: int | None = None completion_tokens = 0 - completion_characters: Optional[int] = None - cache_creation_input_tokens: Optional[int] = None - cache_read_input_tokens: Optional[int] = None + completion_characters: int | None = None + cache_creation_input_tokens: int | None = None + cache_read_input_tokens: int | None = None audio_transcription_file_duration: float = 0.0 - cost_per_token_usage_object: Optional[Usage] = _get_usage_object(completion_response=completion_response) - rerank_billed_units: Optional[RerankBilledUnits] = None + cost_per_token_usage_object: Usage | None = _get_usage_object(completion_response=completion_response) + rerank_billed_units: RerankBilledUnits | None = None # Extract service_tier from optional_params if not provided directly if service_tier is None and optional_params is not None: @@ -1234,7 +1225,7 @@ def completion_cost( isinstance(completion_response, BaseModel) or isinstance(completion_response, dict) ): # tts returns a custom class if isinstance(completion_response, dict): - usage_obj: Optional[Union[dict, Usage]] = completion_response.get("usage", {}) + usage_obj: dict | Usage | None = completion_response.get("usage", {}) else: usage_obj = getattr(completion_response, "usage", {}) if isinstance(usage_obj, BaseModel) and not _is_known_usage_objects(usage_obj=usage_obj): @@ -1258,10 +1249,7 @@ def completion_cost( elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(_usage): tr_usage = TranscriptionUsageObjectTransformation.transform_transcription_usage_object( cast( - Union[ - TranscriptionUsageDurationObject, - TranscriptionUsageTokensObject, - ], + TranscriptionUsageDurationObject | TranscriptionUsageTokensObject, _usage, ) ) @@ -1327,9 +1315,7 @@ def completion_cost( ) # strip the llm provider from the model name -> for image gen cost calculation except Exception as e: verbose_logger.debug( - "litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - {}".format( - str(e) - ) + f"litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - {e!s}" ) if CostCalculatorUtils._call_type_has_image_response(call_type) and isinstance( completion_response, ImageResponse @@ -1348,7 +1334,7 @@ def completion_cost( elif call_type in _VIDEO_CALL_TYPES: ### VIDEO GENERATION COST CALCULATION ### # Extract custom model_info for deployment-specific pricing - _video_model_info: Optional[ModelInfo] = None + _video_model_info: ModelInfo | None = None if custom_pricing and litellm_logging_obj is not None: _litellm_params = getattr(litellm_logging_obj, "litellm_params", None) if _litellm_params is not None: @@ -1356,8 +1342,8 @@ def completion_cost( _video_model_info = _metadata.get("model_info", None) usage_obj = getattr(completion_response, "usage", None) - duration_seconds: Optional[float] = None - video_resolution: Optional[str] = None + duration_seconds: float | None = None + video_resolution: str | None = None if completion_response is not None and usage_obj: # Handle both dict and Pydantic Usage object if isinstance(usage_obj, dict): @@ -1491,10 +1477,7 @@ def completion_cost( ): if cost_per_token_usage_object is None or custom_llm_provider is None: raise ValueError( - "usage object and custom_llm_provider must be provided for realtime stream cost calculation. Got cost_per_token_usage_object={}, custom_llm_provider={}".format( - cost_per_token_usage_object, - custom_llm_provider, - ) + f"usage object and custom_llm_provider must be provided for realtime stream cost calculation. Got cost_per_token_usage_object={cost_per_token_usage_object}, custom_llm_provider={custom_llm_provider}" ) return handle_realtime_stream_cost_calculation( results=completion_response.results, @@ -1577,13 +1560,15 @@ def completion_cost( if completion_response is not None: hidden_params = getattr(completion_response, "_hidden_params", None) or {} hidden_model = hidden_params.get("model") or hidden_params.get("litellm_model_name") - if hidden_model and ( - "model_router" in (hidden_model or "").lower() - or "model-router" in (hidden_model or "").lower() + if ( + hidden_model + and ( + "model_router" in (hidden_model or "").lower() + or "model-router" in (hidden_model or "").lower() + ) + or model_for_additional_costs is None ): model_for_additional_costs = hidden_model - elif model_for_additional_costs is None: - model_for_additional_costs = hidden_model if model_for_additional_costs is None: model_for_additional_costs = model additional_costs = _get_additional_costs( @@ -1639,11 +1624,11 @@ def completion_cost( # Store cost breakdown in logging object if available if litellm_logging_obj is not None: - _reasoning_cost: Optional[float] = None - _cache_read_cost: Optional[float] = None - _cache_creation_cost: Optional[float] = None + _reasoning_cost: float | None = None + _cache_read_cost: float | None = None + _cache_creation_cost: float | None = None if cost_per_token_usage_object is not None and model: - _breakdown_provider: Optional[str] = ( + _breakdown_provider: str | None = ( custom_llm_provider if isinstance(custom_llm_provider, str) else None ) _token_type_breakdown = get_token_type_cost_breakdown( @@ -1677,20 +1662,18 @@ def completion_cost( return _final_cost except Exception as e: verbose_logger.debug( - "litellm.cost_calculator.py::completion_cost() - Error calculating cost for model={} - {}".format( - model, str(e) - ) + f"litellm.cost_calculator.py::completion_cost() - Error calculating cost for model={model} - {e!s}" ) if idx == len(potential_model_names) - 1: raise e - raise Exception("Unable to calculat cost for received potential model names - {}".format(potential_model_names)) + raise Exception(f"Unable to calculat cost for received potential model names - {potential_model_names}") except Exception as e: raise e def get_response_cost_from_hidden_params( - hidden_params: Union[dict, BaseModel], -) -> Optional[float]: + hidden_params: dict | BaseModel, +) -> float | None: if isinstance(hidden_params, BaseModel): _hidden_params_dict = cast(BaseModel, hidden_params).model_dump() else: @@ -1706,22 +1689,20 @@ def get_response_cost_from_hidden_params( def response_cost_calculator( - response_object: Union[ - ModelResponse, - EmbeddingResponse, - ImageResponse, - TranscriptionResponse, - TextCompletionResponse, - HttpxBinaryResponseContent, - RerankResponse, - ResponsesAPIResponse, - LiteLLMRealtimeStreamLoggingObject, - OpenAIModerationResponse, - Response, - SearchResponse, - ], + response_object: ModelResponse + | EmbeddingResponse + | ImageResponse + | TranscriptionResponse + | TextCompletionResponse + | HttpxBinaryResponseContent + | RerankResponse + | ResponsesAPIResponse + | LiteLLMRealtimeStreamLoggingObject + | OpenAIModerationResponse + | Response + | SearchResponse, model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, call_type: Literal[ "embedding", "aembedding", @@ -1743,18 +1724,18 @@ def response_cost_calculator( "asearch", ], optional_params: dict, - cache_hit: Optional[bool] = None, - base_model: Optional[str] = None, - custom_pricing: Optional[bool] = None, + cache_hit: bool | None = None, + base_model: str | None = None, + custom_pricing: bool | None = None, prompt: str = "", - standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None, - litellm_model_name: Optional[str] = None, - router_model_id: Optional[str] = None, - litellm_logging_obj: Optional[LitellmLoggingObject] = None, + standard_built_in_tools_params: StandardBuiltInToolsParams | None = None, + litellm_model_name: str | None = None, + router_model_id: str | None = None, + litellm_logging_obj: LitellmLoggingObject | None = None, ### SERVICE TIER ### - service_tier: Optional[str] = None, # for OpenAI service tier pricing + service_tier: str | None = None, # for OpenAI service tier pricing ### DATA RESIDENCY ### - data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") + data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") ) -> float: """ Returns @@ -1795,9 +1776,9 @@ def response_cost_calculator( def ocr_cost( model: str, - custom_llm_provider: Optional[str], - response: Optional[Any] = None, -) -> Tuple[float, float]: + custom_llm_provider: str | None, + response: Any | None = None, +) -> tuple[float, float]: """ Args: model: str - model name @@ -1821,7 +1802,7 @@ def ocr_cost( raise ValueError("OCR response usage_info is None") try: - model_info: Optional[ModelInfo] = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + model_info: ModelInfo | None = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) except Exception: model_info = None @@ -1832,7 +1813,7 @@ def ocr_cost( if credits is not None and cost_per_credit is not None: return cost_per_credit * credits, 0.0 - ocr_cost_per_page: Optional[float] = None + ocr_cost_per_page: float | None = None if model_info is not None: ocr_cost_per_page = model_info.get("ocr_cost_per_page") @@ -1874,15 +1855,15 @@ def ocr_cost( def vector_store_search_cost( - model: Optional[str], + model: str | None, custom_llm_provider: str, response: VectorStoreSearchResponse, -) -> Tuple[float, float]: +) -> tuple[float, float]: """ Returns - float or None: cost of vector store search """ - api_type: Optional[str] = None + api_type: str | None = None if custom_llm_provider is None: custom_llm_provider = "openai" @@ -1907,9 +1888,9 @@ def vector_store_search_cost( def rerank_cost( model: str, - custom_llm_provider: Optional[str], - billed_units: Optional[RerankBilledUnits] = None, -) -> Tuple[float, float]: + custom_llm_provider: str | None, + billed_units: RerankBilledUnits | None = None, +) -> tuple[float, float]: """ Returns - float or None: cost of response OR none if error. @@ -1925,9 +1906,7 @@ def rerank_cost( ) try: - model_info: Optional[ModelInfo] = litellm.get_model_info( - model=model, custom_llm_provider=custom_llm_provider - ) + model_info: ModelInfo | None = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) except Exception: model_info = None @@ -1941,17 +1920,17 @@ def rerank_cost( raise e -def transcription_cost(model: str, custom_llm_provider: Optional[str], duration: float) -> Tuple[float, float]: +def transcription_cost(model: str, custom_llm_provider: str | None, duration: float) -> tuple[float, float]: return openai_cost_per_second(model=model, custom_llm_provider=custom_llm_provider, duration=duration) def default_image_cost_calculator( model: str, - custom_llm_provider: Optional[str] = None, - quality: Optional[str] = None, - n: Optional[int] = 1, # Default to 1 image - size: Optional[str] = "1024-x-1024", # OpenAI default - optional_params: Optional[dict] = None, + custom_llm_provider: str | None = None, + quality: str | None = None, + n: int | None = 1, # Default to 1 image + size: str | None = "1024-x-1024", # OpenAI default + optional_params: dict | None = None, ) -> float: """ Default image cost calculator for image generation @@ -1978,7 +1957,7 @@ def default_image_cost_calculator( # Build model names for cost lookup base_model_name = f"{size_str}/{model}" - model_name_without_custom_llm_provider: Optional[str] = None + model_name_without_custom_llm_provider: str | None = None if custom_llm_provider and model.startswith(f"{custom_llm_provider}/"): model_name_without_custom_llm_provider = model.replace(f"{custom_llm_provider}/", "") base_model_name = f"{custom_llm_provider}/{size_str}/{model_name_without_custom_llm_provider}" @@ -1993,8 +1972,8 @@ def default_image_cost_calculator( model_with_quality_without_provider = f"{quality}/{model_without_provider}" if quality else model_without_provider # Try model with quality first, fall back to base model name - cost_info: Optional[dict] = None - models_to_check: List[Optional[str]] = [ + cost_info: dict | None = None + models_to_check: list[str | None] = [ model_name_with_quality, base_model_name, model_name_with_v2_quality, @@ -2023,9 +2002,9 @@ def default_image_cost_calculator( def default_video_cost_calculator( model: str, duration_seconds: float, - custom_llm_provider: Optional[str] = None, - model_info: Optional[ModelInfo] = None, - video_resolution: Optional[str] = None, + custom_llm_provider: str | None = None, + model_info: ModelInfo | None = None, + video_resolution: str | None = None, ) -> float: """ Default video cost calculator for video generation @@ -2046,13 +2025,13 @@ def default_video_cost_calculator( Exception: If model pricing not found in cost map """ # Use custom model_info pricing if provided (deployment-specific pricing) - cost_info: Optional[dict] = None + cost_info: dict | None = None if model_info is not None: cost_info = dict(model_info) else: # Build model names for cost lookup base_model_name = model - model_name_without_custom_llm_provider: Optional[str] = None + model_name_without_custom_llm_provider: str | None = None if custom_llm_provider and model.startswith(f"{custom_llm_provider}/"): model_name_without_custom_llm_provider = model.replace(f"{custom_llm_provider}/", "") base_model_name = f"{custom_llm_provider}/{model_name_without_custom_llm_provider}" @@ -2062,7 +2041,7 @@ def default_video_cost_calculator( model_without_provider = model.split("/")[-1] # Try model with provider first, fall back to base model name - models_to_check: List[Optional[str]] = [ + models_to_check: list[str | None] = [ base_model_name, model, model_without_provider, @@ -2101,10 +2080,10 @@ def default_video_cost_calculator( def batch_cost_calculator( usage: Usage, model: str, - custom_llm_provider: Optional[str] = None, - model_info: Optional[ModelInfo] = None, - data_residency: Optional[str] = None, -) -> Tuple[float, float]: + custom_llm_provider: str | None = None, + model_info: ModelInfo | None = None, + data_residency: str | None = None, +) -> tuple[float, float]: """ Calculate the cost of a batch job. @@ -2191,7 +2170,7 @@ def batch_cost_calculator( return total_prompt_cost, total_completion_cost -def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> List[str]: +def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]: field_names = list(type(prompt_tokens_details).model_fields) if getattr(prompt_tokens_details, "cache_write_tokens", None) is None: return field_names @@ -2200,7 +2179,7 @@ def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> List[str] class BaseTokenUsageProcessor: @staticmethod - def combine_usage_objects(usage_objects: List[Usage]) -> Usage: + def combine_usage_objects(usage_objects: list[Usage]) -> Usage: """ Combine multiple Usage objects into a single Usage object, checking model keys for nested values. """ @@ -2272,15 +2251,15 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): @staticmethod def collect_usage_from_realtime_stream_results( results: OpenAIRealtimeStreamList, - ) -> List[Usage]: + ) -> list[Usage]: """ Collect usage from realtime stream results """ - response_done_events: List[OpenAIRealtimeStreamResponseBaseObject] = cast( - List[OpenAIRealtimeStreamResponseBaseObject], + response_done_events: list[OpenAIRealtimeStreamResponseBaseObject] = cast( + list[OpenAIRealtimeStreamResponseBaseObject], [result for result in results if result["type"] == "response.done"], ) - usage_objects: List[Usage] = [] + usage_objects: list[Usage] = [] for result in response_done_events: usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result["response"].get("usage", {}) @@ -2317,8 +2296,8 @@ def handle_realtime_stream_cost_calculation( combined_usage_object: Usage, custom_llm_provider: str, litellm_model_name: str, - data_residency: Optional[str] = None, - litellm_logging_obj: Optional[LitellmLoggingObject] = None, + data_residency: str | None = None, + litellm_logging_obj: LitellmLoggingObject | None = None, ) -> float: """ Handles the cost calculation for realtime stream responses. @@ -2412,7 +2391,7 @@ def handle_realtime_transcription_cost_calculation( def _get_transcription_model_name_from_results( results: OpenAIRealtimeStreamList, -) -> Optional[str]: +) -> str | None: """Resolve the ASR model from a transcription_session.* / session.* event.""" for result in results: if result.get("type") in ( @@ -2431,7 +2410,7 @@ def _get_transcription_model_name_from_results( return None -def _transcription_usage_cost(usage: dict, model_info: Optional[ModelInfo]) -> float: +def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float: if model_info is None: return 0.0 usage_type = usage.get("type") diff --git a/litellm/endpoints/speech/speech_to_completion_bridge/handler.py b/litellm/endpoints/speech/speech_to_completion_bridge/handler.py index f2b443eb7bf..babb1811d1b 100644 --- a/litellm/endpoints/speech/speech_to_completion_bridge/handler.py +++ b/litellm/endpoints/speech/speech_to_completion_bridge/handler.py @@ -2,7 +2,7 @@ Handler for transforming /chat/completions api requests to litellm.responses requests """ -from typing import TYPE_CHECKING, Optional, Union +from typing import TYPE_CHECKING from typing_extensions import TypedDict @@ -14,7 +14,7 @@ if TYPE_CHECKING: class SpeechToCompletionBridgeHandlerInputKwargs(TypedDict): model: str input: str - voice: Optional[Union[str, dict]] + voice: str | dict | None optional_params: dict litellm_params: dict logging_obj: "LiteLLMLoggingObj" @@ -79,7 +79,7 @@ class SpeechToCompletionBridgeHandler: self, model: str, input: str, - voice: Optional[Union[str, dict]], + voice: str | dict | None, optional_params: dict, litellm_params: dict, headers: dict, @@ -120,7 +120,7 @@ class SpeechToCompletionBridgeHandler: model_response=result, ) else: - raise Exception("Unmapped response type. Got type: {}".format(type(result))) + raise Exception(f"Unmapped response type. Got type: {type(result)}") speech_to_completion_bridge_handler = SpeechToCompletionBridgeHandler() diff --git a/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py b/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py index 94de4878b65..2f2861dfc26 100644 --- a/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py +++ b/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Optional, Union, cast +from typing import TYPE_CHECKING, cast from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS @@ -13,7 +13,7 @@ class SpeechToCompletionBridgeTransformationHandler: self, model: str, input: str, - voice: Optional[Union[str, dict]], + voice: str | dict | None, optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/evals/__init__.py b/litellm/evals/__init__.py index 89dfb62b2b7..14311ded659 100644 --- a/litellm/evals/__init__.py +++ b/litellm/evals/__init__.py @@ -18,16 +18,16 @@ from .main import ( ) __all__ = [ - "acreate_eval", - "alist_evals", - "aget_eval", - "aupdate_eval", - "adelete_eval", "acancel_eval", - "create_eval", - "list_evals", - "get_eval", - "update_eval", - "delete_eval", + "acreate_eval", + "adelete_eval", + "aget_eval", + "alist_evals", + "aupdate_eval", "cancel_eval", + "create_eval", + "delete_eval", + "get_eval", + "list_evals", + "update_eval", ] diff --git a/litellm/evals/main.py b/litellm/evals/main.py index 078b9c7002f..0bcccb73cb1 100644 --- a/litellm/evals/main.py +++ b/litellm/evals/main.py @@ -7,7 +7,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -41,15 +41,15 @@ DEFAULT_OPENAI_API_BASE = "https://api.openai.com" @client async def acreate_eval( - data_source_config: Dict[str, Any], - testing_criteria: List[Dict[str, Any]], - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + data_source_config: dict[str, Any], + testing_criteria: list[dict[str, Any]], + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Eval: """ @@ -110,17 +110,17 @@ async def acreate_eval( @client def create_eval( - data_source_config: Dict[str, Any], - testing_criteria: List[Dict[str, Any]], - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + data_source_config: dict[str, Any], + testing_criteria: list[dict[str, Any]], + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Eval, Coroutine[Any, Any, Eval]]: +) -> Eval | Coroutine[Any, Any, Eval]: """ Create a new evaluation @@ -142,7 +142,7 @@ def create_eval( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acreate_eval", False) is True # Get LiteLLM parameters @@ -153,7 +153,7 @@ def create_eval( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -226,15 +226,15 @@ def create_eval( @client async def alist_evals( - limit: Optional[int] = None, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - order_by: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + after: str | None = None, + before: str | None = None, + order: str | None = None, + order_by: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ListEvalsResponse: """ @@ -295,17 +295,17 @@ async def alist_evals( @client def list_evals( - limit: Optional[int] = None, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - order_by: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + after: str | None = None, + before: str | None = None, + order: str | None = None, + order_by: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ListEvalsResponse, Coroutine[Any, Any, ListEvalsResponse]]: +) -> ListEvalsResponse | Coroutine[Any, Any, ListEvalsResponse]: """ List all evaluations @@ -327,7 +327,7 @@ def list_evals( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("alist_evals", False) is True # Get LiteLLM parameters @@ -338,7 +338,7 @@ def list_evals( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -413,10 +413,10 @@ def list_evals( @client async def aget_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Eval: """ @@ -470,12 +470,12 @@ async def aget_eval( @client def get_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Eval, Coroutine[Any, Any, Eval]]: +) -> Eval | Coroutine[Any, Any, Eval]: """ Get an evaluation by ID @@ -493,7 +493,7 @@ def get_eval( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aget_eval", False) is True # Get LiteLLM parameters @@ -504,7 +504,7 @@ def get_eval( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -563,13 +563,13 @@ def get_eval( @client async def aupdate_eval( eval_id: str, - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Eval: """ @@ -629,15 +629,15 @@ async def aupdate_eval( @client def update_eval( eval_id: str, - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Eval, Coroutine[Any, Any, Eval]]: +) -> Eval | Coroutine[Any, Any, Eval]: """ Update an evaluation @@ -658,7 +658,7 @@ def update_eval( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aupdate_eval", False) is True # Get LiteLLM parameters @@ -669,7 +669,7 @@ def update_eval( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -783,10 +783,10 @@ def update_eval( @client async def adelete_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> DeleteEvalResponse: """ @@ -840,12 +840,12 @@ async def adelete_eval( @client def delete_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[DeleteEvalResponse, Coroutine[Any, Any, DeleteEvalResponse]]: +) -> DeleteEvalResponse | Coroutine[Any, Any, DeleteEvalResponse]: """ Delete an evaluation @@ -863,7 +863,7 @@ def delete_eval( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete_eval", False) is True # Get LiteLLM parameters @@ -874,7 +874,7 @@ def delete_eval( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -933,10 +933,10 @@ def delete_eval( @client async def acancel_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> CancelEvalResponse: """ @@ -990,12 +990,12 @@ async def acancel_eval( @client def cancel_eval( eval_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[CancelEvalResponse, Coroutine[Any, Any, CancelEvalResponse]]: +) -> CancelEvalResponse | Coroutine[Any, Any, CancelEvalResponse]: """ Cancel a running evaluation @@ -1013,7 +1013,7 @@ def cancel_eval( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acancel_eval", False) is True # Get LiteLLM parameters @@ -1024,7 +1024,7 @@ def cancel_eval( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1092,14 +1092,14 @@ def cancel_eval( @client async def acreate_run( eval_id: str, - data_source: Dict[str, Any], - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + data_source: dict[str, Any], + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Run: """ @@ -1161,16 +1161,16 @@ async def acreate_run( @client def create_run( eval_id: str, - data_source: Dict[str, Any], - name: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + data_source: dict[str, Any], + name: str | None = None, + metadata: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Run, Coroutine[Any, Any, Run]]: +) -> Run | Coroutine[Any, Any, Run]: """ Create a new run for an evaluation @@ -1192,7 +1192,7 @@ def create_run( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acreate_run", False) is True # Get LiteLLM parameters @@ -1203,7 +1203,7 @@ def create_run( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1276,14 +1276,14 @@ def create_run( @client async def alist_runs( eval_id: str, - limit: Optional[int] = None, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + after: str | None = None, + before: str | None = None, + order: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ListRunsResponse: """ @@ -1345,16 +1345,16 @@ async def alist_runs( @client def list_runs( eval_id: str, - limit: Optional[int] = None, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + after: str | None = None, + before: str | None = None, + order: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ListRunsResponse, Coroutine[Any, Any, ListRunsResponse]]: +) -> ListRunsResponse | Coroutine[Any, Any, ListRunsResponse]: """ List all runs for an evaluation @@ -1376,7 +1376,7 @@ def list_runs( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("alist_runs", False) is True # Get LiteLLM parameters @@ -1387,7 +1387,7 @@ def list_runs( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1462,10 +1462,10 @@ def list_runs( async def aget_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Run: """ @@ -1522,12 +1522,12 @@ async def aget_run( def get_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Run, Coroutine[Any, Any, Run]]: +) -> Run | Coroutine[Any, Any, Run]: """ Get a specific run @@ -1546,7 +1546,7 @@ def get_run( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aget_run", False) is True # Get LiteLLM parameters @@ -1557,7 +1557,7 @@ def get_run( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1618,10 +1618,10 @@ def get_run( async def acancel_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> CancelRunResponse: """ @@ -1678,12 +1678,12 @@ async def acancel_run( def cancel_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[CancelRunResponse, Coroutine[Any, Any, CancelRunResponse]]: +) -> CancelRunResponse | Coroutine[Any, Any, CancelRunResponse]: """ Cancel a running run @@ -1702,7 +1702,7 @@ def cancel_run( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acancel_run", False) is True # Get LiteLLM parameters @@ -1713,7 +1713,7 @@ def cancel_run( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1783,10 +1783,10 @@ def cancel_run( async def adelete_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> RunDeleteResponse: """ @@ -1843,12 +1843,12 @@ async def adelete_run( def delete_run( eval_id: str, run_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[RunDeleteResponse, Coroutine[Any, Any, RunDeleteResponse]]: +) -> RunDeleteResponse | Coroutine[Any, Any, RunDeleteResponse]: """ Delete a run @@ -1867,7 +1867,7 @@ def delete_run( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete_run", False) is True # Get LiteLLM parameters @@ -1878,7 +1878,7 @@ def delete_run( custom_llm_provider = "openai" # Get provider config - evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore + evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore provider=litellm.LlmProviders(custom_llm_provider), ) diff --git a/litellm/exceptions.py b/litellm/exceptions.py index fd0a2afb3e8..0d85c795c7b 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -10,7 +10,7 @@ ## LiteLLM versions of the OpenAI Exception Types import enum -from typing import Any, Dict, Optional, Union +from typing import Any import httpx import openai @@ -85,7 +85,7 @@ _RATE_LIMIT_CATEGORY_VALUES = frozenset(c.value for c in RateLimitErrorCategory) _RATE_LIMIT_TYPE_VALUES = frozenset(t.value for t in RateLimitType) -def validate_rate_limit_category(value: Any) -> Optional[str]: +def validate_rate_limit_category(value: Any) -> str | None: """Return ``value`` only if it matches a known :class:`RateLimitErrorCategory`. Used at duck-typed read sites (StandardLoggingPayload extraction, Prometheus @@ -100,7 +100,7 @@ def validate_rate_limit_category(value: Any) -> Optional[str]: return None -def validate_rate_limit_type(value: Any) -> Optional[str]: +def validate_rate_limit_type(value: Any) -> str | None: """Return ``value`` only if it matches a known :class:`RateLimitType`. See :func:`validate_rate_limit_category` for the rationale. @@ -112,7 +112,7 @@ def validate_rate_limit_type(value: Any) -> Optional[str]: return None -_MINIMAL_ERROR_RESPONSE: Optional[httpx.Response] = None +_MINIMAL_ERROR_RESPONSE: httpx.Response | None = None def _get_minimal_error_response() -> httpx.Response: @@ -132,13 +132,13 @@ class AuthenticationError(openai.AuthenticationError): # type: ignore message, llm_provider, model, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 401 - self.message = "litellm.AuthenticationError: {}".format(message) + self.message = f"litellm.AuthenticationError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -176,13 +176,13 @@ class NotFoundError(openai.NotFoundError): # type: ignore message, model, llm_provider, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 404 - self.message = "litellm.NotFoundError: {}".format(message) + self.message = f"litellm.NotFoundError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -219,14 +219,14 @@ class BadRequestError(openai.BadRequestError): # type: ignore message, model, llm_provider, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, - body: Optional[dict] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, + body: dict | None = None, ): self.status_code = 400 - self.message = "litellm.BadRequestError: {}".format(message) + self.message = f"litellm.BadRequestError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -270,11 +270,11 @@ class ImageFetchError(BadRequestError): message, model=None, llm_provider=None, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, - body: Optional[dict] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, + body: dict | None = None, ): super().__init__( message=message, @@ -295,12 +295,12 @@ class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore model, llm_provider, response: httpx.Response, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 422 - self.message = "litellm.UnprocessableEntityError: {}".format(message) + self.message = f"litellm.UnprocessableEntityError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -333,11 +333,11 @@ class Timeout(openai.APITimeoutError): # type: ignore message, model, llm_provider, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, - headers: Optional[dict] = None, - exception_status_code: Optional[int] = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, + headers: dict | None = None, + exception_status_code: int | None = None, ): request = httpx.Request( method="POST", @@ -345,7 +345,7 @@ class Timeout(openai.APITimeoutError): # type: ignore ) super().__init__(request=request) # Call the base class constructor with the parameters it needs self.status_code = exception_status_code or 408 - self.message = "litellm.Timeout: {}".format(message) + self.message = f"litellm.Timeout: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -378,12 +378,12 @@ class PermissionDeniedError(openai.PermissionDeniedError): # type: ignore llm_provider, model, response: httpx.Response, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 403 - self.message = "litellm.PermissionDeniedError: {}".format(message) + self.message = f"litellm.PermissionDeniedError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -428,17 +428,17 @@ class RateLimitError(openai.RateLimitError): # type: ignore message, llm_provider, model, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, - category: Union[str, RateLimitErrorCategory] = (RateLimitErrorCategory.VENDOR_RATE_LIMIT), - rate_limit_type: Optional[Union[str, RateLimitType]] = None, - headers: Optional[Dict[str, str]] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, + category: str | RateLimitErrorCategory = (RateLimitErrorCategory.VENDOR_RATE_LIMIT), + rate_limit_type: str | RateLimitType | None = None, + headers: dict[str, str] | None = None, detail: Any = None, ): self.status_code = 429 - self.message = "litellm.RateLimitError: {}".format(message) + self.message = f"litellm.RateLimitError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -448,7 +448,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore # Which dimension was exceeded — request count, token count, parallel # requests, budget, max iterations. None when the source didn't # classify the failure (e.g. legacy vendor 429 with no header hints). - self.rate_limit_type: Optional[str] = ( + self.rate_limit_type: str | None = ( rate_limit_type.value if isinstance(rate_limit_type, RateLimitType) else rate_limit_type ) # Headers explicitly attached to the error (e.g. retry-after, @@ -465,7 +465,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore # explicitly want them; only the proxy-supplied `headers=` kwarg # makes it onto `self.headers`. _response_headers = getattr(response, "headers", None) if response is not None else None - self.headers: Optional[Dict[str, str]] = {k: str(v) for k, v in headers.items()} if headers else None + self.headers: dict[str, str] | None = {k: str(v) for k, v in headers.items()} if headers else None # Mirrors FastAPI HTTPException.detail so the same instance can be # serialized through both the ProxyException and HTTPException paths. self.detail = detail if detail is not None else self.message @@ -507,8 +507,8 @@ class ContextWindowExceededError(BadRequestError): # type: ignore message, model, llm_provider, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, ): self.status_code = 400 self.model = model @@ -523,7 +523,7 @@ class ContextWindowExceededError(BadRequestError): # type: ignore ) # Call the base class constructor with the parameters it needs # set after, to make it clear the raised error is a context window exceeded error - self.message = "litellm.ContextWindowExceededError: {}".format(self.message) + self.message = f"litellm.ContextWindowExceededError: {self.message}" def __str__(self): _message = self.message @@ -550,10 +550,10 @@ class RejectedRequestError(BadRequestError): # type: ignore model, llm_provider, request_data: dict, - litellm_debug_info: Optional[str] = None, + litellm_debug_info: str | None = None, ): self.status_code = 400 - self.message = "litellm.RejectedRequestError: {}".format(message) + self.message = f"litellm.RejectedRequestError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -592,13 +592,13 @@ class ContentPolicyViolationError(BadRequestError): # type: ignore message, model, llm_provider, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - provider_specific_fields: Optional[dict] = None, - body: Optional[dict] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + provider_specific_fields: dict | None = None, + body: dict | None = None, ): self.status_code = 400 - self.message = "litellm.ContentPolicyViolationError: {}".format(message) + self.message = f"litellm.ContentPolicyViolationError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -636,13 +636,13 @@ class ServiceUnavailableError(openai.APIStatusError): # type: ignore message, llm_provider, model, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 503 - self.message = "litellm.ServiceUnavailableError: {}".format(message) + self.message = f"litellm.ServiceUnavailableError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -684,13 +684,13 @@ class BadGatewayError(openai.APIStatusError): # type: ignore message, llm_provider, model, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 502 - self.message = "litellm.BadGatewayError: {}".format(message) + self.message = f"litellm.BadGatewayError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -732,13 +732,13 @@ class InternalServerError(openai.InternalServerError): # type: ignore message, llm_provider, model, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 500 - self.message = "litellm.InternalServerError: {}".format(message) + self.message = f"litellm.InternalServerError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -782,13 +782,13 @@ class APIError(openai.APIError): # type: ignore message, llm_provider, model, - request: Optional[httpx.Request] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + request: httpx.Request | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = status_code - self.message = "litellm.APIError: {}".format(message) + self.message = f"litellm.APIError: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -822,12 +822,12 @@ class APIConnectionError(openai.APIConnectionError): # type: ignore message, llm_provider, model, - request: Optional[httpx.Request] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + request: httpx.Request | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): - self.message = "litellm.APIConnectionError: {}".format(message) + self.message = f"litellm.APIConnectionError: {message}" self.llm_provider = llm_provider self.model = model self.status_code = 500 @@ -861,11 +861,11 @@ class APIResponseValidationError(openai.APIResponseValidationError): # type: ig message, llm_provider, model, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): - self.message = "litellm.APIResponseValidationError: {}".format(message) + self.message = f"litellm.APIResponseValidationError: {message}" self.llm_provider = llm_provider self.model = model request = httpx.Request(method="POST", url="https://api.openai.com/v1") @@ -897,9 +897,7 @@ class JSONSchemaValidationError(APIResponseValidationError): self.raw_response = raw_response self.schema = schema self.model = model - message = "litellm.JSONSchemaValidationError: model={}, returned an invalid response={}, for schema={}.\nAccess raw response with `e.raw_response`".format( - model, raw_response, schema - ) + message = f"litellm.JSONSchemaValidationError: model={model}, returned an invalid response={raw_response}, for schema={schema}.\nAccess raw response with `e.raw_response`" self.message = message super().__init__(model=model, message=message, llm_provider=llm_provider) @@ -914,16 +912,16 @@ class UnsupportedParamsError(BadRequestError): def __init__( self, message, - llm_provider: Optional[str] = None, - model: Optional[str] = None, + llm_provider: str | None = None, + model: str | None = None, status_code: int = 400, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = 400 - self.message = "litellm.UnsupportedParamsError: {}".format(message) + self.message = f"litellm.UnsupportedParamsError: {message}" self.model = model self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info @@ -964,10 +962,10 @@ class BudgetExceededError(Exception): self, current_cost: float, max_budget: float, - message: Optional[str] = None, - llm_provider: Optional[str] = None, - entity_type: Optional[str] = None, - entity_id: Optional[str] = None, + message: str | None = None, + llm_provider: str | None = None, + entity_type: str | None = None, + entity_id: str | None = None, ): self.current_cost = current_cost self.max_budget = max_budget @@ -1012,13 +1010,13 @@ class MockException(openai.APIError): message, llm_provider, model, - request: Optional[httpx.Request] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + request: httpx.Request | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, ): self.status_code = status_code - self.message = "litellm.MockException: {}".format(message) + self.message = f"litellm.MockException: {message}" self.llm_provider = llm_provider self.model = model self.litellm_debug_info = litellm_debug_info @@ -1030,7 +1028,7 @@ class MockException(openai.APIError): class LiteLLMUnknownProvider(BadRequestError): - def __init__(self, model: str, custom_llm_provider: Optional[str] = None): + def __init__(self, model: str, custom_llm_provider: str | None = None): self.message = LiteLLMCommonStrings.llm_provider_not_provided.value.format( model=model, custom_llm_provider=custom_llm_provider ) @@ -1043,7 +1041,7 @@ class LiteLLMUnknownProvider(BadRequestError): class GuardrailRaisedException(Exception): def __init__( self, - guardrail_name: Optional[str] = None, + guardrail_name: str | None = None, message: str = "", should_wrap_with_default_message: bool = True, status_code: int = 400, @@ -1059,7 +1057,7 @@ class BlockedPiiEntityError(Exception): def __init__( self, entity_type: str, - guardrail_name: Optional[str] = None, + guardrail_name: str | None = None, status_code: int = 400, ): """ @@ -1078,11 +1076,11 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore message: str, model: str, llm_provider: str, - original_exception: Optional[Exception] = None, - response: Optional[httpx.Response] = None, - litellm_debug_info: Optional[str] = None, - max_retries: Optional[int] = None, - num_retries: Optional[int] = None, + original_exception: Exception | None = None, + response: httpx.Response | None = None, + litellm_debug_info: str | None = None, + max_retries: int | None = None, + num_retries: int | None = None, generated_content: str = "", is_pre_first_chunk: bool = False, ): @@ -1142,7 +1140,7 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore if self.max_retries: _message += f", LiteLLM Max Retries: {self.max_retries}" if self.original_exception: - _message += f" Original exception: {type(self.original_exception).__name__}: {str(self.original_exception)}" + _message += f" Original exception: {type(self.original_exception).__name__}: {self.original_exception!s}" return _message def __repr__(self): @@ -1166,10 +1164,10 @@ class ModifyResponseException(Exception): self, message: str, model: str, - request_data: Dict[str, Any], - guardrail_name: Optional[str] = None, - detection_info: Optional[Dict[str, Any]] = None, - original_response: Optional[Any] = None, + request_data: dict[str, Any], + guardrail_name: str | None = None, + detection_info: dict[str, Any] | None = None, + original_response: Any | None = None, ): self.message = message self.model = model @@ -1201,9 +1199,9 @@ class SensitiveDataRouteException(Exception): self, route_to_model: str, session_id: str, - guardrail_name: Optional[str] = None, - detection_info: Optional[Dict[str, Any]] = None, - message: Optional[str] = None, + guardrail_name: str | None = None, + detection_info: dict[str, Any] | None = None, + message: str | None = None, sticky_session_routing: bool = True, ): self.route_to_model = route_to_model diff --git a/litellm/experimental_mcp_client/__init__.py b/litellm/experimental_mcp_client/__init__.py index 7110d5375e4..5399968ff74 100644 --- a/litellm/experimental_mcp_client/__init__.py +++ b/litellm/experimental_mcp_client/__init__.py @@ -1,3 +1,3 @@ from .tools import call_openai_tool, load_mcp_tools -__all__ = ["load_mcp_tools", "call_openai_tool"] +__all__ = ["call_openai_tool", "load_mcp_tools"] diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 8a5440c8553..72248c4448d 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -8,12 +8,7 @@ import os from collections.abc import Awaitable, Callable, Generator from typing import ( Any, - Dict, - List, - Optional, - Tuple, TypeVar, - Union, ) import httpx @@ -21,7 +16,7 @@ from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParamete from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client -streamable_http_client: Optional[Any] = None +streamable_http_client: Any | None = None try: import mcp.client.streamable_http as streamable_http_module # type: ignore @@ -58,15 +53,15 @@ def to_basic_auth(auth_value: str) -> str: return base64.b64encode(auth_value.encode("utf-8")).decode() -def _strip_header_whitespace(headers: Dict[str, str]) -> Dict[str, str]: +def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]: return { (key.strip() if isinstance(key, str) else key): (value.strip() if isinstance(value, str) else value) for key, value in headers.items() } -def _first_non_cancelled_cause(exc: BaseException) -> Optional[BaseException]: - queue: List[BaseException] = [exc] +def _first_non_cancelled_cause(exc: BaseException) -> BaseException | None: + queue: list[BaseException] = [exc] while queue: current = queue.pop(0) nested = getattr(current, "exceptions", None) @@ -92,13 +87,13 @@ class MCPSigV4Auth(httpx.Auth): def __init__( self, - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, - aws_region_name: Optional[str] = None, - aws_service_name: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_session_name: Optional[str] = None, + aws_access_key_id: str | None = None, + aws_secret_access_key: str | None = None, + aws_session_token: str | None = None, + aws_region_name: str | None = None, + aws_service_name: str | None = None, + aws_role_name: str | None = None, + aws_session_name: str | None = None, ): try: from botocore.credentials import Credentials @@ -140,10 +135,10 @@ class MCPSigV4Auth(httpx.Auth): @staticmethod def _assume_role( aws_role_name: str, - aws_session_name: Optional[str], - aws_access_key_id: Optional[str], - aws_secret_access_key: Optional[str], - aws_session_token: Optional[str], + aws_session_name: str | None, + aws_access_key_id: str | None, + aws_secret_access_key: str | None, + aws_session_token: str | None, aws_region_name: str, ): """Call STS AssumeRole and return temporary credentials.""" @@ -207,47 +202,47 @@ class MCPClient: server_url: str = "", transport_type: MCPTransportType = MCPTransport.http, auth_type: MCPAuthType = None, - auth_value: Optional[Union[str, Dict[str, str]]] = None, - timeout: Optional[float] = None, - stdio_config: Optional[MCPStdioConfig] = None, - extra_headers: Optional[Dict[str, str]] = None, - ssl_verify: Optional[VerifyTypes] = None, - aws_auth: Optional[httpx.Auth] = None, - resolved_auth: Optional[httpx.Auth] = None, - sampling_callback: Optional[Callable] = None, - elicitation_callback: Optional[Callable] = None, - logging_callback: Optional[Callable] = None, + auth_value: str | dict[str, str] | None = None, + timeout: float | None = None, + stdio_config: MCPStdioConfig | None = None, + extra_headers: dict[str, str] | None = None, + ssl_verify: VerifyTypes | None = None, + aws_auth: httpx.Auth | None = None, + resolved_auth: httpx.Auth | None = None, + sampling_callback: Callable | None = None, + elicitation_callback: Callable | None = None, + logging_callback: Callable | None = None, ): self.server_url: str = server_url self.transport_type: MCPTransport = transport_type self.auth_type: MCPAuthType = auth_type self.timeout: float = timeout if timeout is not None else MCP_CLIENT_TIMEOUT - self._mcp_auth_value: Optional[Union[str, Dict[str, str]]] = None - self.stdio_config: Optional[MCPStdioConfig] = stdio_config - self.extra_headers: Optional[Dict[str, str]] = extra_headers - self.ssl_verify: Optional[VerifyTypes] = ssl_verify - self._aws_auth: Optional[httpx.Auth] = aws_auth + self._mcp_auth_value: str | dict[str, str] | None = None + self.stdio_config: MCPStdioConfig | None = stdio_config + self.extra_headers: dict[str, str] | None = extra_headers + self.ssl_verify: VerifyTypes | None = ssl_verify + self._aws_auth: httpx.Auth | None = aws_auth # A pre-resolved httpx.Auth (e.g. from the v2 credential resolver) attached to the # upstream client's auth= slot, taking precedence over the SigV4 aws_auth. - self._resolved_auth: Optional[httpx.Auth] = resolved_auth - self._last_initialize_instructions: Optional[str] = None - self._sampling_callback: Optional[Callable] = sampling_callback - self._elicitation_callback: Optional[Callable] = elicitation_callback - self._logging_callback: Optional[Callable] = logging_callback + self._resolved_auth: httpx.Auth | None = resolved_auth + self._last_initialize_instructions: str | None = None + self._sampling_callback: Callable | None = sampling_callback + self._elicitation_callback: Callable | None = elicitation_callback + self._logging_callback: Callable | None = logging_callback # handle the basic auth value if provided if auth_value: self.update_auth_value(auth_value) def _create_transport_context( self, - ) -> Tuple[Any, Optional[httpx.AsyncClient]]: + ) -> tuple[Any, httpx.AsyncClient | None]: """ Create the appropriate transport context based on transport type. Returns: Tuple of (transport_context, http_client). http_client is only set for HTTP transport and needs cleanup. """ - http_client: Optional[httpx.AsyncClient] = None + http_client: httpx.AsyncClient | None = None if self.transport_type == MCPTransport.stdio: if not self.stdio_config: raise ValueError("stdio_config is required for stdio transport") @@ -285,7 +280,7 @@ class MCPClient: ) return transport_ctx, http_client - def _get_safe_stdio_env(self, provided_env: Optional[Dict[str, str]]) -> Optional[Dict[str, str]]: + def _get_safe_stdio_env(self, provided_env: dict[str, str] | None) -> dict[str, str] | None: """ Return a safe environment for the stdio subprocess. @@ -344,11 +339,11 @@ class MCPClient: user input (elicitation), or send log messages. """ transport = await transport_ctx.__aenter__() - in_flight_error: Optional[BaseException] = None + in_flight_error: BaseException | None = None try: read_stream, write_stream = transport[0], transport[1] # Build session kwargs with optional callbacks - session_kwargs: Dict[str, Any] = {} + session_kwargs: dict[str, Any] = {} if self._sampling_callback is not None: session_kwargs["sampling_callback"] = self._sampling_callback if self._elicitation_callback is not None: @@ -393,7 +388,7 @@ class MCPClient: quiet_on_error demotes the failure line to debug for callers that own the exception (call_tool / list_tools under raise_on_error), so an expected pass-through re-auth does not emit a warning per call; every other caller keeps the operator-visible warning.""" - http_client: Optional[httpx.AsyncClient] = None + http_client: httpx.AsyncClient | None = None try: self._last_initialize_instructions = None transport_ctx, http_client = self._create_transport_context() @@ -409,7 +404,7 @@ class MCPClient: except BaseException as e: verbose_logger.debug(f"Error during http_client cleanup: {e}") - def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): + def update_auth_value(self, mcp_auth_value: str | dict[str, str]): """ Set the authentication header for the MCP client. """ @@ -462,9 +457,9 @@ class MCPClient: def factory( *, - headers: Optional[Dict[str, str]] = None, - timeout: Optional[httpx.Timeout] = None, - auth: Optional[httpx.Auth] = None, + headers: dict[str, str] | None = None, + timeout: httpx.Timeout | None = None, + auth: httpx.Auth | None = None, ) -> httpx.AsyncClient: """Create an httpx.AsyncClient with LiteLLM's SSL configuration.""" # Get unified SSL configuration using the same logic as http_handler.py @@ -485,7 +480,7 @@ class MCPClient: return factory - async def list_tools(self, raise_on_error: bool = False) -> List[MCPTool]: + async def list_tools(self, raise_on_error: bool = False) -> list[MCPTool]: """List available tools from the server. Args: @@ -520,7 +515,7 @@ class MCPClient: _log( f"MCP client list_tools failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" ) @@ -541,14 +536,14 @@ class MCPClient: def error_tool_result(exc: Exception) -> MCPCallToolResult: """The error result ``call_tool`` returns when it swallows a failure (no re-execution).""" return MCPCallToolResult( - content=[TextContent(type="text", text=f"{type(exc).__name__}: {str(exc)}")], + content=[TextContent(type="text", text=f"{type(exc).__name__}: {exc!s}")], isError=True, ) async def call_tool( self, call_tool_request_params: MCPCallToolRequestParams, - host_progress_callback: Optional[Callable] = None, + host_progress_callback: Callable | None = None, raise_on_error: bool = False, ) -> MCPCallToolResult: """ @@ -606,7 +601,7 @@ class MCPClient: _log( f"MCP client call_tool failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Tool: {call_tool_request_params.name}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" @@ -622,7 +617,7 @@ class MCPClient: # Return a default error result instead of raising return self.error_tool_result(e) - async def list_prompts(self) -> List[Prompt]: + async def list_prompts(self) -> list[Prompt]: """List available prompts from the server.""" verbose_logger.debug(f"MCP client listing tools from {self.server_url or 'stdio'}") @@ -645,7 +640,7 @@ class MCPClient: verbose_logger.error( f"MCP client list_prompts failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" ) @@ -686,7 +681,7 @@ class MCPClient: verbose_logger.error( f"MCP client get_prompt failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Prompt: {get_prompt_request_params.name}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" @@ -722,7 +717,7 @@ class MCPClient: verbose_logger.error( f"MCP client list_resources failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" ) @@ -758,7 +753,7 @@ class MCPClient: verbose_logger.error( f"MCP client list_resource_templates failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" ) @@ -796,7 +791,7 @@ class MCPClient: verbose_logger.error( f"MCP client read_resource failed - " f"Error Type: {error_type}, " - f"Error: {str(e)}, " + f"Error: {e!s}, " f"Url: {url}, " f"Server: {self.server_url or 'stdio'}, " f"Transport: {self.transport_type}" diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py index 500d226752b..23ab77f0037 100644 --- a/litellm/experimental_mcp_client/tools.py +++ b/litellm/experimental_mcp_client/tools.py @@ -1,5 +1,5 @@ import json -from typing import Dict, List, Literal, Union +from typing import Literal from mcp import ClientSession from mcp.types import CallToolRequestParams as MCPCallToolRequestParams @@ -92,7 +92,7 @@ def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessages async def load_mcp_tools( session: ClientSession, format: Literal["mcp", "openai"] = "mcp" -) -> Union[List[MCPTool], List[ChatCompletionToolParam]]: +) -> list[MCPTool] | list[ChatCompletionToolParam]: """ Load all available MCP tools @@ -138,7 +138,7 @@ def _get_function_arguments(function: FunctionDefinition) -> dict: def transform_openai_tool_call_request_to_mcp_tool_call_request( - openai_tool: Union[ChatCompletionMessageToolCall, Dict], + openai_tool: ChatCompletionMessageToolCall | dict, ) -> MCPCallToolRequestParams: """Convert an OpenAI ChatCompletionMessageToolCall to an MCP CallToolRequestParams.""" function = openai_tool["function"] diff --git a/litellm/files/main.py b/litellm/files/main.py index 297adfcffaa..e692cdc7c76 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -12,7 +12,7 @@ import uuid as uuid_module from collections.abc import Coroutine from functools import partial from types import MappingProxyType -from typing import Any, Dict, Literal, Optional, Union, cast +from typing import Any, Literal, cast import httpx @@ -70,7 +70,7 @@ base_llm_http_handler = BaseLLMHTTPHandler() def _should_sdk_support_streaming( - custom_llm_provider: Optional[Union[FileContentProvider, str]], + custom_llm_provider: FileContentProvider | str | None, ) -> bool: """ Return whether file content streaming is supported for the provider. @@ -86,7 +86,7 @@ bedrock_files_instance = BedrockFilesHandler() def _add_trusted_model_credentials_to_litellm_params( - litellm_params_dict: Dict[str, Any], kwargs: Dict[str, Any] + litellm_params_dict: dict[str, Any], kwargs: dict[str, Any] ) -> None: trusted_model_credentials = kwargs.get("_litellm_internal_model_credentials") if isinstance(trusted_model_credentials, type(MappingProxyType({}))): @@ -97,10 +97,10 @@ def _add_trusted_model_credentials_to_litellm_params( async def acreate_file( file: FileTypes, purpose: Literal["assistants", "batch", "fine-tune", "messages"], - expires_after: Optional[FileExpiresAfter] = None, + expires_after: FileExpiresAfter | None = None, custom_llm_provider: FileCreateProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> OpenAIFileObject: """ @@ -142,12 +142,12 @@ async def acreate_file( def create_file( file: FileTypes, purpose: Literal["assistants", "batch", "fine-tune", "messages"], - expires_after: Optional[FileExpiresAfter] = None, - custom_llm_provider: Optional[FileCreateProvider] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + expires_after: FileExpiresAfter | None = None, + custom_llm_provider: FileCreateProvider | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[OpenAIFileObject, Coroutine[Any, Any, OpenAIFileObject]]: +) -> OpenAIFileObject | Coroutine[Any, Any, OpenAIFileObject]: """ Files are used to upload documents that can be used with features like Assistants, Fine-tuning, and Batch API. @@ -159,7 +159,7 @@ def create_file( _is_async = kwargs.pop("acreate_file", False) is True optional_params = GenericLiteLLMParams(**kwargs) litellm_params_dict = dict(**kwargs) - logging_obj = cast(Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")) + logging_obj = cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj")) if logging_obj is None: raise ValueError("logging_obj is required") client = kwargs.get("client") @@ -246,9 +246,7 @@ def create_file( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_file'. Only ['openai', 'azure', 'vertex_ai', 'manus', 'anthropic'] are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_file'. Only ['openai', 'azure', 'vertex_ai', 'manus', 'anthropic'] are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -266,8 +264,8 @@ def create_file( async def afile_retrieve( file_id: str, custom_llm_provider: FileRetrieveProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> OpenAIFileObject: """ @@ -307,8 +305,8 @@ async def afile_retrieve( def file_retrieve( file_id: str, custom_llm_provider: FileRetrieveProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> FileObject: """ @@ -412,9 +410,7 @@ def file_retrieve( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'file_retrieve'. Only 'openai', 'azure', 'manus', and 'anthropic' are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'file_retrieve'. Only 'openai', 'azure', 'manus', and 'anthropic' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -437,8 +433,8 @@ def file_retrieve( async def afile_delete( file_id: str, custom_llm_provider: FileDeleteProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> Coroutine[Any, Any, FileObject]: """ @@ -479,10 +475,10 @@ async def afile_delete( @client def file_delete( file_id: str, - model: Optional[str] = None, - custom_llm_provider: Union[FileDeleteProvider, str] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + model: str | None = None, + custom_llm_provider: FileDeleteProvider | str = "openai", + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> FileDeleted: """ @@ -591,9 +587,7 @@ def file_delete( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'file_delete'. Only 'openai', 'azure', 'gemini', 'manus', and 'anthropic' are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'file_delete'. Only 'openai', 'azure', 'gemini', 'manus', and 'anthropic' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -614,9 +608,9 @@ def file_delete( @client async def afile_list( custom_llm_provider: FileListProvider = "openai", - purpose: Optional[str] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + purpose: str | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -655,9 +649,9 @@ async def afile_list( @client def file_list( custom_llm_provider: FileListProvider = "openai", - purpose: Optional[str] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + purpose: str | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -755,9 +749,7 @@ def file_list( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'file_list'. Only 'openai', 'azure', 'manus', and 'anthropic' are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'file_list'. Only 'openai', 'azure', 'manus', and 'anthropic' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -775,12 +767,12 @@ def file_list( async def afile_content( file_id: str, custom_llm_provider: FileContentProvider = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, chunk_size: int = 1024 * 1024, stream: bool = False, **kwargs, -) -> Union[HttpxBinaryResponseContent, FileContentStreamingResult]: +) -> HttpxBinaryResponseContent | FileContentStreamingResult: """ Async: Get file contents @@ -821,19 +813,19 @@ async def afile_content( @client def file_content( file_id: str, - model: Optional[str] = None, - custom_llm_provider: Optional[Union[FileContentProvider, str]] = None, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + model: str | None = None, + custom_llm_provider: FileContentProvider | str | None = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, chunk_size: int = 1024 * 1024, stream: bool = False, **kwargs, -) -> Union[ - HttpxBinaryResponseContent, - FileContentStreamingResult, - Coroutine[Any, Any, HttpxBinaryResponseContent], - Coroutine[Any, Any, FileContentStreamingResult], -]: +) -> ( + HttpxBinaryResponseContent + | FileContentStreamingResult + | Coroutine[Any, Any, HttpxBinaryResponseContent] + | Coroutine[Any, Any, FileContentStreamingResult] +): """ Returns the contents of the specified file. @@ -887,7 +879,7 @@ def file_content( chunk_size=chunk_size, optional_params=optional_params, timeout=timeout, - logging_obj=cast(Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")), + logging_obj=cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj")), _is_async=_is_async, client=client, ) @@ -989,9 +981,7 @@ def file_content( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'file_content'. Supported providers are 'openai', 'azure', 'vertex_ai', 'bedrock', 'manus', 'anthropic'.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'file_content'. Supported providers are 'openai', 'azure', 'vertex_ai', 'bedrock', 'manus', 'anthropic'.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -1008,17 +998,17 @@ def file_content( def file_content_streaming( *, file_id: str, - model: Optional[str], - custom_llm_provider: Optional[Union[FileContentProvider, str]], - extra_headers: Optional[Dict[str, str]], - extra_body: Optional[Dict[str, str]], + model: str | None, + custom_llm_provider: FileContentProvider | str | None, + extra_headers: dict[str, str] | None, + extra_body: dict[str, str] | None, chunk_size: int, optional_params: GenericLiteLLMParams, - timeout: Union[float, httpx.Timeout], - logging_obj: Optional[LiteLLMLoggingObj], + timeout: float | httpx.Timeout, + logging_obj: LiteLLMLoggingObj | None, _is_async: bool, - client: Optional[Any], -) -> Union[FileContentStreamingResult, Coroutine[Any, Any, FileContentStreamingResult]]: + client: Any | None, +) -> FileContentStreamingResult | Coroutine[Any, Any, FileContentStreamingResult]: if logging_obj is not None: logging_obj.model = model or "" logging_obj.model_call_details["model"] = model or "" @@ -1043,8 +1033,8 @@ def file_content_streaming( headers=response.headers, ) - response: Union[FileContentStreamingResult, Coroutine[Any, Any, FileContentStreamingResult]] = ( - FileContentStreamingResult(stream_iterator=iter(()), headers={}) + response: FileContentStreamingResult | Coroutine[Any, Any, FileContentStreamingResult] = FileContentStreamingResult( + stream_iterator=iter(()), headers={} ) if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: openai_creds = get_openai_credentials( @@ -1069,10 +1059,7 @@ def file_content_streaming( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for streaming 'file_content'. Supported providers are {}.".format( - custom_llm_provider, - sorted(OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS), - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for streaming 'file_content'. Supported providers are {sorted(OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS)}.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( diff --git a/litellm/files/streaming.py b/litellm/files/streaming.py index 9422a973e21..7c2b53b395e 100644 --- a/litellm/files/streaming.py +++ b/litellm/files/streaming.py @@ -4,9 +4,7 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, Optional, - Union, cast, ) @@ -29,10 +27,10 @@ class FileContentStreamingResponse: def __init__( self, - stream_iterator: Union[Iterator[bytes], AsyncIterator[bytes]], + stream_iterator: Iterator[bytes] | AsyncIterator[bytes], file_id: str, - model: Optional[str], - custom_llm_provider: Optional[Union[FileContentProvider, str]], + model: str | None, + custom_llm_provider: FileContentProvider | str | None, logging_obj: Optional["LiteLLMLoggingObj"], ) -> None: self.stream_iterator = stream_iterator @@ -40,8 +38,8 @@ class FileContentStreamingResponse: self.model = model self.custom_llm_provider = custom_llm_provider self.logging_obj = logging_obj - self.standard_logging_object: Optional["StandardLoggingPayload"] = None - self._hidden_params: Dict[str, Any] = {} + self.standard_logging_object: StandardLoggingPayload | None = None + self._hidden_params: dict[str, Any] = {} self._logging_completed = False self._close_completed = False self._start_time = ( @@ -94,7 +92,7 @@ class FileContentStreamingResponse: self._close_completed = True self._logging_completed = True stream_to_close = self.stream_iterator - self.stream_iterator = cast(Union[Iterator[bytes], AsyncIterator[bytes]], iter(())) + self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(())) # Shield cleanup from request cancellation so upstream HTTP connections # are released promptly on client disconnects. @@ -113,12 +111,12 @@ class FileContentStreamingResponse: self._close_completed = True self._logging_completed = True stream_to_close = self.stream_iterator - self.stream_iterator = cast(Union[Iterator[bytes], AsyncIterator[bytes]], iter(())) + self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(())) if hasattr(stream_to_close, "close"): cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined] - def _build_logging_response(self) -> Dict[str, str]: + def _build_logging_response(self) -> dict[str, str]: response = { "id": self.file_id, "object": "file.content", @@ -170,7 +168,7 @@ class FileContentStreamingResponse: merged_hidden_params = cast( "StandardLoggingHiddenParams", { - **cast(Dict[str, Any], payload.get("hidden_params") or {}), + **cast(dict[str, Any], payload.get("hidden_params") or {}), **self._hidden_params, }, ) diff --git a/litellm/files/types.py b/litellm/files/types.py index 93feb4676c6..8cadd69f024 100644 --- a/litellm/files/types.py +++ b/litellm/files/types.py @@ -1,9 +1,9 @@ from collections.abc import AsyncIterator, Iterator -from typing import Dict, Literal, NamedTuple, Union +from typing import Literal, NamedTuple FileContentProvider = Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"] class FileContentStreamingResult(NamedTuple): - stream_iterator: Union[Iterator[bytes], AsyncIterator[bytes]] - headers: Dict[str, str] + stream_iterator: Iterator[bytes] | AsyncIterator[bytes] + headers: dict[str, str] diff --git a/litellm/files/utils.py b/litellm/files/utils.py index 3ee4953bfef..3c58533f66a 100644 --- a/litellm/files/utils.py +++ b/litellm/files/utils.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.types.llms.openai import CreateFileRequest from litellm.types.utils import ExtractedFileData @@ -37,7 +35,7 @@ class FilesAPIUtils: ) @staticmethod - def is_batch_jsonl_request(create_file_data: CreateFileRequest, content_type: Optional[str]) -> bool: + def is_batch_jsonl_request(create_file_data: CreateFileRequest, content_type: str | None) -> bool: """ Batch-jsonl check from metadata only, so the body can stay a streamable Path/handle instead of being read into memory. @@ -49,7 +47,7 @@ class FilesAPIUtils: ) @staticmethod - def valid_content_type(content_type: Optional[str]) -> bool: + def valid_content_type(content_type: str | None) -> bool: """ Whether the upload's MIME type is one a batch JSONL file is plausibly sent as (see ``_BATCH_JSONL_CONTENT_TYPES``). diff --git a/litellm/fine_tuning/main.py b/litellm/fine_tuning/main.py index c6ab8a749d0..1987bf6a284 100644 --- a/litellm/fine_tuning/main.py +++ b/litellm/fine_tuning/main.py @@ -13,7 +13,7 @@ import contextvars import os from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, Literal, Optional, Union +from typing import Any, Literal import httpx @@ -36,10 +36,10 @@ vertex_fine_tuning_apis_instance = VertexFineTuningAPI() def _prepare_azure_extra_body( - extra_body: Optional[Dict[str, Any]], - kwargs: Dict[str, Any], - azure_specific_hyperparams: Dict[str, Any], -) -> Dict[str, Any]: + extra_body: dict[str, Any] | None, + kwargs: dict[str, Any], + azure_specific_hyperparams: dict[str, Any], +) -> dict[str, Any]: """ Prepare extra_body for Azure fine-tuning API by combining Azure-specific parameters. @@ -77,14 +77,14 @@ def _prepare_azure_extra_body( async def acreate_fine_tuning_job( model: str, training_file: str, - hyperparameters: Optional[dict] = {}, - suffix: Optional[str] = None, - validation_file: Optional[str] = None, - integrations: Optional[List[str]] = None, - seed: Optional[int] = None, + hyperparameters: dict | None = {}, + suffix: str | None = None, + validation_file: str | None = None, + integrations: List[str] | None = None, + seed: int | None = None, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> LiteLLMFineTuningJob: """ @@ -140,7 +140,7 @@ def _build_fine_tuning_job_data(model, training_file, hyperparameters, suffix, v def _resolve_fine_tuning_timeout( timeout: Any, custom_llm_provider: str, -) -> Union[float, httpx.Timeout]: +) -> float | httpx.Timeout: """Normalise a raw timeout value to a float (seconds) or httpx.Timeout for fine-tuning calls.""" timeout = timeout or 600.0 if isinstance(timeout, httpx.Timeout): @@ -154,16 +154,16 @@ def _resolve_fine_tuning_timeout( def create_fine_tuning_job( model: str, training_file: str, - hyperparameters: Optional[dict] = {}, - suffix: Optional[str] = None, - validation_file: Optional[str] = None, - integrations: Optional[List[str]] = None, - seed: Optional[int] = None, + hyperparameters: dict | None = {}, + suffix: str | None = None, + validation_file: str | None = None, + integrations: List[str] | None = None, + seed: int | None = None, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: +) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: """ Creates a fine-tuning job which begins the process of creating a new model from a given dataset. @@ -315,9 +315,7 @@ def create_fine_tuning_job( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_batch'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_batch'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -336,8 +334,8 @@ def create_fine_tuning_job( async def acancel_fine_tuning_job( fine_tuning_job_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> LiteLLMFineTuningJob: """ @@ -374,10 +372,10 @@ async def acancel_fine_tuning_job( def cancel_fine_tuning_job( fine_tuning_job_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: +) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: """ Immediately cancel a fine-tune job. @@ -469,9 +467,7 @@ def cancel_fine_tuning_job( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_batch'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_batch'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -486,11 +482,11 @@ def cancel_fine_tuning_job( async def alist_fine_tuning_jobs( - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -525,11 +521,11 @@ async def alist_fine_tuning_jobs( def list_fine_tuning_jobs( - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ): """ @@ -627,9 +623,7 @@ def list_fine_tuning_jobs( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'create_batch'. Only 'openai' is supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'create_batch'. Only 'openai' is supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( @@ -647,8 +641,8 @@ def list_fine_tuning_jobs( async def aretrieve_fine_tuning_job( fine_tuning_job_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, ) -> LiteLLMFineTuningJob: """ @@ -685,10 +679,10 @@ async def aretrieve_fine_tuning_job( def retrieve_fine_tuning_job( fine_tuning_job_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, str] | None = None, **kwargs, -) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: +) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: """ Get info about a fine-tuning job. """ @@ -767,9 +761,7 @@ def retrieve_fine_tuning_job( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'retrieve_fine_tuning_job'. Only 'openai' and 'azure' are supported.".format( - custom_llm_provider - ), + message=f"LiteLLM doesn't support {custom_llm_provider} for 'retrieve_fine_tuning_job'. Only 'openai' and 'azure' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( diff --git a/litellm/google_genai/__init__.py b/litellm/google_genai/__init__.py index ca7b547c440..eeff6a5fd65 100644 --- a/litellm/google_genai/__init__.py +++ b/litellm/google_genai/__init__.py @@ -12,8 +12,8 @@ from .main import ( ) __all__ = [ - "generate_content", "agenerate_content", - "generate_content_stream", "agenerate_content_stream", + "generate_content", + "generate_content_stream", ] diff --git a/litellm/google_genai/adapters/__init__.py b/litellm/google_genai/adapters/__init__.py index 6fbe7d95a55..796ddce8831 100644 --- a/litellm/google_genai/adapters/__init__.py +++ b/litellm/google_genai/adapters/__init__.py @@ -13,7 +13,7 @@ from .handler import GenerateContentToCompletionHandler from .transformation import GoogleGenAIAdapter, GoogleGenAIStreamWrapper __all__ = [ + "GenerateContentToCompletionHandler", "GoogleGenAIAdapter", "GoogleGenAIStreamWrapper", - "GenerateContentToCompletionHandler", ] diff --git a/litellm/google_genai/adapters/handler.py b/litellm/google_genai/adapters/handler.py index e6d7486fdb3..573f0633af5 100644 --- a/litellm/google_genai/adapters/handler.py +++ b/litellm/google_genai/adapters/handler.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator, Coroutine -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast import litellm from litellm.types.router import GenericLiteLLMParams @@ -17,12 +17,12 @@ class GenerateContentToCompletionHandler: @staticmethod def _prepare_completion_kwargs( model: str, - contents: Union[List[Dict[str, Any]], Dict[str, Any]], - config: Optional[Dict[str, Any]] = None, + contents: list[dict[str, Any]] | dict[str, Any], + config: dict[str, Any] | None = None, stream: bool = False, - litellm_params: Optional[GenericLiteLLMParams] = None, - extra_kwargs: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + litellm_params: GenericLiteLLMParams | None = None, + extra_kwargs: dict[str, Any] | None = None, + ) -> dict[str, Any]: """Prepare kwargs for litellm.completion/acompletion""" # Transform generate_content request to completion format @@ -34,7 +34,7 @@ class GenerateContentToCompletionHandler: **(extra_kwargs or {}), ) - completion_kwargs: Dict[str, Any] = dict(completion_request) + completion_kwargs: dict[str, Any] = dict(completion_request) # Forward extra_kwargs that should be passed to completion call if extra_kwargs is not None: @@ -53,12 +53,12 @@ class GenerateContentToCompletionHandler: @staticmethod async def async_generate_content_handler( model: str, - contents: Union[List[Dict[str, Any]], Dict[str, Any]], + contents: list[dict[str, Any]] | dict[str, Any], litellm_params: GenericLiteLLMParams, - config: Optional[Dict[str, Any]] = None, + config: dict[str, Any] | None = None, stream: bool = False, **kwargs, - ) -> Union[Dict[str, Any], AsyncIterator[bytes]]: + ) -> dict[str, Any] | AsyncIterator[bytes]: """Handle generate_content call asynchronously using completion adapter""" completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs( @@ -98,22 +98,18 @@ class GenerateContentToCompletionHandler: return generate_content_response except Exception as e: - raise ValueError(f"Error calling litellm.acompletion for generate_content: {str(e)}") + raise ValueError(f"Error calling litellm.acompletion for generate_content: {e!s}") @staticmethod def generate_content_handler( model: str, - contents: Union[List[Dict[str, Any]], Dict[str, Any]], + contents: list[dict[str, Any]] | dict[str, Any], litellm_params: GenericLiteLLMParams, - config: Optional[Dict[str, Any]] = None, + config: dict[str, Any] | None = None, stream: bool = False, _is_async: bool = False, **kwargs, - ) -> Union[ - Dict[str, Any], - AsyncIterator[bytes], - Coroutine[Any, Any, Union[Dict[str, Any], AsyncIterator[bytes]]], - ]: + ) -> dict[str, Any] | AsyncIterator[bytes] | Coroutine[Any, Any, dict[str, Any] | AsyncIterator[bytes]]: """Handle generate_content call using completion adapter""" if _is_async: @@ -163,4 +159,4 @@ class GenerateContentToCompletionHandler: return generate_content_response except Exception as e: - raise ValueError(f"Error calling litellm.completion for generate_content: {str(e)}") + raise ValueError(f"Error calling litellm.completion for generate_content: {e!s}") diff --git a/litellm/google_genai/adapters/transformation.py b/litellm/google_genai/adapters/transformation.py index d15b6f47013..7c2800db07b 100644 --- a/litellm/google_genai/adapters/transformation.py +++ b/litellm/google_genai/adapters/transformation.py @@ -1,6 +1,6 @@ import json from collections.abc import AsyncIterator, Iterator -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast from litellm import verbose_logger from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema @@ -36,7 +36,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper): sent_first_chunk: bool = False # State tracking for accumulating partial tool calls - accumulated_tool_calls: Dict[str, Dict[str, Any]] + accumulated_tool_calls: dict[str, dict[str, Any]] def __init__(self, completion_stream: Any): self.sent_first_chunk = False @@ -108,7 +108,6 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper): f"Name: {tool_call_data['name']}. " f"Partial args: {tool_call_data['arguments']}" ) - pass if parts: final_chunk = { "candidates": [ @@ -178,11 +177,11 @@ class GoogleGenAIAdapter: def translate_generate_content_to_completion( self, model: str, - contents: Union[List[Dict[str, Any]], Dict[str, Any]], - config: Optional[Dict[str, Any]] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + contents: list[dict[str, Any]] | dict[str, Any], + config: dict[str, Any] | None = None, + litellm_params: GenericLiteLLMParams | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform generate_content request to litellm completion format @@ -273,8 +272,8 @@ class GoogleGenAIAdapter: def _add_generic_litellm_params_to_request( self, - completion_request_dict: Dict[str, Any], - litellm_params: Optional[GenericLiteLLMParams] = None, + completion_request_dict: dict[str, Any], + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: """Add generic litellm params to request. e.g add api_base, api_key, api_version, etc. @@ -296,7 +295,7 @@ class GoogleGenAIAdapter: def translate_completion_output_params_streaming( self, completion_stream: Any, - ) -> Union[AsyncIterator[bytes], None]: + ) -> AsyncIterator[bytes] | None: """Transform streaming completion output to Google GenAI format""" google_genai_wrapper = GoogleGenAIStreamWrapper(completion_stream=completion_stream) # Return the SSE-wrapped version for proper event formatting @@ -304,15 +303,15 @@ class GoogleGenAIAdapter: def _transform_google_genai_tools_to_openai( self, - tools: List[Dict[str, Any]], - ) -> List[ChatCompletionToolParam]: + tools: list[dict[str, Any]], + ) -> list[ChatCompletionToolParam]: """Transform Google GenAI tools to OpenAI tools format""" - openai_tools: List[Dict[str, Any]] = [] + openai_tools: list[dict[str, Any]] = [] for tool in tools: if "functionDeclarations" in tool: for func_decl in tool["functionDeclarations"]: - function_chunk: Dict[str, Any] = { + function_chunk: dict[str, Any] = { "name": func_decl.get("name", ""), } @@ -327,12 +326,12 @@ class GoogleGenAIAdapter: # normalize the tool schemas normalized_tools = [normalize_tool_schema(tool) for tool in openai_tools] - return cast(List[ChatCompletionToolParam], normalized_tools) + return cast(list[ChatCompletionToolParam], normalized_tools) def _transform_google_genai_tool_config_to_openai( self, - tool_config: Dict[str, Any], - ) -> Optional[ChatCompletionToolChoiceValues]: + tool_config: dict[str, Any], + ) -> ChatCompletionToolChoiceValues | None: """Transform Google GenAI tool_config to OpenAI tool_choice""" function_calling_config = tool_config.get("functionCallingConfig", {}) mode = function_calling_config.get("mode", "AUTO") @@ -344,11 +343,11 @@ class GoogleGenAIAdapter: def _transform_contents_to_messages( self, - contents: List[Dict[str, Any]], - system_instruction: Optional[Dict[str, Any]] = None, - ) -> List[AllMessageValues]: + contents: list[dict[str, Any]], + system_instruction: dict[str, Any] | None = None, + ) -> list[AllMessageValues]: """Transform Google GenAI contents to OpenAI messages format""" - messages: List[AllMessageValues] = [] + messages: list[AllMessageValues] = [] # Handle system instruction if system_instruction: @@ -362,8 +361,8 @@ class GoogleGenAIAdapter: if role == "user": # Handle user messages with potential function responses - content_parts: List[Union[ChatCompletionTextObject, ChatCompletionImageObject]] = [] - tool_messages: List[ChatCompletionToolMessage] = [] + content_parts: list[ChatCompletionTextObject | ChatCompletionImageObject] = [] + tool_messages: list[ChatCompletionToolMessage] = [] for part in parts: if isinstance(part, dict): @@ -420,7 +419,7 @@ class GoogleGenAIAdapter: elif role == "model": # Handle assistant messages with potential function calls combined_text = "" - tool_calls: List[ChatCompletionAssistantToolCall] = [] + tool_calls: list[ChatCompletionAssistantToolCall] = [] for part in parts: if isinstance(part, dict): @@ -461,7 +460,7 @@ class GoogleGenAIAdapter: def translate_completion_to_generate_content( self, response: ModelResponse, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform litellm completion response to Google GenAI generate_content format @@ -490,7 +489,7 @@ class GoogleGenAIAdapter: parts = [{"text": message_content}] if message_content else [] # Create Google GenAI format response - generate_content_response: Dict[str, Any] = { + generate_content_response: dict[str, Any] = { "candidates": [ { "content": {"parts": parts, "role": "model"}, @@ -522,9 +521,9 @@ class GoogleGenAIAdapter: def translate_streaming_completion_to_generate_content( self, - response: Union[ModelResponse, ModelResponseStream], + response: ModelResponse | ModelResponseStream, wrapper: GoogleGenAIStreamWrapper, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """ Transform streaming litellm completion chunk to Google GenAI generate_content format @@ -560,7 +559,7 @@ class GoogleGenAIAdapter: return None # Create Google GenAI streaming format response - streaming_chunk: Dict[str, Any] = { + streaming_chunk: dict[str, Any] = { "candidates": [ { "content": {"parts": parts, "role": "model"}, @@ -597,9 +596,9 @@ class GoogleGenAIAdapter: def _transform_openai_message_to_google_genai_parts( self, message: Any, - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """Transform OpenAI message to Google GenAI parts format""" - parts: List[Dict[str, Any]] = [] + parts: list[dict[str, Any]] = [] # Add text content if present if hasattr(message, "content") and message.content: @@ -626,14 +625,14 @@ class GoogleGenAIAdapter: def _transform_openai_delta_to_google_genai_parts_with_accumulation( self, delta: Any, wrapper: GoogleGenAIStreamWrapper - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """Transforms OpenAI delta to Google GenAI parts, accumulating streaming tool calls.""" # 1. Initialize wrapper state if it doesn't exist if not hasattr(wrapper, "accumulated_tool_calls"): wrapper.accumulated_tool_calls = {} - parts: List[Dict[str, Any]] = [] + parts: list[dict[str, Any]] = [] if hasattr(delta, "content") and delta.content: parts.append({"text": delta.content}) @@ -699,7 +698,7 @@ class GoogleGenAIAdapter: return parts - def _map_finish_reason(self, finish_reason: Optional[str]) -> str: + def _map_finish_reason(self, finish_reason: str | None) -> str: """Map OpenAI finish reasons to Google GenAI finish reasons""" if not finish_reason: return "STOP" @@ -714,7 +713,7 @@ class GoogleGenAIAdapter: return mapping.get(finish_reason, "STOP") - def _map_usage(self, usage: Any) -> Dict[str, int]: + def _map_usage(self, usage: Any) -> dict[str, int]: """Map OpenAI usage to Google GenAI usage format""" return { "promptTokenCount": getattr(usage, "prompt_tokens", 0) or 0, diff --git a/litellm/google_genai/main.py b/litellm/google_genai/main.py index 5e119b75af3..dbb124a3106 100644 --- a/litellm/google_genai/main.py +++ b/litellm/google_genai/main.py @@ -2,7 +2,7 @@ import asyncio import contextvars from collections.abc import Iterator from functools import partial -from typing import TYPE_CHECKING, Any, ClassVar, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, ClassVar import httpx from pydantic import BaseModel, ConfigDict @@ -52,14 +52,14 @@ class GenerateContentSetupResult(BaseModel): model_config: ClassVar[ConfigDict] = ConfigDict(arbitrary_types_allowed=True) model: str - request_body: Dict[str, Any] + request_body: dict[str, Any] custom_llm_provider: str - generate_content_provider_config: Optional[BaseGoogleGenAIGenerateContentConfig] - generate_content_config_dict: Dict[str, Any] + generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None + generate_content_config_dict: dict[str, Any] native_request_fields: dict[str, object] litellm_params: GenericLiteLLMParams litellm_logging_obj: LiteLLMLoggingObj - litellm_call_id: Optional[str] + litellm_call_id: str | None class GenerateContentHelper: @@ -68,7 +68,7 @@ class GenerateContentHelper: @staticmethod def mock_generate_content_response( mock_response: str = "This is a mock response from Google GenAI generate_content.", - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Mock response for generate_content for testing purposes""" return { "text": mock_response, @@ -91,9 +91,9 @@ class GenerateContentHelper: def setup_generate_content_call( model: str, contents: GenerateContentContentListUnionDict, - config: Optional[GenerateContentConfigDict] = None, - custom_llm_provider: Optional[str] = None, - tools: Optional[ToolConfigDict] = None, + config: GenerateContentConfigDict | None = None, + custom_llm_provider: str | None = None, + tools: ToolConfigDict | None = None, **kwargs, ) -> GenerateContentSetupResult: """ @@ -110,8 +110,8 @@ class GenerateContentHelper: Returns: GenerateContentSetupResult containing all setup information """ - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj") + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) # get llm provider logic litellm_params = GenericLiteLLMParams(**kwargs) @@ -140,7 +140,7 @@ class GenerateContentHelper: litellm_params.custom_llm_provider = custom_llm_provider # get provider config - generate_content_provider_config: Optional[BaseGoogleGenAIGenerateContentConfig] = ( + generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None = ( ProviderConfigManager.get_provider_google_genai_generate_content_config( model=model, provider=litellm.LlmProviders(custom_llm_provider), @@ -235,16 +235,16 @@ def _merge_native_request_fields( async def agenerate_content( model: str, contents: GenerateContentContentListUnionDict, - config: Optional[GenerateContentConfigDict] = None, - tools: Optional[ToolConfigDict] = None, + config: GenerateContentConfigDict | None = None, + tools: ToolConfigDict | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Any: """ @@ -303,16 +303,16 @@ async def agenerate_content( def generate_content( model: str, contents: GenerateContentContentListUnionDict, - config: Optional[GenerateContentConfigDict] = None, - tools: Optional[ToolConfigDict] = None, + config: GenerateContentConfigDict | None = None, + tools: ToolConfigDict | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Any: """ @@ -393,16 +393,16 @@ def generate_content( async def agenerate_content_stream( model: str, contents: GenerateContentContentListUnionDict, - config: Optional[GenerateContentConfigDict] = None, - tools: Optional[ToolConfigDict] = None, + config: GenerateContentConfigDict | None = None, + tools: ToolConfigDict | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Any: """ @@ -488,16 +488,16 @@ async def agenerate_content_stream( def generate_content_stream( model: str, contents: GenerateContentContentListUnionDict, - config: Optional[GenerateContentConfigDict] = None, - tools: Optional[ToolConfigDict] = None, + config: GenerateContentConfigDict | None = None, + tools: ToolConfigDict | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Iterator[Any]: """ diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index 900a171640b..2829699492d 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -1,6 +1,6 @@ import asyncio from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.success_handler import ( @@ -18,12 +18,12 @@ else: GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ = PassThroughEndpointLogging() -def _encode_google_genai_sse_event(event_lines: List[str]) -> bytes: +def _encode_google_genai_sse_event(event_lines: list[str]) -> bytes: return ("\n".join(event_lines) + "\n\n").encode("utf-8") def _next_google_genai_sse_chunk(line_iter) -> bytes: - event_lines: List[str] = [] + event_lines: list[str] = [] while True: try: line = next(line_iter) @@ -39,7 +39,7 @@ def _next_google_genai_sse_chunk(line_iter) -> bytes: async def _anext_google_genai_sse_chunk(line_iter) -> bytes: - event_lines: List[str] = [] + event_lines: list[str] = [] while True: try: line = await line_iter.__anext__() @@ -65,14 +65,14 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: litellm_logging_obj: LiteLLMLoggingObj, request_body: dict, model: str, - hidden_params: Optional[Dict[str, Any]] = None, + hidden_params: dict[str, Any] | None = None, ): self.litellm_logging_obj = litellm_logging_obj self.request_body = request_body self.start_time = datetime.now() - self.collected_chunks: List[bytes] = [] + self.collected_chunks: list[bytes] = [] self.model = model - self._hidden_params: Dict[str, Any] = hidden_params or {} + self._hidden_params: dict[str, Any] = hidden_params or {} async def _handle_async_streaming_logging( self, @@ -111,8 +111,8 @@ class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContent generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig, litellm_metadata: dict, custom_llm_provider: str, - request_body: Optional[dict] = None, - hidden_params: Optional[Dict[str, Any]] = None, + request_body: dict | None = None, + hidden_params: dict[str, Any] | None = None, ): super().__init__( litellm_logging_obj=logging_obj, @@ -162,8 +162,8 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateCo generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig, litellm_metadata: dict, custom_llm_provider: str, - request_body: Optional[dict] = None, - hidden_params: Optional[Dict[str, Any]] = None, + request_body: dict | None = None, + hidden_params: dict[str, Any] | None = None, ): super().__init__( litellm_logging_obj=logging_obj, diff --git a/litellm/images/main.py b/litellm/images/main.py index 99f2ddffeb3..d26c9d54f83 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -6,11 +6,8 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Union, cast, overload, ) @@ -116,7 +113,7 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: # Await normally init_response = await loop.run_in_executor(None, func_with_context) - response: Optional[ImageResponse] = None + response: ImageResponse | None = None if isinstance(init_response, dict): response = ImageResponse(**init_response) elif isinstance(init_response, ImageResponse): ## CACHING SCENARIO @@ -145,17 +142,17 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: @overload def image_generation( prompt: str, - model: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + n: int | None = None, + quality: str | ImageGenerationRequestQuality | None = None, + response_format: str | None = None, + size: str | None = None, + style: str | None = None, + user: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider=None, *, aimg_generation: Literal[True], @@ -169,17 +166,17 @@ def image_generation( @overload def image_generation( prompt: str, - model: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + n: int | None = None, + quality: str | ImageGenerationRequestQuality | None = None, + response_format: str | None = None, + size: str | None = None, + style: str | None = None, + user: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider=None, *, aimg_generation: Literal[False] = False, @@ -193,23 +190,20 @@ def image_generation( @client def image_generation( prompt: str, - model: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + n: int | None = None, + quality: str | ImageGenerationRequestQuality | None = None, + response_format: str | None = None, + size: str | None = None, + style: str | None = None, + user: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, custom_llm_provider=None, **kwargs, -) -> Union[ - ImageResponse, - Coroutine[Any, Any, ImageResponse], -]: +) -> ImageResponse | Coroutine[Any, Any, ImageResponse]: """ Maps the https://api.openai.com/v1/images/generations endpoint. @@ -220,7 +214,7 @@ def image_generation( aimg_generation = kwargs.get("aimg_generation", False) litellm_call_id = kwargs.get("litellm_call_id", None) logger_fn = kwargs.get("logger_fn", None) - mock_response: Optional[str] = kwargs.get("mock_response", None) # type: ignore + mock_response: str | None = kwargs.get("mock_response", None) # type: ignore proxy_server_request = kwargs.get("proxy_server_request", None) azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) model_info = kwargs.get("model_info", None) @@ -233,7 +227,7 @@ def image_generation( if extra_headers is not None: headers.update(extra_headers) model_response: ImageResponse = litellm.utils.ImageResponse() - dynamic_api_key: Optional[str] = None + dynamic_api_key: str | None = None if model is not None or custom_llm_provider is not None: model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( model=model, # type: ignore @@ -267,7 +261,7 @@ def image_generation( k: v for k, v in kwargs.items() if k not in default_params } # model-specific params - pass them straight to the model/provider - image_generation_config: Optional[BaseImageGenerationConfig] = None + image_generation_config: BaseImageGenerationConfig | None = None if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values(): image_generation_config = ProviderConfigManager.get_provider_image_generation_config( model=base_model or model, @@ -474,7 +468,7 @@ def image_generation( if extra_headers is not None: optional_params["extra_headers"] = extra_headers # Forward OpenAI organization if present (set by proxy pre-call utils) - organization: Optional[str] = kwargs.get("organization", None) + organization: str | None = kwargs.get("organization", None) model_response = openai_chat_completions.image_generation( model=model, prompt=prompt, @@ -506,7 +500,7 @@ def image_generation( ) elif custom_llm_provider in litellm._custom_providers: # Assume custom LLM provider # Get the Custom Handler - custom_handler: Optional[CustomLLM] = None + custom_handler: CustomLLM | None = None for item in litellm.custom_provider_map: if item["provider"] == custom_llm_provider: custom_handler = item["custom_handler"] @@ -516,7 +510,7 @@ def image_generation( ## ROUTE LLM CALL ## if aimg_generation is True: - async_custom_client: Optional[AsyncHTTPHandler] = None + async_custom_client: AsyncHTTPHandler | None = None if client is not None and isinstance(client, AsyncHTTPHandler): async_custom_client = client @@ -533,7 +527,7 @@ def image_generation( client=async_custom_client, ) else: - custom_client: Optional[HTTPHandler] = None + custom_client: HTTPHandler | None = None if client is not None and isinstance(client, HTTPHandler): custom_client = client @@ -619,8 +613,8 @@ def image_variation( model: str = "dall-e-2", # set to dall-e-2 by default - like OpenAI. n: int = 1, response_format: Literal["url", "b64_json"] = "url", - size: Optional[str] = None, - user: Optional[str] = None, + size: str | None = None, + user: str | None = None, **kwargs, ) -> ImageResponse: # get non-default params @@ -648,7 +642,7 @@ def image_variation( ) model_response = ImageResponse() - response: Optional[ImageResponse] = None + response: ImageResponse | None = None provider_config = ProviderConfigManager.get_provider_model_info( model=model or "", # openai defaults to dall-e-2 @@ -711,25 +705,25 @@ def image_variation( @client def image_edit( - image: Optional[Union[FileTypes, List[FileTypes]]] = None, - prompt: Optional[str] = None, - model: Optional[str] = None, - mask: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + image: FileTypes | list[FileTypes] | None = None, + prompt: str | None = None, + model: str | None = None, + mask: str | None = None, + n: int | None = None, + quality: str | ImageGenerationRequestQuality | None = None, + response_format: str | None = None, + size: str | None = None, + user: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse]]: +) -> ImageResponse | Coroutine[Any, Any, ImageResponse]: """ Maps the image edit functionality, similar to OpenAI's images/edits endpoint. """ @@ -759,7 +753,7 @@ def image_edit( k: v for k, v in kwargs.items() if k not in default_params } # model-specific params - pass them straight to the model/provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) model_info = kwargs.get("model_info", None) metadata = kwargs.get("metadata", {}) _is_async = kwargs.pop("async_call", False) is True @@ -768,7 +762,7 @@ def image_edit( images = image if isinstance(image, list) else ([image] if image is not None else []) headers_from_kwargs = kwargs.get("headers") - merged_extra_headers: Dict[str, Any] = {} + merged_extra_headers: dict[str, Any] = {} if isinstance(headers_from_kwargs, dict): merged_extra_headers.update(headers_from_kwargs) if isinstance(extra_headers, dict): @@ -786,7 +780,7 @@ def image_edit( # Check for custom provider if custom_llm_provider in litellm._custom_providers: - custom_handler: Optional[CustomLLM] = None + custom_handler: CustomLLM | None = None for item in litellm.custom_provider_map: if item["provider"] == custom_llm_provider: custom_handler = item["custom_handler"] @@ -797,7 +791,7 @@ def image_edit( model_response = ImageResponse() if _is_async: - async_custom_client: Optional[AsyncHTTPHandler] = None + async_custom_client: AsyncHTTPHandler | None = None if kwargs.get("client") is not None and isinstance(kwargs.get("client"), AsyncHTTPHandler): async_custom_client = kwargs.get("client") @@ -814,7 +808,7 @@ def image_edit( client=async_custom_client, ) else: - custom_client: Optional[HTTPHandler] = None + custom_client: HTTPHandler | None = None if kwargs.get("client") is not None and isinstance(kwargs.get("client"), HTTPHandler): custom_client = kwargs.get("client") @@ -832,11 +826,9 @@ def image_edit( ) # get provider config - image_edit_provider_config: Optional[BaseImageEditConfig] = ( - ProviderConfigManager.get_provider_image_edit_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) + image_edit_provider_config: BaseImageEditConfig | None = ProviderConfigManager.get_provider_image_edit_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), ) if image_edit_provider_config is None: @@ -848,7 +840,7 @@ def image_edit( _get_ImageEditRequestUtils().get_requested_image_edit_optional_param(local_vars) ) # Get optional parameters for the responses API - image_edit_request_params: Dict = _get_ImageEditRequestUtils().get_optional_params_image_edit( + image_edit_request_params: dict = _get_ImageEditRequestUtils().get_optional_params_image_edit( model=model, image_edit_provider_config=image_edit_provider_config, image_edit_optional_params=image_edit_optional_params, @@ -952,23 +944,23 @@ def image_edit( @client async def aimage_edit( - image: Union[FileTypes, List[FileTypes]], + image: FileTypes | list[FileTypes], model: str, prompt: str, - mask: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + mask: str | None = None, + n: int | None = None, + quality: str | ImageGenerationRequestQuality | None = None, + response_format: str | None = None, + size: str | None = None, + user: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ImageResponse: """ diff --git a/litellm/images/utils.py b/litellm/images/utils.py index f0d4c985c01..e906cb5094b 100644 --- a/litellm/images/utils.py +++ b/litellm/images/utils.py @@ -1,5 +1,5 @@ from io import BufferedReader, BytesIO -from typing import Any, Dict, List, Optional, cast, get_type_hints +from typing import Any, cast, get_type_hints import litellm from litellm.litellm_core_utils.token_counter import get_image_type @@ -14,9 +14,9 @@ class ImageEditRequestUtils: model: str, image_edit_provider_config: BaseImageEditConfig, image_edit_optional_params: ImageEditOptionalRequestParams, - drop_params: Optional[bool] = None, - additional_drop_params: Optional[List[str]] = None, - ) -> Dict: + drop_params: bool | None = None, + additional_drop_params: list[str] | None = None, + ) -> dict: """ Get optional parameters for the image edit API. @@ -61,7 +61,7 @@ class ImageEditRequestUtils: @staticmethod def get_requested_image_edit_optional_param( - params: Dict[str, Any], + params: dict[str, Any], ) -> ImageEditOptionalRequestParams: """ Filter parameters to only include those defined in ImageEditOptionalRequestParams. diff --git a/litellm/integrations/SlackAlerting/batching_handler.py b/litellm/integrations/SlackAlerting/batching_handler.py index 42f4f562422..e5a60640ee2 100644 --- a/litellm/integrations/SlackAlerting/batching_handler.py +++ b/litellm/integrations/SlackAlerting/batching_handler.py @@ -70,6 +70,6 @@ async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item, count) if response.status_code != 200: verbose_proxy_logger.debug(f"Error sending slack alert to url={item['url']}. Error={response.text}") except Exception as e: - verbose_proxy_logger.debug(f"Error sending slack alert: {str(e)}") + verbose_proxy_logger.debug(f"Error sending slack alert: {e!s}") finally: _print_alerting_payload_warning(payload, slackAlertingInstance=slackAlertingInstance) diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index 2a19ec0b7fa..50700774ea6 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -10,12 +10,10 @@ class BaseBudgetAlertType(ABC): @abstractmethod def get_event_message(self) -> str: """Return the event message for this alert type""" - pass @abstractmethod def get_id(self, user_info: CallInfo) -> str: """Return the ID to use for caching/tracking this alert""" - pass class ProxyBudgetAlert(BaseBudgetAlertType): diff --git a/litellm/integrations/SlackAlerting/hanging_request_check.py b/litellm/integrations/SlackAlerting/hanging_request_check.py index 136b6583f38..55dff2fde1f 100644 --- a/litellm/integrations/SlackAlerting/hanging_request_check.py +++ b/litellm/integrations/SlackAlerting/hanging_request_check.py @@ -9,7 +9,7 @@ Notes: import asyncio import time -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_proxy_logger @@ -49,7 +49,7 @@ class AlertingHangingRequestCheck: async def add_request_to_hanging_request_check( self, - request_data: Optional[dict] = None, + request_data: dict | None = None, ): """ Add a request to the hanging request cache. This is the list of request_ids that gets periodicall checked for hanging requests @@ -59,7 +59,7 @@ class AlertingHangingRequestCheck: request_metadata = get_litellm_metadata_from_kwargs(kwargs=request_data) model = request_data.get("model", "") - api_base: Optional[str] = None + api_base: str | None = None if request_data.get("deployment", None) is not None and isinstance(request_data["deployment"], dict): api_base = litellm.get_api_base( @@ -101,7 +101,7 @@ class AlertingHangingRequestCheck: ) for request_id in hanging_requests: - hanging_request_data: Optional[HangingRequestData] = await self.hanging_request_cache.async_get_cache( + hanging_request_data: HangingRequestData | None = await self.hanging_request_cache.async_get_cache( key=request_id, ) @@ -112,7 +112,7 @@ class AlertingHangingRequestCheck: continue request_status = await proxy_logging_obj.internal_usage_cache.async_get_cache( - key="request_status:{}".format(hanging_request_data.request_id), + key=f"request_status:{hanging_request_data.request_id}", litellm_parent_otel_span=None, local_only=True, ) diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index e93c650ed97..4378b2f754e 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -6,7 +6,7 @@ import os import random import time from datetime import timedelta -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Literal from openai import APIError @@ -61,16 +61,15 @@ class SlackAlerting(CustomBatchLogger): # Class variables or attributes def __init__( self, - internal_usage_cache: Optional[DualCache] = None, - alerting_threshold: Optional[float] = None, # threshold for slow / hanging llm responses (in seconds) - alerting: Optional[List] = [], - alert_types: List[AlertType] = DEFAULT_ALERT_TYPES, - alert_to_webhook_url: Optional[ - Dict[AlertType, Union[List[str], str]] - ] = None, # if user wants to separate alerts to diff channels + internal_usage_cache: DualCache | None = None, + alerting_threshold: float | None = None, # threshold for slow / hanging llm responses (in seconds) + alerting: list | None = [], + alert_types: list[AlertType] = DEFAULT_ALERT_TYPES, + alert_to_webhook_url: dict[AlertType, list[str] | str] + | None = None, # if user wants to separate alerts to diff channels alerting_args={}, - default_webhook_url: Optional[str] = None, - alert_type_config: Optional[Dict[str, dict]] = None, + default_webhook_url: str | None = None, + alert_type_config: dict[str, dict] | None = None, **kwargs, ): if alerting_threshold is None: @@ -89,23 +88,23 @@ class SlackAlerting(CustomBatchLogger): self.hanging_request_check = AlertingHangingRequestCheck( slack_alerting_object=self, ) - self.alert_type_config: Dict[str, AlertTypeConfig] = {} + self.alert_type_config: dict[str, AlertTypeConfig] = {} if alert_type_config: for key, val in alert_type_config.items(): self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val - self.digest_buckets: Dict[str, DigestEntry] = {} + self.digest_buckets: dict[str, DigestEntry] = {} self.digest_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) def update_values( self, - alerting: Optional[List] = None, - alerting_threshold: Optional[float] = None, - alert_types: Optional[List[AlertType]] = None, - alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]] = None, - alerting_args: Optional[Dict] = None, - llm_router: Optional[Router] = None, - alert_type_config: Optional[Dict[str, dict]] = None, + alerting: list | None = None, + alerting_threshold: float | None = None, + alert_types: list[AlertType] | None = None, + alert_to_webhook_url: dict[AlertType, list[str] | str] | None = None, + alerting_args: dict | None = None, + llm_router: Router | None = None, + alert_type_config: dict[str, dict] | None = None, ): if alerting is not None: self.alerting = alerting @@ -134,9 +133,7 @@ class SlackAlerting(CustomBatchLogger): if llm_router is not None: self.llm_router = llm_router - def _prepare_outage_value_for_cache( - self, outage_value: Union[dict, ProviderRegionOutageModel, OutageModel] - ) -> dict: + def _prepare_outage_value_for_cache(self, outage_value: dict | ProviderRegionOutageModel | OutageModel) -> dict: """ Helper method to prepare outage value for Redis caching. Converts set objects to lists for JSON serialization. @@ -148,7 +145,7 @@ class SlackAlerting(CustomBatchLogger): cache_value["deployment_ids"] = list(cache_value["deployment_ids"]) return cache_value - def _restore_outage_value_from_cache(self, outage_value: Optional[dict]) -> Optional[dict]: + def _restore_outage_value_from_cache(self, outage_value: dict | None) -> dict | None: """ Helper method to restore outage value after retrieving from cache. Converts list objects back to sets for proper handling. @@ -210,7 +207,7 @@ class SlackAlerting(CustomBatchLogger): _deployment_latencies = metadata["_latency_per_deployment"] if len(_deployment_latencies) == 0: return None - _deployment_latency_map: Optional[dict] = None + _deployment_latency_map: dict | None = None try: # try sorting deployments by latency _deployment_latencies = sorted(_deployment_latencies.items(), key=lambda x: x[1]) @@ -290,10 +287,7 @@ class SlackAlerting(CustomBatchLogger): ## FAILED REQUESTS ## if deployment_metrics.failed_request: await self.internal_usage_cache.async_increment_cache( - key="{}:{}".format( - deployment_metrics.id, - SlackAlertingCacheKeys.failed_requests_key.value, - ), + key=f"{deployment_metrics.id}:{SlackAlertingCacheKeys.failed_requests_key.value}", value=1, parent_otel_span=None, # no attached request, this is a background operation ) @@ -303,7 +297,7 @@ class SlackAlerting(CustomBatchLogger): ## LATENCY ## if deployment_metrics.latency_per_output_token is not None: await self.internal_usage_cache.async_increment_cache( - key="{}:{}".format(deployment_metrics.id, SlackAlertingCacheKeys.latency_key.value), + key=f"{deployment_metrics.id}:{SlackAlertingCacheKeys.latency_key.value}", value=deployment_metrics.latency_per_output_token, parent_otel_span=None, # no attached request, this is a background operation ) @@ -333,8 +327,8 @@ class SlackAlerting(CustomBatchLogger): ids = router.get_model_ids() # get keys - failed_request_keys = ["{}:{}".format(id, SlackAlertingCacheKeys.failed_requests_key.value) for id in ids] - latency_keys = ["{}:{}".format(id, SlackAlertingCacheKeys.latency_key.value) for id in ids] + failed_request_keys = [f"{id}:{SlackAlertingCacheKeys.failed_requests_key.value}" for id in ids] + latency_keys = [f"{id}:{SlackAlertingCacheKeys.latency_key.value}" for id in ids] combined_metrics_keys = failed_request_keys + latency_keys # reduce cache calls @@ -445,7 +439,7 @@ class SlackAlerting(CustomBatchLogger): async def response_taking_too_long( self, - request_data: Optional[dict] = None, + request_data: dict | None = None, ): if self.alerting is None or self.alert_types is None: return @@ -471,7 +465,7 @@ class SlackAlerting(CustomBatchLogger): _cache: DualCache = self.internal_usage_cache message = "Failed Tracking Cost for " + error_message - _cache_key = "budget_alerts:failed_tracking:{}".format(failing_model) + _cache_key = f"budget_alerts:failed_tracking:{failing_model}" result = await _cache.async_get_cache(key=_cache_key) if result is None: await self.send_alert( @@ -528,16 +522,11 @@ class SlackAlerting(CustomBatchLogger): event_message = budget_alert_class.get_event_message() # Set default event unless we're in projected_limit_exceeded - event: Optional[ - Literal[ - "budget_crossed", - "threshold_crossed", - "projected_limit_exceeded", - "soft_budget_crossed", - ] - ] = "projected_limit_exceeded" if type == "projected_limit_exceeded" else None + event: ( + Literal["budget_crossed", "threshold_crossed", "projected_limit_exceeded", "soft_budget_crossed"] | None + ) = "projected_limit_exceeded" if type == "projected_limit_exceeded" else None - webhook_event: Optional[WebhookEvent] = None + webhook_event: WebhookEvent | None = None # percent of max_budget left to spend if user_info.max_budget is None and user_info.soft_budget is None: @@ -552,7 +541,7 @@ class SlackAlerting(CustomBatchLogger): # send alert if event is not None and user_info.event_group is not None: - _cache_key = "budget_alerts:{}:{}".format(event, _id) + _cache_key = f"budget_alerts:{event}:{_id}" result = await _cache.async_get_cache(key=_cache_key) if result is None: webhook_event = WebhookEvent( @@ -579,24 +568,10 @@ class SlackAlerting(CustomBatchLogger): def _get_event_and_event_message( self, user_info: CallInfo, - event: Optional[ - Literal[ - "budget_crossed", - "threshold_crossed", - "soft_budget_crossed", - "projected_limit_exceeded", - ] - ], + event: Literal["budget_crossed", "threshold_crossed", "soft_budget_crossed", "projected_limit_exceeded"] | None, event_message: str, - ) -> Tuple[ - Optional[ - Literal[ - "budget_crossed", - "threshold_crossed", - "soft_budget_crossed", - "projected_limit_exceeded", - ] - ], + ) -> tuple[ + Literal["budget_crossed", "threshold_crossed", "soft_budget_crossed", "projected_limit_exceeded"] | None, str, ]: """ @@ -642,7 +617,7 @@ class SlackAlerting(CustomBatchLogger): """ percent_left: float = 0.0 current_spend: float = user_info.spend - max_budget: Optional[float] = user_info.max_budget + max_budget: float | None = user_info.max_budget if max_budget is None: return percent_left if max_budget <= 0: @@ -666,11 +641,11 @@ class SlackAlerting(CustomBatchLogger): async def customer_spend_alert( self, - token: Optional[str], - key_alias: Optional[str], - end_user_id: Optional[str], - response_cost: Optional[float], - max_budget: Optional[float], + token: str | None, + key_alias: str | None, + end_user_id: str | None, + response_cost: float | None, + max_budget: float | None, ): if ( self.alerting is not None @@ -693,12 +668,12 @@ class SlackAlerting(CustomBatchLogger): projected_spend=None, event="spend_tracked", event_group=Litellm_EntityType.END_USER, - event_message="Customer spend tracked. Customer={}, spend={}".format(end_user_id, response_cost), + event_message=f"Customer spend tracked. Customer={end_user_id}, spend={response_cost}", ) await self.send_webhook_alert(webhook_event=event) - def _count_outage_alerts(self, alerts: List[int]) -> str: + def _count_outage_alerts(self, alerts: list[int]) -> str: """ Parameters: - alerts: List[int] -> list of error codes (either 408 or 500+) @@ -718,7 +693,7 @@ class SlackAlerting(CustomBatchLogger): error_msg = "" for key, value in error_breakdown.items(): if value > 0: - error_msg += "\n{}: {}\n".format(key, value) + error_msg += f"\n{key}: {value}\n" return error_msg @@ -728,7 +703,7 @@ class SlackAlerting(CustomBatchLogger): key: Literal["Model", "Region"], key_val: str, provider: str, - api_base: Optional[str], + api_base: str | None, outage_value: BaseOutageModel, ) -> str: """Format an alert message for slack""" @@ -788,9 +763,7 @@ class SlackAlerting(CustomBatchLogger): ### UNIQUE CACHE KEY ### cache_key = provider + region_name - outage_value: Optional[ProviderRegionOutageModel] = await self.internal_usage_cache.async_get_cache( - key=cache_key - ) + outage_value: ProviderRegionOutageModel | None = await self.internal_usage_cache.async_get_cache(key=cache_key) # Convert deployment_ids back to set if it was stored as a list if outage_value is not None: @@ -911,7 +884,7 @@ class SlackAlerting(CustomBatchLogger): max_alerts_size = 10 """ try: - outage_value: Optional[OutageModel] = await self.internal_usage_cache.async_get_cache(key=deployment_id) # type: ignore + outage_value: OutageModel | None = await self.internal_usage_cache.async_get_cache(key=deployment_id) # type: ignore if ( getattr(exception, "status_code", None) is None or ( @@ -1024,7 +997,7 @@ class SlackAlerting(CustomBatchLogger): for k, v in model_info.items(): if k == "input_cost_per_token" or k == "output_cost_per_token": # when converting to string it should not be 1.63e-06 - v = "{:.8f}".format(v) + v = f"{v:.8f}" model_info_str += f"{k}: {v}\n" @@ -1105,15 +1078,14 @@ Model Info: async def _check_if_using_premium_email_feature( self, premium_user: bool, - email_logo_url: Optional[str] = None, - email_support_contact: Optional[str] = None, + email_logo_url: str | None = None, + email_support_contact: str | None = None, ): from litellm.proxy.proxy_server import CommonProxyErrors, premium_user if premium_user is not True: if email_logo_url is not None or email_support_contact is not None: raise ValueError(f"Trying to Customize Email Alerting\n {CommonProxyErrors.not_premium_user.value}") - return async def send_key_created_or_user_invited_email(self, webhook_event: WebhookEvent) -> bool: try: @@ -1274,9 +1246,9 @@ Model Info: level: Literal["Low", "Medium", "High"], alert_type: AlertType, alerting_metadata: dict, - user_info: Optional[WebhookEvent] = None, - request_model: Optional[str] = None, - api_base: Optional[str] = None, + user_info: WebhookEvent | None = None, + request_model: str | None = None, + api_base: str | None = None, **kwargs, ): """ @@ -1323,7 +1295,7 @@ Model Info: if _atc is not None and _atc.digest: # Resolve webhook URL for this alert type (needed for digest entry) if self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url: - _digest_webhook: Optional[Union[str, List[str]]] = self.alert_to_webhook_url[alert_type] + _digest_webhook: str | list[str] | None = self.alert_to_webhook_url[alert_type] elif self.default_webhook_url is not None: _digest_webhook = self.default_webhook_url else: @@ -1376,7 +1348,7 @@ Model Info: # check if we find the slack webhook url in self.alert_to_webhook_url if self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url: - slack_webhook_url: Optional[Union[str, List[str]]] = self.alert_to_webhook_url[alert_type] + slack_webhook_url: str | list[str] | None = self.alert_to_webhook_url[alert_type] elif self.default_webhook_url is not None: slack_webhook_url = self.default_webhook_url else: @@ -1431,7 +1403,7 @@ Model Info: from datetime import datetime now = datetime.now() - flushed_keys: List[str] = [] + flushed_keys: list[str] = [] async with self.digest_lock: for key, entry in self.digest_buckets.items(): @@ -1495,7 +1467,7 @@ Model Info: try: await self._flush_digest_buckets() except Exception as e: - verbose_proxy_logger.debug(f"Error flushing digest buckets: {str(e)}") + verbose_proxy_logger.debug(f"Error flushing digest buckets: {e!s}") await self.flush_queue() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -1530,9 +1502,8 @@ Model Info: ) except Exception as e: verbose_proxy_logger.error( - f"[Non-Blocking Error] Slack Alerting: Got error in logging LLM deployment latency: {str(e)}" + f"[Non-Blocking Error] Slack Alerting: Got error in logging LLM deployment latency: {e!s}" ) - pass async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """Log failure + deployment latency""" @@ -1551,7 +1522,7 @@ Model Info: ) ) except Exception as e: - verbose_logger.debug(f"Exception raises -{str(e)}") + verbose_logger.debug(f"Exception raises -{e!s}") if isinstance(kwargs.get("exception", ""), APIError): if "outage_alerts" in self.alert_types: @@ -1601,7 +1572,7 @@ Model Info: return report_sent_bool - async def _run_scheduled_daily_report(self, llm_router: Optional[Any] = None): + async def _run_scheduled_daily_report(self, llm_router: Any | None = None): """ If 'daily_reports' enabled @@ -1785,8 +1756,6 @@ Model Info: except Exception as e: verbose_proxy_logger.error("Error sending weekly spend report %s", e) - pass - async def send_virtual_key_event_slack( self, key_event: VirtualKeyEvent, @@ -1830,9 +1799,7 @@ Model Info: except Exception as e: verbose_proxy_logger.error("Error sending send_virtual_key_event_slack %s", e) - return - - async def _request_is_completed(self, request_data: Optional[dict]) -> bool: + async def _request_is_completed(self, request_data: dict | None) -> bool: """ Returns True if the request is completed - either as a success or failure """ @@ -1842,8 +1809,8 @@ Model Info: if request_data.get("litellm_status", "") != "success" and request_data.get("litellm_status", "") != "fail": ## CHECK IF CACHE IS UPDATED litellm_call_id = request_data.get("litellm_call_id", "") - status: Optional[str] = await self.internal_usage_cache.async_get_cache( - key="request_status:{}".format(litellm_call_id), local_only=True + status: str | None = await self.internal_usage_cache.async_get_cache( + key=f"request_status:{litellm_call_id}", local_only=True ) if status is not None and (status == "success" or status == "fail"): return True diff --git a/litellm/integrations/SlackAlerting/utils.py b/litellm/integrations/SlackAlerting/utils.py index 4424bedba81..9587e0ae78b 100644 --- a/litellm/integrations/SlackAlerting/utils.py +++ b/litellm/integrations/SlackAlerting/utils.py @@ -3,7 +3,7 @@ Utils used for slack alerting """ import asyncio -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import litellm from litellm.proxy._types import AlertType @@ -18,8 +18,8 @@ else: def process_slack_alerting_variables( - alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]], -) -> Optional[Dict[AlertType, Union[List[str], str]]]: + alert_to_webhook_url: dict[AlertType, list[str] | str] | None, +) -> dict[AlertType, list[str] | str] | None: """ process alert_to_webhook_url - check if any urls are set as os.environ/SLACK_WEBHOOK_URL_1 read env var and set the correct value @@ -29,7 +29,7 @@ def process_slack_alerting_variables( for alert_type, webhook_urls in alert_to_webhook_url.items(): if isinstance(webhook_urls, list): - _webhook_values: List[str] = [] + _webhook_values: list[str] = [] for webhook_url in webhook_urls: if "os.environ/" in webhook_url: _env_value = get_secret(secret_name=webhook_url) @@ -56,8 +56,8 @@ def process_slack_alerting_variables( async def _add_langfuse_trace_id_to_alert( - request_data: Optional[dict] = None, -) -> Optional[str]: + request_data: dict | None = None, +) -> str | None: """ Returns langfuse trace url @@ -73,7 +73,7 @@ async def _add_langfuse_trace_id_to_alert( ######################################################### if request_data is not None and request_data.get("litellm_logging_obj", None) is not None: - trace_id: Optional[str] = None + trace_id: str | None = None litellm_logging_obj: Logging = request_data["litellm_logging_obj"] for _ in range(3): diff --git a/litellm/integrations/additional_logging_utils.py b/litellm/integrations/additional_logging_utils.py index 59319140a18..3f79a8ac007 100644 --- a/litellm/integrations/additional_logging_utils.py +++ b/litellm/integrations/additional_logging_utils.py @@ -7,7 +7,6 @@ Base class for Additional Logging Utils for CustomLoggers from abc import ABC, abstractmethod from datetime import datetime -from typing import Optional from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus @@ -21,15 +20,14 @@ class AdditionalLoggingUtils(ABC): """ Check if the service is healthy """ - pass @abstractmethod async def get_request_response_payload( self, request_id: str, - start_time_utc: Optional[datetime], - end_time_utc: Optional[datetime], - ) -> Optional[dict]: + start_time_utc: datetime | None, + end_time_utc: datetime | None, + ) -> dict | None: """ Get the request and response payload for a given `request_id` """ diff --git a/litellm/integrations/agentops/agentops.py b/litellm/integrations/agentops/agentops.py index c60e5cb0e2a..5295d8bf2be 100644 --- a/litellm/integrations/agentops/agentops.py +++ b/litellm/integrations/agentops/agentops.py @@ -4,7 +4,8 @@ AgentOps integration for LiteLLM - Provides OpenTelemetry tracing for LLM calls import os from dataclasses import dataclass -from typing import Optional, Dict, Any +from typing import Any + from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig from litellm.llms.custom_httpx.http_handler import _get_httpx_client @@ -12,9 +13,9 @@ from litellm.llms.custom_httpx.http_handler import _get_httpx_client @dataclass class AgentOpsConfig: endpoint: str = "https://otlp.agentops.cloud/v1/traces" - api_key: Optional[str] = None - service_name: Optional[str] = None - deployment_environment: Optional[str] = None + api_key: str | None = None + service_name: str | None = None + deployment_environment: str | None = None auth_endpoint: str = "https://api.agentops.ai/v3/auth/token" @classmethod @@ -47,7 +48,7 @@ class AgentOps(OpenTelemetry): def __init__( self, - config: Optional[AgentOpsConfig] = None, + config: AgentOpsConfig | None = None, ): if config is None: config = AgentOpsConfig.from_env() @@ -82,7 +83,7 @@ class AgentOps(OpenTelemetry): self.resource_attributes = resource_attrs - def _fetch_auth_token(self, api_key: str, auth_endpoint: str) -> Dict[str, Any]: + def _fetch_auth_token(self, api_key: str, auth_endpoint: str) -> dict[str, Any]: """ Fetch JWT authentication token from AgentOps API diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index faedf8ae1a3..751c8c01aae 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -10,7 +10,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and """ import copy -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -39,17 +39,17 @@ class AnthropicCacheControlHook(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Apply cache control directives based on specified injection points. @@ -59,7 +59,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): - non_default_params: dict - params with any global cache controls """ # Extract cache control injection points - injection_points: List[CacheControlInjectionPoint] = non_default_params.pop( + injection_points: list[CacheControlInjectionPoint] = non_default_params.pop( "cache_control_injection_points", [] ) if not injection_points: @@ -69,8 +69,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = copy.deepcopy(messages) # Separate message-level and non-message-level injection points - message_points: List[CacheControlMessageInjectionPoint] = [] - remaining_points: List[CacheControlInjectionPoint] = [] + message_points: list[CacheControlMessageInjectionPoint] = [] + remaining_points: list[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": message_points.append(cast(CacheControlMessageInjectionPoint, point)) @@ -99,10 +99,10 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _apply_message_injections( - points: List[CacheControlMessageInjectionPoint], - messages: List[AllMessageValues], + points: list[CacheControlMessageInjectionPoint], + messages: list[AllMessageValues], max_blocks: int, - ) -> List[AllMessageValues]: + ) -> list[AllMessageValues]: """Apply message-level cache control injection points in order. Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control @@ -151,11 +151,11 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _resolve_target_indices( - point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] - ) -> List[int]: + point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues] + ) -> list[int]: """Resolve which message indices an injection point targets.""" - _targetted_index: Optional[Union[int, str]] = point.get("index", None) - targetted_index: Optional[int] = None + _targetted_index: int | str | None = point.get("index", None) + targetted_index: int | None = None if isinstance(_targetted_index, str): try: targetted_index = int(_targetted_index) @@ -232,10 +232,10 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def apply_to_anthropic_messages_request( - messages: List[Dict], + messages: list[dict], system: str | list | None, - injection_points: List[CacheControlInjectionPoint], - ) -> Tuple[List[Dict], str | list | None, List[CacheControlInjectionPoint]]: + injection_points: list[CacheControlInjectionPoint], + ) -> tuple[list[dict], str | list | None, list[CacheControlInjectionPoint]]: """Apply cache control injection for the Anthropic-native v1/messages endpoint. Returns (messages, system, remaining_non_message_points). @@ -243,12 +243,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): if not injection_points: return messages, system, [] - processed_messages: List[Dict] = copy.deepcopy(messages) + processed_messages: list[dict] = copy.deepcopy(messages) processed_system = copy.deepcopy(system) if system is not None else None - message_points: List[CacheControlMessageInjectionPoint] = [] - system_points: List[CacheControlMessageInjectionPoint] = [] - remaining_points: List[CacheControlInjectionPoint] = [] + message_points: list[CacheControlMessageInjectionPoint] = [] + system_points: list[CacheControlMessageInjectionPoint] = [] + remaining_points: list[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": @@ -292,7 +292,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = AnthropicCacheControlHook._apply_message_injections( points=message_points, - messages=cast(List[AllMessageValues], processed_messages), + messages=cast(list[AllMessageValues], processed_messages), max_blocks=max_blocks - used_blocks, ) @@ -462,13 +462,13 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def maybe_inject_cache_control( - messages: List[Dict], + messages: list[dict], system: str | list | None, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], model: str | None = None, custom_llm_provider: str | None = None, tools: list[dict] | None = None, - ) -> Tuple[List[Dict], str | list | None]: + ) -> tuple[list[dict], str | list | None]: """Extract cache_control_injection_points from kwargs and apply if present. Configured points stand down entirely when the client already marked @@ -515,8 +515,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """Always return False since this is not a true prompt management system.""" @@ -524,12 +524,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return PromptManagementClient( @@ -542,12 +542,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return self._compile_prompt_helper( @@ -562,19 +562,19 @@ class AnthropicCacheControlHook(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """Async version - delegates to sync since no async operations needed.""" return self.get_chat_completion_prompt( model=model, @@ -591,15 +591,15 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) @staticmethod - def should_use_anthropic_cache_control_hook(non_default_params: Dict) -> bool: + def should_use_anthropic_cache_control_hook(non_default_params: dict) -> bool: if non_default_params.get("cache_control_injection_points", None): return True return False @staticmethod def get_custom_logger_for_anthropic_cache_control_hook( - non_default_params: Dict, - ) -> Optional[CustomLogger]: + non_default_params: dict, + ) -> CustomLogger | None: from litellm.litellm_core_utils.litellm_logging import ( _init_custom_logger_compatible_class, ) diff --git a/litellm/integrations/argilla.py b/litellm/integrations/argilla.py index a86b6f9e388..d41291f9f98 100644 --- a/litellm/integrations/argilla.py +++ b/litellm/integrations/argilla.py @@ -7,7 +7,7 @@ import json import os import random import types -from typing import Any, Dict, List, Optional +from typing import Any import httpx from pydantic import BaseModel # type: ignore @@ -41,9 +41,9 @@ def is_serializable(value): class ArgillaLogger(CustomBatchLogger): def __init__( self, - argilla_api_key: Optional[str] = None, - argilla_dataset_name: Optional[str] = None, - argilla_base_url: Optional[str] = None, + argilla_api_key: str | None = None, + argilla_dataset_name: str | None = None, + argilla_base_url: str | None = None, **kwargs, ): if litellm.argilla_transformation_object is None: @@ -69,7 +69,7 @@ class ArgillaLogger(CustomBatchLogger): self.flush_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) - def validate_argilla_transformation_object(self, argilla_transformation_object: Dict[str, Any]): + def validate_argilla_transformation_object(self, argilla_transformation_object: dict[str, Any]): if not isinstance(argilla_transformation_object, dict): raise Exception("'argilla_transformation_object' must be a dictionary, to log your payload to Argilla.") @@ -81,9 +81,9 @@ class ArgillaLogger(CustomBatchLogger): def get_credentials_from_env( self, - argilla_api_key: Optional[str], - argilla_dataset_name: Optional[str], - argilla_base_url: Optional[str], + argilla_api_key: str | None, + argilla_dataset_name: str | None, + argilla_base_url: str | None, ) -> ArgillaCredentialsObject: _credentials_api_key = argilla_api_key or os.getenv("ARGILLA_API_KEY") if _credentials_api_key is None: @@ -115,7 +115,7 @@ class ArgillaLogger(CustomBatchLogger): ARGILLA_DATASET_NAME=_credentials_dataset_name, ) - def get_chat_messages(self, payload: StandardLoggingPayload) -> List[Dict[str, Any]]: + def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, Any]]: payload_messages = payload.get("messages", None) if payload_messages is None: @@ -141,10 +141,10 @@ class ArgillaLogger(CustomBatchLogger): else: raise Exception(f"Invalid response format: {response}") - def _prepare_log_data(self, kwargs, response_obj, start_time, end_time) -> Optional[ArgillaItem]: + def _prepare_log_data(self, kwargs, response_obj, start_time, end_time) -> ArgillaItem | None: try: # Ensure everything in the payload is converted to str - payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if payload is None: raise Exception("Error logging request payload. Payload=none.") @@ -204,9 +204,7 @@ class ArgillaLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.debug( @@ -233,9 +231,7 @@ class ArgillaLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.debug( @@ -243,7 +239,7 @@ class ArgillaLogger(CustomBatchLogger): kwargs, response_obj, ) - payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) data = self._prepare_log_data(kwargs, response_obj, start_time, end_time) @@ -276,7 +272,7 @@ class ArgillaLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format(sampling_rate, random_sample) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.info("Langsmith Failure Event Logging!") diff --git a/litellm/integrations/arize/__init__.py b/litellm/integrations/arize/__init__.py index ab2627801e6..24271a9b926 100644 --- a/litellm/integrations/arize/__init__.py +++ b/litellm/integrations/arize/__init__.py @@ -1,16 +1,16 @@ import os -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: - from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.integrations.custom_prompt_management import CustomPromptManagement + from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .arize_phoenix_prompt_manager import ArizePhoenixPromptManager # Global instances -global_arize_config: Optional[dict] = None +global_arize_config: dict | None = None def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement": diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 44fd7a0d01a..032f490860d 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -1,5 +1,5 @@ import json -from typing import TYPE_CHECKING, Any, Dict, Optional, Type +from typing import TYPE_CHECKING, Any from typing_extensions import override @@ -31,7 +31,7 @@ from litellm.integrations._types.open_inference import ( class ArizeOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @override - def set_messages(span: "Span", kwargs: Dict[str, Any]): + def set_messages(span: "Span", kwargs: dict[str, Any]): messages = kwargs.get("messages") # for /chat/completions @@ -302,7 +302,7 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): ) -def _infer_open_inference_span_kind(call_type: Optional[str]) -> str: +def _infer_open_inference_span_kind(call_type: str | None) -> str: """ Map LiteLLM call types to OpenInference span kinds. """ @@ -360,7 +360,7 @@ def _infer_open_inference_span_kind(call_type: Optional[str]) -> str: return OpenInferenceSpanKindValues.UNKNOWN.value -def _set_tool_attributes(span: "Span", optional_tools: Optional[list], metadata_tools: Optional[list]): +def _set_tool_attributes(span: "Span", optional_tools: list | None, metadata_tools: list | None): """set tool attributes on span from optional_params or tool call metadata""" if optional_tools: for idx, tool in enumerate(optional_tools): @@ -408,7 +408,7 @@ def _set_tool_attributes(span: "Span", optional_tools: Optional[list], metadata_ ) -def set_attributes(span: "Span", kwargs, response_obj, attributes: Type[BaseLLMObsOTELAttributes]): +def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMObsOTELAttributes]): """ Populates span with OpenInference-compliant LLM attributes for Arize and Phoenix tracing. """ @@ -427,7 +427,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: Type[BaseLLMO try: optional_params = _sanitize_optional_params(kwargs.get("optional_params")) litellm_params = kwargs.get("litellm_params", {}) or {} - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") @@ -482,19 +482,19 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: Type[BaseLLMO ) -def _sanitize_optional_params(optional_params: Optional[dict]) -> dict: +def _sanitize_optional_params(optional_params: dict | None) -> dict: if not isinstance(optional_params, dict): return {} optional_params.pop("secret_fields", None) return optional_params -def _set_metadata_attributes(span: "Span", metadata: Optional[Any], span_attrs) -> None: +def _set_metadata_attributes(span: "Span", metadata: Any | None, span_attrs) -> None: if metadata is not None: safe_set_attribute(span, span_attrs.METADATA, safe_dumps(metadata)) -def _extract_metadata_tools(metadata: Optional[Any]) -> Optional[list]: +def _extract_metadata_tools(metadata: Any | None) -> list | None: if not isinstance(metadata, dict): return None llm_obj = metadata.get("llm") @@ -503,7 +503,7 @@ def _extract_metadata_tools(metadata: Optional[Any]) -> Optional[list]: return None -def _extract_optional_tools(optional_params: dict) -> Optional[list]: +def _extract_optional_tools(optional_params: dict) -> list | None: return optional_params.get("tools") if isinstance(optional_params, dict) else None @@ -544,7 +544,7 @@ def _set_request_attributes( safe_set_attribute(span, "llm.response.model", response_obj.get("model")) -def _set_model_params(span: "Span", model_params: Optional[dict], span_attrs) -> None: +def _set_model_params(span: "Span", model_params: dict | None, span_attrs) -> None: if not model_params: return @@ -606,7 +606,7 @@ def _coerce_response_obj_for_attrs(response_obj): return response_obj -def _coerce_text(value) -> Optional[str]: +def _coerce_text(value) -> str | None: """Best-effort text extraction from a message-content value. Returns None when no textual portion can be derived. Handles: @@ -650,7 +650,7 @@ def _to_plain_dict(value): return value -def _get_tool_calls(message) -> Optional[list]: +def _get_tool_calls(message) -> list | None: """Return ``message.tool_calls`` only when it's a non-empty list. Works for dicts and Pydantic message objects via ``_safe_get``. @@ -659,7 +659,7 @@ def _get_tool_calls(message) -> Optional[list]: return tool_calls if isinstance(tool_calls, list) and tool_calls else None -def _normalize_tool_call(raw_tc) -> Optional[Dict[str, Any]]: +def _normalize_tool_call(raw_tc) -> dict[str, Any] | None: """Normalize a single tool_call (dict or Pydantic) into a stable shape: {"id": str|None, "type": str, "function": {"name": str|None, "arguments": str|None}} @@ -879,7 +879,7 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None: safe_set_attribute(span, "llm.response.cost", cost_value) -def _is_passthrough_call_type(call_type: Optional[str]) -> bool: +def _is_passthrough_call_type(call_type: str | None) -> bool: if not call_type: return False lowered = str(call_type).lower() diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py index e5fdb231933..9d743659135 100644 --- a/litellm/integrations/arize/arize.py +++ b/litellm/integrations/arize/arize.py @@ -6,7 +6,7 @@ this file has Arize ai specific helper functions import os from datetime import datetime -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from litellm.integrations.arize import _utils from litellm.integrations.arize._utils import ArizeOTELAttributes @@ -61,16 +61,13 @@ class ArizeLogger(OpenTelemetry): ``open_telemetry_logger``. That attribute is reserved for the primary ``otel`` callback which handles proxy-level parent spans. """ - pass - def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): + def set_attributes(self, span: Span, kwargs, response_obj: Any | None): ArizeLogger.set_arize_attributes(span, kwargs, response_obj) - return @staticmethod def set_arize_attributes(span: Span, kwargs, response_obj): _utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes) - return @staticmethod def get_arize_config() -> ArizeConfig: @@ -116,25 +113,23 @@ class ArizeLogger(OpenTelemetry): async def async_service_success_hook( self, payload: ServiceLoggerPayload, - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[datetime, float]] = None, - event_metadata: Optional[dict] = None, + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: datetime | float | None = None, + event_metadata: dict | None = None, ): """Arize is used mainly for LLM I/O tracing, sending router+caching metrics adds bloat to arize logs""" - pass async def async_service_failure_hook( self, payload: ServiceLoggerPayload, - error: Optional[str] = "", - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[float, datetime]] = None, - event_metadata: Optional[dict] = None, + error: str | None = "", + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: float | datetime | None = None, + event_metadata: dict | None = None, ): """Arize is used mainly for LLM I/O tracing, sending router+caching metrics adds bloat to arize logs""" - pass # def create_litellm_proxy_request_started_span( # self, @@ -174,12 +169,12 @@ class ArizeLogger(OpenTelemetry): except Exception as e: return { "status": "unhealthy", - "error_message": f"Arize health check failed: {str(e)}", + "error_message": f"Arize health check failed: {e!s}", } def construct_dynamic_otel_headers( self, standard_callback_dynamic_params: StandardCallbackDynamicParams - ) -> Optional[dict]: + ) -> dict | None: """ Construct dynamic Arize headers from standard callback dynamic params diff --git a/litellm/integrations/arize/arize_phoenix.py b/litellm/integrations/arize/arize_phoenix.py index db7aed1a71c..db698dd6b77 100644 --- a/litellm/integrations/arize/arize_phoenix.py +++ b/litellm/integrations/arize/arize_phoenix.py @@ -1,7 +1,7 @@ import os import threading from collections import OrderedDict -from typing import TYPE_CHECKING, Any, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union from litellm._logging import verbose_logger from litellm.integrations.arize import _utils @@ -12,8 +12,7 @@ if TYPE_CHECKING: from opentelemetry.sdk.trace import TracerProvider from opentelemetry.sdk.trace.export import SpanProcessor from opentelemetry.trace import Span as _Span - from opentelemetry.trace import SpanKind - from opentelemetry.trace import Tracer + from opentelemetry.trace import SpanKind, Tracer from litellm.integrations.opentelemetry import OpenTelemetry as _OpenTelemetry from litellm.integrations.opentelemetry import ( @@ -176,7 +175,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore self._project_providers[project_name] = new_provider return new_provider.get_tracer(LITELLM_TRACER_NAME) - def _resolve_tracer_for_kwargs(self, kwargs: dict) -> Tuple[str, Tracer]: + def _resolve_tracer_for_kwargs(self, kwargs: dict) -> tuple[str, Tracer]: """Resolve project name once and return the matching tracer.""" project_name = self._resolve_project_name(kwargs) return project_name, self._get_tracer_for(project_name) @@ -193,19 +192,16 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore ``open_telemetry_logger``. That attribute is reserved for the primary ``otel`` callback which handles proxy-level parent spans. """ - pass - def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): + def set_attributes(self, span: Span, kwargs, response_obj: Any | None): ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj) - return @staticmethod def set_arize_phoenix_attributes(span: Span, kwargs, response_obj): _utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes) - return @staticmethod - def _normalize_project_name(name: Optional[str]) -> Optional[str]: + def _normalize_project_name(name: str | None) -> str | None: if name is None: return None normalized = str(name).strip() @@ -236,7 +232,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore return isinstance(litellm_params, dict) and bool(litellm_params.get("proxy_server_request")) @staticmethod - def _project_from_metadata_dict(metadata: dict, metadata_key: str, *, proxy_mode: bool) -> Optional[str]: + def _project_from_metadata_dict(metadata: dict, metadata_key: str, *, proxy_mode: bool) -> str | None: """ Read a Phoenix project field from proxy/SDK metadata. @@ -255,7 +251,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore return None @staticmethod - def _metadata_project_from_kwargs(kwargs: dict, metadata_key: str) -> Optional[str]: + def _metadata_project_from_kwargs(kwargs: dict, metadata_key: str) -> str | None: proxy_mode = ArizePhoenixLogger._is_proxy_request(kwargs) for metadata in ArizePhoenixLogger._iter_metadata_dicts_from_kwargs(kwargs): project = ArizePhoenixLogger._project_from_metadata_dict(metadata, metadata_key, proxy_mode=proxy_mode) @@ -288,7 +284,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore return "default" - def _get_phoenix_context(self, kwargs, tracer: Optional[Tracer] = None): + def _get_phoenix_context(self, kwargs, tracer: Tracer | None = None): """ Build a trace context for Phoenix's dedicated TracerProvider. diff --git a/litellm/integrations/arize/arize_phoenix_client.py b/litellm/integrations/arize/arize_phoenix_client.py index 7c0715d2e1e..6f1787fae9e 100644 --- a/litellm/integrations/arize/arize_phoenix_client.py +++ b/litellm/integrations/arize/arize_phoenix_client.py @@ -3,7 +3,7 @@ Arize Phoenix API client for fetching prompt versions from Arize Phoenix. """ import urllib.parse -from typing import Any, Dict, Optional +from typing import Any from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -27,7 +27,7 @@ class ArizePhoenixClient: - Direct API base URL configuration """ - def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None): + def __init__(self, api_key: str | None = None, api_base: str | None = None): """ Initialize the Arize Phoenix client. @@ -53,7 +53,7 @@ class ArizePhoenixClient: # Initialize HTTPHandler self.http_handler = HTTPHandler(disable_default_headers=True) - def get_prompt_version(self, prompt_version_id: str) -> Optional[Dict[str, Any]]: + def get_prompt_version(self, prompt_version_id: str) -> dict[str, Any] | None: """ Fetch a prompt version from Arize Phoenix. diff --git a/litellm/integrations/arize/arize_phoenix_prompt_manager.py b/litellm/integrations/arize/arize_phoenix_prompt_manager.py index 4053b725a0f..ca74835e167 100644 --- a/litellm/integrations/arize/arize_phoenix_prompt_manager.py +++ b/litellm/integrations/arize/arize_phoenix_prompt_manager.py @@ -3,7 +3,7 @@ Arize Phoenix prompt manager that integrates with LiteLLM's prompt management sy Fetches prompt versions from Arize Phoenix and provides workspace-based access control. """ -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any from jinja2 import DictLoader, select_autoescape from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -28,9 +28,9 @@ class ArizePhoenixPromptTemplate: def __init__( self, template_id: str, - messages: List[Dict[str, Any]], - metadata: Dict[str, Any], - model: Optional[str] = None, + messages: list[dict[str, Any]], + metadata: dict[str, Any], + model: str | None = None, ): self.template_id = template_id self.messages = messages @@ -61,14 +61,14 @@ class ArizePhoenixTemplateManager: def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - prompt_id: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + prompt_id: str | None = None, ): self.api_key = api_key self.api_base = api_base self.prompt_id = prompt_id - self.prompts: Dict[str, ArizePhoenixPromptTemplate] = {} + self.prompts: dict[str, ArizePhoenixPromptTemplate] = {} self.arize_client = ArizePhoenixClient(api_key=self.api_key, api_base=self.api_base) # Templates fetched from Arize Phoenix come from external workspace @@ -107,7 +107,7 @@ class ArizePhoenixTemplateManager: except Exception as e: raise Exception(f"Failed to load prompt version '{prompt_version_id}' from Arize Phoenix: {e}") - def _parse_prompt_data(self, data: Dict[str, Any], prompt_version_id: str) -> ArizePhoenixPromptTemplate: + def _parse_prompt_data(self, data: dict[str, Any], prompt_version_id: str) -> ArizePhoenixPromptTemplate: """Parse Arize Phoenix prompt data and extract messages and metadata.""" template_data = data.get("template", {}) messages = template_data.get("messages", []) @@ -146,13 +146,13 @@ class ArizePhoenixTemplateManager: metadata=metadata, ) - def render_template(self, template_id: str, variables: Optional[Dict[str, Any]] = None) -> List[AllMessageValues]: + def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> list[AllMessageValues]: """Render a template with the given variables and return formatted messages.""" if template_id not in self.prompts: raise ValueError(f"Template '{template_id}' not found") template = self.prompts[template_id] - rendered_messages: List[AllMessageValues] = [] + rendered_messages: list[AllMessageValues] = [] for message in template.messages: role = message.get("role", "user") @@ -180,11 +180,11 @@ class ArizePhoenixTemplateManager: return rendered_messages - def get_template(self, template_id: str) -> Optional[ArizePhoenixPromptTemplate]: + def get_template(self, template_id: str) -> ArizePhoenixPromptTemplate | None: """Get a template by ID.""" return self.prompts.get(template_id) - def list_templates(self) -> List[str]: + def list_templates(self) -> list[str]: """List all available template IDs.""" return list(self.prompts.keys()) @@ -215,16 +215,16 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - prompt_id: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + prompt_id: str | None = None, **kwargs, ): super().__init__(**kwargs) self.api_key = api_key self.api_base = api_base self.prompt_id = prompt_id - self._prompt_manager: Optional[ArizePhoenixTemplateManager] = None + self._prompt_manager: ArizePhoenixTemplateManager | None = None @property def integration_name(self) -> str: @@ -245,8 +245,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def get_prompt_template( self, prompt_id: str, - prompt_variables: Optional[Dict[str, Any]] = None, - ) -> Tuple[List[AllMessageValues], Dict[str, Any]]: + prompt_variables: dict[str, Any] | None = None, + ) -> tuple[list[AllMessageValues], dict[str, Any]]: """ Get a prompt template and render it with variables. @@ -289,14 +289,14 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def pre_call_hook( self, - user_id: Optional[str], - messages: List[AllMessageValues], - function_call: Optional[Union[Dict[str, Any], str]] = None, - litellm_params: Optional[Dict[str, Any]] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[Dict[str, Any]] = None, + user_id: str | None, + messages: list[AllMessageValues], + function_call: dict[str, Any] | str | None = None, + litellm_params: dict[str, Any] | None = None, + prompt_id: str | None = None, + prompt_variables: dict[str, Any] | None = None, **kwargs, - ) -> Tuple[List[AllMessageValues], Optional[Dict[str, Any]]]: + ) -> tuple[list[AllMessageValues], dict[str, Any] | None]: """ Pre-call hook that processes the prompt template before making the LLM call. """ @@ -342,7 +342,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement): litellm._logging.verbose_proxy_logger.error(f"Error in Arize Phoenix prompt pre_call_hook: {e}") return messages, litellm_params - def get_available_prompts(self) -> List[str]: + def get_available_prompts(self) -> list[str]: """Get list of available prompt IDs.""" return self.prompt_manager.list_templates() @@ -354,8 +354,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """ @@ -368,12 +368,12 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Compile an Arize Phoenix prompt template into a PromptManagementClient structure. @@ -422,12 +422,12 @@ class ArizePhoenixPromptManager(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Async version of compile prompt helper. Since Arize Phoenix operations are synchronous, @@ -447,17 +447,17 @@ class ArizePhoenixPromptManager(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Get chat completion prompt from Arize Phoenix and return processed model, messages, and parameters. """ diff --git a/litellm/integrations/athina.py b/litellm/integrations/athina.py index d1bf8e68624..f57c4c8b545 100644 --- a/litellm/integrations/athina.py +++ b/litellm/integrations/athina.py @@ -82,4 +82,3 @@ class AthinaLogger: print_verbose(f"Athina Logger Succeeded - {response.text}") except Exception as e: print_verbose(f"Athina Logger Error - {e}, Stack trace: {traceback.format_exc()}") - pass diff --git a/litellm/integrations/azure_sentinel/azure_sentinel.py b/litellm/integrations/azure_sentinel/azure_sentinel.py index 5f8afe58cb0..f0200b75c43 100644 --- a/litellm/integrations/azure_sentinel/azure_sentinel.py +++ b/litellm/integrations/azure_sentinel/azure_sentinel.py @@ -16,7 +16,6 @@ import asyncio import os import time import traceback -from typing import List, Optional, Union from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -35,13 +34,13 @@ class AzureSentinelLogger(CustomBatchLogger): def __init__( self, - dcr_immutable_id: Optional[str] = None, - stream_name: Optional[str] = None, - endpoint: Optional[str] = None, - tenant_id: Optional[str] = None, - client_id: Optional[str] = None, - client_secret: Optional[str] = None, - audit_stream_name: Optional[str] = None, + dcr_immutable_id: str | None = None, + stream_name: str | None = None, + endpoint: str | None = None, + tenant_id: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + audit_stream_name: str | None = None, **kwargs, ): """ @@ -120,14 +119,14 @@ class AzureSentinelLogger(CustomBatchLogger): # OAuth2 scope for Azure Monitor self.oauth_scope = "https://monitor.azure.com/.default" - self.oauth_token: Optional[str] = None - self.oauth_token_expires_at: Optional[float] = None + self.oauth_token: str | None = None + self.oauth_token_expires_at: float | None = None self.flush_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) asyncio.create_task(self.periodic_flush()) - self.log_queue: List[StandardLoggingPayload] = [] - self.audit_log_queue: List[StandardAuditLogPayload] = [] + self.log_queue: list[StandardLoggingPayload] = [] + self.audit_log_queue: list[StandardAuditLogPayload] = [] @staticmethod def _build_api_endpoint(endpoint: str, dcr_immutable_id: str, stream_name: str) -> str: @@ -204,8 +203,7 @@ class AzureSentinelLogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Azure Sentinel Layer Error - {e!s}\n{traceback.format_exc()}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """ @@ -235,8 +233,7 @@ class AzureSentinelLogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Azure Sentinel Layer Error - {e!s}\n{traceback.format_exc()}") async def async_log_audit_log_event(self, audit_log: StandardAuditLogPayload) -> None: """ @@ -259,8 +256,7 @@ class AzureSentinelLogger(CustomBatchLogger): await self.async_send_audit_batch() except Exception as e: - verbose_logger.exception(f"Azure Sentinel Audit Log Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Azure Sentinel Audit Log Layer Error - {e!s}\n{traceback.format_exc()}") async def async_send_batch(self): """ @@ -287,7 +283,7 @@ class AzureSentinelLogger(CustomBatchLogger): async def _async_send_batch_to_api( self, - log_queue: List[Union[StandardLoggingPayload, StandardAuditLogPayload]], + log_queue: list[StandardLoggingPayload | StandardAuditLogPayload], api_endpoint: str, log_type: str, ) -> None: @@ -327,7 +323,7 @@ class AzureSentinelLogger(CustomBatchLogger): ) except Exception as e: - verbose_logger.exception(f"Azure Sentinel Error sending batch API - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"Azure Sentinel Error sending batch API - {e!s}\n{traceback.format_exc()}") finally: log_queue.clear() diff --git a/litellm/integrations/azure_storage/azure_storage.py b/litellm/integrations/azure_storage/azure_storage.py index 5ccd1a86bff..bbd6e9698bb 100644 --- a/litellm/integrations/azure_storage/azure_storage.py +++ b/litellm/integrations/azure_storage/azure_storage.py @@ -1,20 +1,19 @@ import asyncio import os import time -from litellm._uuid import uuid from datetime import datetime, timedelta -from typing import List, Optional from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.constants import _DEFAULT_TTL_FOR_HTTPX_CLIENTS, AZURE_STORAGE_MSFT_VERSION from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.azure.common_utils import get_azure_ad_token_from_entra_id from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.types.utils import StandardLoggingPayload @@ -30,7 +29,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): self.tenant_id = os.getenv("AZURE_STORAGE_TENANT_ID") self.client_id = os.getenv("AZURE_STORAGE_CLIENT_ID") self.client_secret = os.getenv("AZURE_STORAGE_CLIENT_SECRET") - self.azure_storage_account_key: Optional[str] = os.getenv("AZURE_STORAGE_ACCOUNT_KEY") + self.azure_storage_account_key: str | None = os.getenv("AZURE_STORAGE_ACCOUNT_KEY") # Required Env Variables for Azure Storage _azure_storage_account_name = os.getenv("AZURE_STORAGE_ACCOUNT_NAME") @@ -43,19 +42,19 @@ class AzureBlobStorageLogger(CustomBatchLogger): self.azure_storage_file_system: str = _azure_storage_file_system self._service_client = None # Time that the azure service client expires, in order to reset the connection pool and keep it fresh - self._service_client_timeout: Optional[float] = None + self._service_client_timeout: float | None = None # Internal variables used for Token based authentication - self.azure_auth_token: Optional[str] = None # the Azure AD token to use for Azure Storage API requests - self.token_expiry: Optional[datetime] = None # the expiry time of the currentAzure AD token + self.azure_auth_token: str | None = None # the Azure AD token to use for Azure Storage API requests + self.token_expiry: datetime | None = None # the expiry time of the currentAzure AD token asyncio.create_task(self.periodic_flush()) self.flush_lock = asyncio.Lock() - self.log_queue: List[StandardLoggingPayload] = [] + self.log_queue: list[StandardLoggingPayload] = [] super().__init__(**kwargs, flush_lock=self.flush_lock) except Exception as e: verbose_logger.exception( - f"AzureBlobStorageLogger: Got exception on init AzureBlobStorageLogger client {str(e)}" + f"AzureBlobStorageLogger: Got exception on init AzureBlobStorageLogger client {e!s}" ) raise e @@ -72,7 +71,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): "AzureBlobStorageLogger: Logging - Enters logging function for model %s", kwargs, ) - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_payload is not set") @@ -80,8 +79,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): self.log_queue.append(standard_logging_payload) except Exception as e: - verbose_logger.exception(f"AzureBlobStorageLogger Layer Error - {str(e)}") - pass + verbose_logger.exception(f"AzureBlobStorageLogger Layer Error - {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """ @@ -96,15 +94,14 @@ class AzureBlobStorageLogger(CustomBatchLogger): "AzureBlobStorageLogger: Logging - Enters logging function for model %s", kwargs, ) - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_payload is not set") self.log_queue.append(standard_logging_payload) except Exception as e: - verbose_logger.exception(f"AzureBlobStorageLogger Layer Error - {str(e)}") - pass + verbose_logger.exception(f"AzureBlobStorageLogger Layer Error - {e!s}") async def async_send_batch(self): """ @@ -127,7 +124,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): await self.async_upload_payload_to_azure_blob_storage(payload=payload) except Exception as e: - verbose_logger.exception(f"AzureBlobStorageLogger Error sending batch API - {str(e)}") + verbose_logger.exception(f"AzureBlobStorageLogger Error sending batch API - {e!s}") async def async_upload_payload_to_azure_blob_storage(self, payload: StandardLoggingPayload): """ @@ -156,7 +153,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): verbose_logger.debug(f"Successfully uploaded log to Azure Blob Storage: {filename}") except Exception as e: - verbose_logger.exception(f"Error uploading to Azure Blob Storage: {str(e)}") + verbose_logger.exception(f"Error uploading to Azure Blob Storage: {e!s}") raise e async def _create_file(self, client: AsyncHTTPHandler, base_url: str): @@ -172,7 +169,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): response.raise_for_status() verbose_logger.debug("Successfully created file resource") except Exception as e: - verbose_logger.exception(f"Error creating file resource: {str(e)}") + verbose_logger.exception(f"Error creating file resource: {e!s}") raise async def _append_data(self, client: AsyncHTTPHandler, base_url: str, json_payload: str): @@ -192,7 +189,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): response.raise_for_status() verbose_logger.debug("Successfully appended data") except Exception as e: - verbose_logger.exception(f"Error appending data: {str(e)}") + verbose_logger.exception(f"Error appending data: {e!s}") raise async def _flush_data(self, client: AsyncHTTPHandler, base_url: str, position: int): @@ -208,7 +205,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): response.raise_for_status() verbose_logger.debug("Successfully flushed data") except Exception as e: - verbose_logger.exception(f"Error flushing data: {str(e)}") + verbose_logger.exception(f"Error flushing data: {e!s}") raise ####### Helper methods to managing Authentication to Azure Storage ####### @@ -236,9 +233,9 @@ class AzureBlobStorageLogger(CustomBatchLogger): def get_azure_ad_token_from_azure_storage( self, - tenant_id: Optional[str], - client_id: Optional[str], - client_secret: Optional[str], + tenant_id: str | None, + client_id: str | None, + client_secret: str | None, ) -> str: """ Gets Azure AD token to use for Azure Storage API requests @@ -348,4 +345,4 @@ class AzureBlobStorageLogger(CustomBatchLogger): verbose_logger.debug(f"Successfully uploaded and wrote to {today}/{file_name}") except Exception as e: - verbose_logger.exception(f"Error occurred: {str(e)}") + verbose_logger.exception(f"Error occurred: {e!s}") diff --git a/litellm/integrations/bitbucket/__init__.py b/litellm/integrations/bitbucket/__init__.py index 2b9bd568e32..28f645597e1 100644 --- a/litellm/integrations/bitbucket/__init__.py +++ b/litellm/integrations/bitbucket/__init__.py @@ -1,16 +1,17 @@ -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: - from .bitbucket_prompt_manager import BitBucketPromptManager - from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.integrations.custom_prompt_management import CustomPromptManagement + from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec + + from .bitbucket_prompt_manager import BitBucketPromptManager from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .bitbucket_prompt_manager import BitBucketPromptManager # Global instances -global_bitbucket_config: Optional[dict] = None +global_bitbucket_config: dict | None = None def set_global_bitbucket_config(config: dict) -> None: @@ -57,6 +58,6 @@ prompt_initializer_registry = { # Export public API __all__ = [ "BitBucketPromptManager", - "set_global_bitbucket_config", "global_bitbucket_config", + "set_global_bitbucket_config", ] diff --git a/litellm/integrations/bitbucket/bitbucket_client.py b/litellm/integrations/bitbucket/bitbucket_client.py index c02d56811a7..756d1bed80c 100644 --- a/litellm/integrations/bitbucket/bitbucket_client.py +++ b/litellm/integrations/bitbucket/bitbucket_client.py @@ -4,7 +4,7 @@ BitBucket API client for fetching .prompt files from BitBucket repositories. import base64 import urllib.parse -from typing import Any, Dict, List, Optional +from typing import Any from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -31,7 +31,7 @@ class BitBucketClient: - Branch-specific file fetching """ - def __init__(self, config: Dict[str, Any]): + def __init__(self, config: dict[str, Any]): """ Initialize the BitBucket client. @@ -74,7 +74,7 @@ class BitBucketClient: # Initialize HTTPHandler self.http_handler = HTTPHandler() - def get_file_content(self, file_path: str) -> Optional[str]: + def get_file_content(self, file_path: str) -> str | None: """ Fetch the content of a file from the BitBucket repository. @@ -117,7 +117,7 @@ class BitBucketClient: else: raise Exception(f"Error fetching file '{file_path}': {e}") - def list_files(self, directory_path: str = "", file_extension: str = ".prompt") -> List[str]: + def list_files(self, directory_path: str = "", file_extension: str = ".prompt") -> list[str]: """ List files in a directory with a specific extension. @@ -162,7 +162,7 @@ class BitBucketClient: else: raise Exception(f"Error listing files in '{directory_path}': {e}") - def get_repository_info(self) -> Dict[str, Any]: + def get_repository_info(self) -> dict[str, Any]: """ Get information about the repository. @@ -191,7 +191,7 @@ class BitBucketClient: except Exception: return False - def get_branches(self) -> List[Dict[str, Any]]: + def get_branches(self) -> list[dict[str, Any]]: """ Get list of branches in the repository. @@ -209,7 +209,7 @@ class BitBucketClient: except Exception as e: raise Exception(f"Failed to get branches: {e}") - def get_file_metadata(self, file_path: str) -> Optional[Dict[str, Any]]: + def get_file_metadata(self, file_path: str) -> dict[str, Any] | None: """ Get metadata about a file (size, last modified, etc.). diff --git a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py index 6dca4d76c04..c76466b2f40 100644 --- a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py +++ b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py @@ -3,7 +3,7 @@ BitBucket prompt manager that integrates with LiteLLM's prompt management system Fetches .prompt files from BitBucket repositories and provides team-based access control. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from jinja2 import DictLoader, select_autoescape from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -34,8 +34,8 @@ class BitBucketPromptTemplate: self, template_id: str, content: str, - metadata: Dict[str, Any], - model: Optional[str] = None, + metadata: dict[str, Any], + model: str | None = None, ): self.template_id = template_id self.content = content @@ -65,12 +65,12 @@ class BitBucketTemplateManager: def __init__( self, - bitbucket_config: Dict[str, Any], - prompt_id: Optional[str] = None, + bitbucket_config: dict[str, Any], + prompt_id: str | None = None, ): self.bitbucket_config = bitbucket_config self.prompt_id = prompt_id - self.prompts: Dict[str, BitBucketPromptTemplate] = {} + self.prompts: dict[str, BitBucketPromptTemplate] = {} self.bitbucket_client = BitBucketClient(bitbucket_config) # Templates fetched from a BitBucket repo are not trustworthy: @@ -123,7 +123,7 @@ class BitBucketTemplateManager: template_content = content # Parse YAML frontmatter - metadata: Dict[str, Any] = {} + metadata: dict[str, Any] = {} if frontmatter_str: try: import yaml @@ -141,9 +141,9 @@ class BitBucketTemplateManager: metadata=metadata, ) - def _parse_yaml_basic(self, yaml_str: str) -> Dict[str, Any]: + def _parse_yaml_basic(self, yaml_str: str) -> dict[str, Any]: """Basic YAML parser for simple cases when PyYAML is not available.""" - result: Dict[str, Any] = {} + result: dict[str, Any] = {} for line in yaml_str.split("\n"): line = line.strip() if ":" in line and not line.startswith("#"): @@ -162,7 +162,7 @@ class BitBucketTemplateManager: result[key] = value.strip("\"'") return result - def render_template(self, template_id: str, variables: Optional[Dict[str, Any]] = None) -> str: + def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str: """Render a template with the given variables.""" if template_id not in self.prompts: raise ValueError(f"Template '{template_id}' not found") @@ -172,11 +172,11 @@ class BitBucketTemplateManager: return jinja_template.render(**(variables or {})) - def get_template(self, template_id: str) -> Optional[BitBucketPromptTemplate]: + def get_template(self, template_id: str) -> BitBucketPromptTemplate | None: """Get a template by ID.""" return self.prompts.get(template_id) - def list_templates(self) -> List[str]: + def list_templates(self) -> list[str]: """List all available template IDs.""" return list(self.prompts.keys()) @@ -209,12 +209,12 @@ class BitBucketPromptManager(CustomPromptManagement): def __init__( self, - bitbucket_config: Dict[str, Any], - prompt_id: Optional[str] = None, + bitbucket_config: dict[str, Any], + prompt_id: str | None = None, ): self.bitbucket_config = bitbucket_config self.prompt_id = prompt_id - self._prompt_manager: Optional[BitBucketTemplateManager] = None + self._prompt_manager: BitBucketTemplateManager | None = None @property def integration_name(self) -> str: @@ -234,8 +234,8 @@ class BitBucketPromptManager(CustomPromptManagement): def get_prompt_template( self, prompt_id: str, - prompt_variables: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + prompt_variables: dict[str, Any] | None = None, + ) -> tuple[str, dict[str, Any]]: """ Get a prompt template and render it with variables. @@ -265,14 +265,14 @@ class BitBucketPromptManager(CustomPromptManagement): def pre_call_hook( self, - user_id: Optional[str], - messages: List[AllMessageValues], - function_call: Optional[Union[Dict[str, Any], str]] = None, - litellm_params: Optional[Dict[str, Any]] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[Dict[str, Any]] = None, + user_id: str | None, + messages: list[AllMessageValues], + function_call: dict[str, Any] | str | None = None, + litellm_params: dict[str, Any] | None = None, + prompt_id: str | None = None, + prompt_variables: dict[str, Any] | None = None, **kwargs, - ) -> Tuple[List[AllMessageValues], Optional[Dict[str, Any]]]: + ) -> tuple[list[AllMessageValues], dict[str, Any] | None]: """ Pre-call hook that processes the prompt template before making the LLM call. """ @@ -289,7 +289,7 @@ class BitBucketPromptManager(CustomPromptManagement): # Merge with existing messages if parsed_messages: # If we have parsed messages, use them instead of the original messages - final_messages: List[AllMessageValues] = parsed_messages + final_messages: list[AllMessageValues] = parsed_messages else: # If no messages were parsed, prepend the prompt to existing messages final_messages = [ @@ -323,7 +323,7 @@ class BitBucketPromptManager(CustomPromptManagement): litellm._logging.verbose_proxy_logger.error(f"Error in BitBucket prompt pre_call_hook: {e}") return messages, litellm_params - def _parse_prompt_to_messages(self, prompt_content: str) -> List[AllMessageValues]: + def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]: """ Parse prompt content into a list of messages. Handles both simple prompts and multi-role conversations. @@ -385,13 +385,13 @@ class BitBucketPromptManager(CustomPromptManagement): def post_call_hook( self, - user_id: Optional[str], + user_id: str | None, response: Any, - input_messages: List[AllMessageValues], - function_call: Optional[Union[Dict[str, Any], str]] = None, - litellm_params: Optional[Dict[str, Any]] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[Dict[str, Any]] = None, + input_messages: list[AllMessageValues], + function_call: dict[str, Any] | str | None = None, + litellm_params: dict[str, Any] | None = None, + prompt_id: str | None = None, + prompt_variables: dict[str, Any] | None = None, **kwargs, ) -> Any: """ @@ -399,7 +399,7 @@ class BitBucketPromptManager(CustomPromptManagement): """ return response - def get_available_prompts(self) -> List[str]: + def get_available_prompts(self) -> list[str]: """Get list of available prompt IDs.""" return self.prompt_manager.list_templates() @@ -411,8 +411,8 @@ class BitBucketPromptManager(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """ @@ -425,12 +425,12 @@ class BitBucketPromptManager(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Compile a BitBucket prompt template into a PromptManagementClient structure. @@ -483,12 +483,12 @@ class BitBucketPromptManager(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Async version of compile prompt helper. Since BitBucket operations use sync client, @@ -509,17 +509,17 @@ class BitBucketPromptManager(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Get chat completion prompt from BitBucket and return processed model, messages, and parameters. """ @@ -539,19 +539,19 @@ class BitBucketPromptManager(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Async version - delegates to PromptManagementBase async implementation. """ diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index 686c37d3e17..a4f3335809a 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -3,15 +3,14 @@ import os from datetime import datetime -from typing import Dict, Optional import httpx import litellm from litellm import verbose_logger from litellm.integrations.braintrust_mock_client import ( - should_use_braintrust_mock, create_mock_braintrust_client, + should_use_braintrust_mock, ) from litellm.integrations.custom_logger import CustomLogger from litellm.llms.custom_httpx.http_handler import ( @@ -34,7 +33,7 @@ def get_utc_datetime(): class BraintrustLogger(CustomLogger): - def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> None: + def __init__(self, api_key: str | None = None, api_base: str | None = None) -> None: super().__init__() self.is_mock_mode = should_use_braintrust_mock() if self.is_mock_mode: @@ -48,11 +47,11 @@ class BraintrustLogger(CustomLogger): "Authorization": "Bearer " + self.api_key, "Content-Type": "application/json", } - self._project_id_cache: Dict[str, str] = {} # Cache mapping project names to IDs + self._project_id_cache: dict[str, str] = {} # Cache mapping project names to IDs self.global_braintrust_http_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) self.global_braintrust_sync_http_handler = HTTPHandler() - def validate_environment(self, api_key: Optional[str]): + def validate_environment(self, api_key: str | None): """ Expects BRAINTRUST_API_KEY @@ -64,7 +63,7 @@ class BraintrustLogger(CustomLogger): missing_keys.append("BRAINTRUST_API_KEY") if len(missing_keys) > 0: - raise Exception("Missing keys={} in environment.".format(missing_keys)) + raise Exception(f"Missing keys={missing_keys} in environment.") def get_project_id_sync(self, project_name: str) -> str: """ @@ -180,7 +179,7 @@ class BraintrustLogger(CustomLogger): cost = kwargs.get("response_cost", None) - metrics: Optional[dict] = None + metrics: dict | None = None usage_obj = getattr(response_obj, "usage", None) if usage_obj and isinstance(usage_obj, litellm.Usage): litellm.utils.get_logging_id(start_time, response_obj) @@ -305,7 +304,7 @@ class BraintrustLogger(CustomLogger): cost = kwargs.get("response_cost", None) - metrics: Optional[dict] = None + metrics: dict | None = None usage_obj = getattr(response_obj, "usage", None) if usage_obj and isinstance(usage_obj, litellm.Usage): litellm.utils.get_logging_id(start_time, response_obj) diff --git a/litellm/integrations/cloudzero/cloudzero.py b/litellm/integrations/cloudzero/cloudzero.py index 121b1dc6967..e6faf4a6a62 100644 --- a/litellm/integrations/cloudzero/cloudzero.py +++ b/litellm/integrations/cloudzero/cloudzero.py @@ -1,6 +1,6 @@ import os from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Optional, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_logger @@ -25,9 +25,9 @@ class CloudZeroLogger(CustomLogger): def __init__( self, - api_key: Optional[str] = None, - connection_id: Optional[str] = None, - timezone: Optional[str] = None, + api_key: str | None = None, + connection_id: str | None = None, + timezone: str | None = None, **kwargs, ): """Initialize CloudZero logger with configuration from parameters or environment variables.""" @@ -92,10 +92,10 @@ class CloudZeroLogger(CustomLogger): async def export_usage_data( self, - limit: Optional[int] = None, + limit: int | None = None, operation: str = "replace_hourly", - start_time_utc: Optional[datetime] = None, - end_time_utc: Optional[datetime] = None, + start_time_utc: datetime | None = None, + end_time_utc: datetime | None = None, ): """ Exports the usage data to CloudZero. @@ -153,10 +153,10 @@ class CloudZeroLogger(CustomLogger): verbose_logger.debug(f"CloudZero Logger: Successfully exported {len(cbf_data)} records to CloudZero") except Exception as e: - verbose_logger.error(f"CloudZero Logger: Error exporting usage data: {str(e)}") + verbose_logger.error(f"CloudZero Logger: Error exporting usage data: {e!s}") raise - async def dry_run_export_usage_data(self, limit: Optional[int] = 10000): + async def dry_run_export_usage_data(self, limit: int | None = 10000): """ Returns the data that would be exported to CloudZero without actually sending it. @@ -244,8 +244,8 @@ class CloudZeroLogger(CustomLogger): } except Exception as e: - verbose_logger.error(f"CloudZero Logger: Error in dry run export: {str(e)}") - verbose_logger.error(f"CloudZero Dry Run Error: {str(e)}") + verbose_logger.error(f"CloudZero Logger: Error in dry run export: {e!s}") + verbose_logger.error(f"CloudZero Dry Run Error: {e!s}") raise def _display_cbf_data_on_screen(self, cbf_data): @@ -346,7 +346,7 @@ class CloudZeroLogger(CustomLogger): from litellm.constants import CLOUDZERO_EXPORT_INTERVAL_MINUTES from litellm.integrations.custom_logger import CustomLogger - prometheus_loggers: List[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( + prometheus_loggers: list[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( callback_type=CloudZeroLogger ) # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them diff --git a/litellm/integrations/cloudzero/cz_resource_names.py b/litellm/integrations/cloudzero/cz_resource_names.py index 15cb66002f7..6147002bd7c 100644 --- a/litellm/integrations/cloudzero/cz_resource_names.py +++ b/litellm/integrations/cloudzero/cz_resource_names.py @@ -34,7 +34,6 @@ class CZRNGenerator: def __init__(self): """Initialize CZRN generator.""" - pass def create_from_litellm_data(self, row: dict[str, Any]) -> str: """Create a CZRN from LiteLLM daily spend data. diff --git a/litellm/integrations/cloudzero/cz_stream_api.py b/litellm/integrations/cloudzero/cz_stream_api.py index 47d6f7474a2..2a2507011e4 100644 --- a/litellm/integrations/cloudzero/cz_stream_api.py +++ b/litellm/integrations/cloudzero/cz_stream_api.py @@ -20,7 +20,7 @@ import zoneinfo from datetime import datetime, timezone -from typing import Any, Optional, Union +from typing import Any import httpx import polars as pl @@ -30,7 +30,7 @@ from rich.console import Console class CloudZeroStreamer: """Stream CBF data to CloudZero AnyCost API with proper batching and timezone handling.""" - def __init__(self, api_key: str, connection_id: str, user_timezone: Optional[str] = None): + def __init__(self, api_key: str, connection_id: str, user_timezone: str | None = None): """Initialize CloudZero streamer with credentials.""" self.api_key = api_key self.connection_id = connection_id @@ -38,7 +38,7 @@ class CloudZeroStreamer: self.console = Console() # Set timezone - default to UTC - self.user_timezone: Union[zoneinfo.ZoneInfo, timezone] + self.user_timezone: zoneinfo.ZoneInfo | timezone if user_timezone: try: self.user_timezone = zoneinfo.ZoneInfo(user_timezone) @@ -75,7 +75,7 @@ class CloudZeroStreamer: self.console.print("[red]Error: Missing 'time/usage_start' column for date grouping[/red]") return {} - timestamp_str: Optional[str] = None + timestamp_str: str | None = None for row in data.iter_rows(named=True): try: # Parse the timestamp and convert to UTC @@ -205,7 +205,7 @@ class CloudZeroStreamer: return payload - def _convert_cbf_to_api_format(self, row: dict[str, Any]) -> Optional[dict[str, Any]]: + def _convert_cbf_to_api_format(self, row: dict[str, Any]) -> dict[str, Any] | None: """Convert CBF row to CloudZero API format - keeping CBF field names as CloudZero expects them.""" try: # CloudZero expects CBF format field names directly, not converted names diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 71929398103..16fb99517ae 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -19,7 +19,7 @@ """Database connection and data extraction for LiteLLM.""" from datetime import datetime -from typing import Any, Optional, List +from typing import Any import polars as pl @@ -39,9 +39,9 @@ class LiteLLMDatabase: async def get_usage_data( self, - limit: Optional[int] = None, - start_time_utc: Optional[datetime] = None, - end_time_utc: Optional[datetime] = None, + limit: int | None = None, + start_time_utc: datetime | None = None, + end_time_utc: datetime | None = None, ) -> pl.DataFrame: """Retrieve usage data from LiteLLM daily user spend table.""" client = self._ensure_prisma_client() @@ -80,7 +80,7 @@ class LiteLLMDatabase: ORDER BY dus.date DESC, dus.created_at DESC """ - params: List[Any] = [ + params: list[Any] = [ start_time_utc, end_time_utc, ] @@ -98,4 +98,4 @@ class LiteLLMDatabase: # This prevents schema mismatch errors when data types vary across rows return pl.DataFrame(db_response, infer_schema_length=None) except Exception as e: - raise Exception(f"Error retrieving usage data: {str(e)}") + raise Exception(f"Error retrieving usage data: {e!s}") diff --git a/litellm/integrations/cloudzero/transform.py b/litellm/integrations/cloudzero/transform.py index c72001aee1a..3acfd3d8451 100644 --- a/litellm/integrations/cloudzero/transform.py +++ b/litellm/integrations/cloudzero/transform.py @@ -19,7 +19,7 @@ """Transform LiteLLM data to CloudZero AnyCost CBF format.""" from datetime import datetime -from typing import Any, Optional +from typing import Any import polars as pl @@ -187,7 +187,7 @@ class CBFTransformer: return CBFRecord(cbf_record) - def _parse_date(self, date_str) -> Optional[datetime]: + def _parse_date(self, date_str) -> datetime | None: """Parse date string from daily spend tables (e.g., '2025-04-19').""" if date_str is None: return None diff --git a/litellm/integrations/code_interpreter_interception/handler.py b/litellm/integrations/code_interpreter_interception/handler.py index 759b2be3a84..db34f00b051 100644 --- a/litellm/integrations/code_interpreter_interception/handler.py +++ b/litellm/integrations/code_interpreter_interception/handler.py @@ -11,19 +11,19 @@ import time import uuid from typing import Any, Literal, TypedDict, cast -import litellm from pydantic import ValidationError +import litellm from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger from litellm.types.integrations.code_interpreter_interception import ( CodeInterpreterInterceptionConfig, ) from litellm.types.integrations.custom_logger import ( - AgenticLoopPlan, - AgenticLoopRequestPatch, CHAT_COMPLETION_AGENTIC_SURFACE, NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + AgenticLoopPlan, + AgenticLoopRequestPatch, is_interception_internal_key, ) from litellm.types.llms.openai import ( diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index f0a696aa1e1..93765001c94 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -7,7 +7,7 @@ litellm_content_retrieve tool calls server-side via the typed agentic loop plan. import time import uuid -from typing import Any, Dict, List, Optional, Tuple, cast +from typing import Any, cast from litellm._logging import verbose_logger from litellm.compression import compress @@ -76,9 +76,9 @@ class CompressionInterceptionLogger(CustomLogger): self, enabled: bool = True, compression_trigger: int = 200_000, - compression_target: Optional[int] = None, - embedding_model: Optional[str] = None, - embedding_model_params: Optional[Dict[str, Any]] = None, + compression_target: int | None = None, + embedding_model: str | None = None, + embedding_model_params: dict[str, Any] | None = None, ): super().__init__() self.enabled = enabled @@ -86,7 +86,7 @@ class CompressionInterceptionLogger(CustomLogger): self.compression_target = compression_target self.embedding_model = embedding_model self.embedding_model_params = embedding_model_params - self._compression_cache_by_call_id: Dict[str, Tuple[Dict[str, str], float]] = {} + self._compression_cache_by_call_id: dict[str, tuple[dict[str, str], float]] = {} @classmethod def from_config_yaml(cls, config: CompressionInterceptionConfig) -> "CompressionInterceptionLogger": @@ -100,8 +100,8 @@ class CompressionInterceptionLogger(CustomLogger): @staticmethod def initialize_from_proxy_config( - litellm_settings: Dict[str, Any], - callback_specific_params: Dict[str, Any], + litellm_settings: dict[str, Any], + callback_specific_params: dict[str, Any], ) -> "CompressionInterceptionLogger": compression_params: CompressionInterceptionConfig = {} if "compression_interception_params" in litellm_settings: @@ -115,9 +115,7 @@ class CompressionInterceptionLogger(CustomLogger): ) return CompressionInterceptionLogger.from_config_yaml(compression_params) - async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: if not self.enabled: return None if call_type is not None and call_type != CallTypes.anthropic_messages: @@ -145,9 +143,9 @@ class CompressionInterceptionLogger(CustomLogger): embedding_model_params=self.embedding_model_params, ) - cache = cast(Dict[str, str], compressed.get("cache", {})) - skip_reason = cast(Optional[str], compressed.get("compression_skipped_reason")) - compressed_tools = cast(List[Dict[str, Any]], compressed.get("tools", [])) + cache = cast(dict[str, str], compressed.get("cache", {})) + skip_reason = cast(str | None, compressed.get("compression_skipped_reason")) + compressed_tools = cast(list[dict[str, Any]], compressed.get("tools", [])) # Only mutate kwargs when compression actually produced a result. # If compression was a no-op (below trigger, invalid tool sequence, etc.), @@ -158,10 +156,10 @@ class CompressionInterceptionLogger(CustomLogger): kwargs["messages"] = compressed["messages"] if compressed_tools: kwargs["tools"] = self._merge_tools( - existing_tools=cast(Optional[List[Dict[str, Any]]], kwargs.get("tools")), + existing_tools=cast(list[dict[str, Any]] | None, kwargs.get("tools")), compressed_tools=compressed_tools, ) - call_id = cast(Optional[str], kwargs.get("litellm_call_id")) + call_id = cast(str | None, kwargs.get("litellm_call_id")) if not call_id: call_id = str(uuid.uuid4()) kwargs["litellm_call_id"] = call_id @@ -193,12 +191,12 @@ class CompressionInterceptionLogger(CustomLogger): self, response: Any, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Tuple[bool, Dict]: + kwargs: dict, + ) -> tuple[bool, dict]: if not self.enabled: return False, {} if not self._has_retrieval_tool(tools): @@ -216,19 +214,19 @@ class CompressionInterceptionLogger(CustomLogger): async def async_build_agentic_loop_plan( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: Any, anthropic_messages_provider_config: Any, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: Any, stream: bool, - kwargs: Dict, + kwargs: dict, ) -> AgenticLoopPlan: self._prune_expired_cache() - tool_calls = cast(List[Dict[str, Any]], tools.get("tool_calls", [])) - thinking_blocks = cast(List[Dict[str, Any]], tools.get("thinking_blocks", [])) + tool_calls = cast(list[dict[str, Any]], tools.get("tool_calls", [])) + thinking_blocks = cast(list[dict[str, Any]], tools.get("thinking_blocks", [])) call_id = self._resolve_call_id(logging_obj=logging_obj, kwargs=kwargs) cache = self._get_cache(call_id=call_id) @@ -261,7 +259,7 @@ class CompressionInterceptionLogger(CustomLogger): follow_up_messages = messages + [assistant_message, user_message] max_tokens = cast( - Optional[int], + int | None, anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get("max_tokens"), ) optional_params_without_max_tokens = { @@ -298,7 +296,7 @@ class CompressionInterceptionLogger(CustomLogger): if now - created_at <= _CACHE_TTL_SECONDS } - def _get_cache(self, call_id: Optional[str]) -> Dict[str, str]: + def _get_cache(self, call_id: str | None) -> dict[str, str]: if not call_id: return {} cache_entry = self._compression_cache_by_call_id.get(call_id) @@ -306,15 +304,15 @@ class CompressionInterceptionLogger(CustomLogger): return {} return cache_entry[0] - def _resolve_call_id(self, logging_obj: Any, kwargs: Dict[str, Any]) -> Optional[str]: + def _resolve_call_id(self, logging_obj: Any, kwargs: dict[str, Any]) -> str | None: if logging_obj is not None: logging_call_id = getattr(logging_obj, "litellm_call_id", None) if isinstance(logging_call_id, str) and logging_call_id: return logging_call_id kwargs_call_id = kwargs.get("litellm_call_id") - return cast(Optional[str], kwargs_call_id if isinstance(kwargs_call_id, str) else None) + return cast(str | None, kwargs_call_id if isinstance(kwargs_call_id, str) else None) - def _resolve_retrieval_content(self, tool_call: Dict[str, Any], cache: Dict[str, str]) -> str: + def _resolve_retrieval_content(self, tool_call: dict[str, Any], cache: dict[str, str]) -> str: raw_input = tool_call.get("input", {}) key = "" if isinstance(raw_input, dict): @@ -325,7 +323,7 @@ class CompressionInterceptionLogger(CustomLogger): return cache[key] return f"[compressed content key '{key}' not found]" - def _extract_retrieval_tool_calls(self, response: Any) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: + def _extract_retrieval_tool_calls(self, response: Any) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: if isinstance(response, dict): content = response.get("content", []) else: @@ -334,8 +332,8 @@ class CompressionInterceptionLogger(CustomLogger): if not isinstance(content, list): return [], [] - tool_calls: List[Dict[str, Any]] = [] - thinking_blocks: List[Dict[str, Any]] = [] + tool_calls: list[dict[str, Any]] = [] + thinking_blocks: list[dict[str, Any]] = [] for block in content: if isinstance(block, dict): @@ -382,7 +380,7 @@ class CompressionInterceptionLogger(CustomLogger): return tool_calls, thinking_blocks - def _prepare_followup_kwargs(self, kwargs: Dict[str, Any]) -> Dict[str, Any]: + def _prepare_followup_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]: internal_keys = {"litellm_logging_obj"} return { k: v for k, v in kwargs.items() if not k.startswith("_compression_interception") and k not in internal_keys @@ -404,9 +402,9 @@ class CompressionInterceptionLogger(CustomLogger): def _merge_tools( self, - existing_tools: Optional[List[Dict[str, Any]]], - compressed_tools: List[Dict[str, Any]], - ) -> List[Dict[str, Any]]: + existing_tools: list[dict[str, Any]] | None, + compressed_tools: list[dict[str, Any]], + ) -> list[dict[str, Any]]: merged = list(existing_tools or []) if self._has_retrieval_tool(merged): return merged diff --git a/litellm/integrations/custom_batch_logger.py b/litellm/integrations/custom_batch_logger.py index aded12fa399..98a8e4ba739 100644 --- a/litellm/integrations/custom_batch_logger.py +++ b/litellm/integrations/custom_batch_logger.py @@ -6,7 +6,6 @@ Use this if you want your logs to be stored in memory and flushed periodically. import asyncio import time -from typing import List, Optional import litellm from litellm._logging import verbose_logger @@ -25,10 +24,10 @@ class CustomBatchLogger(CustomLogger): def __init__( self, - flush_lock: Optional[asyncio.Lock] = None, - batch_size: Optional[int] = None, - flush_interval: Optional[int] = None, - max_queue_size: Optional[int] = None, + flush_lock: asyncio.Lock | None = None, + batch_size: int | None = None, + flush_interval: int | None = None, + max_queue_size: int | None = None, **kwargs, ) -> None: """ @@ -36,7 +35,7 @@ class CustomBatchLogger(CustomLogger): flush_lock (Optional[asyncio.Lock], optional): Lock to use when flushing the queue. Defaults to None. Only used for custom loggers that do batching max_queue_size (Optional[int], optional): Maximum number of events to retain in ``log_queue``. When the limit is exceeded (e.g. because the send destination is unreachable and events are preserved for retry), the oldest events are dropped. Defaults to ``DEFAULT_MAX_QUEUE_SIZE``. """ - self.log_queue: List = [] + self.log_queue: list = [] self.flush_interval = flush_interval or litellm.DEFAULT_FLUSH_INTERVAL_SECONDS self.batch_size: int = batch_size or litellm.DEFAULT_BATCH_SIZE self.last_flush_time = time.time() diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index fbe70c03d74..743c539c36f 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -7,23 +7,19 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, - Dict, - List, Literal, Optional, - Type, - Union, get_args, ) from litellm._logging import verbose_logger +from litellm.caching import DualCache +from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, get_or_create_metadata_bucket, redact_nested_match_and_regex_keys, ) -from litellm.caching import DualCache -from litellm.integrations.custom_logger import CustomLogger from litellm.secret_managers.main import str_to_bool from litellm.types.guardrails import ( DynamicGuardrailParams, @@ -90,7 +86,7 @@ def _strict_guardrail_modes_enabled() -> bool: return True if parsed is None else parsed -def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]: +def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None: """Extract session_id from request data (litellm_session_id or metadata).""" session_id = request_data.get("litellm_session_id") if session_id: @@ -117,18 +113,18 @@ class CustomGuardrail(CustomLogger): def __init__( self, - guardrail_name: Optional[str] = None, - supported_event_hooks: Optional[List[GuardrailEventHooks]] = None, - event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None, + guardrail_name: str | None = None, + supported_event_hooks: list[GuardrailEventHooks] | None = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, mask_request_content: bool = False, mask_response_content: bool = False, - violation_message_template: Optional[str] = None, - end_session_after_n_fails: Optional[int] = None, - on_violation: Optional[str] = None, - realtime_violation_message: Optional[str] = None, - on_sensitive_data: Optional[str] = None, - sensitive_data_route_to_model: Optional[str] = None, + violation_message_template: str | None = None, + end_session_after_n_fails: int | None = None, + on_violation: str | None = None, + realtime_violation_message: str | None = None, + on_sensitive_data: str | None = None, + sensitive_data_route_to_model: str | None = None, sticky_session_routing: bool = True, run_in_parallel: bool = False, only_scan_new_messages: bool = False, @@ -156,16 +152,16 @@ class CustomGuardrail(CustomLogger): """ self.guardrail_name = guardrail_name self.supported_event_hooks = supported_event_hooks - self.event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = event_hook + self.event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = event_hook self.default_on: bool = default_on self.mask_request_content: bool = mask_request_content self.mask_response_content: bool = mask_response_content - self.violation_message_template: Optional[str] = violation_message_template - self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails - self.on_violation: Optional[str] = on_violation - self.realtime_violation_message: Optional[str] = realtime_violation_message - self.on_sensitive_data: Optional[str] = on_sensitive_data - self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model + self.violation_message_template: str | None = violation_message_template + self.end_session_after_n_fails: int | None = end_session_after_n_fails + self.on_violation: str | None = on_violation + self.realtime_violation_message: str | None = realtime_violation_message + self.on_sensitive_data: str | None = on_sensitive_data + self.sensitive_data_route_to_model: str | None = sensitive_data_route_to_model self.sticky_session_routing: bool = sticky_session_routing self.run_in_parallel: bool = run_in_parallel self.only_scan_new_messages: bool = only_scan_new_messages @@ -185,13 +181,13 @@ class CustomGuardrail(CustomLogger): ) super().__init__(**kwargs) - def render_violation_message(self, default: str, context: Optional[Dict[str, Any]] = None) -> str: + def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str: """Return a custom violation message if template is configured.""" if not self.violation_message_template: return default - format_context: Dict[str, Any] = {"default_message": default} + format_context: dict[str, Any] = {"default_message": default} if context: format_context.update(context) try: @@ -207,8 +203,8 @@ class CustomGuardrail(CustomLogger): def raise_passthrough_exception( self, violation_message: str, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Raise a passthrough exception for guardrail violations. @@ -251,8 +247,8 @@ class CustomGuardrail(CustomLogger): def raise_sensitive_data_route_exception( self, route_to_model: str, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Raise an exception to reroute the request to a different model. @@ -289,7 +285,7 @@ class CustomGuardrail(CustomLogger): sticky_session_routing=self.sticky_session_routing, ) - def _get_session_id_from_request_data(self, request_data: Dict[str, Any]) -> Optional[str]: + def _get_session_id_from_request_data(self, request_data: dict[str, Any]) -> str | None: """Extract session_id from request data.""" return get_session_id_from_request_data(request_data) @@ -396,8 +392,8 @@ class CustomGuardrail(CustomLogger): def handle_sensitive_data_detection( self, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Handle sensitive data detection based on guardrail configuration. @@ -439,7 +435,7 @@ class CustomGuardrail(CustomLogger): ) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Returns the config model for the guardrail @@ -448,7 +444,7 @@ class CustomGuardrail(CustomLogger): return None @classmethod - def get_supported_event_hooks(cls) -> Optional[List[GuardrailEventHooks]]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks] | None: """ Returns the event hooks this guardrail supports, for the UI to render. @@ -461,12 +457,12 @@ class CustomGuardrail(CustomLogger): def _validate_event_hook( self, - event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]], - supported_event_hooks: List[GuardrailEventHooks], + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None, + supported_event_hooks: list[GuardrailEventHooks], ) -> None: def _validate_event_hook_list_is_in_supported_event_hooks( - event_hook: Union[List[GuardrailEventHooks], List[str]], - supported_event_hooks: List[GuardrailEventHooks], + event_hook: list[GuardrailEventHooks] | list[str], + supported_event_hooks: list[GuardrailEventHooks], ) -> None: for hook in event_hook: if isinstance(hook, str): @@ -517,7 +513,7 @@ class CustomGuardrail(CustomLogger): key_meta = meta.get("user_api_key_metadata") or key_meta return {**team_meta, **key_meta} - def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: + def get_disable_global_guardrail(self, data: dict) -> bool | None: """ Returns True if the global guardrail should be disabled. @@ -526,7 +522,7 @@ class CustomGuardrail(CustomLogger): """ return self._get_admin_metadata(data).get("disable_global_guardrails", False) - def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]: + def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> list[str]: """ Returns the list of global guardrail names the team/key has opted out of. @@ -557,7 +553,7 @@ class CustomGuardrail(CustomLogger): return True raise - def get_guardrail_from_metadata(self, data: dict) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]: + def get_guardrail_from_metadata(self, data: dict) -> list[str] | list[dict[str, DynamicGuardrailParams]]: """ Returns the guardrail(s) to be run from the metadata or root """ @@ -578,7 +574,7 @@ class CustomGuardrail(CustomLogger): def _guardrail_is_in_requested_guardrails( self, - requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]], + requested_guardrails: list[str] | list[dict[str, DynamicGuardrailParams]], ) -> bool: for _guardrail in requested_guardrails: if isinstance(_guardrail, dict): @@ -590,13 +586,13 @@ class CustomGuardrail(CustomLogger): return False - def _pre_call_marker(self) -> Optional[str]: + def _pre_call_marker(self) -> str | None: name = self.guardrail_name if not name: return None return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}" - def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None: + def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None: """ Record that this guardrail's ``async_pre_call_hook`` already ran for this request, so the deployment-level hook does not run it a second time. @@ -621,7 +617,7 @@ class CustomGuardrail(CustomLogger): return data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]} - def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool: + def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool: marker = self._pre_call_marker() if marker is None: return False @@ -649,9 +645,7 @@ class CustomGuardrail(CustomLogger): ) from e return unified_guardrail - async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: from litellm.proxy._types import UserAPIKeyAuth # should run guardrail @@ -694,8 +688,8 @@ class CustomGuardrail(CustomLogger): self, request_data: dict, response: LLMResponseTypes, - call_type: Optional[CallTypes], - ) -> Optional[LLMResponseTypes]: + call_type: CallTypes | None, + ) -> LLMResponseTypes | None: """ Allow modifying / reviewing the response just after it's received from the deployment. """ @@ -877,16 +871,16 @@ class CustomGuardrail(CustomLogger): def add_standard_logging_guardrail_information_to_request_data( self, - guardrail_json_response: Union[Exception, str, dict, List[dict]], + guardrail_json_response: Exception | str | dict | list[dict], request_data: dict, guardrail_status: GuardrailStatus, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - masked_entity_count: Optional[Dict[str, int]] = None, - guardrail_provider: Optional[str] = None, - event_type: Optional[GuardrailEventHooks] = None, - tracing_detail: Optional[GuardrailTracingDetail] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + masked_entity_count: dict[str, int] | None = None, + guardrail_provider: str | None = None, + event_type: GuardrailEventHooks | None = None, + tracing_detail: GuardrailTracingDetail | None = None, ) -> None: """ Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc. @@ -901,7 +895,7 @@ class CustomGuardrail(CustomLogger): from litellm.types.utils import GuardrailMode # Use event_type if provided, otherwise fall back to self.event_hook - guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]] + guardrail_mode: GuardrailEventHooks | GuardrailMode | list[GuardrailEventHooks] if event_type is not None: guardrail_mode = event_type elif isinstance(self.event_hook, Mode): @@ -1009,13 +1003,13 @@ class CustomGuardrail(CustomLogger): def _process_response( self, - response: Optional[Dict], + response: dict | None, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, - original_inputs: Optional[Dict] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, + original_inputs: dict | None = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -1023,7 +1017,7 @@ class CustomGuardrail(CustomLogger): This gets logged on downsteam Langfuse, DataDog, etc. """ # Convert None to empty dict to satisfy type requirements - guardrail_response: Union[Dict[str, Any], str] = {} if response is None else response + guardrail_response: dict[str, Any] | str = {} if response is None else response # For apply_guardrail functions in custom_code_guardrail scenario, # simplify the logged response to "allow", "deny", or "mask" @@ -1090,10 +1084,10 @@ class CustomGuardrail(CustomLogger): self, e: Exception, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -1105,7 +1099,7 @@ class CustomGuardrail(CustomLogger): ) # For custom_code_guardrail scenario, log as "deny" instead of full exception # Check if this is from custom_code_guardrail by checking the class name - guardrail_response: Union[Exception, str] = e + guardrail_response: Exception | str = e if "CustomCodeGuardrail" in self.__class__.__name__: guardrail_response = "deny" @@ -1120,7 +1114,7 @@ class CustomGuardrail(CustomLogger): ) raise e - def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool: + def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool: """ Compare original inputs with response to determine if content was modified. @@ -1165,8 +1159,8 @@ class CustomGuardrail(CustomLogger): setattr(self, key, value) def get_guardrails_messages_for_call_type( - self, call_type: CallTypes, data: Optional[dict] = None - ) -> Optional[List[AllMessageValues]]: + self, call_type: CallTypes, data: dict | None = None + ) -> list[AllMessageValues] | None: """ Returns the messages for the given call type and data """ @@ -1205,7 +1199,7 @@ class CustomGuardrail(CustomLogger): input=input_data, responses_api_request=data, ) - return cast(List[AllMessageValues], messages) + return cast(list[AllMessageValues], messages) return None @@ -1281,7 +1275,7 @@ def log_guardrail_information(func): def _infer_event_type_from_function_name( func_name: str, - ) -> Optional[GuardrailEventHooks]: + ) -> GuardrailEventHooks | None: """Infer the actual event type from the function name""" if func_name == "async_pre_call_hook": return GuardrailEventHooks.pre_call diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index d2f0acd15d4..971d53ffec4 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -6,10 +6,7 @@ from collections.abc import AsyncGenerator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Optional, - Tuple, Union, ) @@ -82,10 +79,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac """ self.message_logging = message_logging self.turn_off_message_logging = turn_off_message_logging - pass @staticmethod - def get_callback_env_vars(callback_name: Optional[str] = None) -> List[str]: + def get_callback_env_vars(callback_name: str | None = None) -> list[str]: """ Return the environment variables associated with a given callback name as defined in the proxy callback registry. @@ -145,7 +141,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_log_pre_api_call(self, model, messages, kwargs): pass - async def async_pre_request_hook(self, model: str, messages: List, kwargs: Dict) -> Optional[Dict]: + async def async_pre_request_hook(self, model: str, messages: list, kwargs: dict) -> dict | None: """ Hook called before making the API request to allow modifying request parameters. @@ -169,7 +165,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac return kwargs ``` """ - pass async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): pass @@ -179,26 +174,25 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_log_audit_log_event(self, audit_log: "StandardAuditLogPayload"): """Called when an audit log is created. Override in subclasses to handle.""" - pass #### PROMPT MANAGEMENT HOOKS #### async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Returns: - model: str - the model to use (can be pulled from prompt management tool) @@ -210,17 +204,17 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Returns: - model: str - the model to use (can be pulled from prompt management tool) @@ -237,11 +231,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - ) -> Optional[PreRoutingHookResponse]: + request_kwargs: dict, + messages: list[dict[str, Any]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + ) -> PreRoutingHookResponse | None: """ This hook is called before the routing decision is made. @@ -252,16 +246,14 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, + ) -> list[dict]: return healthy_deployments - async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: """ Allow modifying the request just before it's sent to the deployment. @@ -269,41 +261,38 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac Used in managed_files.py """ + + async def async_pre_call_check(self, deployment: dict, parent_otel_span: Span | None) -> dict | None: pass - async def async_pre_call_check(self, deployment: dict, parent_otel_span: Optional[Span]) -> Optional[dict]: - pass - - def pre_call_check(self, deployment: dict) -> Optional[dict]: + def pre_call_check(self, deployment: dict) -> dict | None: pass async def async_post_call_success_deployment_hook( self, request_data: dict, response: LLMResponseTypes, - call_type: Optional[CallTypes], - ) -> Optional[LLMResponseTypes]: + call_type: CallTypes | None, + ) -> LLMResponseTypes | None: """ Allow modifying / reviewing the response just after it's received from the deployment. """ - pass async def async_post_call_streaming_deployment_hook( self, request_data: dict, response_chunk: Any, - call_type: Optional[CallTypes], - ) -> Optional[Any]: + call_type: CallTypes | None, + ) -> Any | None: """ Allow modifying streaming chunks just before they're returned to the user. This is called for each streaming chunk in the response. """ - pass #### Fallback Events - router/proxy only #### async def log_model_group_rate_limit_error( - self, exception: Exception, original_model_group: Optional[str], kwargs: dict + self, exception: Exception, original_model_group: str | None, kwargs: dict ): pass @@ -315,33 +304,30 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac #### ADAPTERS #### Allow calling 100+ LLMs in custom format - https://github.com/BerriAI/litellm/pulls - def translate_completion_input_params(self, kwargs) -> Optional[ChatCompletionRequest]: + def translate_completion_input_params(self, kwargs) -> ChatCompletionRequest | None: """ Translates the input params, from the provider's native format to the litellm.completion() format. """ - pass - def translate_completion_output_params(self, response: ModelResponse) -> Optional[BaseModel]: + def translate_completion_output_params(self, response: ModelResponse) -> BaseModel | None: """ Translates the output params, from the OpenAI format to the custom format. """ - pass def translate_completion_output_params_streaming( self, completion_stream: Any - ) -> Optional[AdapterCompletionStreamWrapper]: + ) -> AdapterCompletionStreamWrapper | None: """ Translates the streaming chunk, from the OpenAI format to the custom format. """ - pass ### DATASET HOOKS #### - currently only used for Argilla async def async_dataset_hook( self, logged_item: ArgillaItem, - standard_logging_payload: Optional[StandardLoggingPayload], - ) -> Optional[ArgillaItem]: + standard_logging_payload: StandardLoggingPayload | None, + ) -> ArgillaItem | None: """ - Decide if the result should be logged to Argilla. - Modify the result before logging to Argilla. @@ -360,9 +346,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac cache: "DualCache", data: dict, call_type: CallTypesLiteral, - ) -> Optional[ - Union[Exception, str, dict] - ]: # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm + ) -> ( + Exception | str | dict | None + ): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm pass async def async_post_call_response_headers_hook( @@ -370,9 +356,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any, - request_headers: Optional[Dict[str, str]] = None, - litellm_call_info: Optional[Dict[str, Any]] = None, - ) -> Optional[Dict[str, str]]: + request_headers: dict[str, str] | None = None, + litellm_call_info: dict[str, Any] | None = None, + ) -> dict[str, str] | None: """ Called after an LLM API call (success or failure) to allow injecting custom HTTP response headers. @@ -398,7 +384,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, + traceback_str: str | None = None, ) -> Optional["HTTPException"]: """ Called after an LLM API call fails. Can return or raise HTTPException to transform error responses. @@ -413,7 +399,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac - Optional[HTTPException]: Return an HTTPException to transform the error response sent to the client. Return None to use the original exception. """ - pass async def async_post_call_success_hook( self, @@ -423,11 +408,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac ) -> Any: pass - async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """For masking logged request/response. Return a modified version of the request/result.""" return kwargs, result - def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """For masking logged request/response. Return a modified version of the request/result.""" return kwargs, result @@ -493,7 +478,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac ) except Exception: print_verbose(f"Custom Logger Error - {traceback.format_exc()}") - pass async def async_log_event(self, kwargs, response_obj, start_time, end_time, print_verbose, callback_func): # Method definition @@ -507,7 +491,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac ) except Exception: print_verbose(f"Custom Logger Error - {traceback.format_exc()}") - pass ######################################################### # MCP TOOL CALL HOOKS @@ -515,7 +498,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_post_mcp_tool_call_hook( self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time - ) -> Optional[MCPPostCallResponseObject]: + ) -> MCPPostCallResponseObject | None: """ This log gets called after the MCP tool call is made. @@ -537,12 +520,12 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac self, response: Any, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Tuple[bool, Dict]: + kwargs: dict, + ) -> tuple[bool, dict]: """ Hook to determine if agentic loop should be executed. @@ -593,15 +576,15 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_run_agentic_loop( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: Any, anthropic_messages_provider_config: Any, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> Any: """ Hook to execute agentic loop based on context from should_run hook. @@ -659,19 +642,18 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac return final_response """ - pass async def async_build_agentic_loop_plan( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: Any, anthropic_messages_provider_config: Any, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> AgenticLoopPlan: """ Build a typed rerun plan for Anthropic Messages agentic loops. @@ -685,7 +667,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac self, response: Any, plan: AgenticLoopPlan, - kwargs: Dict, + kwargs: dict, ) -> Any: """ Post-process the response returned by the agentic-loop follow-up call. @@ -718,18 +700,18 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac Default does nothing. """ - return None + return async def async_should_run_chat_completion_agentic_loop( self, response: Any, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Tuple[bool, Dict]: + kwargs: dict, + ) -> tuple[bool, dict]: """ Hook to determine if chat completion agentic loop should be executed. """ @@ -737,30 +719,29 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_run_chat_completion_agentic_loop( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: Any, - optional_params: Dict, + optional_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> Any: """ Hook to execute chat completion agentic loop based on context from should_run hook. """ - pass async def async_build_chat_completion_agentic_loop_plan( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: Any, - optional_params: Dict, + optional_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> AgenticLoopPlan: """ Build a typed rerun plan for chat-completions agentic loops. @@ -823,7 +804,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac else text ) - def _select_metadata_field(self, request_kwargs: Optional[Dict] = None) -> Optional[str]: + def _select_metadata_field(self, request_kwargs: dict | None = None) -> str | None: """ Select the metadata field to use for logging @@ -838,7 +819,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac return LITELLM_METADATA_FIELD return OLD_LITELLM_METADATA_FIELD - def redact_standard_logging_payload_from_model_call_details(self, model_call_details: Dict) -> Dict: + def redact_standard_logging_payload_from_model_call_details(self, model_call_details: dict) -> dict: """ Redacts or excludes fields from StandardLoggingPayload before callbacks receive it. @@ -856,7 +837,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac from litellm import Choices, Message, ModelResponse turn_off_message_logging: bool = getattr(self, "turn_off_message_logging", False) - excluded_fields: Optional[List[str]] = getattr(litellm, "standard_logging_payload_excluded_fields", None) + excluded_fields: list[str] | None = getattr(litellm, "standard_logging_payload_excluded_fields", None) # Early return if no processing needed if turn_off_message_logging is False and not excluded_fields: @@ -915,11 +896,10 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def get_proxy_server_request_from_cold_storage_with_object_key( self, object_key: str, - ) -> Optional[dict]: + ) -> dict | None: """ Get the proxy server request from cold storage using the object key directly. """ - pass def handle_callback_failure(self, callback_name: str): """ @@ -947,7 +927,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac except Exception as e: from litellm._logging import verbose_logger - verbose_logger.debug(f"Error in handle_callback_failure for {callback_name}: {str(e)}") + verbose_logger.debug(f"Error in handle_callback_failure for {callback_name}: {e!s}") async def _strip_base64_from_messages( self, @@ -965,7 +945,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac • Recursively redact inline base64 blobs in *any* string field, at any depth. """ raw_messages: Any = payload.get("messages", []) - messages: List[Any] = raw_messages if isinstance(raw_messages, list) else [] + messages: list[Any] = raw_messages if isinstance(raw_messages, list) else [] verbose_logger.debug(f"[CustomLogger] Stripping base64 from {len(messages)} messages") if messages: @@ -997,7 +977,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac • Recursively redact inline base64 blobs in *any* string field, at any depth. """ raw_messages: Any = payload.get("messages", []) - messages: List[Any] = raw_messages if isinstance(raw_messages, list) else [] + messages: list[Any] = raw_messages if isinstance(raw_messages, list) else [] verbose_logger.debug(f"[CustomLogger] Stripping base64 from {len(messages)} messages") if messages: @@ -1049,16 +1029,16 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac def _process_messages( self, - messages: List[Any], + messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER, - ) -> List[Dict[str, Any]]: - filtered_messages: List[Dict[str, Any]] = [] + ) -> list[dict[str, Any]]: + filtered_messages: list[dict[str, Any]] = [] for msg in messages: if not isinstance(msg, dict): continue contents: Any = msg.get("content") if isinstance(contents, list): - cleaned: List[Any] = [] + cleaned: list[Any] = [] for c in contents: if self._should_keep_content(content=c): cleaned.append(self._redact_base64(value=c, max_depth=max_depth)) diff --git a/litellm/integrations/custom_prompt_management.py b/litellm/integrations/custom_prompt_management.py index fbca1867793..7078416b7a8 100644 --- a/litellm/integrations/custom_prompt_management.py +++ b/litellm/integrations/custom_prompt_management.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Tuple - from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prompt_management_base import ( PromptManagementBase, @@ -13,8 +11,8 @@ from litellm.types.utils import StandardCallbackDynamicParams class CustomPromptManagement(CustomLogger, PromptManagementBase): def __init__( self, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, **kwargs, ): self.ignore_prompt_manager_model = ignore_prompt_manager_model @@ -23,17 +21,17 @@ class CustomPromptManagement(CustomLogger, PromptManagementBase): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Returns: - model: str - the model to use (can be pulled from prompt management tool) @@ -48,30 +46,30 @@ class CustomPromptManagement(CustomLogger, PromptManagementBase): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: return True def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: raise NotImplementedError("Custom prompt management does not support compile prompt helper") async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: raise NotImplementedError("Custom prompt management does not support async compile prompt helper") diff --git a/litellm/integrations/custom_secret_manager.py b/litellm/integrations/custom_secret_manager.py index a1bb7b00d92..8cb7f02b798 100644 --- a/litellm/integrations/custom_secret_manager.py +++ b/litellm/integrations/custom_secret_manager.py @@ -37,7 +37,7 @@ Usage: """ from abc import abstractmethod -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -87,7 +87,7 @@ class CustomSecretManager(BaseSecretManager): def __init__( self, - secret_manager_name: Optional[str] = None, + secret_manager_name: str | None = None, **kwargs, ): """ @@ -106,9 +106,9 @@ class CustomSecretManager(BaseSecretManager): async def async_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Asynchronously read a secret from your custom secret manager. @@ -123,15 +123,14 @@ class CustomSecretManager(BaseSecretManager): Raises: Exception: If there's an error reading the secret """ - pass @abstractmethod def sync_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Synchronously read a secret from your custom secret manager. @@ -146,17 +145,16 @@ class CustomSecretManager(BaseSecretManager): Raises: Exception: If there's an error reading the secret """ - pass async def async_write_secret( self, secret_name: str, secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None, - ) -> Dict[str, Any]: + description: str | None = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + tags: dict | list | None = None, + ) -> dict[str, Any]: """ Asynchronously write a secret to your custom secret manager. @@ -185,9 +183,9 @@ class CustomSecretManager(BaseSecretManager): async def async_delete_secret( self, secret_name: str, - recovery_window_in_days: Optional[int] = 7, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + recovery_window_in_days: int | None = 7, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Asynchronously delete a secret from your custom secret manager. @@ -227,7 +225,7 @@ class CustomSecretManager(BaseSecretManager): verbose_logger.debug("No environment validation configured for custom secret manager") return True - async def async_health_check(self, timeout: Optional[Union[float, httpx.Timeout]] = None) -> bool: + async def async_health_check(self, timeout: float | httpx.Timeout | None = None) -> bool: """ Perform a health check on your secret manager. diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index 28d3f330fcd..047d69c9c9c 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -20,7 +20,7 @@ import time import traceback from collections.abc import Sequence from datetime import datetime as datetimeObj -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx from httpx import Response @@ -94,10 +94,10 @@ class DataDogLogger( # Class variables or attributes def __init__( self, - dd_api_key: Optional[str] = None, - dd_site: Optional[str] = None, - dd_agent_host: Optional[str] = None, - dd_agent_port: Optional[str] = None, + dd_api_key: str | None = None, + dd_site: str | None = None, + dd_agent_host: str | None = None, + dd_agent_port: str | None = None, allow_env_credentials: bool = True, **kwargs, ): @@ -171,20 +171,20 @@ class DataDogLogger( batch_size=_resolve_dd_batch_size(), ) except Exception as e: - verbose_logger.exception(f"Datadog: Got exception on init Datadog client {str(e)}") + verbose_logger.exception(f"Datadog: Got exception on init Datadog client {e!s}") raise e - def _get_datadog_params(self) -> Dict: + def _get_datadog_params(self) -> dict: """ Get the datadog_params from litellm.datadog_params These are params specific to initializing the DataDogLogger e.g. turn_off_message_logging """ - dict_datadog_params: Dict = {} + dict_datadog_params: dict = {} if litellm.datadog_params is not None: if isinstance(litellm.datadog_params, DatadogInitParams): dict_datadog_params = litellm.datadog_params.model_dump() - elif isinstance(litellm.datadog_params, Dict): + elif isinstance(litellm.datadog_params, dict): # only allow params that are of DatadogInitParams dict_datadog_params = DatadogInitParams(**litellm.datadog_params).model_dump() return dict_datadog_params @@ -192,8 +192,8 @@ class DataDogLogger( def _configure_dd_agent( self, dd_agent_host: str, - dd_agent_port: Optional[str] = None, - dd_api_key: Optional[str] = None, + dd_agent_port: str | None = None, + dd_api_key: str | None = None, allow_env_credentials: bool = True, ) -> None: """ @@ -214,8 +214,8 @@ class DataDogLogger( def _configure_dd_direct_api( self, - dd_api_key: Optional[str] = None, - dd_site: Optional[str] = None, + dd_api_key: str | None = None, + dd_site: str | None = None, allow_env_credentials: bool = True, ) -> None: """ @@ -257,8 +257,7 @@ class DataDogLogger( await self._log_async_event(kwargs, response_obj, start_time, end_time) except Exception as e: - verbose_logger.exception(f"Datadog Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Datadog Layer Error - {e!s}\n{traceback.format_exc()}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: @@ -266,16 +265,15 @@ class DataDogLogger( await self._log_async_event(kwargs, response_obj, start_time, end_time) except Exception as e: - verbose_logger.exception(f"Datadog Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Datadog Layer Error - {e!s}\n{traceback.format_exc()}") async def async_post_call_failure_hook( self, request_data: dict, original_exception: Exception, user_api_key_dict: Any, - traceback_str: Optional[str] = None, - ) -> Optional[Any]: + traceback_str: str | None = None, + ) -> Any | None: """ Log proxy-level failures (e.g. 401 auth, DB connection errors) to Datadog. @@ -294,12 +292,12 @@ class DataDogLogger( traceback_str=traceback_str, ) _code = error_information.get("error_code") or "" - status_code: Optional[int] = None + status_code: int | None = None if _code and str(_code).strip().isdigit(): status_code = int(_code) # Use project-standard sanitized user context when running in proxy - user_context: Dict[str, Any] = {} + user_context: dict[str, Any] = {} try: from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, @@ -342,7 +340,7 @@ class DataDogLogger( if len(self.log_queue) >= self.batch_size: await self.flush_queue() except Exception as e: - verbose_logger.exception(f"Datadog: async_post_call_failure_hook - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"Datadog: async_post_call_failure_hook - {e!s}\n{traceback.format_exc()}") return None async def async_send_batch(self): @@ -382,9 +380,9 @@ class DataDogLogger( except Exception as e: self.log_queue = batch_to_send + self.log_queue - verbose_logger.exception(f"Datadog Error sending batch API - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"Datadog Error sending batch API - {e!s}\n{traceback.format_exc()}") - async def _send_with_413_split(self, batch: List) -> List: + async def _send_with_413_split(self, batch: list) -> list: """ Send a batch, halving any sub-batch that exceeds Datadog's intake limits before sending, and halving again on a 413 (payload too large) response, since Datadog @@ -397,7 +395,7 @@ class DataDogLogger( that could not be delivered because of a non-413 (transient) error, so the caller re-queues only those and never the events already accepted by Datadog. """ - pending: List[List] = [batch] + pending: list[list] = [batch] while pending: chunk = pending.pop() if not chunk: @@ -413,7 +411,7 @@ class DataDogLogger( if isinstance(e, MaskedHTTPStatusError) and e.status_code == 413: response = e.response else: - verbose_logger.exception(f"Datadog Error sending batch API - {str(e)}") + verbose_logger.exception(f"Datadog Error sending batch API - {e!s}") return self._undelivered(chunk, pending) if response.status_code == 413: @@ -442,7 +440,7 @@ class DataDogLogger( return [] @staticmethod - def _undelivered(chunk: List, pending: List[List]) -> List: + def _undelivered(chunk: list, pending: list[list]) -> list: return chunk + [event for remaining in reversed(pending) for event in remaining] @staticmethod @@ -517,9 +515,7 @@ class DataDogLogger( ) except Exception as e: - verbose_logger.exception(f"Datadog Layer Error - {str(e)}\n{traceback.format_exc()}") - pass - pass + verbose_logger.exception(f"Datadog Layer Error - {e!s}\n{traceback.format_exc()}") async def _log_async_event(self, kwargs, response_obj, start_time, end_time): dd_payload = self.create_datadog_logging_payload( @@ -557,7 +553,7 @@ class DataDogLogger( def create_datadog_logging_payload( self, - kwargs: Union[dict, Any], + kwargs: dict | Any, response_obj: Any, start_time: datetime.datetime, end_time: datetime.datetime, @@ -575,7 +571,7 @@ class DataDogLogger( DatadogPayload: defined in types.py """ - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: raise ValueError("standard_logging_object not found in kwargs") @@ -592,7 +588,7 @@ class DataDogLogger( ) return dd_payload - async def async_send_compressed_data(self, data: List) -> Response: + async def async_send_compressed_data(self, data: list) -> Response: """ Async helper to send compressed data to datadog self.intake_url @@ -628,11 +624,11 @@ class DataDogLogger( async def async_service_failure_hook( self, payload: ServiceLoggerPayload, - error: Optional[str] = "", - parent_otel_span: Optional[Any] = None, - start_time: Optional[Union[datetimeObj, float]] = None, - end_time: Optional[Union[float, datetimeObj]] = None, - event_metadata: Optional[dict] = None, + error: str | None = "", + parent_otel_span: Any | None = None, + start_time: datetimeObj | float | None = None, + end_time: float | datetimeObj | None = None, + event_metadata: dict | None = None, ): """ Logs failures from Redis, Postgres (Adjacent systems), as 'WARNING' on DataDog @@ -658,16 +654,15 @@ class DataDogLogger( except Exception as e: verbose_logger.exception(f"Datadog: Logger - Exception in async_service_failure_hook: {e}") - pass async def async_service_success_hook( self, payload: ServiceLoggerPayload, - error: Optional[str] = "", - parent_otel_span: Optional[Any] = None, - start_time: Optional[Union[datetimeObj, float]] = None, - end_time: Optional[Union[float, datetimeObj]] = None, - event_metadata: Optional[dict] = None, + error: str | None = "", + parent_otel_span: Any | None = None, + start_time: datetimeObj | float | None = None, + end_time: float | datetimeObj | None = None, + event_metadata: dict | None = None, ): """ Logs success from Redis, Postgres (Adjacent systems), as 'INFO' on DataDog @@ -701,7 +696,7 @@ class DataDogLogger( def _create_v0_logging_payload( self, - kwargs: Union[dict, Any], + kwargs: dict | Any, response_obj: Any, start_time: datetime.datetime, end_time: datetime.datetime, @@ -800,7 +795,7 @@ class DataDogLogger( except Exception: verbose_logger.exception("Datadog: Failed to attach trace context to payload") - def _get_active_trace_context(self) -> Optional[Dict[str, str]]: + def _get_active_trace_context(self) -> dict[str, str] | None: try: current_span = None current_span_fn = getattr(tracer, "current_span", None) @@ -820,7 +815,7 @@ class DataDogLogger( return None span_id = getattr(current_span, "span_id", None) - trace_context: Dict[str, str] = {"trace_id": str(trace_id)} + trace_context: dict[str, str] = {"trace_id": str(trace_id)} if span_id is not None: trace_context["span_id"] = str(span_id) return trace_context @@ -863,7 +858,7 @@ class DataDogLogger( async def get_request_response_payload( self, request_id: str, - start_time_utc: Optional[datetimeObj], - end_time_utc: Optional[datetimeObj], - ) -> Optional[dict]: + start_time_utc: datetimeObj | None, + end_time_utc: datetimeObj | None, + ) -> dict | None: pass diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index 714a50eb2f2..7b22f4658f2 100644 --- a/litellm/integrations/datadog/datadog_cost_management.py +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -2,7 +2,7 @@ import asyncio import os import time from datetime import datetime -from typing import Any, Dict, List, Optional, Tuple, cast +from typing import Any, cast from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -44,8 +44,8 @@ _RESERVED_TAG_KEYS: frozenset = frozenset( class DatadogCostManagementLogger(CustomBatchLogger): - def __init__(self, cost_tag_keys: Optional[List[str]] = None, **kwargs): - self.cost_tag_keys: List[str] = list(cost_tag_keys) if cost_tag_keys else [] + def __init__(self, cost_tag_keys: list[str] | None = None, **kwargs): + self.cost_tag_keys: list[str] = list(cost_tag_keys) if cost_tag_keys else [] self.dd_api_key = os.getenv("DD_API_KEY") self.dd_app_key = os.getenv("DD_APP_KEY") self.dd_site = os.getenv("DD_SITE", "datadoghq.com") @@ -71,7 +71,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -84,7 +84,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"Datadog Cost Management: Error in async_log_success_event: {str(e)}") + verbose_logger.exception(f"Datadog Cost Management: Error in async_log_success_event: {e!s}") async def async_send_batch(self): if not self.log_queue: @@ -104,14 +104,14 @@ class DatadogCostManagementLogger(CustomBatchLogger): await self._upload_to_datadog(aggregated_entries) except Exception as e: self.log_queue = batch_to_send + self.log_queue - verbose_logger.exception(f"Datadog Cost Management: Error in async_send_batch: {str(e)}") + verbose_logger.exception(f"Datadog Cost Management: Error in async_send_batch: {e!s}") - def _aggregate_costs(self, logs: List[StandardLoggingPayload]) -> List[DatadogFOCUSCostEntry]: + def _aggregate_costs(self, logs: list[StandardLoggingPayload]) -> list[DatadogFOCUSCostEntry]: """ Aggregates costs by Provider, Model, and Date. Returns a list of DatadogFOCUSCostEntry. """ - aggregator: Dict[Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry] = {} + aggregator: dict[tuple[str, str, str, tuple[tuple[str, str], ...]], DatadogFOCUSCostEntry] = {} for log in logs: try: @@ -164,8 +164,8 @@ class DatadogCostManagementLogger(CustomBatchLogger): return list(aggregator.values()) - def _extract_tags(self, log: StandardLoggingPayload) -> Dict[str, str]: - tags: Dict[str, str] = { + def _extract_tags(self, log: StandardLoggingPayload) -> dict[str, str]: + tags: dict[str, str] = { "env": get_datadog_env(), "service": get_datadog_service(), "host": get_datadog_hostname(), @@ -180,7 +180,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): # cast because StandardLoggingMetadata is a TypedDict; we iterate it # as a generic mapping below. - metadata: Dict[str, Any] = cast(Dict[str, Any], log.get("metadata") or {}) + metadata: dict[str, Any] = cast(dict[str, Any], log.get("metadata") or {}) # Backwards-compat: team/user/model_group preserved regardless of allowlist. if metadata.get("user_api_key_alias"): @@ -220,7 +220,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): return tags @staticmethod - def _set_custom_tag(tags: Dict[str, str], key: str, value: str) -> None: + def _set_custom_tag(tags: dict[str, str], key: str, value: str) -> None: if key in _RESERVED_TAG_KEYS: verbose_logger.debug( "Datadog Cost Management: dropping user-supplied tag %r=%r — " @@ -232,11 +232,11 @@ class DatadogCostManagementLogger(CustomBatchLogger): tags[key] = value @staticmethod - def _add_tag(tags: Dict[str, str], key: str, value: Any) -> None: + def _add_tag(tags: dict[str, str], key: str, value: Any) -> None: if value: tags[key] = str(value) - async def _upload_to_datadog(self, payload: List[Dict]): + async def _upload_to_datadog(self, payload: list[dict]): if not self.dd_api_key or not self.dd_app_key: return diff --git a/litellm/integrations/datadog/datadog_handler.py b/litellm/integrations/datadog/datadog_handler.py index b6bb2b57037..6a86803ed46 100644 --- a/litellm/integrations/datadog/datadog_handler.py +++ b/litellm/integrations/datadog/datadog_handler.py @@ -3,7 +3,6 @@ from __future__ import annotations import os -from typing import List, Optional from litellm.types.utils import StandardLoggingPayload @@ -20,7 +19,7 @@ def get_datadog_hostname() -> str: return os.getenv("HOSTNAME", "") -def get_datadog_base_url_from_env() -> Optional[str]: +def get_datadog_base_url_from_env() -> str | None: """ Get base URL override from common DD_BASE_URL env var. This is useful for testing or custom endpoints. @@ -37,8 +36,8 @@ def get_datadog_pod_name() -> str: def get_datadog_tags( - standard_logging_object: Optional[StandardLoggingPayload] = None, -) -> List[str]: + standard_logging_object: StandardLoggingPayload | None = None, +) -> list[str]: """Build Datadog tags as a list of individual tag strings. Returns a list of "key:value" strings suitable for Datadog LLM Observability @@ -54,7 +53,7 @@ def get_datadog_tags( "POD_NAME": get_datadog_pod_name(), } - tags: List[str] = [f"{k}:{v}" for k, v in base_tags.items()] + tags: list[str] = [f"{k}:{v}" for k, v in base_tags.items()] if standard_logging_object: request_tags = standard_logging_object.get("request_tags", []) or [] diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 1078f05165a..e10071cb083 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -9,23 +9,23 @@ API Reference: https://docs.datadoghq.com/llm_observability/setup/api/?tab=examp import asyncio import json import os -from litellm._uuid import uuid from datetime import datetime -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal import httpx import litellm from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.integrations.custom_batch_logger import CustomBatchLogger -from litellm.integrations.datadog.datadog_mock_client import ( - should_use_datadog_mock, - create_mock_datadog_client, -) from litellm.integrations.datadog.datadog_handler import ( + get_datadog_base_url_from_env, get_datadog_service, get_datadog_tags, - get_datadog_base_url_from_env, +) +from litellm.integrations.datadog.datadog_mock_client import ( + create_mock_datadog_client, + should_use_datadog_mock, ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -80,7 +80,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): asyncio.create_task(self.periodic_flush()) self.flush_lock = asyncio.Lock() - self.log_queue: List[LLMObsPayload] = [] + self.log_queue: list[LLMObsPayload] = [] ######################################################### # Handle datadog_llm_observability_params set as litellm.datadog_llm_observability_params @@ -89,7 +89,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): kwargs.update(dict_datadog_llm_obs_params) CustomBatchLogger.__init__(self, **kwargs, flush_lock=self.flush_lock) except Exception as e: - verbose_logger.exception(f"DataDogLLMObs: Error initializing - {str(e)}") + verbose_logger.exception(f"DataDogLLMObs: Error initializing - {e!s}") raise e def _configure_dd_agent(self, dd_agent_host: str): @@ -118,17 +118,17 @@ class DataDogLLMObsLogger(CustomBatchLogger): self.intake_url = f"https://api.{self.DD_SITE}/api/intake/llm-obs/v1/trace/spans" - def _get_datadog_llm_obs_params(self) -> Dict: + def _get_datadog_llm_obs_params(self) -> dict: """ Get the datadog_llm_observability_params from litellm.datadog_llm_observability_params These are params specific to initializing the DataDogLLMObsLogger e.g. turn_off_message_logging """ - dict_datadog_llm_obs_params: Dict = {} + dict_datadog_llm_obs_params: dict = {} if litellm.datadog_llm_observability_params is not None: if isinstance(litellm.datadog_llm_observability_params, DatadogLLMObsInitParams): dict_datadog_llm_obs_params = litellm.datadog_llm_observability_params.model_dump() - elif isinstance(litellm.datadog_llm_observability_params, Dict): + elif isinstance(litellm.datadog_llm_observability_params, dict): # only allow params that are of DatadogLLMObsInitParams dict_datadog_llm_obs_params = DatadogLLMObsInitParams( **litellm.datadog_llm_observability_params @@ -145,7 +145,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): if len(self.log_queue) >= self.batch_size: await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"DataDogLLMObs: Error logging success event - {str(e)}") + verbose_logger.exception(f"DataDogLLMObs: Error logging success event - {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: @@ -157,7 +157,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): if len(self.log_queue) >= self.batch_size: await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"DataDogLLMObs: Error logging failure event - {str(e)}") + verbose_logger.exception(f"DataDogLLMObs: Error logging failure event - {e!s}") async def async_send_batch(self): try: @@ -214,10 +214,10 @@ class DataDogLLMObsLogger(CustomBatchLogger): except httpx.HTTPStatusError as e: verbose_logger.exception(f"DataDogLLMObs: Error sending batch - {e.response.text}") except Exception as e: - verbose_logger.exception(f"DataDogLLMObs: Error sending batch - {str(e)}") + verbose_logger.exception(f"DataDogLLMObs: Error sending batch - {e!s}") - def create_llm_obs_payload(self, kwargs: Dict, start_time: datetime, end_time: datetime) -> LLMObsPayload: - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + def create_llm_obs_payload(self, kwargs: dict, start_time: datetime, end_time: datetime) -> LLMObsPayload: + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise Exception("DataDogLLMObs: standard_logging_object is not set") @@ -236,7 +236,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): error_info = self._assemble_error_info(standard_logging_payload) - metadata_parent_id: Optional[str] = None + metadata_parent_id: str | None = None if isinstance(metadata, dict): metadata_parent_id = metadata.get("parent_id") @@ -276,7 +276,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): return payload - def _get_apm_trace_id(self) -> Optional[str]: + def _get_apm_trace_id(self) -> str | None: """Retrieve the current APM trace ID if available.""" try: current_span_fn = getattr(tracer, "current_span", None) @@ -290,16 +290,16 @@ class DataDogLLMObsLogger(CustomBatchLogger): pass return None - def _assemble_error_info(self, standard_logging_payload: StandardLoggingPayload) -> Optional[DDLLMObsError]: + def _assemble_error_info(self, standard_logging_payload: StandardLoggingPayload) -> DDLLMObsError | None: """ Assemble error information for failure cases according to DD LLM Obs API spec """ # Handle error information for failure cases according to DD LLM Obs API spec - error_info: Optional[DDLLMObsError] = None + error_info: DDLLMObsError | None = None if standard_logging_payload.get("status") == "failure": # Try to get structured error information first - error_information: Optional[StandardLoggingPayloadErrorInformation] = standard_logging_payload.get( + error_information: StandardLoggingPayloadErrorInformation | None = standard_logging_payload.get( "error_information" ) @@ -321,9 +321,9 @@ class DataDogLLMObsLogger(CustomBatchLogger): For non streaming calls, CompletionStartTime is time we get the response back """ - start_time: Optional[float] = standard_logging_payload.get("startTime") - completion_start_time: Optional[float] = standard_logging_payload.get("completionStartTime") - end_time: Optional[float] = standard_logging_payload.get("endTime") + start_time: float | None = standard_logging_payload.get("startTime") + completion_start_time: float | None = standard_logging_payload.get("completionStartTime") + end_time: float | None = standard_logging_payload.get("endTime") if completion_start_time is not None and start_time is not None: return completion_start_time - start_time @@ -333,8 +333,8 @@ class DataDogLLMObsLogger(CustomBatchLogger): return 0.0 def _get_response_messages( - self, standard_logging_payload: StandardLoggingPayload, call_type: Optional[str] - ) -> List[Any]: + self, standard_logging_payload: StandardLoggingPayload, call_type: str | None + ) -> list[Any]: """ Get the messages from the response object @@ -382,7 +382,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): return [] def _get_datadog_span_kind( - self, call_type: Optional[str], parent_id: Optional[str] = None + self, call_type: str | None, parent_id: str | None = None ) -> Literal["llm", "tool", "task", "embedding", "retrieval"]: """ Map liteLLM call_type to appropriate DataDog LLM Observability span kind. @@ -484,7 +484,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): # Default fallback for unknown or passthrough operations return "llm" - def _ensure_string_content(self, messages: Optional[Union[str, List[Any], Dict[Any, Any]]]) -> List[Any]: + def _ensure_string_content(self, messages: str | list[Any] | dict[Any, Any] | None) -> list[Any]: if messages is None: return [] if isinstance(messages, str): @@ -495,11 +495,11 @@ class DataDogLLMObsLogger(CustomBatchLogger): return [str(messages.get("content", ""))] return [] - def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> Dict[str, Any]: + def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]: """ Fields to track in DD LLM Observability metadata from litellm standard logging payload """ - _metadata: Dict[str, Any] = { + _metadata: dict[str, Any] = { "model_name": standard_logging_payload.get("model", "unknown"), "model_provider": standard_logging_payload.get("custom_llm_provider", "unknown"), "id": standard_logging_payload.get("id", "unknown"), @@ -549,13 +549,13 @@ class DataDogLLMObsLogger(CustomBatchLogger): latency_metrics["litellm_overhead_time_ms"] = litellm_overhead_ms # Guardrail overhead latency - guardrail_info: Optional[list[StandardLoggingGuardrailInformation]] = standard_logging_payload.get( + guardrail_info: list[StandardLoggingGuardrailInformation] | None = standard_logging_payload.get( "guardrail_information" ) if guardrail_info is not None: total_duration = 0.0 for info in guardrail_info: - _guardrail_duration_seconds: Optional[float] = info.get("duration") + _guardrail_duration_seconds: float | None = info.get("duration") if _guardrail_duration_seconds is not None: total_duration += float(_guardrail_duration_seconds) @@ -647,7 +647,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): return spend_metrics - def _process_input_messages_preserving_tool_calls(self, messages: List[Any]) -> List[Dict[str, Any]]: + def _process_input_messages_preserving_tool_calls(self, messages: list[Any]) -> list[dict[str, Any]]: """ Process input messages while preserving tool_calls and tool message types. @@ -671,13 +671,13 @@ class DataDogLLMObsLogger(CustomBatchLogger): return processed @staticmethod - def _tool_calls_kv_pair(tool_calls: List[Dict[str, Any]]) -> Dict[str, Any]: + def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, Any]: """ Extract tool call information into key-value pairs for Datadog metadata. Similar to OpenTelemetry's implementation but adapted for Datadog's format. """ - kv_pairs: Dict[str, Any] = {} + kv_pairs: dict[str, Any] = {} for idx, tool_call in enumerate(tool_calls): try: # Extract tool call ID @@ -707,16 +707,16 @@ class DataDogLLMObsLogger(CustomBatchLogger): kv_pairs[f"tool_calls.{idx}.function.arguments"] = json.dumps(function_arguments) except (KeyError, TypeError, ValueError) as e: - verbose_logger.debug(f"DataDogLLMObs: Error processing tool call {idx}: {str(e)}") + verbose_logger.debug(f"DataDogLLMObs: Error processing tool call {idx}: {e!s}") continue return kv_pairs - def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> Dict[str, Any]: + def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]: """ Extract tool call information from both input messages and response for Datadog metadata. """ - tool_call_metadata: Dict[str, Any] = {} + tool_call_metadata: dict[str, Any] = {} try: # Extract tool calls from input messages @@ -747,6 +747,6 @@ class DataDogLLMObsLogger(CustomBatchLogger): tool_call_metadata[f"output_{key}"] = value except Exception as e: - verbose_logger.debug(f"DataDogLLMObs: Error extracting tool call metadata: {str(e)}") + verbose_logger.debug(f"DataDogLLMObs: Error extracting tool call metadata: {e!s}") return tool_call_metadata diff --git a/litellm/integrations/datadog/datadog_metrics.py b/litellm/integrations/datadog/datadog_metrics.py index b1e4bc73e77..3fbd0f917dc 100644 --- a/litellm/integrations/datadog/datadog_metrics.py +++ b/litellm/integrations/datadog/datadog_metrics.py @@ -3,7 +3,6 @@ import gzip import os import time from datetime import datetime -from typing import List, Optional, Union from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -60,8 +59,8 @@ class DatadogMetricsLogger(CustomBatchLogger): def _extract_tags( self, log: StandardLoggingPayload, - status_code: Optional[Union[str, int]] = None, - ) -> List[str]: + status_code: str | int | None = None, + ) -> list[str]: """ Builds the list of tags for a Datadog metric point """ @@ -105,7 +104,7 @@ class DatadogMetricsLogger(CustomBatchLogger): self, log: StandardLoggingPayload, kwargs: dict, - status_code: Union[str, int] = "200", + status_code: str | int = "200", ): """ Extracts latencies and appends Datadog metric series to the queue @@ -170,7 +169,7 @@ class DatadogMetricsLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -181,11 +180,11 @@ class DatadogMetricsLogger(CustomBatchLogger): await self.flush_queue() except Exception as e: - verbose_logger.exception(f"Datadog Metrics: Error in async_log_success_event: {str(e)}") + verbose_logger.exception(f"Datadog Metrics: Error in async_log_success_event: {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -203,7 +202,7 @@ class DatadogMetricsLogger(CustomBatchLogger): await self.flush_queue() except Exception as e: - verbose_logger.exception(f"Datadog Metrics: Error in async_log_failure_event: {str(e)}") + verbose_logger.exception(f"Datadog Metrics: Error in async_log_failure_event: {e!s}") async def async_send_batch(self): if not self.log_queue: @@ -215,7 +214,7 @@ class DatadogMetricsLogger(CustomBatchLogger): try: await self._upload_to_datadog(payload_data) except Exception as e: - verbose_logger.exception(f"Datadog Metrics: Error in async_send_batch: {str(e)}") + verbose_logger.exception(f"Datadog Metrics: Error in async_send_batch: {e!s}") raise async def _upload_to_datadog(self, payload: DatadogMetricsPayload): @@ -280,7 +279,7 @@ class DatadogMetricsLogger(CustomBatchLogger): async def get_request_response_payload( self, request_id: str, - start_time_utc: Optional[datetime], - end_time_utc: Optional[datetime], - ) -> Optional[dict]: + start_time_utc: datetime | None, + end_time_utc: datetime | None, + ) -> dict | None: pass diff --git a/litellm/integrations/datadog/datadog_team_handler.py b/litellm/integrations/datadog/datadog_team_handler.py index cae954f753c..53eebe7c505 100644 --- a/litellm/integrations/datadog/datadog_team_handler.py +++ b/litellm/integrations/datadog/datadog_team_handler.py @@ -5,7 +5,7 @@ Used to get the DataDogLogger for a given request. Handles Key/Team Based Datadog Logging, following the same pattern as LangFuseHandler. """ -from typing import TYPE_CHECKING, Any, Dict, Optional, TypedDict +from typing import TYPE_CHECKING, Any, TypedDict from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams @@ -19,10 +19,10 @@ else: class DatadogLoggingConfig(TypedDict): - dd_api_key: Optional[str] - dd_site: Optional[str] - dd_agent_host: Optional[str] - dd_agent_port: Optional[str] + dd_api_key: str | None + dd_site: str | None + dd_agent_host: str | None + dd_agent_port: str | None class DataDogHandler: @@ -63,7 +63,7 @@ class DataDogHandler: @staticmethod def _create_datadog_logger_from_credentials( - credentials: Dict, + credentials: dict, in_memory_dynamic_logger_cache: DynamicLoggingCache, ) -> DataDogLogger: """ diff --git a/litellm/integrations/deepeval/api.py b/litellm/integrations/deepeval/api.py index fccc5970433..adca8928df4 100644 --- a/litellm/integrations/deepeval/api.py +++ b/litellm/integrations/deepeval/api.py @@ -1,7 +1,9 @@ # duplicate -> https://github.com/confident-ai/deepeval/blob/main/deepeval/confident/api.py import logging -import httpx from enum import Enum + +import httpx + from litellm._logging import verbose_logger DEEPEVAL_BASE_URL = "https://deepeval.confident-ai.com" diff --git a/litellm/integrations/deepeval/deepeval.py b/litellm/integrations/deepeval/deepeval.py index 90c1d8eedce..e194d351b8c 100644 --- a/litellm/integrations/deepeval/deepeval.py +++ b/litellm/integrations/deepeval/deepeval.py @@ -1,4 +1,6 @@ import os + +from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.deepeval.api import Api, Endpoints, HttpMethods @@ -12,7 +14,6 @@ from litellm.integrations.deepeval.utils import ( to_zod_compatible_iso, validate_environment, ) -from litellm._logging import verbose_logger # This file includes the custom callbacks for LiteLLM Proxy diff --git a/litellm/integrations/deepeval/types.py b/litellm/integrations/deepeval/types.py index afaf4436db9..c86d01d0468 100644 --- a/litellm/integrations/deepeval/types.py +++ b/litellm/integrations/deepeval/types.py @@ -1,7 +1,8 @@ # Duplicate -> https://github.com/confident-ai/deepeval/blob/main/deepeval/tracing/api.py from enum import Enum -from typing import Any, ClassVar, Dict, List, Optional, Union, Literal -from pydantic import BaseModel, Field, ConfigDict +from typing import Any, ClassVar, Literal + +from pydantic import BaseModel, ConfigDict, Field class SpanApiType(Enum): @@ -24,37 +25,37 @@ class BaseApiSpan(BaseModel): model_config: ClassVar[ConfigDict] = ConfigDict(use_enum_values=True) uuid: str - name: Optional[str] = None + name: str | None = None status: TraceSpanApiStatus type: SpanApiType trace_uuid: str = Field(alias="traceUuid") - parent_uuid: Optional[str] = Field(None, alias="parentUuid") + parent_uuid: str | None = Field(None, alias="parentUuid") start_time: str = Field(alias="startTime") end_time: str = Field(alias="endTime") - input: Optional[Union[Dict, list, str]] = None - output: Optional[Union[Dict, list, str]] = None - error: Optional[str] = None + input: dict | list | str | None = None + output: dict | list | str | None = None + error: str | None = None # llm - model: Optional[str] = None - input_token_count: Optional[int] = Field(None, alias="inputTokenCount") - output_token_count: Optional[int] = Field(None, alias="outputTokenCount") - cost_per_input_token: Optional[float] = Field(None, alias="costPerInputToken") - cost_per_output_token: Optional[float] = Field(None, alias="costPerOutputToken") + model: str | None = None + input_token_count: int | None = Field(None, alias="inputTokenCount") + output_token_count: int | None = Field(None, alias="outputTokenCount") + cost_per_input_token: float | None = Field(None, alias="costPerInputToken") + cost_per_output_token: float | None = Field(None, alias="costPerOutputToken") class TraceApi(BaseModel): uuid: str - base_spans: List[BaseApiSpan] = Field(alias="baseSpans") - agent_spans: List[BaseApiSpan] = Field(alias="agentSpans") - llm_spans: List[BaseApiSpan] = Field(alias="llmSpans") - retriever_spans: List[BaseApiSpan] = Field(alias="retrieverSpans") - tool_spans: List[BaseApiSpan] = Field(alias="toolSpans") + base_spans: list[BaseApiSpan] = Field(alias="baseSpans") + agent_spans: list[BaseApiSpan] = Field(alias="agentSpans") + llm_spans: list[BaseApiSpan] = Field(alias="llmSpans") + retriever_spans: list[BaseApiSpan] = Field(alias="retrieverSpans") + tool_spans: list[BaseApiSpan] = Field(alias="toolSpans") start_time: str = Field(alias="startTime") end_time: str = Field(alias="endTime") - metadata: Optional[Dict[str, Any]] = Field(None) - tags: Optional[List[str]] = Field(None) - environment: Optional[str] = Field(None) + metadata: dict[str, Any] | None = Field(None) + tags: list[str] | None = Field(None) + environment: str | None = Field(None) class Environment(Enum): diff --git a/litellm/integrations/deepeval/utils.py b/litellm/integrations/deepeval/utils.py index 3df9aceb241..9d65b509fe2 100644 --- a/litellm/integrations/deepeval/utils.py +++ b/litellm/integrations/deepeval/utils.py @@ -1,4 +1,5 @@ from datetime import datetime, timezone + from litellm.integrations.deepeval.types import Environment diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py index 8432d50e32b..b254c0315af 100644 --- a/litellm/integrations/dotprompt/__init__.py +++ b/litellm/integrations/dotprompt/__init__.py @@ -1,16 +1,17 @@ from typing import TYPE_CHECKING, Optional if TYPE_CHECKING: - from .prompt_manager import PromptManager, PromptTemplate - from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.integrations.custom_prompt_management import CustomPromptManagement + from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec + + from .prompt_manager import PromptManager, PromptTemplate from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .dotprompt_manager import DotpromptManager # Global instances -global_prompt_directory: Optional[str] = None +global_prompt_directory: str | None = None global_prompt_manager: Optional["PromptManager"] = None @@ -80,10 +81,10 @@ prompt_initializer_registry = { # Export public API __all__ = [ - "PromptManager", "DotpromptManager", + "PromptManager", "PromptTemplate", - "set_global_prompt_directory", "global_prompt_directory", "global_prompt_manager", + "set_global_prompt_directory", ] diff --git a/litellm/integrations/dotprompt/dotprompt_manager.py b/litellm/integrations/dotprompt/dotprompt_manager.py index 3ba9efd68b7..588ef442378 100644 --- a/litellm/integrations/dotprompt/dotprompt_manager.py +++ b/litellm/integrations/dotprompt/dotprompt_manager.py @@ -4,7 +4,7 @@ Builds on top of PromptManagementBase to provide .prompt file support. """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.integrations.prompt_management_base import PromptManagementClient @@ -42,10 +42,10 @@ class DotpromptManager(CustomPromptManagement): def __init__( self, - prompt_directory: Optional[str] = None, - prompt_file: Optional[str] = None, - prompt_data: Optional[Union[dict, str]] = None, - prompt_id: Optional[str] = None, + prompt_directory: str | None = None, + prompt_file: str | None = None, + prompt_data: dict | str | None = None, + prompt_id: str | None = None, ): import litellm @@ -56,7 +56,7 @@ class DotpromptManager(CustomPromptManagement): else: self.prompt_data = prompt_data or {} - self._prompt_manager: Optional[PromptManager] = None + self._prompt_manager: PromptManager | None = None self.prompt_file = prompt_file self.prompt_id = prompt_id @@ -84,8 +84,8 @@ class DotpromptManager(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """ @@ -103,12 +103,12 @@ class DotpromptManager(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Compile a .prompt file into a PromptManagementClient structure. @@ -159,12 +159,12 @@ class DotpromptManager(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Async version of compile prompt helper. Since dotprompt operations are synchronous, @@ -185,17 +185,17 @@ class DotpromptManager(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: from litellm.integrations.prompt_management_base import PromptManagementBase return PromptManagementBase.get_chat_completion_prompt( @@ -214,19 +214,19 @@ class DotpromptManager(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Async version - delegates to PromptManagementBase async implementation. """ @@ -249,7 +249,7 @@ class DotpromptManager(CustomPromptManagement): ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params, ) - def _convert_to_messages(self, rendered_content: str) -> List[AllMessageValues]: + def _convert_to_messages(self, rendered_content: str) -> list[AllMessageValues]: """ Convert rendered prompt content to chat messages. @@ -339,20 +339,20 @@ class DotpromptManager(CustomPromptManagement): if self._prompt_manager: self._prompt_manager.reload_prompts() - def add_prompt_from_json(self, prompt_id: str, json_data: Dict[str, Any]) -> None: + def add_prompt_from_json(self, prompt_id: str, json_data: dict[str, Any]) -> None: """Add a prompt from JSON data.""" content = json_data.get("content", "") metadata = json_data.get("metadata", {}) self.prompt_manager.add_prompt(prompt_id, content, metadata) - def load_prompts_from_json(self, prompts_data: Dict[str, Dict[str, Any]]) -> None: + def load_prompts_from_json(self, prompts_data: dict[str, dict[str, Any]]) -> None: """Load multiple prompts from JSON data.""" self.prompt_manager.load_prompts_from_json_data(prompts_data) - def get_prompts_as_json(self) -> Dict[str, Dict[str, Any]]: + def get_prompts_as_json(self) -> dict[str, dict[str, Any]]: """Get all prompts in JSON format.""" return self.prompt_manager.get_all_prompts_as_json() - def convert_prompt_file_to_json(self, file_path: str) -> Dict[str, Any]: + def convert_prompt_file_to_json(self, file_path: str) -> dict[str, Any]: """Convert a .prompt file to JSON format.""" return self.prompt_manager.prompt_file_to_json(file_path) diff --git a/litellm/integrations/dotprompt/prompt_manager.py b/litellm/integrations/dotprompt/prompt_manager.py index dd198ba1272..5bfe63e0f41 100644 --- a/litellm/integrations/dotprompt/prompt_manager.py +++ b/litellm/integrations/dotprompt/prompt_manager.py @@ -4,7 +4,7 @@ Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/d import re from pathlib import Path -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any import yaml from jinja2 import DictLoader, select_autoescape @@ -17,8 +17,8 @@ class PromptTemplate: def __init__( self, content: str, - metadata: Optional[Dict[str, Any]] = None, - template_id: Optional[str] = None, + metadata: dict[str, Any] | None = None, + template_id: str | None = None, ): self.content = content self.metadata = metadata or {} @@ -52,13 +52,13 @@ class PromptManager: def __init__( self, - prompt_id: Optional[str] = None, - prompt_directory: Optional[str] = None, - prompt_data: Optional[Dict[str, Dict[str, Any]]] = None, - prompt_file: Optional[str] = None, + prompt_id: str | None = None, + prompt_directory: str | None = None, + prompt_data: dict[str, dict[str, Any]] | None = None, + prompt_file: str | None = None, ): self.prompt_directory = Path(prompt_directory) if prompt_directory else None - self.prompts: Dict[str, PromptTemplate] = {} + self.prompts: dict[str, PromptTemplate] = {} self.prompt_file = prompt_file # Sandboxed env: templates can come from user input via /prompts/test, # so we must block access to unsafe Python attributes and mutation of @@ -107,7 +107,7 @@ class PromptManager: # Optional: print(f"Error loading prompt file {prompt_file}") pass - def _load_prompts_from_json(self, prompt_data: Dict[str, Dict[str, Any]], prompt_id: Optional[str] = None) -> None: + def _load_prompts_from_json(self, prompt_data: dict[str, dict[str, Any]], prompt_id: str | None = None) -> None: """Load prompts from JSON data structure. Expected format: @@ -143,7 +143,7 @@ class PromptManager: # Optional: print(f"Error loading prompt from JSON: {prompt_id}") pass - def _load_prompt_file(self, file_path: Union[str, Path], prompt_id: str) -> PromptTemplate: + def _load_prompt_file(self, file_path: str | Path, prompt_id: str) -> PromptTemplate: """Load and parse a single .prompt file.""" if isinstance(file_path, str): file_path = Path(file_path) @@ -159,7 +159,7 @@ class PromptManager: template_id=prompt_id, ) - def _parse_frontmatter(self, content: str) -> Tuple[Dict[str, Any], str]: + def _parse_frontmatter(self, content: str) -> tuple[dict[str, Any], str]: """Parse YAML frontmatter from prompt content.""" # Match YAML frontmatter between --- delimiters frontmatter_pattern = r"^---\s*\n(.*?)\n---\s*\n(.*)$" @@ -183,8 +183,8 @@ class PromptManager: def render( self, prompt_id: str, - prompt_variables: Optional[Dict[str, Any]] = None, - version: Optional[int] = None, + prompt_variables: dict[str, Any] | None = None, + version: int | None = None, ) -> str: """ Render a prompt template with the given variables. @@ -223,7 +223,7 @@ class PromptManager: except Exception as e: raise ValueError(f"Error rendering template '{prompt_id}': {e}") - def _validate_input(self, variables: Dict[str, Any], schema: Dict[str, Any]) -> None: + def _validate_input(self, variables: dict[str, Any], schema: dict[str, Any]) -> None: """Basic validation of input variables against schema.""" for field_name, field_type in schema.items(): if field_name in variables: @@ -236,9 +236,9 @@ class PromptManager: f"expected {getattr(expected_type, '__name__', str(expected_type))}, got {type(value).__name__}" ) - def _get_python_type(self, schema_type: str) -> Union[type, tuple]: + def _get_python_type(self, schema_type: str) -> type | tuple: """Convert schema type string to Python type.""" - type_mapping: Dict[str, Union[type, tuple]] = { + type_mapping: dict[str, type | tuple] = { "string": str, "str": str, "number": (int, float), @@ -255,7 +255,7 @@ class PromptManager: return type_mapping.get(schema_type.lower(), str) # type: ignore - def get_prompt(self, prompt_id: str, version: Optional[int] = None) -> Optional[PromptTemplate]: + def get_prompt(self, prompt_id: str, version: int | None = None) -> PromptTemplate | None: """ Get a prompt template by ID and optional version. @@ -275,11 +275,11 @@ class PromptManager: # Fall back to base prompt_id return self.prompts.get(prompt_id) - def list_prompts(self) -> List[str]: + def list_prompts(self) -> list[str]: """Get a list of all available prompt IDs.""" return list(self.prompts.keys()) - def get_prompt_metadata(self, prompt_id: str) -> Optional[Dict[str, Any]]: + def get_prompt_metadata(self, prompt_id: str) -> dict[str, Any] | None: """Get metadata for a specific prompt.""" template = self.prompts.get(prompt_id) return template.metadata if template else None @@ -290,12 +290,12 @@ class PromptManager: if self.prompt_directory: self._load_prompts() - def add_prompt(self, prompt_id: str, content: str, metadata: Optional[Dict[str, Any]] = None) -> None: + def add_prompt(self, prompt_id: str, content: str, metadata: dict[str, Any] | None = None) -> None: """Add a prompt template programmatically.""" template = PromptTemplate(content=content, metadata=metadata or {}, template_id=prompt_id) self.prompts[prompt_id] = template - def prompt_file_to_json(self, file_path: Union[str, Path]) -> Dict[str, Any]: + def prompt_file_to_json(self, file_path: str | Path) -> dict[str, Any]: """Convert a .prompt file to JSON format. Args: @@ -312,7 +312,7 @@ class PromptManager: return {"content": template_content.strip(), "metadata": frontmatter} - def json_to_prompt_file(self, prompt_data: Dict[str, Any]) -> str: + def json_to_prompt_file(self, prompt_data: dict[str, Any]) -> str: """Convert JSON prompt data to .prompt file format. Args: @@ -335,7 +335,7 @@ class PromptManager: return f"---\n{frontmatter_yaml}---\n{content}" - def get_all_prompts_as_json(self) -> Dict[str, Dict[str, Any]]: + def get_all_prompts_as_json(self) -> dict[str, dict[str, Any]]: """Get all loaded prompts in JSON format. Returns: @@ -349,6 +349,6 @@ class PromptManager: } return result - def load_prompts_from_json_data(self, prompt_data: Dict[str, Dict[str, Any]]) -> None: + def load_prompts_from_json_data(self, prompt_data: dict[str, dict[str, Any]]) -> None: """Load additional prompts from JSON data (merges with existing prompts).""" self._load_prompts_from_json(prompt_data) diff --git a/litellm/integrations/dynamodb.py b/litellm/integrations/dynamodb.py index ab76fa3c8bd..5826a06b0ec 100644 --- a/litellm/integrations/dynamodb.py +++ b/litellm/integrations/dynamodb.py @@ -3,10 +3,10 @@ import os import traceback -from litellm._uuid import uuid from typing import Any import litellm +from litellm._uuid import uuid class DyanmoDBLogger: @@ -70,10 +70,9 @@ class DyanmoDBLogger: # Assuming log_data is a dictionary with log information response = table.put_item(Item=payload) - print_verbose(f"Response from DynamoDB:{str(response)}") + print_verbose(f"Response from DynamoDB:{response!s}") print_verbose(f"DynamoDB Layer Logging - final response object: {response_obj}") return response except Exception: print_verbose(f"DynamoDB Layer Error - {traceback.format_exc()}") - pass diff --git a/litellm/integrations/email_alerting.py b/litellm/integrations/email_alerting.py index 35d63a691f9..92a56eaaf75 100644 --- a/litellm/integrations/email_alerting.py +++ b/litellm/integrations/email_alerting.py @@ -3,7 +3,6 @@ Functions for sending Email Alerts """ import os -from typing import List, Optional from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.proxy._types import WebhookEvent @@ -14,7 +13,7 @@ LITELLM_LOGO_URL = "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" LITELLM_SUPPORT_CONTACT = "support@berri.ai" -async def get_all_team_member_emails(team_id: Optional[str] = None) -> list: +async def get_all_team_member_emails(team_id: str | None = None) -> list: verbose_logger.debug("Email Alerting: Getting all team members for team_id=%s", team_id) if team_id is None: return [] @@ -38,7 +37,7 @@ async def get_all_team_member_emails(team_id: Optional[str] = None) -> list: team_id, _team_members, ) - _team_member_user_ids: List[str] = [] + _team_member_user_ids: list[str] = [] for member in _team_members: if member and isinstance(member, dict): _user_id = member.get("user_id") diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py index 3ae3f6b53ac..0dd8bc5bf38 100644 --- a/litellm/integrations/focus/database.py +++ b/litellm/integrations/focus/database.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any import polars as pl @@ -24,9 +24,9 @@ class FocusLiteLLMDatabase: async def get_usage_data( self, *, - limit: Optional[int] = None, - start_time_utc: Optional[datetime] = None, - end_time_utc: Optional[datetime] = None, + limit: int | None = None, + start_time_utc: datetime | None = None, + end_time_utc: datetime | None = None, ) -> pl.DataFrame: """Return usage data for the requested window.""" client = self._ensure_prisma_client() @@ -100,7 +100,7 @@ class FocusLiteLLMDatabase: except Exception as exc: raise RuntimeError(f"Error retrieving usage data: {exc}") from exc - async def get_table_info(self) -> Dict[str, Any]: + async def get_table_info(self) -> dict[str, Any]: """Return metadata about the spend table for diagnostics.""" client = self._ensure_prisma_client() diff --git a/litellm/integrations/focus/destinations/__init__.py b/litellm/integrations/focus/destinations/__init__.py index 21945c9b457..932184cf485 100644 --- a/litellm/integrations/focus/destinations/__init__.py +++ b/litellm/integrations/focus/destinations/__init__.py @@ -3,16 +3,16 @@ from .base import FocusDestination, FocusTimeWindow from .factory import FocusDestinationFactory from .gcs_destination import FocusGCSDestination -from .s3_destination import FocusS3Destination from .mavvrik_destination import FocusMavvrikDestination +from .s3_destination import FocusS3Destination from .vantage_destination import FocusVantageDestination __all__ = [ "FocusDestination", "FocusDestinationFactory", "FocusGCSDestination", - "FocusTimeWindow", - "FocusS3Destination", "FocusMavvrikDestination", + "FocusS3Destination", + "FocusTimeWindow", "FocusVantageDestination", ] diff --git a/litellm/integrations/focus/destinations/factory.py b/litellm/integrations/focus/destinations/factory.py index 3d79046bf6c..6e807d009ef 100644 --- a/litellm/integrations/focus/destinations/factory.py +++ b/litellm/integrations/focus/destinations/factory.py @@ -3,12 +3,12 @@ from __future__ import annotations import os -from typing import Any, Dict, Optional +from typing import Any from .base import FocusDestination from .gcs_destination import FocusGCSDestination -from .s3_destination import FocusS3Destination from .mavvrik_destination import FocusMavvrikDestination +from .s3_destination import FocusS3Destination from .vantage_destination import FocusVantageDestination @@ -20,7 +20,7 @@ class FocusDestinationFactory: *, provider: str, prefix: str, - config: Optional[Dict[str, Any]] = None, + config: dict[str, Any] | None = None, ) -> FocusDestination: """Return a destination implementation for the requested provider.""" provider_lower = provider.lower() @@ -39,8 +39,8 @@ class FocusDestinationFactory: def _resolve_config( *, provider: str, - overrides: Dict[str, Any], - ) -> Dict[str, Any]: + overrides: dict[str, Any], + ) -> dict[str, Any]: if provider == "s3": resolved = { "bucket_name": overrides.get("bucket_name") or os.getenv("FOCUS_S3_BUCKET_NAME"), diff --git a/litellm/integrations/focus/destinations/gcs_destination.py b/litellm/integrations/focus/destinations/gcs_destination.py index e4525ccd267..898b58d952a 100644 --- a/litellm/integrations/focus/destinations/gcs_destination.py +++ b/litellm/integrations/focus/destinations/gcs_destination.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import timezone -from typing import Any, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase @@ -21,7 +21,7 @@ class FocusGCSDestination(GCSBucketBase, FocusDestination): self, *, prefix: str, - config: Optional[dict[str, Any]] = None, + config: dict[str, Any] | None = None, ) -> None: config = config or {} bucket_name = config.get("bucket_name") diff --git a/litellm/integrations/focus/destinations/mavvrik_destination.py b/litellm/integrations/focus/destinations/mavvrik_destination.py index cf500a71b52..385d0f8ce3b 100644 --- a/litellm/integrations/focus/destinations/mavvrik_destination.py +++ b/litellm/integrations/focus/destinations/mavvrik_destination.py @@ -9,7 +9,7 @@ Flow: from __future__ import annotations import gzip -from typing import Any, Optional +from typing import Any from urllib.parse import urlparse from litellm._logging import verbose_logger @@ -56,7 +56,7 @@ class FocusMavvrikDestination(FocusDestination): self, *, prefix: str, - config: Optional[dict[str, Any]] = None, + config: dict[str, Any] | None = None, ) -> None: config = config or {} api_key = config.get("api_key") @@ -100,7 +100,7 @@ class FocusMavvrikDestination(FocusDestination): def _auth_headers(self) -> dict[str, str]: return {"Content-Type": "application/json", "x-api-key": self.api_key} - async def _ensure_registered(self) -> Optional[int]: + async def _ensure_registered(self) -> int | None: """POST agent endpoint to register/initialize the connector (once per instance). Returns metricsMarker from the Mavvrik response — the last date index @@ -264,7 +264,7 @@ class FocusMavvrikDestination(FocusDestination): ) verbose_logger.debug("Mavvrik FOCUS destination: metricsMarker advanced to %s", date_epoch) - async def get_metrics_marker(self) -> Optional[int]: + async def get_metrics_marker(self) -> int | None: """Register with Mavvrik and return the current metricsMarker. Always calls the Mavvrik register API — unlike deliver() which skips diff --git a/litellm/integrations/focus/destinations/s3_destination.py b/litellm/integrations/focus/destinations/s3_destination.py index c6d5554b438..28896102bb8 100644 --- a/litellm/integrations/focus/destinations/s3_destination.py +++ b/litellm/integrations/focus/destinations/s3_destination.py @@ -4,7 +4,7 @@ from __future__ import annotations import asyncio from datetime import timezone -from typing import Any, Optional +from typing import Any import boto3 @@ -18,7 +18,7 @@ class FocusS3Destination(FocusDestination): self, *, prefix: str, - config: Optional[dict[str, Any]] = None, + config: dict[str, Any] | None = None, ) -> None: config = config or {} bucket_name = config.get("bucket_name") diff --git a/litellm/integrations/focus/destinations/vantage_destination.py b/litellm/integrations/focus/destinations/vantage_destination.py index ffd37aa195b..41363752ab6 100644 --- a/litellm/integrations/focus/destinations/vantage_destination.py +++ b/litellm/integrations/focus/destinations/vantage_destination.py @@ -4,7 +4,7 @@ from __future__ import annotations import csv import io -from typing import Any, Optional +from typing import Any import httpx # noqa: F401 - used at runtime (AsyncClient, HTTPStatusError) @@ -94,7 +94,7 @@ class FocusVantageDestination(FocusDestination): self, *, prefix: str, - config: Optional[dict[str, Any]] = None, + config: dict[str, Any] | None = None, ) -> None: config = config or {} api_key = config.get("api_key") @@ -173,7 +173,7 @@ class FocusVantageDestination(FocusDestination): header = lines[0] data_lines = [line for line in lines[1:] if line.strip()] - first_error: Optional[Exception] = None + first_error: Exception | None = None batch_num = 0 for start in range(0, len(data_lines), VANTAGE_MAX_ROWS_PER_UPLOAD): batch_lines = data_lines[start : start + VANTAGE_MAX_ROWS_PER_UPLOAD] @@ -214,7 +214,7 @@ class FocusVantageDestination(FocusDestination): current_size = len(header) + 1 # header + newline sub_batch = 0 header_size = len(header) + 1 - first_error: Optional[Exception] = None + first_error: Exception | None = None for line in data_lines: line_size = len(line) + 1 # line + newline diff --git a/litellm/integrations/focus/export_engine.py b/litellm/integrations/focus/export_engine.py index 67ae6bcc3d0..197c5ad341c 100644 --- a/litellm/integrations/focus/export_engine.py +++ b/litellm/integrations/focus/export_engine.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import Any, Dict, Optional +from typing import Any import polars as pl @@ -23,7 +23,7 @@ class FocusExportEngine: provider: str, export_format: str, prefix: str, - destination_config: Optional[dict[str, Any]] = None, + destination_config: dict[str, Any] | None = None, ) -> None: self.provider = provider self.export_format = export_format @@ -44,7 +44,7 @@ class FocusExportEngine: return FocusParquetSerializer() raise NotImplementedError(f"Export format '{self.export_format}' not supported. Use 'parquet' or 'csv'.") - async def dry_run_export_usage_data(self, limit: Optional[int]) -> Dict[str, Any]: + async def dry_run_export_usage_data(self, limit: int | None) -> dict[str, Any]: data = await self._database.get_usage_data(limit=limit) normalized = self._transformer.transform(data) @@ -68,7 +68,7 @@ class FocusExportEngine: async def export_all( self, *, - limit: Optional[int], + limit: int | None, ) -> None: """Export all available data without time-window filtering.""" data = await self._database.get_usage_data(limit=limit) @@ -96,7 +96,7 @@ class FocusExportEngine: self, *, window: FocusTimeWindow, - limit: Optional[int], + limit: int | None, ) -> None: data = await self._database.get_usage_data( limit=limit, diff --git a/litellm/integrations/focus/focus_logger.py b/litellm/integrations/focus/focus_logger.py index ac6f1f7af1f..60353cd4aca 100644 --- a/litellm/integrations/focus/focus_logger.py +++ b/litellm/integrations/focus/focus_logger.py @@ -4,7 +4,7 @@ from __future__ import annotations import os from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_logger @@ -14,6 +14,7 @@ from .destinations import FocusTimeWindow if TYPE_CHECKING: from apscheduler.schedulers.asyncio import AsyncIOScheduler + from .export_engine import FocusExportEngine else: AsyncIOScheduler = Any @@ -28,13 +29,13 @@ class FocusLogger(CustomLogger): def __init__( self, *, - provider: Optional[str] = None, - export_format: Optional[str] = None, - frequency: Optional[str] = None, - cron_offset_minute: Optional[int] = None, - interval_seconds: Optional[int] = None, - prefix: Optional[str] = None, - destination_config: Optional[dict[str, Any]] = None, + provider: str | None = None, + export_format: str | None = None, + frequency: str | None = None, + cron_offset_minute: int | None = None, + interval_seconds: int | None = None, + prefix: str | None = None, + destination_config: dict[str, Any] | None = None, **kwargs: Any, ) -> None: super().__init__(**kwargs) @@ -45,7 +46,7 @@ class FocusLogger(CustomLogger): cron_offset_minute if cron_offset_minute is not None else int(os.getenv("FOCUS_CRON_OFFSET", "5")) ) raw_interval = interval_seconds if interval_seconds is not None else os.getenv("FOCUS_INTERVAL_SECONDS") - self.interval_seconds: Optional[int] = None + self.interval_seconds: int | None = None if raw_interval is not None: try: self.interval_seconds = int(raw_interval) @@ -58,9 +59,9 @@ class FocusLogger(CustomLogger): self.prefix: str = prefix if prefix is not None else (env_prefix if env_prefix else "focus_exports") self._destination_config = destination_config - self._engine: Optional["FocusExportEngine"] = None + self._engine: FocusExportEngine | None = None - def _ensure_engine(self) -> "FocusExportEngine": + def _ensure_engine(self) -> FocusExportEngine: """Instantiate the heavy export engine lazily.""" if self._engine is None: from .export_engine import FocusExportEngine @@ -76,9 +77,9 @@ class FocusLogger(CustomLogger): async def export_usage_data( self, *, - limit: Optional[int] = None, - start_time_utc: Optional[datetime] = None, - end_time_utc: Optional[datetime] = None, + limit: int | None = None, + start_time_utc: datetime | None = None, + end_time_utc: datetime | None = None, ) -> None: """Public hook to trigger export immediately. @@ -101,7 +102,7 @@ class FocusLogger(CustomLogger): # No time bounds → export all available data await self._export_all(limit=limit) - async def dry_run_export_usage_data(self, limit: Optional[int] = DEFAULT_DRY_RUN_LIMIT) -> dict[str, Any]: + async def dry_run_export_usage_data(self, limit: int | None = DEFAULT_DRY_RUN_LIMIT) -> dict[str, Any]: """Return transformed data without uploading.""" engine = self._ensure_engine() return await engine.dry_run_export_usage_data(limit=limit) @@ -136,7 +137,7 @@ class FocusLogger(CustomLogger): # Use exact type match to exclude subclasses like VantageLogger, # which have their own dedicated scheduling method. - focus_loggers: List[CustomLogger] = [ + focus_loggers: list[CustomLogger] = [ cb for cb in litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=FocusLogger) if type(cb) is FocusLogger @@ -152,7 +153,7 @@ class FocusLogger(CustomLogger): **trigger_kwargs, ) - def _build_scheduler_trigger(self) -> Dict[str, Any]: + def _build_scheduler_trigger(self) -> dict[str, Any]: """Return scheduler configuration for the selected frequency.""" if self.frequency == "interval": seconds = self.interval_seconds or 60 @@ -178,7 +179,7 @@ class FocusLogger(CustomLogger): async def _export_all( self, *, - limit: Optional[int], + limit: int | None, ) -> None: """Export all available data without a time window filter.""" engine = self._ensure_engine() @@ -188,7 +189,7 @@ class FocusLogger(CustomLogger): self, *, window: FocusTimeWindow, - limit: Optional[int], + limit: int | None, ) -> None: engine = self._ensure_engine() await engine.export_window(window=window, limit=limit) diff --git a/litellm/integrations/focus/serializers/__init__.py b/litellm/integrations/focus/serializers/__init__.py index bdbf5204540..7e15ca7388e 100644 --- a/litellm/integrations/focus/serializers/__init__.py +++ b/litellm/integrations/focus/serializers/__init__.py @@ -4,4 +4,4 @@ from .base import FocusSerializer from .csv import FocusCsvSerializer from .parquet import FocusParquetSerializer -__all__ = ["FocusSerializer", "FocusCsvSerializer", "FocusParquetSerializer"] +__all__ = ["FocusCsvSerializer", "FocusParquetSerializer", "FocusSerializer"] diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index 0ec6d496689..0180af51992 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -5,7 +5,7 @@ import os import re import uuid from datetime import datetime, timezone -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, cast import httpx from pydantic import BaseModel, Field @@ -17,16 +17,16 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, get_content_from_model_response, ) -from litellm.types.llms.openai import ( - AllMessageValues, - HttpxBinaryResponseContent, - ResponsesAPIResponse, -) from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus +from litellm.types.llms.openai import ( + AllMessageValues, + HttpxBinaryResponseContent, + ResponsesAPIResponse, +) GALILEO_CLOUD_API_BASE_URL = "https://api.galileo.ai" # Cap the in-memory buffer so persistent flush failures (e.g. Galileo @@ -44,22 +44,22 @@ class LLMResponse(BaseModel): num_input_tokens: int num_output_tokens: int num_total_tokens: int - cost: Optional[float] = Field( + cost: float | None = Field( default=None, description="Total cost of the LLM call in USD as computed by LiteLLM.", ) - output_logprobs: Optional[Dict[str, Any]] = Field( + output_logprobs: dict[str, Any] | None = Field( default=None, description="Optional. When available, logprobs are used to compute Uncertainty.", ) created_at: str = Field(..., description='timestamp constructed in "%Y-%m-%dT%H:%M:%S" format') - tags: Optional[List[str]] = None - user_metadata: Optional[Dict[str, Any]] = None + tags: list[str] | None = None + user_metadata: dict[str, Any] | None = None class GalileoObserve(CustomLogger): def __init__(self) -> None: - self.in_memory_records: List[dict] = [] + self.in_memory_records: list[dict] = [] self.batch_size = 1 self.api_key = os.getenv("GALILEO_API_KEY") self.project_id = os.getenv("GALILEO_PROJECT_ID") @@ -70,11 +70,11 @@ class GalileoObserve(CustomLogger): if self.api_key and not self.base_url: self.base_url = GALILEO_CLOUD_API_BASE_URL self.use_v2_api = bool(self.api_key) - self.headers: Optional[Dict[str, str]] = None + self.headers: dict[str, str] | None = None self.async_httpx_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) @staticmethod - def _normalize_base_url(base_url: Optional[str]) -> Optional[str]: + def _normalize_base_url(base_url: str | None) -> str | None: if base_url: return base_url.rstrip("/") return None @@ -128,7 +128,7 @@ class GalileoObserve(CustomLogger): except Exception as e: return IntegrationHealthCheckStatus( status="unhealthy", - error_message=f"Galileo health check failed: {str(e)}", + error_message=f"Galileo health check failed: {e!s}", ) async def async_set_galileo_headers(self) -> None: @@ -176,7 +176,7 @@ class GalileoObserve(CustomLogger): return False @staticmethod - def _galileo_input_messages(messages: Optional[Any], input_text: str) -> List[Dict[str, str]]: + def _galileo_input_messages(messages: Any | None, input_text: str) -> list[dict[str, str]]: if isinstance(messages, dict): messages = messages.get("messages") if not messages: @@ -184,7 +184,7 @@ class GalileoObserve(CustomLogger): if not isinstance(messages, list): return [{"role": "user", "content": input_text}] - galileo_messages: List[Dict[str, str]] = [] + galileo_messages: list[dict[str, str]] = [] for message in messages: if not isinstance(message, dict): continue @@ -207,7 +207,7 @@ class GalileoObserve(CustomLogger): return datetime.now().astimezone().tzinfo or timezone.utc @staticmethod - def _format_created_at(dt: Union[datetime, Any]) -> str: + def _format_created_at(dt: datetime | Any) -> str: """Serialize timestamps as UTC ISO-8601 for Galileo.""" if not isinstance(dt, datetime): return str(dt) @@ -226,13 +226,13 @@ class GalileoObserve(CustomLogger): return created_at @staticmethod - def _token_metrics_from_record(record: Dict[str, Any]) -> Dict[str, Any]: + def _token_metrics_from_record(record: dict[str, Any]) -> dict[str, Any]: num_input_tokens = int(record.get("num_input_tokens") or 0) num_output_tokens = int(record.get("num_output_tokens") or 0) num_total_tokens = int(record.get("num_total_tokens") or 0) if num_total_tokens == 0 and (num_input_tokens or num_output_tokens): num_total_tokens = num_input_tokens + num_output_tokens - metrics: Dict[str, Any] = { + metrics: dict[str, Any] = { "num_input_tokens": num_input_tokens, "num_output_tokens": num_output_tokens, "num_total_tokens": num_total_tokens, @@ -244,14 +244,14 @@ class GalileoObserve(CustomLogger): @staticmethod def _record_to_v2_span( - record: Dict[str, Any], + record: dict[str, Any], *, trace_id: str, span_id: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: created_at = GalileoObserve._normalize_created_at(record.get("created_at", "")) - span: Dict[str, Any] = { + span: dict[str, Any] = { "type": "llm", "id": span_id, "trace_id": trace_id, @@ -275,7 +275,7 @@ class GalileoObserve(CustomLogger): return span @staticmethod - def _record_to_v2_trace(record: Dict[str, Any]) -> Dict[str, Any]: + def _record_to_v2_trace(record: dict[str, Any]) -> dict[str, Any]: trace_id = str(uuid.uuid4()) span_id = str(uuid.uuid4()) created_at = GalileoObserve._normalize_created_at(record.get("created_at", "")) @@ -295,8 +295,8 @@ class GalileoObserve(CustomLogger): "spans": [GalileoObserve._record_to_v2_span(record, trace_id=trace_id, span_id=span_id)], } - def _build_traces_payload(self, records: List[dict]) -> Dict[str, Any]: - payload: Dict[str, Any] = { + def _build_traces_payload(self, records: list[dict]) -> dict[str, Any]: + payload: dict[str, Any] = { "traces": [self._record_to_v2_trace(record) for record in records], "logging_method": "api_direct", "reliable": False, @@ -306,7 +306,7 @@ class GalileoObserve(CustomLogger): payload["log_stream_id"] = self.log_stream_id return payload - def _get_ingest_request(self) -> Optional[Tuple[str, Dict[str, Any]]]: + def _get_ingest_request(self) -> tuple[str, dict[str, Any]] | None: if not self.base_url or not self.project_id: return None @@ -330,10 +330,10 @@ class GalileoObserve(CustomLogger): ) @staticmethod - def _redact_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]: + def _redact_headers(headers: dict[str, str] | None) -> dict[str, str]: if not headers: return {} - redacted: Dict[str, str] = {} + redacted: dict[str, str] = {} for key, value in headers.items(): if key.lower() in {"authorization", "galileo-api-key"} and value: redacted[key] = f"{value[:8]}...{value[-4:]}" if len(value) > 12 else "***" @@ -355,8 +355,8 @@ class GalileoObserve(CustomLogger): ) @staticmethod - def _log_v2_payload_validation(payload: Dict[str, Any]) -> None: - missing_fields: List[str] = [] + def _log_v2_payload_validation(payload: dict[str, Any]) -> None: + missing_fields: list[str] = [] traces = payload.get("traces", []) if not traces: missing_fields.append("traces") @@ -384,7 +384,7 @@ class GalileoObserve(CustomLogger): missing_fields, ) - def _log_flush_payload(self, url: str, payload: Dict[str, Any]) -> None: + def _log_flush_payload(self, url: str, payload: dict[str, Any]) -> None: traces = payload.get("traces", []) verbose_logger.debug( "Galileo Logger flush URL: %s trace_count=%s", @@ -415,9 +415,9 @@ class GalileoObserve(CustomLogger): pass @staticmethod - def _build_prompt(kwargs: Dict[str, Any]) -> Dict[str, Any]: + def _build_prompt(kwargs: dict[str, Any]) -> dict[str, Any]: optional_params = kwargs.get("optional_params", {}) or {} - prompt: Dict[str, Any] = {"messages": kwargs.get("messages")} + prompt: dict[str, Any] = {"messages": kwargs.get("messages")} if optional_params.get("functions") is not None: prompt["functions"] = optional_params["functions"] if optional_params.get("tools") is not None: @@ -439,7 +439,7 @@ class GalileoObserve(CustomLogger): return json.dumps(value, default=_json_default) @staticmethod - def _prompt_to_input_text(prompt: Dict[str, Any]) -> str: + def _prompt_to_input_text(prompt: dict[str, Any]) -> str: messages = prompt.get("messages") if messages is not None: text = GalileoObserve._input_text_from_messages(messages) @@ -462,7 +462,7 @@ class GalileoObserve(CustomLogger): @staticmethod def _get_text_completion_content_for_galileo( response_obj: litellm.TextCompletionResponse, - ) -> Optional[str]: + ) -> str | None: if response_obj.choices and len(response_obj.choices) > 0: return response_obj.choices[0].text return None @@ -476,17 +476,17 @@ class GalileoObserve(CustomLogger): return None @staticmethod - def _langfuse_style_rerank_prompt(kwargs: Dict[str, Any]) -> Dict[str, Any]: + def _langfuse_style_rerank_prompt(kwargs: dict[str, Any]) -> dict[str, Any]: """Match Langfuse rerank input: prompt = {"messages": kwargs.get("messages")}.""" return {"messages": kwargs.get("messages")} def _get_galileo_input_output_content( self, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], response_obj: Any, level: str = "DEFAULT", - status_message: Optional[str] = None, - ) -> Tuple[str, str, Any]: + status_message: str | None = None, + ) -> tuple[str, str, Any]: """ Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest. @@ -582,7 +582,7 @@ class GalileoObserve(CustomLogger): return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or [] - def get_output_str_from_response(self, response_obj: Any, kwargs: Dict[str, Any]) -> str: + def get_output_str_from_response(self, response_obj: Any, kwargs: dict[str, Any]) -> str: _, output_text, _ = self._get_galileo_input_output_content(kwargs=kwargs, response_obj=response_obj) return output_text @@ -635,7 +635,7 @@ class GalileoObserve(CustomLogger): ) return - slo: Optional[Dict[str, Any]] = kwargs.get("standard_logging_object") + slo: dict[str, Any] | None = kwargs.get("standard_logging_object") if slo is None: verbose_logger.debug("Galileo Logger: no standard_logging_object in kwargs, skipping") return diff --git a/litellm/integrations/gcs_bucket/gcs_bucket.py b/litellm/integrations/gcs_bucket/gcs_bucket.py index c2e0ad64586..552e078cb60 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket.py @@ -3,11 +3,11 @@ import hashlib import json import os import time -from litellm._uuid import uuid from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.constants import LITELLM_ASYNCIO_QUEUE_MAXSIZE from litellm.integrations.additional_logging_utils import AdditionalLoggingUtils from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase @@ -26,7 +26,7 @@ else: class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): - def __init__(self, bucket_name: Optional[str] = None) -> None: + def __init__(self, bucket_name: str | None = None) -> None: from litellm.proxy.proxy_server import premium_user super().__init__(bucket_name=bucket_name) @@ -67,7 +67,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): kwargs, response_obj, ) - logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") # When queue is at maxsize, flush immediately to make room (no blocking, no data dropped) @@ -76,7 +76,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): await self.log_queue.put(GCSLogQueueItem(payload=logging_payload, kwargs=kwargs, response_obj=response_obj)) except Exception as e: - verbose_logger.exception(f"GCS Bucket logging error: {str(e)}") + verbose_logger.exception(f"GCS Bucket logging error: {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: @@ -86,7 +86,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): response_obj, ) - logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") # When queue is at maxsize, flush immediately to make room (no blocking, no data dropped) @@ -95,9 +95,9 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): await self.log_queue.put(GCSLogQueueItem(payload=logging_payload, kwargs=kwargs, response_obj=response_obj)) except Exception as e: - verbose_logger.exception(f"GCS Bucket logging error: {str(e)}") + verbose_logger.exception(f"GCS Bucket logging error: {e!s}") - def _drain_queue_batch(self) -> List[GCSLogQueueItem]: + def _drain_queue_batch(self) -> list[GCSLogQueueItem]: """ Drain items from the queue (non-blocking), respecting batch_size limit. @@ -106,7 +106,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): Returns: List of items to process, up to batch_size items """ - items_to_process: List[GCSLogQueueItem] = [] + items_to_process: list[GCSLogQueueItem] = [] while len(items_to_process) < self.batch_size: try: items_to_process.append(self.log_queue.get_nowait()) @@ -121,7 +121,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): """ return f"{date_str}/batch-{batch_id}.ndjson" - def _get_config_key(self, kwargs: Dict[str, Any]) -> str: + def _get_config_key(self, kwargs: dict[str, Any]) -> str: """ Extract a synchronous grouping key from kwargs to group items by GCS config. This allows us to batch items with the same bucket/credentials together. @@ -151,14 +151,14 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): hash_obj = hashlib.sha256(config_key.encode("utf-8")) return f"config-{hash_obj.hexdigest()[:8]}" - def _group_items_by_config(self, items: List[GCSLogQueueItem]) -> Dict[str, List[GCSLogQueueItem]]: + def _group_items_by_config(self, items: list[GCSLogQueueItem]) -> dict[str, list[GCSLogQueueItem]]: """ Group items by their GCS config (bucket + credentials). This ensures items with different configs are processed separately. Returns a dict mapping config_key -> list of items with that config. """ - grouped: Dict[str, List[GCSLogQueueItem]] = {} + grouped: dict[str, list[GCSLogQueueItem]] = {} for item in items: config_key = self._get_config_key(item["kwargs"]) if config_key not in grouped: @@ -166,7 +166,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): grouped[config_key].append(item) return grouped - def _combine_payloads_to_ndjson(self, items: List[GCSLogQueueItem]) -> str: + def _combine_payloads_to_ndjson(self, items: list[GCSLogQueueItem]) -> str: """ Combine multiple log payloads into newline-delimited JSON (NDJSON) format. Each line is a valid JSON object representing one log entry. @@ -178,7 +178,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): lines.append(json_line) return "\n".join(lines) - async def _send_grouped_batch(self, items: List[GCSLogQueueItem], config_key: str) -> Tuple[int, int]: + async def _send_grouped_batch(self, items: list[GCSLogQueueItem], config_key: str) -> tuple[int, int]: """ Send a batch of items that share the same GCS config. @@ -218,10 +218,10 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): except Exception as e: success_count = 0 error_count = len(items) - verbose_logger.exception(f"GCS Bucket error logging batch payload to GCS bucket: {str(e)}") + verbose_logger.exception(f"GCS Bucket error logging batch payload to GCS bucket: {e!s}") return (success_count, error_count) - async def _send_individual_logs(self, items: List[GCSLogQueueItem]) -> None: + async def _send_individual_logs(self, items: list[GCSLogQueueItem]) -> None: """ Send each log individually as separate GCS objects (legacy behavior). This is used when GCS_USE_BATCHED_LOGGING is disabled. @@ -255,7 +255,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): logging_payload=item["payload"], ) except Exception as e: - verbose_logger.exception(f"GCS Bucket error logging individual payload to GCS bucket: {str(e)}") + verbose_logger.exception(f"GCS Bucket error logging individual payload to GCS bucket: {e!s}") async def async_send_batch(self): """ @@ -279,7 +279,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): else: await self._send_individual_logs(items_to_process) - def _get_object_name(self, kwargs: Dict, logging_payload: StandardLoggingPayload, response_obj: Any) -> str: + def _get_object_name(self, kwargs: dict, logging_payload: StandardLoggingPayload, response_obj: Any) -> str: """ Get the object name to use for the current payload """ @@ -307,9 +307,9 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): async def get_request_response_payload( self, request_id: str, - start_time_utc: Optional[datetime], - end_time_utc: Optional[datetime], - ) -> Optional[dict]: + start_time_utc: datetime | None, + end_time_utc: datetime | None, + ) -> dict | None: """ Get the request and response payload for a given `request_id` Tries current day, next day, and previous day until it finds the payload @@ -336,7 +336,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): loaded_response = json.loads(response) return loaded_response except Exception as e: - verbose_logger.debug(f"Failed to fetch payload for date {date_str}: {str(e)}") + verbose_logger.debug(f"Failed to fetch payload for date {date_str}: {e!s}") continue return None diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_base.py b/litellm/integrations/gcs_bucket/gcs_bucket_base.py index 0eabf16cff9..58a099d78fa 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket_base.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket_base.py @@ -1,16 +1,14 @@ import json import os -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union - -from litellm.integrations.gcs_bucket.gcs_bucket_mock_client import ( - should_use_gcs_mock, - create_mock_gcs_client, - mock_vertex_auth_methods, -) - +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.integrations.gcs_bucket.gcs_bucket_mock_client import ( + create_mock_gcs_client, + mock_vertex_auth_methods, + should_use_gcs_mock, +) from litellm.litellm_core_utils.cloud_storage_security import ( encode_gcs_object_name_for_url, split_configured_cloud_bucket_name, @@ -30,7 +28,7 @@ IAM_AUTH_KEY = "IAM_AUTH" class GCSBucketBase(CustomBatchLogger): - def __init__(self, bucket_name: Optional[str] = None, **kwargs) -> None: + def __init__(self, bucket_name: str | None = None, **kwargs) -> None: self.is_mock_mode = should_use_gcs_mock() if self.is_mock_mode: @@ -40,16 +38,16 @@ class GCSBucketBase(CustomBatchLogger): self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) _path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT") _bucket_name = bucket_name or os.getenv("GCS_BUCKET_NAME") - self.path_service_account_json: Optional[str] = _path_service_account - self.BUCKET_NAME: Optional[str] = _bucket_name - self.vertex_instances: Dict[str, VertexBase] = {} + self.path_service_account_json: str | None = _path_service_account + self.BUCKET_NAME: str | None = _bucket_name + self.vertex_instances: dict[str, VertexBase] = {} super().__init__(**kwargs) async def construct_request_headers( self, - service_account_json: Optional[str], - vertex_instance: Optional[VertexBase] = None, - ) -> Dict[str, str]: + service_account_json: str | None, + vertex_instance: VertexBase | None = None, + ) -> dict[str, str]: from litellm import vertex_chat_completion if vertex_instance is None: @@ -80,7 +78,7 @@ class GCSBucketBase(CustomBatchLogger): return headers - def sync_construct_request_headers(self) -> Dict[str, str]: + def sync_construct_request_headers(self) -> dict[str, str]: """ Construct request headers for GCS API calls """ @@ -120,7 +118,7 @@ class GCSBucketBase(CustomBatchLogger): self, bucket_name: str, object_name: str, - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Handles when the user passes a bucket name with a folder postfix @@ -137,7 +135,7 @@ class GCSBucketBase(CustomBatchLogger): return bucket_name, object_name return bucket_name, object_name - async def get_gcs_logging_config(self, kwargs: Optional[Dict[str, Any]] = {}) -> GCSLoggingConfig: + async def get_gcs_logging_config(self, kwargs: dict[str, Any] | None = {}) -> GCSLoggingConfig: """ This function is used to get the GCS logging config for the GCS Bucket Logger. It checks if the dynamic parameters are provided in the kwargs and uses them to get the GCS logging config. @@ -146,20 +144,18 @@ class GCSBucketBase(CustomBatchLogger): if kwargs is None: kwargs = {} - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params", None ) bucket_name: str - path_service_account: Optional[str] + path_service_account: str | None if standard_callback_dynamic_params is not None: verbose_logger.debug("Using dynamic GCS logging") verbose_logger.debug("standard_callback_dynamic_params: %s", standard_callback_dynamic_params) - _bucket_name: Optional[str] = ( - standard_callback_dynamic_params.get("gcs_bucket_name", None) or self.BUCKET_NAME - ) - _path_service_account: Optional[str] = ( + _bucket_name: str | None = standard_callback_dynamic_params.get("gcs_bucket_name", None) or self.BUCKET_NAME + _path_service_account: str | None = ( standard_callback_dynamic_params.get("gcs_path_service_account", None) or self.path_service_account_json ) @@ -186,7 +182,7 @@ class GCSBucketBase(CustomBatchLogger): path_service_account=path_service_account, ) - async def get_or_create_vertex_instance(self, credentials: Optional[str]) -> VertexBase: + async def get_or_create_vertex_instance(self, credentials: str | None) -> VertexBase: """ This function is used to get the Vertex instance for the GCS Bucket Logger. It checks if the Vertex instance is already created and cached, if not it creates a new instance and caches it. @@ -204,7 +200,7 @@ class GCSBucketBase(CustomBatchLogger): self.vertex_instances[_in_memory_key] = vertex_instance return self.vertex_instances[_in_memory_key] - def _get_in_memory_key_for_vertex_instance(self, credentials: Optional[str]) -> str: + def _get_in_memory_key_for_vertex_instance(self, credentials: str | None) -> str: """ Returns key to use for caching the Vertex instance in-memory. @@ -297,10 +293,10 @@ class GCSBucketBase(CustomBatchLogger): async def _log_json_data_on_gcs( self, - headers: Dict[str, str], + headers: dict[str, str], bucket_name: str, object_name: str, - logging_payload: Union[StandardLoggingPayload, str], + logging_payload: StandardLoggingPayload | str, ): """ Helper function to make POST request to GCS Bucket in the specified bucket. diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py b/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py index fae7ddaf536..86cf8617dd5 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py @@ -13,8 +13,8 @@ import asyncio from litellm._logging import verbose_logger from litellm.integrations.mock_client_factory import ( MockClientConfig, - create_mock_client_factory, MockResponse, + create_mock_client_factory, ) # Use factory for POST handler diff --git a/litellm/integrations/gcs_pubsub/pub_sub.py b/litellm/integrations/gcs_pubsub/pub_sub.py index c1bccb0b390..6ade70ab6d6 100644 --- a/litellm/integrations/gcs_pubsub/pub_sub.py +++ b/litellm/integrations/gcs_pubsub/pub_sub.py @@ -10,7 +10,7 @@ import asyncio import json import os import traceback -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any from litellm.types.utils import StandardLoggingPayload @@ -31,9 +31,9 @@ from litellm.llms.custom_httpx.http_handler import ( class GcsPubSubLogger(CustomBatchLogger): def __init__( self, - project_id: Optional[str] = None, - topic_id: Optional[str] = None, - credentials_path: Optional[str] = None, + project_id: str | None = None, + topic_id: str | None = None, + credentials_path: str | None = None, **kwargs, ): """ @@ -60,9 +60,9 @@ class GcsPubSubLogger(CustomBatchLogger): self.flush_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) asyncio.create_task(self.periodic_flush()) - self.log_queue: List[Union[SpendLogsPayload, StandardLoggingPayload]] = [] + self.log_queue: list[SpendLogsPayload | StandardLoggingPayload] = [] - async def construct_request_headers(self) -> Dict[str, str]: + async def construct_request_headers(self) -> dict[str, str]: """Construct authorization headers using Vertex AI auth""" from litellm import vertex_chat_completion @@ -132,8 +132,7 @@ class GcsPubSubLogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"PubSub Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"PubSub Layer Error - {e!s}\n{traceback.format_exc()}") async def async_send_batch(self): """ @@ -149,13 +148,11 @@ class GcsPubSubLogger(CustomBatchLogger): await self.publish_message(message) except Exception as e: - verbose_logger.exception(f"PubSub Error sending batch - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"PubSub Error sending batch - {e!s}\n{traceback.format_exc()}") finally: self.log_queue.clear() - async def publish_message( - self, message: Union[SpendLogsPayload, StandardLoggingPayload] - ) -> Optional[Dict[str, Any]]: + async def publish_message(self, message: SpendLogsPayload | StandardLoggingPayload) -> dict[str, Any] | None: """ Publish message to Google Cloud Pub/Sub using REST API diff --git a/litellm/integrations/generic_api/generic_api_callback.py b/litellm/integrations/generic_api/generic_api_callback.py index da6009c3a94..a524755540e 100644 --- a/litellm/integrations/generic_api/generic_api_callback.py +++ b/litellm/integrations/generic_api/generic_api_callback.py @@ -11,9 +11,10 @@ import json import os import re import traceback -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal import httpx + import litellm from litellm._logging import verbose_logger from litellm._uuid import uuid @@ -29,7 +30,7 @@ API_EVENT_TYPES = Literal["llm_api_success", "llm_api_failure"] LOG_FORMAT_TYPES = Literal["json_array", "ndjson", "single"] -def load_compatible_callbacks() -> Dict: +def load_compatible_callbacks() -> dict: """ Load the generic_api_compatible_callbacks.json file @@ -41,7 +42,7 @@ def load_compatible_callbacks() -> Dict: with open(json_path, "r") as f: return json.load(f) except Exception as e: - verbose_logger.warning(f"Error loading generic_api_compatible_callbacks.json: {str(e)}") + verbose_logger.warning(f"Error loading generic_api_compatible_callbacks.json: {e!s}") return {} @@ -59,7 +60,7 @@ def is_callback_compatible(callback_name: str) -> bool: return callback_name in compatible_callbacks -def get_callback_config(callback_name: str) -> Optional[Dict]: +def get_callback_config(callback_name: str) -> dict | None: """ Get the configuration for a specific callback @@ -95,14 +96,14 @@ def substitute_env_variables(value: str) -> str: class GenericAPILogger(CustomBatchLogger): def __init__( self, - endpoint: Optional[str] = None, - headers: Optional[dict] = None, - event_types: Optional[List[API_EVENT_TYPES]] = None, - callback_name: Optional[str] = None, - log_format: Optional[LOG_FORMAT_TYPES] = None, + endpoint: str | None = None, + headers: dict | None = None, + event_types: list[API_EVENT_TYPES] | None = None, + callback_name: str | None = None, + log_format: LOG_FORMAT_TYPES | None = None, max_retries: int = 0, retry_delay: float = 1.0, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ): """ @@ -157,10 +158,10 @@ class GenericAPILogger(CustomBatchLogger): "endpoint not set for GenericAPILogger, GENERIC_LOGGER_ENDPOINT not found in environment variables" ) - self.headers: Dict = self._get_headers(headers) + self.headers: dict = self._get_headers(headers) self.endpoint: str = endpoint - self.event_types: Optional[List[API_EVENT_TYPES]] = event_types - self.callback_name: Optional[str] = callback_name + self.event_types: list[API_EVENT_TYPES] | None = event_types + self.callback_name: str | None = callback_name self.max_retries = max(0, int(max_retries or 0)) retry_delay_value = 0.0 if retry_delay is None else retry_delay self.retry_delay = max(0.0, float(retry_delay_value)) @@ -185,9 +186,9 @@ class GenericAPILogger(CustomBatchLogger): self.flush_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) asyncio.create_task(self.periodic_flush()) - self.log_queue: List[Union[Dict, StandardLoggingPayload]] = [] + self.log_queue: list[dict | StandardLoggingPayload] = [] - def _get_headers(self, headers: Optional[dict] = None): + def _get_headers(self, headers: dict | None = None): """ Get headers for the Generic API Logger @@ -213,7 +214,7 @@ class GenericAPILogger(CustomBatchLogger): key, value = item.split("=", 1) headers_dict[key.strip()] = value.strip() except Exception as e: - verbose_logger.warning(f"Error parsing headers from environment variables: {str(e)}") + verbose_logger.warning(f"Error parsing headers from environment variables: {e!s}") # 2. Update with litellm generic headers if available if litellm.generic_logger_headers: @@ -242,7 +243,7 @@ class GenericAPILogger(CustomBatchLogger): await asyncio.sleep(delay) async def _post_with_retries(self, data: str) -> httpx.Response: - post_kwargs: Dict[str, Any] = { + post_kwargs: dict[str, Any] = { "url": self.endpoint, "headers": self.headers, "data": data, @@ -307,8 +308,7 @@ class GenericAPILogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"Generic API Logger Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Generic API Logger Error - {e!s}\n{traceback.format_exc()}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """ @@ -339,7 +339,7 @@ class GenericAPILogger(CustomBatchLogger): await self.async_send_batch() except Exception as e: - verbose_logger.exception(f"Generic API Logger Error - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"Generic API Logger Error - {e!s}\n{traceback.format_exc()}") async def async_send_batch(self): """ @@ -395,7 +395,7 @@ class GenericAPILogger(CustomBatchLogger): ) except Exception as e: - verbose_logger.exception(f"Generic API Logger Error sending batch - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"Generic API Logger Error sending batch - {e!s}\n{traceback.format_exc()}") finally: self.log_queue.clear() diff --git a/litellm/integrations/generic_prompt_management/__init__.py b/litellm/integrations/generic_prompt_management/__init__.py index 44c61aa5f50..1d2d6dfa70c 100644 --- a/litellm/integrations/generic_prompt_management/__init__.py +++ b/litellm/integrations/generic_prompt_management/__init__.py @@ -1,18 +1,19 @@ """Generic prompt management integration for LiteLLM.""" -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: - from .generic_prompt_manager import GenericPromptManager - from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.integrations.custom_prompt_management import CustomPromptManagement + from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec + + from .generic_prompt_manager import GenericPromptManager from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .generic_prompt_manager import GenericPromptManager # Global instances -global_generic_prompt_config: Optional[dict] = None +global_generic_prompt_config: dict | None = None def set_global_generic_prompt_config(config: dict) -> None: @@ -72,7 +73,7 @@ prompt_initializer_registry = { # Export public API __all__ = [ "GenericPromptManager", - "set_global_generic_prompt_config", "global_generic_prompt_config", "prompt_initializer_registry", + "set_global_generic_prompt_config", ] diff --git a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py index f9837efdde2..18cbe4e90a2 100644 --- a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py +++ b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py @@ -4,7 +4,7 @@ Fetches prompts from any API that implements the /beta/litellm_prompt_management """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -54,10 +54,10 @@ class GenericPromptManager(CustomPromptManagement): def __init__( self, api_base: str, - api_key: Optional[str] = None, + api_key: str | None = None, timeout: int = 30, - prompt_id: Optional[str] = None, - additional_provider_specific_query_params: Optional[Dict[str, Any]] = None, + prompt_id: str | None = None, + additional_provider_specific_query_params: dict[str, Any] | None = None, **kwargs, ): """ @@ -75,14 +75,14 @@ class GenericPromptManager(CustomPromptManagement): self.timeout = timeout self.prompt_id = prompt_id self.additional_provider_specific_query_params = additional_provider_specific_query_params - self._prompt_cache: Dict[str, PromptManagementClient] = {} + self._prompt_cache: dict[str, PromptManagementClient] = {} @property def integration_name(self) -> str: """Integration name used in model names like 'generic_prompt/gpt-4'.""" return "generic_prompt" - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """Get HTTP headers for API requests.""" headers = { "Content-Type": "application/json", @@ -92,7 +92,7 @@ class GenericPromptManager(CustomPromptManagement): headers["Authorization"] = f"Bearer {self.api_key}" return headers - def _fetch_prompt_from_api(self, prompt_id: Optional[str], prompt_spec: Optional[PromptSpec]) -> Dict[str, Any]: + def _fetch_prompt_from_api(self, prompt_id: str | None, prompt_spec: PromptSpec | None) -> dict[str, Any]: """ Fetch a prompt from the API. @@ -130,8 +130,8 @@ class GenericPromptManager(CustomPromptManagement): raise Exception(f"Failed to parse prompt response for '{prompt_id}': {e}") async def async_fetch_prompt_from_api( - self, prompt_id: Optional[str], prompt_spec: Optional[PromptSpec] - ) -> Dict[str, Any]: + self, prompt_id: str | None, prompt_spec: PromptSpec | None + ) -> dict[str, Any]: """ Fetch a prompt from the API asynchronously. """ @@ -167,9 +167,9 @@ class GenericPromptManager(CustomPromptManagement): def _parse_api_response( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - api_response: Dict[str, Any], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + api_response: dict[str, Any], ) -> PromptManagementClient: """ Parse the API response into a PromptManagementClient structure. @@ -205,8 +205,8 @@ class GenericPromptManager(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """ @@ -223,19 +223,19 @@ class GenericPromptManager(CustomPromptManagement): def _get_cache_key( self, - prompt_id: Optional[str], - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_id: str | None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> str: return f"{prompt_id}:{prompt_label}:{prompt_version}" def _common_caching_logic( self, - prompt_id: Optional[str], - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - prompt_variables: Optional[dict] = None, - ) -> Optional[PromptManagementClient]: + prompt_id: str | None, + prompt_label: str | None = None, + prompt_version: int | None = None, + prompt_variables: dict | None = None, + ) -> PromptManagementClient | None: """ Common caching logic for the prompt manager. """ @@ -251,12 +251,12 @@ class GenericPromptManager(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Compile a prompt template into a PromptManagementClient structure. @@ -307,12 +307,12 @@ class GenericPromptManager(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: # Check cache first cached_prompt = self._common_caching_logic( @@ -349,7 +349,7 @@ class GenericPromptManager(CustomPromptManagement): def _apply_variables( self, prompt_client: PromptManagementClient, - variables: Dict[str, Any], + variables: dict[str, Any], ) -> PromptManagementClient: """ Apply variables to the prompt template. @@ -364,7 +364,7 @@ class GenericPromptManager(CustomPromptManagement): Updated PromptManagementClient with variables applied """ # Create a copy of the prompt template with variables applied - updated_messages: List[AllMessageValues] = [] + updated_messages: list[AllMessageValues] = [] for message in prompt_client["prompt_template"]: updated_message = dict(message) # type: ignore if "content" in updated_message and isinstance(updated_message["content"], str): @@ -386,19 +386,19 @@ class GenericPromptManager(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: "LiteLLMLoggingObj", - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Get chat completion prompt and return processed model, messages, and parameters. """ @@ -432,17 +432,17 @@ class GenericPromptManager(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Get chat completion prompt and return processed model, messages, and parameters. """ diff --git a/litellm/integrations/gitlab/__init__.py b/litellm/integrations/gitlab/__init__.py index f06c28c5001..3c37e68db58 100644 --- a/litellm/integrations/gitlab/__init__.py +++ b/litellm/integrations/gitlab/__init__.py @@ -1,17 +1,18 @@ -from typing import TYPE_CHECKING, Optional, Dict, Any +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from .gitlab_prompt_manager import GitLabPromptManager - from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec from litellm.integrations.custom_prompt_management import CustomPromptManagement + from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec + + from .gitlab_prompt_manager import GitLabPromptManager -from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from litellm.integrations.custom_prompt_management import CustomPromptManagement -from litellm.types.prompts.init_prompts import PromptSpec, PromptLiteLLMParams -from .gitlab_prompt_manager import GitLabPromptManager, GitLabPromptCache +from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec, SupportedPromptIntegrations + +from .gitlab_prompt_manager import GitLabPromptCache, GitLabPromptManager # Global instances -global_gitlab_config: Optional[dict] = None +global_gitlab_config: dict | None = None def set_global_gitlab_config(config: dict) -> None: @@ -65,8 +66,8 @@ def _gitlab_prompt_initializer( # You can store arbitrary integration-specific config on PromptLiteLLMParams. # If your dataclass doesn't have these attributes, add them or put inside # `litellm_params.extra` and pull them from there. - gitlab_config: Dict[str, Any] = getattr(litellm_params, "gitlab_config", None) or {} - git_ref: Optional[str] = getattr(litellm_params, "git_ref", None) + gitlab_config: dict[str, Any] = getattr(litellm_params, "gitlab_config", None) or {} + git_ref: str | None = getattr(litellm_params, "git_ref", None) if not gitlab_config: raise ValueError("gitlab_config is required for gitlab prompt integration") @@ -85,8 +86,8 @@ prompt_initializer_registry = { # Export public API __all__ = [ - "GitLabPromptManager", "GitLabPromptCache", - "set_global_gitlab_config", + "GitLabPromptManager", "global_gitlab_config", + "set_global_gitlab_config", ] diff --git a/litellm/integrations/gitlab/gitlab_client.py b/litellm/integrations/gitlab/gitlab_client.py index ca366274ccd..3cd5198f3b6 100644 --- a/litellm/integrations/gitlab/gitlab_client.py +++ b/litellm/integrations/gitlab/gitlab_client.py @@ -4,7 +4,7 @@ Now supports selecting a tag via `config["tag"]`; falls back to branch ("main"). """ import base64 -from typing import Any, Dict, List, Optional +from typing import Any from urllib.parse import quote from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -22,7 +22,7 @@ class GitLabClient: - Directory listing via the repository tree API """ - def __init__(self, config: Dict[str, Any]): + def __init__(self, config: dict[str, Any]): """ Initialize the GitLab client. @@ -76,12 +76,12 @@ class GitLabClient: # Core helpers # ------------------------ - def _file_raw_url(self, file_path: str, *, ref: Optional[str] = None) -> str: + def _file_raw_url(self, file_path: str, *, ref: str | None = None) -> str: file_enc = quote(file_path, safe="") ref_q = quote(ref or self.ref, safe="") return f"{self.base_url}/projects/{self._project_enc}/repository/files/{file_enc}/raw?ref={ref_q}" - def _file_json_url(self, file_path: str, *, ref: Optional[str] = None) -> str: + def _file_json_url(self, file_path: str, *, ref: str | None = None) -> str: file_enc = quote(file_path, safe="") ref_q = quote(ref or self.ref, safe="") return f"{self.base_url}/projects/{self._project_enc}/repository/files/{file_enc}?ref={ref_q}" @@ -91,7 +91,7 @@ class GitLabClient: directory_path: str = "", recursive: bool = False, *, - ref: Optional[str] = None, + ref: str | None = None, ) -> str: path_q = f"&path={quote(directory_path, safe='')}" if directory_path else "" rec_q = "&recursive=true" if recursive else "" @@ -108,7 +108,7 @@ class GitLabClient: raise ValueError("ref must be a non-empty string") self.ref = ref - def get_file_content(self, file_path: str, *, ref: Optional[str] = None) -> Optional[str]: + def get_file_content(self, file_path: str, *, ref: str | None = None) -> str | None: """ Fetch the content of a file from the GitLab repository at the given ref (tag, branch, or commit SHA). If `ref` is None, uses self.ref. @@ -149,7 +149,7 @@ class GitLabClient: raise Exception("Authentication failed. Check your GitLab token and auth_method.") raise Exception(f"Failed to fetch file '{file_path}': {e}") - def _get_file_content_via_json(self, file_path: str, *, ref: Optional[str] = None) -> Optional[str]: + def _get_file_content_via_json(self, file_path: str, *, ref: str | None = None) -> str | None: """ Fallback for get_file_content(): use the JSON file API which returns base64 content. """ @@ -186,8 +186,8 @@ class GitLabClient: file_extension: str = ".prompt", recursive: bool = False, *, - ref: Optional[str] = None, - ) -> List[str]: + ref: str | None = None, + ) -> list[str]: """ List files in a directory with a specific extension using the repository tree API. @@ -209,7 +209,7 @@ class GitLabClient: resp.raise_for_status() data = resp.json() or [] - files: List[str] = [] + files: list[str] = [] for item in data: if item.get("type") == "blob": file_path = item.get("path", "") @@ -229,7 +229,7 @@ class GitLabClient: raise Exception("Authentication failed. Check your GitLab token and auth_method.") raise Exception(f"Failed to list files in '{directory_path}': {e}") - def get_repository_info(self) -> Dict[str, Any]: + def get_repository_info(self) -> dict[str, Any]: """Get information about the project/repository.""" url = f"{self.base_url}/projects/{self._project_enc}" try: @@ -247,7 +247,7 @@ class GitLabClient: except Exception: return False - def get_branches(self) -> List[Dict[str, Any]]: + def get_branches(self) -> list[dict[str, Any]]: """Get list of branches in the repository.""" url = f"{self.base_url}/projects/{self._project_enc}/repository/branches" try: @@ -258,7 +258,7 @@ class GitLabClient: except Exception as e: raise Exception(f"Failed to get branches: {e}") - def get_file_metadata(self, file_path: str, *, ref: Optional[str] = None) -> Optional[Dict[str, Any]]: + def get_file_metadata(self, file_path: str, *, ref: str | None = None) -> dict[str, Any] | None: """ Get minimal metadata about a file via RAW endpoint headers at a given ref. diff --git a/litellm/integrations/gitlab/gitlab_prompt_manager.py b/litellm/integrations/gitlab/gitlab_prompt_manager.py index 4896f95f398..54e0a3ad02e 100644 --- a/litellm/integrations/gitlab/gitlab_prompt_manager.py +++ b/litellm/integrations/gitlab/gitlab_prompt_manager.py @@ -2,7 +2,7 @@ GitLab prompt manager with configurable prompts folder. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from jinja2 import DictLoader, select_autoescape from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -44,8 +44,8 @@ class GitLabPromptTemplate: self, template_id: str, content: str, - metadata: Dict[str, Any], - model: Optional[str] = None, + metadata: dict[str, Any], + model: str | None = None, ): self.template_id = template_id self.content = content @@ -69,14 +69,14 @@ class GitLabTemplateManager: def __init__( self, - gitlab_config: Dict[str, Any], - prompt_id: Optional[str] = None, - ref: Optional[str] = None, - gitlab_client: Optional[GitLabClient] = None, + gitlab_config: dict[str, Any], + prompt_id: str | None = None, + ref: str | None = None, + gitlab_client: GitLabClient | None = None, ): self.gitlab_config = dict(gitlab_config) self.prompt_id = prompt_id - self.prompts: Dict[str, GitLabPromptTemplate] = {} + self.prompts: dict[str, GitLabPromptTemplate] = {} self.gitlab_client = gitlab_client or GitLabClient(self.gitlab_config) if ref: @@ -124,13 +124,12 @@ class GitLabTemplateManager: path = repo_path.strip("/") if self.prompts_path and path.startswith(self.prompts_path.strip("/") + "/"): path = path[len(self.prompts_path.strip("/")) + 1 :] - if path.endswith(".prompt"): - path = path[: -len(".prompt")] + path = path.removesuffix(".prompt") return encode_prompt_id(path) # ---------- loading ---------- - def _load_prompt_from_gitlab(self, prompt_id: str, *, ref: Optional[str] = None) -> None: + def _load_prompt_from_gitlab(self, prompt_id: str, *, ref: str | None = None) -> None: """Load a specific .prompt file from GitLab (scoped under prompts_path if set).""" try: # prompt_id = decode_prompt_id(prompt_id) @@ -142,12 +141,12 @@ class GitLabTemplateManager: except Exception as e: raise Exception(f"Failed to load prompt '{encode_prompt_id(prompt_id)}' from GitLab: {e}") - def load_all_prompts(self, *, recursive: bool = True) -> List[str]: + def load_all_prompts(self, *, recursive: bool = True) -> list[str]: """ Eagerly load all .prompt files from prompts_path. Returns loaded IDs. """ files = self.list_templates(recursive=recursive) - loaded: List[str] = [] + loaded: list[str] = [] for pid in files: if pid not in self.prompts: self._load_prompt_from_gitlab(pid) @@ -169,7 +168,7 @@ class GitLabTemplateManager: frontmatter_str = "" template_content = content - metadata: Dict[str, Any] = {} + metadata: dict[str, Any] = {} if frontmatter_str: try: import yaml @@ -186,8 +185,8 @@ class GitLabTemplateManager: metadata=metadata, ) - def _parse_yaml_basic(self, yaml_str: str) -> Dict[str, Any]: - result: Dict[str, Any] = {} + def _parse_yaml_basic(self, yaml_str: str) -> dict[str, Any]: + result: dict[str, Any] = {} for line in yaml_str.split("\n"): line = line.strip() if ":" in line and not line.startswith("#"): @@ -207,17 +206,17 @@ class GitLabTemplateManager: result[key] = value.strip("\"'") return result - def render_template(self, template_id: str, variables: Optional[Dict[str, Any]] = None) -> str: + def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str: if template_id not in self.prompts: raise ValueError(f"Template '{template_id}' not found") template = self.prompts[template_id] jinja_template = self.jinja_env.from_string(template.content) return jinja_template.render(**(variables or {})) - def get_template(self, template_id: str) -> Optional[GitLabPromptTemplate]: + def get_template(self, template_id: str) -> GitLabPromptTemplate | None: return self.prompts.get(template_id) - def list_templates(self, *, recursive: bool = True) -> List[str]: + def list_templates(self, *, recursive: bool = True) -> list[str]: """ List available prompt IDs under prompts_path (no extension). Compatible with both list_files signatures: @@ -232,7 +231,7 @@ class GitLabTemplateManager: recursive=recursive, ) base = self.prompts_path.strip("/") - out: List[str] = [] + out: list[str] = [] for p in files or []: path = str(p).strip("/") if base and not path.startswith(base + "/"): @@ -279,14 +278,14 @@ class GitLabPromptManager(CustomPromptManagement): def __init__( self, - gitlab_config: Dict[str, Any], - prompt_id: Optional[str] = None, - ref: Optional[str] = None, # tag/branch/SHA override - gitlab_client: Optional[GitLabClient] = None, + gitlab_config: dict[str, Any], + prompt_id: str | None = None, + ref: str | None = None, # tag/branch/SHA override + gitlab_client: GitLabClient | None = None, ): self.gitlab_config = gitlab_config self.prompt_id = prompt_id - self._prompt_manager: Optional[GitLabTemplateManager] = None + self._prompt_manager: GitLabTemplateManager | None = None self._ref_override = ref self._injected_gitlab_client = gitlab_client if self.prompt_id: @@ -314,10 +313,10 @@ class GitLabPromptManager(CustomPromptManagement): def get_prompt_template( self, prompt_id: str, - prompt_variables: Optional[Dict[str, Any]] = None, + prompt_variables: dict[str, Any] | None = None, *, - ref: Optional[str] = None, - ) -> Tuple[str, Dict[str, Any]]: + ref: str | None = None, + ) -> tuple[str, dict[str, Any]]: if prompt_id not in self.prompt_manager.prompts: self.prompt_manager._load_prompt_from_gitlab(prompt_id, ref=ref) @@ -337,15 +336,15 @@ class GitLabPromptManager(CustomPromptManagement): def pre_call_hook( self, - user_id: Optional[str], - messages: List[AllMessageValues], - function_call: Optional[Union[Dict[str, Any], str]] = None, - litellm_params: Optional[Dict[str, Any]] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[Dict[str, Any]] = None, - prompt_version: Optional[str] = None, + user_id: str | None, + messages: list[AllMessageValues], + function_call: dict[str, Any] | str | None = None, + litellm_params: dict[str, Any] | None = None, + prompt_id: str | None = None, + prompt_variables: dict[str, Any] | None = None, + prompt_version: str | None = None, **kwargs, - ) -> Tuple[List[AllMessageValues], Optional[Dict[str, Any]]]: + ) -> tuple[list[AllMessageValues], dict[str, Any] | None]: if not prompt_id: return messages, litellm_params try: @@ -356,7 +355,7 @@ class GitLabPromptManager(CustomPromptManagement): parsed_messages = self._parse_prompt_to_messages(rendered_prompt) if parsed_messages: - final_messages: List[AllMessageValues] = parsed_messages + final_messages: list[AllMessageValues] = parsed_messages else: final_messages = [{"role": "user", "content": rendered_prompt}] + messages # type: ignore @@ -383,11 +382,11 @@ class GitLabPromptManager(CustomPromptManagement): litellm._logging.verbose_proxy_logger.error(f"Error in GitLab prompt pre_call_hook: {e}") return messages, litellm_params - def _parse_prompt_to_messages(self, prompt_content: str) -> List[AllMessageValues]: - messages: List[AllMessageValues] = [] + def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]: + messages: list[AllMessageValues] = [] lines = prompt_content.strip().split("\n") - current_role: Optional[str] = None - current_content: List[str] = [] + current_role: str | None = None + current_content: list[str] = [] for raw in lines: line = raw.strip() @@ -435,18 +434,18 @@ class GitLabPromptManager(CustomPromptManagement): def post_call_hook( self, - user_id: Optional[str], + user_id: str | None, response: Any, - input_messages: List[AllMessageValues], - function_call: Optional[Union[Dict[str, Any], str]] = None, - litellm_params: Optional[Dict[str, Any]] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[Dict[str, Any]] = None, + input_messages: list[AllMessageValues], + function_call: dict[str, Any] | str | None = None, + litellm_params: dict[str, Any] | None = None, + prompt_id: str | None = None, + prompt_variables: dict[str, Any] | None = None, **kwargs, ) -> Any: return response - def get_available_prompts(self) -> List[str]: + def get_available_prompts(self) -> list[str]: """ Return prompt IDs. Prefer already-loaded templates in memory to avoid unnecessary network calls (and to make tests deterministic). @@ -466,20 +465,20 @@ class GitLabPromptManager(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: return prompt_id is not None def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: if prompt_id is None: raise ValueError("prompt_id is required for GitLab prompt manager") @@ -499,7 +498,7 @@ class GitLabPromptManager(CustomPromptManagement): messages = self._parse_prompt_to_messages(rendered_prompt) template_model = prompt_metadata.get("model") - optional_params: Dict[str, Any] = {} + optional_params: dict[str, Any] = {} for param in [ "temperature", "max_tokens", @@ -522,12 +521,12 @@ class GitLabPromptManager(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """ Async version of compile prompt helper. Since GitLab operations use sync client, @@ -548,17 +547,17 @@ class GitLabPromptManager(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: return PromptManagementBase.get_chat_completion_prompt( self, model, @@ -575,19 +574,19 @@ class GitLabPromptManager(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Async version - delegates to PromptManagementBase async implementation. """ @@ -644,10 +643,10 @@ class GitLabPromptCache: def __init__( self, - gitlab_config: Dict[str, Any], + gitlab_config: dict[str, Any], *, - ref: Optional[str] = None, - gitlab_client: Optional[GitLabClient] = None, + ref: str | None = None, + gitlab_client: GitLabClient | None = None, ) -> None: # Build a PromptManager (which internally builds TemplateManager + Client) self.prompt_manager = GitLabPromptManager( @@ -659,14 +658,14 @@ class GitLabPromptCache: self.template_manager: GitLabTemplateManager = self.prompt_manager.prompt_manager # In-memory stores - self._by_file: Dict[str, Dict[str, Any]] = {} - self._by_id: Dict[str, Dict[str, Any]] = {} + self._by_file: dict[str, dict[str, Any]] = {} + self._by_id: dict[str, dict[str, Any]] = {} # ------------------------- # Public API # ------------------------- - def load_all(self, *, recursive: bool = True) -> Dict[str, Dict[str, Any]]: + def load_all(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]: """ Scan GitLab for all .prompt files under prompts_path, load and parse each, and return the mapping of repo file path -> JSON-like dict. @@ -696,25 +695,25 @@ class GitLabPromptCache: return self._by_id - def reload(self, *, recursive: bool = True) -> Dict[str, Dict[str, Any]]: + def reload(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]: """Clear the cache and re-load from GitLab.""" self._by_file.clear() self._by_id.clear() return self.load_all(recursive=recursive) - def list_files(self) -> List[str]: + def list_files(self) -> list[str]: """Return the repo file paths currently cached.""" return list(self._by_file.keys()) - def list_ids(self) -> List[str]: + def list_ids(self) -> list[str]: """Return the template IDs (relative to prompts_path, without extension) currently cached.""" return list(self._by_id.keys()) - def get_by_file(self, file_path: str) -> Optional[Dict[str, Any]]: + def get_by_file(self, file_path: str) -> dict[str, Any] | None: """Get a cached prompt JSON by repo file path.""" return self._by_file.get(file_path) - def get_by_id(self, prompt_id: str) -> Optional[Dict[str, Any]]: + def get_by_id(self, prompt_id: str) -> dict[str, Any] | None: """Get a cached prompt JSON by prompt ID (relative to prompts_path).""" if prompt_id in self._by_id: return self._by_id[prompt_id] @@ -729,7 +728,7 @@ class GitLabPromptCache: # Internals # ------------------------- - def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> Dict[str, Any]: + def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> dict[str, Any]: """ Normalize a GitLabPromptTemplate into a JSON-like dict that is easy to serialize. """ diff --git a/litellm/integrations/greenscale.py b/litellm/integrations/greenscale.py index e2aca361010..08beaa79a51 100644 --- a/litellm/integrations/greenscale.py +++ b/litellm/integrations/greenscale.py @@ -57,4 +57,3 @@ class GreenscaleLogger: print_verbose(f"Greenscale Logger Succeeded - {response.text}") except Exception as e: print_verbose(f"Greenscale Logger Error - {e}, Stack trace: {traceback.format_exc()}") - pass diff --git a/litellm/integrations/helicone.py b/litellm/integrations/helicone.py index 21e9479491e..e67ab9fa93b 100644 --- a/litellm/integrations/helicone.py +++ b/litellm/integrations/helicone.py @@ -6,8 +6,8 @@ import traceback import litellm from litellm._logging import verbose_logger from litellm.integrations.helicone_mock_client import ( - should_use_helicone_mock, create_mock_helicone_client, + should_use_helicone_mock, ) @@ -36,8 +36,7 @@ class HeliconeLogger: self.provider_url = "https://api.openai.com/v1" self.key = os.getenv("HELICONE_API_KEY") self.api_base = os.getenv("HELICONE_API_BASE") or "https://api.hconeai.com" - if self.api_base.endswith("/"): - self.api_base = self.api_base[:-1] + self.api_base = self.api_base.removesuffix("/") def claude_mapping(self, model, messages, response_obj): from anthropic import AI_PROMPT, HUMAN_PROMPT @@ -201,4 +200,3 @@ class HeliconeLogger: print_verbose(f"Helicone Logging - Error {response.text}") except Exception: print_verbose(f"Helicone Logging Error - {traceback.format_exc()}") - pass diff --git a/litellm/integrations/humanloop.py b/litellm/integrations/humanloop.py index 2a5cb70baee..57f99d0bbf6 100644 --- a/litellm/integrations/humanloop.py +++ b/litellm/integrations/humanloop.py @@ -4,7 +4,7 @@ Humanloop integration https://humanloop.com/ """ -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, cast import httpx from typing_extensions import TypedDict @@ -22,9 +22,9 @@ from .custom_logger import CustomLogger class PromptManagementClient(TypedDict): prompt_id: str - prompt_template: List[AllMessageValues] - model: Optional[str] - optional_params: Optional[Dict[str, Any]] + prompt_template: list[AllMessageValues] + model: str | None + optional_params: dict[str, Any] | None class HumanLoopPromptManager(DualCache): @@ -32,12 +32,12 @@ class HumanLoopPromptManager(DualCache): def integration_name(self): return "humanloop" - def _get_prompt_from_id_cache(self, humanloop_prompt_id: str) -> Optional[PromptManagementClient]: - return cast(Optional[PromptManagementClient], self.get_cache(key=humanloop_prompt_id)) + def _get_prompt_from_id_cache(self, humanloop_prompt_id: str) -> PromptManagementClient | None: + return cast(PromptManagementClient | None, self.get_cache(key=humanloop_prompt_id)) def _compile_prompt_helper( - self, prompt_template: List[AllMessageValues], prompt_variables: Dict[str, Any] - ) -> List[AllMessageValues]: + self, prompt_template: list[AllMessageValues], prompt_variables: dict[str, Any] + ) -> list[AllMessageValues]: """ Helper function to compile the prompt by substituting variables in the template. @@ -48,7 +48,7 @@ class HumanLoopPromptManager(DualCache): Returns: list: A list of dictionaries with variables substituted. """ - compiled_prompts: List[AllMessageValues] = [] + compiled_prompts: list[AllMessageValues] = [] for template in prompt_template: tc = template.get("content") @@ -63,7 +63,7 @@ class HumanLoopPromptManager(DualCache): def _get_prompt_from_id_api(self, humanloop_prompt_id: str, humanloop_api_key: str) -> PromptManagementClient: client = _get_httpx_client() - base_url = "https://api.humanloop.com/v5/prompts/{}".format(humanloop_prompt_id) + base_url = f"https://api.humanloop.com/v5/prompts/{humanloop_prompt_id}" response = client.get( url=base_url, @@ -93,7 +93,7 @@ class HumanLoopPromptManager(DualCache): optional_params[k] = v return PromptManagementClient( prompt_id=humanloop_prompt_id, - prompt_template=cast(List[AllMessageValues], template_messages), + prompt_template=cast(list[AllMessageValues], template_messages), model=template_model, optional_params=optional_params, ) @@ -111,10 +111,10 @@ class HumanLoopPromptManager(DualCache): def compile_prompt( self, - prompt_template: List[AllMessageValues], - prompt_variables: Optional[dict], - ) -> List[AllMessageValues]: - compiled_prompt: Optional[Union[str, list]] = None + prompt_template: list[AllMessageValues], + prompt_variables: dict | None, + ) -> list[AllMessageValues]: + compiled_prompt: str | list | None = None if prompt_variables is None: prompt_variables = {} @@ -130,7 +130,7 @@ class HumanLoopPromptManager(DualCache): if prompt_management_client["model"] is not None: return prompt_management_client["model"] else: - return model.replace("{}/".format(self.integration_name), "") + return model.replace(f"{self.integration_name}/", "") prompt_manager = HumanLoopPromptManager() @@ -140,19 +140,19 @@ class HumanloopLogger(CustomLogger): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[ + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[ str, - List[AllMessageValues], + list[AllMessageValues], dict, ]: humanloop_api_key = dynamic_callback_params.get("humanloop_api_key") or get_secret_str("HUMANLOOP_API_KEY") diff --git a/litellm/integrations/lago.py b/litellm/integrations/lago.py index 0052e04644d..3186f1bf58b 100644 --- a/litellm/integrations/lago.py +++ b/litellm/integrations/lago.py @@ -3,13 +3,13 @@ import json import os -from litellm._uuid import uuid -from typing import Literal, Optional +from typing import Literal import httpx import litellm from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.integrations.custom_logger import CustomLogger from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, @@ -58,7 +58,7 @@ class LagoLogger(CustomLogger): missing_keys.append("LAGO_API_EVENT_CODE") if len(missing_keys) > 0: - raise Exception("Missing keys={} in environment.".format(missing_keys)) + raise Exception(f"Missing keys={missing_keys} in environment.") def _common_logic(self, kwargs: dict, response_obj) -> dict: response_obj.get("id", kwargs.get("litellm_call_id")) @@ -84,7 +84,7 @@ class LagoLogger(CustomLogger): litellm_params["metadata"].get("user_api_key_org_id", None) charge_by: Literal["end_user_id", "team_id", "user_id"] = "end_user_id" - external_customer_id: Optional[str] = None + external_customer_id: str | None = None if os.getenv("LAGO_API_CHARGE_BY", None) is not None and isinstance(os.environ["LAGO_API_CHARGE_BY"], str): if os.environ["LAGO_API_CHARGE_BY"] in [ @@ -105,9 +105,7 @@ class LagoLogger(CustomLogger): if external_customer_id is None: raise Exception( - "External Customer ID is not set. Charge_by={}. User_id={}. End_user_id={}. Team_id={}".format( - charge_by, user_id, end_user_id, team_id - ) + f"External Customer ID is not set. Charge_by={charge_by}. User_id={user_id}. End_user_id={end_user_id}. Team_id={team_id}" ) returned_val = { @@ -119,13 +117,13 @@ class LagoLogger(CustomLogger): } } - verbose_logger.debug("\033[91mLogged Lago Object:\n{}\033[0m\n".format(returned_val)) + verbose_logger.debug(f"\033[91mLogged Lago Object:\n{returned_val}\033[0m\n") return returned_val def log_success_event(self, kwargs, response_obj, start_time, end_time): _url = os.getenv("LAGO_API_BASE") assert _url is not None and isinstance(_url, str), ( - "LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={}".format(_url) + f"LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={_url}" ) if _url.endswith("/"): _url += "api/v1/events" @@ -137,7 +135,7 @@ class LagoLogger(CustomLogger): _data = self._common_logic(kwargs=kwargs, response_obj=response_obj) _headers = { "Content-Type": "application/json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } try: @@ -159,7 +157,7 @@ class LagoLogger(CustomLogger): verbose_logger.debug("ENTERS LAGO CALLBACK") _url = os.getenv("LAGO_API_BASE") assert _url is not None and isinstance(_url, str), ( - "LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={}".format(_url) + f"LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={_url}" ) if _url.endswith("/"): _url += "api/v1/events" @@ -171,12 +169,12 @@ class LagoLogger(CustomLogger): _data = self._common_logic(kwargs=kwargs, response_obj=response_obj) _headers = { "Content-Type": "application/json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } except Exception as e: raise e - response: Optional[httpx.Response] = None + response: httpx.Response | None = None try: response = await self.async_http_handler.post( url=_url, diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 86da03a078b..3fb50e07b01 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -7,11 +7,6 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Union, cast, ) @@ -241,23 +236,21 @@ class LangFuseLogger: def log_event_on_langfuse( self, kwargs: dict, - response_obj: Union[ - None, - dict, - EmbeddingResponse, - ModelResponse, - TextCompletionResponse, - ImageResponse, - TranscriptionResponse, - RerankResponse, - HttpxBinaryResponseContent, - ResponsesAPIResponse, - ], - start_time: Optional[datetime] = None, - end_time: Optional[datetime] = None, - user_id: Optional[str] = None, + response_obj: None + | dict + | EmbeddingResponse + | ModelResponse + | TextCompletionResponse + | ImageResponse + | TranscriptionResponse + | RerankResponse + | HttpxBinaryResponseContent + | ResponsesAPIResponse, + start_time: datetime | None = None, + end_time: datetime | None = None, + user_id: str | None = None, level: str = "DEFAULT", - status_message: Optional[str] = None, + status_message: str | None = None, ) -> dict: """ Logs a success or error event on Langfuse @@ -337,28 +330,26 @@ class LangFuseLogger: return {"trace_id": trace_id, "generation_id": generation_id} except Exception as e: - verbose_logger.exception("Langfuse Layer Error(): Exception occured - {}".format(str(e))) + verbose_logger.exception(f"Langfuse Layer Error(): Exception occured - {e!s}") return {"trace_id": None, "generation_id": None} def _get_langfuse_input_output_content( self, kwargs: dict, - response_obj: Union[ - None, - dict, - EmbeddingResponse, - ModelResponse, - TextCompletionResponse, - ImageResponse, - TranscriptionResponse, - RerankResponse, - HttpxBinaryResponseContent, - ResponsesAPIResponse, - ], + response_obj: None + | dict + | EmbeddingResponse + | ModelResponse + | TextCompletionResponse + | ImageResponse + | TranscriptionResponse + | RerankResponse + | HttpxBinaryResponseContent + | ResponsesAPIResponse, prompt: dict, level: str, - status_message: Optional[str], - ) -> Tuple[Optional[dict], Optional[Union[str, dict, list]]]: + status_message: str | None, + ) -> tuple[dict | None, str | dict | list | None]: """ Get the input and output content for Langfuse logging @@ -374,7 +365,7 @@ class LangFuseLogger: output: The output content for Langfuse logging """ input = None - output: Optional[Union[str, dict, List[Any]]] = None + output: str | dict | list[Any] | None = None if level == "ERROR" and status_message is not None and isinstance(status_message, str): input = prompt output = status_message @@ -461,7 +452,7 @@ class LangFuseLogger: ) ) - custom_llm_provider = cast(Optional[str], kwargs.get("custom_llm_provider")) + custom_llm_provider = cast(str | None, kwargs.get("custom_llm_provider")) model_name = reconstruct_model_name(kwargs.get("model", ""), custom_llm_provider, metadata) trace.generation( @@ -483,24 +474,24 @@ class LangFuseLogger: def _log_langfuse_v2( self, - user_id: Optional[str], + user_id: str | None, metadata: dict, litellm_params: dict, - output: Optional[Union[str, dict, list]], - start_time: Optional[datetime], - end_time: Optional[datetime], + output: str | dict | list | None, + start_time: datetime | None, + end_time: datetime | None, kwargs: dict, optional_params: dict, - input: Optional[dict], + input: dict | None, response_obj, level: str, - litellm_call_id: Optional[str], + litellm_call_id: str | None, ) -> tuple: verbose_logger.debug("Langfuse Layer Logging - logging to langfuse v2") try: - standard_logging_object: Optional[StandardLoggingPayload] = cast( - Optional[StandardLoggingPayload], + standard_logging_object: StandardLoggingPayload | None = cast( + StandardLoggingPayload | None, kwargs.get("standard_logging_object", None), ) tags = ( @@ -511,19 +502,19 @@ class LangFuseLogger: if standard_logging_object is None: end_user_id = None - prompt_management_metadata: Optional[StandardLoggingPromptManagementMetadata] = None + prompt_management_metadata: StandardLoggingPromptManagementMetadata | None = None else: end_user_id = standard_logging_object["metadata"].get("user_api_key_end_user_id", None) prompt_management_metadata = cast( - Optional[StandardLoggingPromptManagementMetadata], + StandardLoggingPromptManagementMetadata | None, standard_logging_object["metadata"].get("prompt_management_metadata", None), ) # Clean Metadata before logging - never log raw metadata # the raw metadata can contain circular references which leads to infinite recursion # we clean out all extra litellm metadata params before logging - clean_metadata: Dict[str, Any] = {} + clean_metadata: dict[str, Any] = {} if prompt_management_metadata is not None: clean_metadata["prompt_management_metadata"] = prompt_management_metadata if isinstance(metadata, dict): @@ -551,12 +542,12 @@ class LangFuseLogger: tags = self.add_default_langfuse_tags(tags=tags, kwargs=kwargs, metadata=metadata) session_id = clean_metadata.pop("session_id", None) - trace_name = cast(Optional[str], clean_metadata.pop("trace_name", None)) + trace_name = cast(str | None, clean_metadata.pop("trace_name", None)) trace_id = clean_metadata.pop("trace_id", None) # Use standard_logging_object.trace_id if available (when trace_id from metadata is None) # This allows standard trace_id to be used when provided in standard_logging_object if trace_id is None and standard_logging_object is not None: - trace_id = cast(Optional[str], standard_logging_object.get("trace_id")) + trace_id = cast(str | None, standard_logging_object.get("trace_id")) # Fallback to litellm_call_id if no trace_id found if trace_id is None: trace_id = kwargs.get("litellm_trace_id") or litellm_call_id @@ -588,7 +579,7 @@ class LangFuseLogger: trace_name = f"litellm-{kwargs.get('call_type', 'completion')}" if existing_trace_id is not None: - trace_params: Dict[str, Any] = {"id": existing_trace_id} + trace_params: dict[str, Any] = {"id": existing_trace_id} # Update the following keys for this trace for metadata_param_key in update_trace_keys: @@ -731,7 +722,7 @@ class LangFuseLogger: # if `generation_name` is None, use sensible default values # If using litellm proxy user `key_alias` if not None # If `key_alias` is None, just log `litellm-{call_type}` as the generation name - _user_api_key_alias = cast(Optional[str], clean_metadata.get("user_api_key_alias", None)) + _user_api_key_alias = cast(str | None, clean_metadata.get("user_api_key_alias", None)) generation_name = f"litellm-{cast(str, kwargs.get('call_type', 'completion'))}" if _user_api_key_alias is not None: generation_name = f"litellm:{_user_api_key_alias}" @@ -744,7 +735,7 @@ class LangFuseLogger: if system_fingerprint is not None: optional_params["system_fingerprint"] = system_fingerprint - custom_llm_provider = cast(Optional[str], kwargs.get("custom_llm_provider")) + custom_llm_provider = cast(str | None, kwargs.get("custom_llm_provider")) model_name = reconstruct_model_name(kwargs.get("model", ""), custom_llm_provider, metadata) generation_params = { @@ -837,8 +828,8 @@ class LangFuseLogger: @staticmethod def _get_langfuse_tags( - standard_logging_object: Optional[StandardLoggingPayload], - ) -> List[str]: + standard_logging_object: StandardLoggingPayload | None, + ) -> list[str]: if standard_logging_object is None: return [] return standard_logging_object.get("request_tags", []) or [] @@ -933,7 +924,7 @@ class LangFuseLogger: def _log_guardrail_information_as_span( self, trace: StatefulTraceClient, - standard_logging_object: Optional[StandardLoggingPayload], + standard_logging_object: StandardLoggingPayload | None, ): """ Log guardrail information as a span @@ -982,7 +973,7 @@ class LangFuseLogger: def _add_prompt_to_generation_params( generation_params: dict, clean_metadata: dict, - prompt_management_metadata: Optional[StandardLoggingPromptManagementMetadata], + prompt_management_metadata: StandardLoggingPromptManagementMetadata | None, langfuse_client: Any, ) -> dict: from langfuse import Langfuse @@ -1045,7 +1036,6 @@ def _add_prompt_to_generation_params( generation_params["prompt"] = langfuse_client.get_prompt(prompt_management_metadata["prompt_id"]) except Exception as e: verbose_logger.debug(f"[Non-blocking] Langfuse Logger: Error getting prompt client for logging: {e}") - pass else: generation_params["prompt"] = user_prompt diff --git a/litellm/integrations/langfuse/langfuse_handler.py b/litellm/integrations/langfuse/langfuse_handler.py index b1d083bd7d4..507a46f4948 100644 --- a/litellm/integrations/langfuse/langfuse_handler.py +++ b/litellm/integrations/langfuse/langfuse_handler.py @@ -6,7 +6,7 @@ Used to get the LangFuseLogger for a given request Handles Key/Team Based Langfuse Logging """ -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams @@ -23,7 +23,7 @@ class LangFuseHandler: def get_langfuse_logger_for_request( standard_callback_dynamic_params: StandardCallbackDynamicParams, in_memory_dynamic_logger_cache: DynamicLoggingCache, - globalLangfuseLogger: Optional[LangFuseLogger] = None, + globalLangfuseLogger: LangFuseLogger | None = None, ) -> LangFuseLogger: """ This function is used to get the LangFuseLogger for a given request @@ -35,7 +35,7 @@ class LangFuseHandler: 2. If dynamic credentials are not passed return the globalLangfuseLogger """ - temp_langfuse_logger: Optional[LangFuseLogger] = globalLangfuseLogger + temp_langfuse_logger: LangFuseLogger | None = globalLangfuseLogger if LangFuseHandler._dynamic_langfuse_credentials_are_passed(standard_callback_dynamic_params) is False: return LangFuseHandler._return_global_langfuse_logger( globalLangfuseLogger=globalLangfuseLogger, @@ -65,7 +65,7 @@ class LangFuseHandler: @staticmethod def _return_global_langfuse_logger( - globalLangfuseLogger: Optional[LangFuseLogger], + globalLangfuseLogger: LangFuseLogger | None, in_memory_dynamic_logger_cache: DynamicLoggingCache, ) -> LangFuseLogger: """ @@ -79,7 +79,7 @@ class LangFuseHandler: if globalLangfuseLogger is not None: return globalLangfuseLogger - credentials_dict: Dict[ + credentials_dict: dict[ str, Any ] = {} # the global langfuse logger uses Environment Variables, there are no dynamic credentials globalLangfuseLogger = in_memory_dynamic_logger_cache.get_cache( @@ -95,7 +95,7 @@ class LangFuseHandler: @staticmethod def _create_langfuse_logger_from_credentials( - credentials: Dict, + credentials: dict, in_memory_dynamic_logger_cache: DynamicLoggingCache, ) -> LangFuseLogger: """ @@ -120,7 +120,7 @@ class LangFuseHandler: @staticmethod def get_dynamic_langfuse_logging_config( standard_callback_dynamic_params: StandardCallbackDynamicParams, - globalLangfuseLogger: Optional[LangFuseLogger] = None, + globalLangfuseLogger: LangFuseLogger | None = None, ) -> LangfuseLoggingConfig: """ This function is used to get the Langfuse logging config to use for a given request. diff --git a/litellm/integrations/langfuse/langfuse_mock_client.py b/litellm/integrations/langfuse/langfuse_mock_client.py index b7862274f62..0fcd899706e 100644 --- a/litellm/integrations/langfuse/langfuse_mock_client.py +++ b/litellm/integrations/langfuse/langfuse_mock_client.py @@ -9,6 +9,7 @@ Usage: """ import httpx + from litellm.integrations.mock_client_factory import ( MockClientConfig, create_mock_client_factory, diff --git a/litellm/integrations/langfuse/langfuse_otel.py b/litellm/integrations/langfuse/langfuse_otel.py index d464d55453d..143362f3468 100644 --- a/litellm/integrations/langfuse/langfuse_otel.py +++ b/litellm/integrations/langfuse/langfuse_otel.py @@ -2,7 +2,7 @@ import base64 import json import os from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Optional, Union from litellm._logging import verbose_logger from litellm.integrations.arize import _utils @@ -51,7 +51,6 @@ class LangfuseOtelLogger(OpenTelemetry): # Set Langfuse specific attributes ######################################################### LangfuseOtelLogger._set_langfuse_specific_attributes(span=span, kwargs=kwargs, response_obj=response_obj) - return @staticmethod def _extract_langfuse_metadata(kwargs: dict) -> dict: @@ -245,7 +244,7 @@ class LangfuseOtelLogger(OpenTelemetry): LangfuseOtelLogger._set_observation_output(span=span, response_obj=response_obj) @staticmethod - def _get_langfuse_otel_host() -> Optional[str]: + def _get_langfuse_otel_host() -> str | None: """ Returns the Langfuse OTEL host based on environment variables. @@ -307,7 +306,7 @@ class LangfuseOtelLogger(OpenTelemetry): @staticmethod def _build_langfuse_otel_config( - public_key: str, secret_key: str, langfuse_host: Optional[str] + public_key: str, secret_key: str, langfuse_host: str | None ) -> "OpenTelemetryConfig": """ Builds an OTLP HTTP config pointing at the Langfuse OTEL endpoint for the @@ -343,7 +342,7 @@ class LangfuseOtelLogger(OpenTelemetry): return f"Basic {auth_header}" @staticmethod - def _build_langfuse_otel_headers(auth_header: str) -> Dict[str, str]: + def _build_langfuse_otel_headers(auth_header: str) -> dict[str, str]: """ Build the OTLP header set Langfuse expects. @@ -356,7 +355,7 @@ class LangfuseOtelLogger(OpenTelemetry): } @staticmethod - def _format_otel_headers(headers: Dict[str, str]) -> str: + def _format_otel_headers(headers: dict[str, str]) -> str: """ Serialize a header mapping into the comma-separated OTLP header string """ @@ -364,7 +363,7 @@ class LangfuseOtelLogger(OpenTelemetry): def construct_dynamic_otel_headers( self, standard_callback_dynamic_params: StandardCallbackDynamicParams - ) -> Optional[dict]: + ) -> dict | None: """ Construct dynamic Langfuse headers from standard callback dynamic params @@ -413,7 +412,7 @@ class LangfuseOtelLogger(OpenTelemetry): self, start_time: datetime, headers: dict, - ) -> Optional[Span]: + ) -> Span | None: """ Override to prevent creating empty proxy request spans. @@ -429,10 +428,8 @@ class LangfuseOtelLogger(OpenTelemetry): """ Langfuse should not receive service success logs. """ - pass async def async_service_failure_hook(self, *args, **kwargs): """ Langfuse should not receive service failure logs. """ - pass diff --git a/litellm/integrations/langfuse/langfuse_otel_attributes.py b/litellm/integrations/langfuse/langfuse_otel_attributes.py index 46bfc21968f..6bf24fab79f 100644 --- a/litellm/integrations/langfuse/langfuse_otel_attributes.py +++ b/litellm/integrations/langfuse/langfuse_otel_attributes.py @@ -5,7 +5,7 @@ Relevant Issue: https://github.com/BerriAI/litellm/issues/13764 """ import json -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any from pydantic import BaseModel from typing_extensions import override @@ -29,20 +29,18 @@ if TYPE_CHECKING: def get_output_content_by_type( - response_obj: Union[ - None, - dict, - EmbeddingResponse, - ModelResponse, - TextCompletionResponse, - ImageResponse, - TranscriptionResponse, - RerankResponse, - HttpxBinaryResponseContent, - ResponsesAPIResponse, - list, - ], - kwargs: Optional[Dict[str, Any]] = None, + response_obj: None + | dict + | EmbeddingResponse + | ModelResponse + | TextCompletionResponse + | ImageResponse + | TranscriptionResponse + | RerankResponse + | HttpxBinaryResponseContent + | ResponsesAPIResponse + | list, + kwargs: dict[str, Any] | None = None, ) -> str: """ Extract output content from response objects based on their type. @@ -83,7 +81,7 @@ def get_output_content_by_type( class LangfuseLLMObsOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @override - def set_messages(span: "Span", kwargs: Dict[str, Any]): + def set_messages(span: "Span", kwargs: dict[str, Any]): prompt = {"messages": kwargs.get("messages")} optional_params = kwargs.get("optional_params", {}) functions = optional_params.get("functions") diff --git a/litellm/integrations/langfuse/langfuse_prompt_management.py b/litellm/integrations/langfuse/langfuse_prompt_management.py index a2b760aa31d..9a5ee49bd0d 100644 --- a/litellm/integrations/langfuse/langfuse_prompt_management.py +++ b/litellm/integrations/langfuse/langfuse_prompt_management.py @@ -4,7 +4,7 @@ Call Hook for LiteLLM Proxy which allows Langfuse prompt management. import os from functools import lru_cache -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, TypeAlias, Union, cast +from typing import TYPE_CHECKING, Any, Literal, TypeAlias, Union, cast from packaging.version import Version @@ -140,8 +140,8 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge self, langfuse_prompt_id: str, langfuse_client: LangfuseClass, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PROMPT_CLIENT: prompt_client = langfuse_client.get_prompt(langfuse_prompt_id, label=prompt_label, version=prompt_version) @@ -150,10 +150,10 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge def _compile_prompt( self, langfuse_prompt_client: PROMPT_CLIENT, - langfuse_prompt_variables: Optional[dict], - call_type: Union[Literal["completion"], Literal["text_completion"]], - ) -> List[AllMessageValues]: - compiled_prompt: Optional[Union[str, list]] = None + langfuse_prompt_variables: dict | None, + call_type: Literal["completion", "text_completion"], + ) -> list[AllMessageValues]: + compiled_prompt: str | list | None = None if langfuse_prompt_variables is None: langfuse_prompt_variables = {} @@ -163,7 +163,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge if isinstance(compiled_prompt, str): compiled_prompt = [ChatCompletionSystemMessage(role="system", content=compiled_prompt)] else: - compiled_prompt = cast(List[AllMessageValues], compiled_prompt) + compiled_prompt = cast(list[AllMessageValues], compiled_prompt) return compiled_prompt @@ -178,21 +178,21 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[ + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[ str, - List[AllMessageValues], + list[AllMessageValues], dict, ]: return self.get_chat_completion_prompt( @@ -211,8 +211,8 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: if prompt_id is None: @@ -232,12 +232,12 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: if prompt_id is None: raise ValueError("prompt_id is required for Langfuse prompt management") @@ -277,12 +277,12 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: return self._compile_prompt_helper( prompt_id=prompt_id, @@ -317,7 +317,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge except Exception as e: from litellm._logging import verbose_logger - verbose_logger.exception(f"Langfuse Layer Error - Exception occurred while logging success event: {str(e)}") + verbose_logger.exception(f"Langfuse Layer Error - Exception occurred while logging success event: {e!s}") self.handle_callback_failure(callback_name="langfuse") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): @@ -329,7 +329,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache, ) standard_logging_object = cast( - Optional[StandardLoggingPayload], + StandardLoggingPayload | None, kwargs.get("standard_logging_object", None), ) status_message = str(kwargs.get("exception", "Unknown error")) @@ -347,5 +347,5 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge except Exception as e: from litellm._logging import verbose_logger - verbose_logger.exception(f"Langfuse Layer Error - Exception occurred while logging failure event: {str(e)}") + verbose_logger.exception(f"Langfuse Layer Error - Exception occurred while logging failure event: {e!s}") self.handle_callback_failure(callback_name="langfuse") diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 565ea833768..1f5d3179fb3 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -6,7 +6,7 @@ import random import traceback import types from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any import httpx from pydantic import BaseModel # type: ignore @@ -41,11 +41,11 @@ def is_serializable(value): class LangsmithLogger(CustomBatchLogger): def __init__( self, - langsmith_api_key: Optional[str] = None, - langsmith_project: Optional[str] = None, - langsmith_base_url: Optional[str] = None, - langsmith_sampling_rate: Optional[float] = None, - langsmith_tenant_id: Optional[str] = None, + langsmith_api_key: str | None = None, + langsmith_project: str | None = None, + langsmith_base_url: str | None = None, + langsmith_sampling_rate: float | None = None, + langsmith_tenant_id: str | None = None, **kwargs, ): self.flush_lock = asyncio.Lock() @@ -74,10 +74,10 @@ class LangsmithLogger(CustomBatchLogger): if _batch_size: self.batch_size = int(_batch_size) - self.log_queue: List[LangsmithQueueObject] = [] - self._flush_task: Optional[asyncio.Task[Any]] = self._start_periodic_flush_task() + self.log_queue: list[LangsmithQueueObject] = [] + self._flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task() - def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]: + def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None: """Start the periodic flush task only when an event loop is already running.""" try: loop = asyncio.get_running_loop() @@ -96,10 +96,10 @@ class LangsmithLogger(CustomBatchLogger): def get_credentials_from_env( self, - langsmith_api_key: Optional[str] = None, - langsmith_project: Optional[str] = None, - langsmith_base_url: Optional[str] = None, - langsmith_tenant_id: Optional[str] = None, + langsmith_api_key: str | None = None, + langsmith_project: str | None = None, + langsmith_base_url: str | None = None, + langsmith_tenant_id: str | None = None, allow_env_credentials: bool = True, ) -> LangsmithCredentialsObject: if allow_env_credentials is False and langsmith_base_url is not None: @@ -142,7 +142,7 @@ class LangsmithLogger(CustomBatchLogger): redacted["requester_metadata"] = redact_user_api_key_info(metadata=nested) return redacted - def _build_extra_metadata(self, metadata: Dict): + def _build_extra_metadata(self, metadata: dict): extra_metadata = dict(metadata) requester_metadata = extra_metadata.get("requester_metadata") if requester_metadata and isinstance(requester_metadata, dict): @@ -152,9 +152,9 @@ class LangsmithLogger(CustomBatchLogger): return self._redact_metadata(extra_metadata) - def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> Dict[str, Any]: + def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> dict[str, Any]: response = payload["response"] - outputs: Dict[str, Any] + outputs: dict[str, Any] if isinstance(response, dict): outputs = {**response} else: @@ -167,7 +167,7 @@ class LangsmithLogger(CustomBatchLogger): } return outputs - def _ensure_required_ids(self, data: dict, run_id: Optional[str]): + def _ensure_required_ids(self, data: dict, run_id: str | None): if "id" not in data or data["id"] is None: run_id = str(uuid.uuid4()) data["id"] = run_id @@ -197,7 +197,7 @@ class LangsmithLogger(CustomBatchLogger): f"Langsmith Logging - project_name: {fields['project_name']}, run_name {fields['run_name']}" ) - payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if payload is None: raise Exception("Error logging request payload. Payload=none.") @@ -244,9 +244,7 @@ class LangsmithLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.debug( @@ -284,9 +282,7 @@ class LangsmithLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.debug( @@ -325,9 +321,7 @@ class LangsmithLogger(CustomBatchLogger): random_sample = random.random() if random_sample > sampling_rate: verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) + f"Skipping Langsmith logging. Sampling rate={sampling_rate}, random_sample={random_sample}" ) return # Skip logging verbose_logger.info("Langsmith Failure Event Logging!") @@ -393,7 +387,7 @@ class LangsmithLogger(CustomBatchLogger): async def _log_batch_on_langsmith( self, credentials: LangsmithCredentialsObject, - queue_objects: List[LangsmithQueueObject], + queue_objects: list[LangsmithQueueObject], ): """ Logs a batch of runs to Langsmith @@ -439,9 +433,9 @@ class LangsmithLogger(CustomBatchLogger): except Exception: verbose_logger.exception(f"Langsmith Layer Error - {traceback.format_exc()}") - def _group_batches_by_credentials(self) -> Dict[CredentialsKey, BatchGroup]: + def _group_batches_by_credentials(self) -> dict[CredentialsKey, BatchGroup]: """Groups queue objects by credentials using a proper key structure""" - log_queue_by_credentials: Dict[CredentialsKey, BatchGroup] = {} + log_queue_by_credentials: dict[CredentialsKey, BatchGroup] = {} for queue_object in self.log_queue: credentials = queue_object["credentials"] @@ -467,8 +461,8 @@ class LangsmithLogger(CustomBatchLogger): return log_queue_by_credentials - def _get_sampling_rate_to_use_for_request(self, kwargs: Dict[str, Any]) -> float: - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + def _get_sampling_rate_to_use_for_request(self, kwargs: dict[str, Any]) -> float: + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params", None ) sampling_rate: float = self.sampling_rate @@ -478,7 +472,7 @@ class LangsmithLogger(CustomBatchLogger): sampling_rate = float(_sampling_rate) return sampling_rate - def _get_credentials_to_use_for_request(self, kwargs: Dict[str, Any]) -> LangsmithCredentialsObject: + def _get_credentials_to_use_for_request(self, kwargs: dict[str, Any]) -> LangsmithCredentialsObject: """ Handles key/team based logging @@ -486,7 +480,7 @@ class LangsmithLogger(CustomBatchLogger): Otherwise, use the default credentials. """ - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params", None ) if standard_callback_dynamic_params is not None: diff --git a/litellm/integrations/levo/levo.py b/litellm/integrations/levo/levo.py index a865944485c..ba16e1ee7d8 100644 --- a/litellm/integrations/levo/levo.py +++ b/litellm/integrations/levo/levo.py @@ -1,5 +1,5 @@ import os -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from litellm.integrations.opentelemetry import OpenTelemetry @@ -25,7 +25,7 @@ class LevoConfig: def __init__( self, - otlp_auth_headers: Optional[str], + otlp_auth_headers: str | None, protocol: Protocol, endpoint: str, ): diff --git a/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py b/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py index 85d209da5b1..44242e09f4c 100644 --- a/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py +++ b/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py @@ -5,8 +5,6 @@ When model is litellm_agent/gpt-3.5-turbo, this hook replaces it with gpt-3.5-tu before the completion call, similar to langfuse/model resolution. """ -from typing import Dict, List, Optional, Tuple - from litellm.integrations.custom_logger import CustomLogger from litellm.types.llms.openai import AllMessageValues from litellm.types.prompts.init_prompts import PromptSpec @@ -25,17 +23,17 @@ class LiteLLMAgentModelResolver(CustomLogger): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Strip litellm_agent/ prefix from model name. @@ -50,19 +48,19 @@ class LiteLLMAgentModelResolver(CustomLogger): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: object, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """Async delegate to get_chat_completion_prompt.""" return self.get_chat_completion_prompt( model=model, diff --git a/litellm/integrations/literal_ai.py b/litellm/integrations/literal_ai.py index c8c931eb667..a54fdcf4dbc 100644 --- a/litellm/integrations/literal_ai.py +++ b/litellm/integrations/literal_ai.py @@ -2,12 +2,11 @@ # This file contains the LiteralAILogger class which is used to log steps to the LiteralAI observability platform. import asyncio import os -from litellm._uuid import uuid -from typing import List, Optional import httpx from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, @@ -162,7 +161,7 @@ class LiteralAILogger(CustomBatchLogger): verbose_logger.exception("Literal AI Layer Error") def _prepare_log_data(self, kwargs, response_obj, start_time, end_time) -> dict: - logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") @@ -172,7 +171,7 @@ class LiteralAILogger(CustomBatchLogger): settings = logging_payload["model_parameters"] messages = logging_payload["messages"] response = logging_payload["response"] - choices: List = [] + choices: list = [] if isinstance(response, dict) and "choices" in response: choices = response["choices"] message_completion = choices[0]["message"] if choices else None diff --git a/litellm/integrations/logfire_logger.py b/litellm/integrations/logfire_logger.py index c92dfff2934..78735c47e5b 100644 --- a/litellm/integrations/logfire_logger.py +++ b/litellm/integrations/logfire_logger.py @@ -3,19 +3,19 @@ import os import traceback -from litellm._uuid import uuid from enum import Enum -from typing import Any, Dict, NamedTuple +from typing import Any, NamedTuple from typing_extensions import LiteralString from litellm._logging import print_verbose, verbose_logger +from litellm._uuid import uuid from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info class SpanConfig(NamedTuple): message_template: LiteralString - span_data: Dict[str, Any] + span_data: dict[str, Any] class LogfireLevel(str, Enum): @@ -35,7 +35,7 @@ class LogfireLogger: if logfire.DEFAULT_LOGFIRE_INSTANCE.config.send_to_logfire: logfire.configure(token=os.getenv("LOGFIRE_TOKEN")) except Exception as e: - print_verbose(f"Got exception on init logfire client {str(e)}") + print_verbose(f"Got exception on init logfire client {e!s}") raise e def _get_span_config(self, payload) -> SpanConfig: @@ -159,5 +159,4 @@ class LogfireLogger: print_verbose(f"Logfire Layer Logging - final response object: {response_obj}") except Exception as e: - verbose_logger.debug(f"Logfire Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.debug(f"Logfire Layer Error - {e!s}\n{traceback.format_exc()}") diff --git a/litellm/integrations/lunary.py b/litellm/integrations/lunary.py index aaf5751cb79..0ec4cf34875 100644 --- a/litellm/integrations/lunary.py +++ b/litellm/integrations/lunary.py @@ -172,4 +172,3 @@ class LunaryLogger: except Exception: print_verbose(f"Lunary Logging Error - {traceback.format_exc()}") - pass diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py index 26b2f32f32f..2fb8e557da4 100644 --- a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -21,7 +21,7 @@ from __future__ import annotations import os from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_proxy_logger @@ -36,8 +36,8 @@ else: def _parse_metrics_marker( - marker: Optional[object], -) -> Optional[datetime]: + marker: object | None, +) -> datetime | None: """Parse metricsMarker from Mavvrik register response into a UTC datetime. Handles both formats Mavvrik may return: @@ -70,7 +70,7 @@ def _parse_metrics_marker( return None -def _is_empty_metrics_marker(marker: Optional[object]) -> bool: +def _is_empty_metrics_marker(marker: object | None) -> bool: if marker is None: return True if isinstance(marker, (int, float)): @@ -105,13 +105,13 @@ class MavvrikFocusLogger(FocusLogger): **kwargs, ) raw = os.getenv("MAVVRIK_FOCUS_MAX_ROWS") - self._max_rows: Optional[int] = int(raw) if raw else 500_000 + self._max_rows: int | None = int(raw) if raw else 500_000 async def _export_window( self, *, window: FocusTimeWindow, - limit: Optional[int], + limit: int | None, ) -> None: """Export with Mavvrik row cap applied when no explicit limit is passed.""" effective_limit = limit if limit is not None else self._max_rows @@ -249,7 +249,7 @@ class MavvrikFocusLogger(FocusLogger): scheduler: AsyncIOScheduler, ) -> None: """Register the Mavvrik FOCUS export job on the provided scheduler.""" - loggers: List[MavvrikFocusLogger] = [ + loggers: list[MavvrikFocusLogger] = [ cb for cb in litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=MavvrikFocusLogger) if type(cb) is MavvrikFocusLogger diff --git a/litellm/integrations/mlflow.py b/litellm/integrations/mlflow.py index 1952c95eac9..7ce5b3d3eea 100644 --- a/litellm/integrations/mlflow.py +++ b/litellm/integrations/mlflow.py @@ -1,6 +1,6 @@ import json import threading -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -53,8 +53,10 @@ class MlflowLogger(CustomLogger): def _extract_and_set_chat_attributes(self, span, kwargs, response_obj): try: - from mlflow.tracing.utils import set_span_chat_messages # type: ignore - from mlflow.tracing.utils import set_span_chat_tools # type: ignore + from mlflow.tracing.utils import ( + set_span_chat_messages, # type: ignore + set_span_chat_tools, # type: ignore + ) except ImportError: return @@ -185,7 +187,7 @@ class MlflowLogger(CustomLogger): "call_type": kwargs.get("call_type"), "model": kwargs.get("model"), } - standard_obj: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_obj: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_obj: attributes.update( { @@ -215,7 +217,7 @@ class MlflowLogger(CustomLogger): ) return attributes - def _get_span_type(self, call_type: Optional[str]) -> str: + def _get_span_type(self, call_type: str | None) -> str: from mlflow.entities import SpanType if call_type in ["completion", "acompletion"]: diff --git a/litellm/integrations/mock_client_factory.py b/litellm/integrations/mock_client_factory.py index 76c0ac03b7b..2aab1792618 100644 --- a/litellm/integrations/mock_client_factory.py +++ b/litellm/integrations/mock_client_factory.py @@ -6,12 +6,13 @@ API calls and return successful mock responses, allowing full code execution wit making actual network calls. """ -import httpx -import json import asyncio -from datetime import timedelta -from typing import Dict, Optional, List, cast +import json from dataclasses import dataclass +from datetime import timedelta +from typing import cast + +import httpx from litellm._logging import verbose_logger @@ -24,8 +25,8 @@ class MockClientConfig: env_var: str # e.g., "GCS_MOCK", "LANGFUSE_MOCK" default_latency_ms: int = 100 # Default mock latency in milliseconds default_status_code: int = 200 # Default HTTP status code - default_json_data: Optional[Dict] = None # Default JSON response data - url_matchers: Optional[List[str]] = None # List of strings to match in URLs (e.g., ["storage.googleapis.com"]) + default_json_data: dict | None = None # Default JSON response data + url_matchers: list[str] | None = None # List of strings to match in URLs (e.g., ["storage.googleapis.com"]) patch_async_handler: bool = True # Whether to patch AsyncHTTPHandler.post patch_sync_client: bool = False # Whether to patch httpx.Client.post patch_http_handler: bool = False # Whether to patch HTTPHandler.post (for sync calls that use HTTPHandler) @@ -42,8 +43,8 @@ class MockResponse: def __init__( self, status_code: int = 200, - json_data: Optional[Dict] = None, - url: Optional[str] = None, + json_data: dict | None = None, + url: str | None = None, elapsed_seconds: float = 0.0, ): self.status_code = status_code @@ -67,7 +68,7 @@ class MockResponse: """Return response content.""" return self._content - def json(self) -> Dict: + def json(self) -> dict: """Return JSON response data.""" return self._json_data @@ -81,7 +82,7 @@ class MockResponse: raise Exception(f"HTTP {self.status_code}") -def _is_url_match(url, matchers: List[str]) -> bool: +def _is_url_match(url, matchers: list[str]) -> bool: """Check if URL matches any of the provided matchers.""" try: parsed_url = httpx.URL(url) if isinstance(url, str) else url @@ -125,7 +126,7 @@ def create_mock_client_factory(config: MockClientConfig): # Create URL matcher function def _is_mock_url(url) -> bool: # url_matchers is guaranteed to be a list after __post_init__ - return _is_url_match(url, cast(List[str], config.url_matchers)) + return _is_url_match(url, cast(list[str], config.url_matchers)) # Create async handler mock async def _mock_async_handler_post( @@ -265,6 +266,7 @@ def create_mock_client_factory(config: MockClientConfig): def should_use_mock() -> bool: """Determine if mock mode should be enabled.""" import os + from litellm.secret_managers.main import str_to_bool mock_mode = os.getenv(config.env_var, "false") diff --git a/litellm/integrations/newrelic/newrelic.py b/litellm/integrations/newrelic/newrelic.py index 3c2bed60ef4..5511cb06174 100644 --- a/litellm/integrations/newrelic/newrelic.py +++ b/litellm/integrations/newrelic/newrelic.py @@ -47,15 +47,15 @@ import os import threading import time import uuid -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any import litellm from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.redact_messages import should_redact_message_logging -from litellm.types.integrations.newrelic import NewRelicInitParams from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus -from litellm.types.utils import ModelResponse, Message, StandardLoggingPayload +from litellm.types.integrations.newrelic import NewRelicInitParams +from litellm.types.utils import Message, ModelResponse, StandardLoggingPayload try: import newrelic.agent as _newrelic_agent @@ -123,17 +123,17 @@ class NewRelicLogger(CustomLogger): verbose_logger.error(f"Failed to initialize New Relic agent: {e}. Integration will be disabled.") self.enabled = False - def _get_newrelic_params(self) -> Dict: + def _get_newrelic_params(self) -> dict: """ Get the newrelic_params from litellm.newrelic_params These are params specific to initializing the NewRelicLogger e.g. turn_off_message_logging """ - dict_newrelic_params: Dict = {} + dict_newrelic_params: dict = {} if litellm.newrelic_params is not None: if isinstance(litellm.newrelic_params, NewRelicInitParams): dict_newrelic_params = litellm.newrelic_params.model_dump() - elif isinstance(litellm.newrelic_params, Dict): + elif isinstance(litellm.newrelic_params, dict): # only allow params that are of NewRelicInitParams dict_newrelic_params = NewRelicInitParams(**litellm.newrelic_params).model_dump() return dict_newrelic_params @@ -246,8 +246,8 @@ class NewRelicLogger(CustomLogger): def _get_trace_context( self, - kwargs: Dict, - standard_logging_object: Optional[StandardLoggingPayload] = None, + kwargs: dict, + standard_logging_object: StandardLoggingPayload | None = None, ) -> str: """ Get the New Relic trace ID for AI monitoring events. @@ -273,7 +273,7 @@ class NewRelicLogger(CustomLogger): Returns: trace_id: always a non-empty string. """ - trace_id: Optional[str] = None + trace_id: str | None = None try: litellm_params = kwargs.get("litellm_params") or {} metadata = litellm_params.get("metadata") or {} @@ -306,7 +306,7 @@ class NewRelicLogger(CustomLogger): return trace_id - def _extract_completion_id(self, kwargs: Dict, response_obj: ModelResponse) -> str: + def _extract_completion_id(self, kwargs: dict, response_obj: ModelResponse) -> str: """ Extract completion ID from kwargs or response_obj, or generate one. """ @@ -326,8 +326,8 @@ class NewRelicLogger(CustomLogger): def _get_vendor( self, - kwargs: Dict, - standard_logging_object: Optional[StandardLoggingPayload] = None, + kwargs: dict, + standard_logging_object: StandardLoggingPayload | None = None, ) -> str: """Extract vendor/provider, preferring StandardLoggingPayload.""" if standard_logging_object: @@ -339,10 +339,10 @@ class NewRelicLogger(CustomLogger): def _get_model_names( self, - kwargs: Dict, + kwargs: dict, response_obj: ModelResponse, - standard_logging_object: Optional[StandardLoggingPayload] = None, - ) -> Tuple[str, str]: + standard_logging_object: StandardLoggingPayload | None = None, + ) -> tuple[str, str]: """ Extract request and response model names, preferring StandardLoggingPayload for the request model. @@ -363,8 +363,8 @@ class NewRelicLogger(CustomLogger): def _extract_usage( self, response_obj: ModelResponse, - standard_logging_object: Optional[StandardLoggingPayload] = None, - ) -> Dict[str, int]: + standard_logging_object: StandardLoggingPayload | None = None, + ) -> dict[str, int]: """Extract usage statistics, preferring StandardLoggingPayload.""" if standard_logging_object: prompt = standard_logging_object.get("prompt_tokens") @@ -406,11 +406,11 @@ class NewRelicLogger(CustomLogger): def _get_duration( self, - kwargs: Dict, + kwargs: dict, start_time: Any, end_time: Any, - standard_logging_object: Optional[StandardLoggingPayload] = None, - ) -> Optional[float]: + standard_logging_object: StandardLoggingPayload | None = None, + ) -> float | None: """ Extract duration in milliseconds. @@ -435,9 +435,9 @@ class NewRelicLogger(CustomLogger): def _get_request_params( self, - kwargs: Dict, - standard_logging_object: Optional[StandardLoggingPayload] = None, - ) -> Dict[str, Any]: + kwargs: dict, + standard_logging_object: StandardLoggingPayload | None = None, + ) -> dict[str, Any]: """ Extract request parameters like temperature and max_tokens, preferring StandardLoggingPayload.model_parameters. @@ -461,7 +461,7 @@ class NewRelicLogger(CustomLogger): return params - def _extract_message_content(self, message: Union[Message, Dict]) -> str: + def _extract_message_content(self, message: Message | dict) -> str: """ Extract content from a message, handling various formats. @@ -496,12 +496,12 @@ class NewRelicLogger(CustomLogger): def _extract_all_messages( self, - kwargs: Dict, + kwargs: dict, response_obj: ModelResponse, response_model: str, vendor: str, - standard_logging_object: Optional[StandardLoggingPayload] = None, - ) -> List[Dict[str, Any]]: + standard_logging_object: StandardLoggingPayload | None = None, + ) -> list[dict[str, Any]]: """ Extract all messages (request + response) with sequence numbers and timestamps. @@ -590,15 +590,15 @@ class NewRelicLogger(CustomLogger): def _record_summary_event( self, request_id: str, - trace_id: Optional[str], + trace_id: str | None, request_model: str, response_model: str, vendor: str, finish_reason: str, num_messages: int, - usage: Dict[str, int], - duration: Optional[float] = None, - request_params: Optional[Dict[str, Any]] = None, + usage: dict[str, int], + duration: float | None = None, + request_params: dict[str, Any] | None = None, ): """Record LlmChatCompletionSummary event to New Relic.""" try: @@ -645,8 +645,8 @@ class NewRelicLogger(CustomLogger): self, request_id: str, llm_response_id: str, - trace_id: Optional[str], - messages: List[Dict[str, Any]], + trace_id: str | None, + messages: list[dict[str, Any]], ): """Record LlmChatCompletionMessage events to New Relic. @@ -719,10 +719,10 @@ class NewRelicLogger(CustomLogger): def _process_success( self, - kwargs: Dict, + kwargs: dict, response_obj: ModelResponse, - start_time: Optional[float] = None, - end_time: Optional[float] = None, + start_time: float | None = None, + end_time: float | None = None, ): """ Core logic for processing successful LLM calls. @@ -736,7 +736,7 @@ class NewRelicLogger(CustomLogger): self._check_and_emit_periodic_metric() # Use StandardLoggingPayload where available for normalized, pre-computed values - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") # Get trace context trace_id = self._get_trace_context(kwargs, standard_logging_object) @@ -832,11 +832,9 @@ class NewRelicLogger(CustomLogger): def log_pre_api_call(self, model, messages, kwargs): """Unused per spec.""" - pass def log_post_api_call(self, kwargs, response_obj, start_time, end_time): """Unused per spec.""" - pass def log_success_event(self, kwargs, response_obj, start_time, end_time): """ diff --git a/litellm/integrations/openmeter.py b/litellm/integrations/openmeter.py index e9cc68a7841..af957f3bbd0 100644 --- a/litellm/integrations/openmeter.py +++ b/litellm/integrations/openmeter.py @@ -45,7 +45,7 @@ class OpenMeterLogger(CustomLogger): missing_keys.append("OPENMETER_API_KEY") if len(missing_keys) > 0: - raise Exception("Missing keys={} in environment.".format(missing_keys)) + raise Exception(f"Missing keys={missing_keys} in environment.") def _common_logic(self, kwargs: dict, response_obj): call_id = response_obj.get("id", kwargs.get("litellm_call_id")) @@ -107,7 +107,7 @@ class OpenMeterLogger(CustomLogger): _data = self._common_logic(kwargs=kwargs, response_obj=response_obj) _headers = { "Content-Type": "application/cloudevents+json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } try: @@ -133,7 +133,7 @@ class OpenMeterLogger(CustomLogger): _data = self._common_logic(kwargs=kwargs, response_obj=response_obj) _headers = { "Content-Type": "application/cloudevents+json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } try: diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 60eb960be4e..5342894a04f 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -4,12 +4,6 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - FrozenSet, - List, - Optional, - Set, - Tuple, Union, cast, ) @@ -95,7 +89,7 @@ _VALID_CAPTURE_MODES = { CAPTURE_MODE_SPAN_AND_EVENT, } -METRIC_METADATA_KEYS: Tuple[str, ...] = ( +METRIC_METADATA_KEYS: tuple[str, ...] = ( "user_api_key_hash", "user_api_key_alias", "user_api_key_team_id", @@ -115,7 +109,7 @@ METRIC_METADATA_KEYS: Tuple[str, ...] = ( TOKEN_TYPE_ATTRIBUTE: str = "gen_ai.token.type" -VALID_METRIC_ATTRIBUTE_NAMES: FrozenSet[str] = frozenset( +VALID_METRIC_ATTRIBUTE_NAMES: frozenset[str] = frozenset( ( "gen_ai.operation.name", "gen_ai.provider.name", @@ -130,8 +124,8 @@ VALID_METRIC_ATTRIBUTE_NAMES: FrozenSet[str] = frozenset( @dataclass(frozen=True) class OTELMetricAttributeFilter: - include_list: Optional[List[str]] = None - exclude_list: Optional[List[str]] = None + include_list: list[str] | None = None + exclude_list: list[str] | None = None def _build_metric_attribute_filter(value: Any) -> OTELMetricAttributeFilter: @@ -149,8 +143,8 @@ def _build_metric_attribute_filter(value: Any) -> OTELMetricAttributeFilter: def _resolve_metric_attribute_filter( - attributes: Optional[OTELMetricAttributeFilter], -) -> Tuple[Optional[FrozenSet[str]], Optional[FrozenSet[str]]]: + attributes: OTELMetricAttributeFilter | None, +) -> tuple[frozenset[str] | None, frozenset[str] | None]: if attributes is None: return None, None include = attributes.include_list or None @@ -173,7 +167,7 @@ def _resolve_metric_attribute_filter( ) -def _normalize_team_metadata_keys(value: Any) -> List[str]: +def _normalize_team_metadata_keys(value: Any) -> list[str]: """Coerce a team-metadata allowlist from a list or comma-separated string. config.yaml passes a YAML list; an env var passes a comma-separated string. @@ -218,28 +212,28 @@ def _freeze_for_dedupe(value: object, _depth: int = 0) -> HashableScope: @dataclass class OpenTelemetryConfig: - exporter: Union[str, SpanExporter] = "console" - endpoint: Optional[str] = None - headers: Optional[str] = None + exporter: str | SpanExporter = "console" + endpoint: str | None = None + headers: str | None = None enable_metrics: bool = False enable_events: bool = False - service_name: Optional[str] = None - deployment_environment: Optional[str] = None - model_id: Optional[str] = None - ignore_context_propagation: Optional[bool] = None + service_name: str | None = None + deployment_environment: str | None = None + model_id: str | None = None + ignore_context_propagation: bool | None = None # When True, create a private TracerProvider instead of reusing or setting the global one. skip_set_global: bool = False # Programmatic override for OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT. # One of NO_CONTENT, SPAN_ONLY, EVENT_ONLY, SPAN_AND_EVENT (or "true" as legacy alias). - capture_message_content: Optional[str] = None - semconv_stability_opt_in: Set[OTELSemconvCategory] = field(default_factory=set) + capture_message_content: str | None = None + semconv_stability_opt_in: set[OTELSemconvCategory] = field(default_factory=set) # Sub-keys of the team's free-form metadata stamped onto the inference span # under ``litellm.team.metadata``. Empty by default so none of a team's # metadata leaves the process until explicitly allowlisted. - baggage_team_metadata_keys: List[str] = field(default_factory=list) + baggage_team_metadata_keys: list[str] = field(default_factory=list) # Prometheus-style include/exclude control over which attributes are stamped # on emitted metrics, to cap metric cardinality. - attributes: Optional[OTELMetricAttributeFilter] = None + attributes: OTELMetricAttributeFilter | None = None def __post_init__(self) -> None: # If endpoint is specified but exporter is still the default "console", @@ -305,12 +299,12 @@ class OpenTelemetryConfig: class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def __init__( self, - config: Optional[OpenTelemetryConfig] = None, - callback_name: Optional[str] = None, + config: OpenTelemetryConfig | None = None, + callback_name: str | None = None, # injection points for testing - tracer_provider: Optional[Any] = None, - logger_provider: Optional[Any] = None, - meter_provider: Optional[Any] = None, + tracer_provider: Any | None = None, + logger_provider: Any | None = None, + meter_provider: Any | None = None, **kwargs, ): team_metadata_keys_override = kwargs.pop("baggage_team_metadata_keys", None) @@ -328,15 +322,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): # callback_settings.otel.attributes after this logger is constructed, so # reading it now would miss it. An explicit config is validated eagerly so # a bad config still fails at startup. - self._metric_attr_include: Optional[FrozenSet[str]] = None - self._metric_attr_exclude: Optional[FrozenSet[str]] = None + self._metric_attr_include: frozenset[str] | None = None + self._metric_attr_exclude: frozenset[str] | None = None self._metric_attr_filter_resolved = False if config.attributes is not None: self._ensure_metric_attribute_filter() self.OTEL_EXPORTER = self.config.exporter self.OTEL_ENDPOINT = self.config.endpoint self.OTEL_HEADERS = self.config.headers - self._tracer_provider_cache: Dict[str, Any] = {} + self._tracer_provider_cache: dict[str, Any] = {} self._init_tracing(tracer_provider) _debug_otel = str(os.getenv("DEBUG_OTEL", "False")).lower() @@ -366,7 +360,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): """Create an OpenTelemetry Resource using config-driven defaults.""" from opentelemetry.sdk.resources import OTELResourceDetector, Resource - base_attributes: Dict[str, Optional[str]] = { + base_attributes: dict[str, str | None] = { "service.name": config.service_name, "deployment.environment": config.deployment_environment, "model_id": config.model_id or config.service_name, @@ -486,7 +480,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): # langfuse_otel relies on the Langfuse SDK's providers; don't overwrite them. return self.config.skip_set_global or (hasattr(self, "callback_name") and self.callback_name == "langfuse_otel") - def _compute_capture_mode_from_init_state(self) -> Optional[str]: + def _compute_capture_mode_from_init_state(self) -> str | None: """Sample explicit settings at init. Returns the resolved mode or None if nothing explicit is set (in which case the legacy ``self.message_logging`` flag is consulted dynamically per request). @@ -672,10 +666,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): async def async_service_success_hook( self, payload: ServiceLoggerPayload, - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[datetime, float]] = None, - event_metadata: Optional[dict] = None, + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: datetime | float | None = None, + event_metadata: dict | None = None, ): from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -731,11 +725,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): async def async_service_failure_hook( self, payload: ServiceLoggerPayload, - error: Optional[str] = "", - parent_otel_span: Optional[Span] = None, - start_time: Optional[Union[datetime, float]] = None, - end_time: Optional[Union[float, datetime]] = None, - event_metadata: Optional[dict] = None, + error: str | None = "", + parent_otel_span: Span | None = None, + start_time: datetime | float | None = None, + end_time: float | datetime | None = None, + event_metadata: dict | None = None, ): from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -797,7 +791,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, + traceback_str: str | None = None, ): from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -883,7 +877,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _emit_guardrail_spans_from_request_data( self, request_data: dict, - parent_span: Optional[Any], + parent_span: Any | None, ) -> None: """Emit ``guardrail`` spans from the request's proxy-internal metadata bucket (``standard_logging_guardrail_information``). @@ -909,7 +903,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): # kwargs["litellm_params"]["metadata"]["_otel_internal"]. Pass the # SAME metadata dict the proxy populated so _handle_failure and # this hook see the same dedupe markers. - kwargs: Dict[str, Any] = { + kwargs: dict[str, Any] = { "litellm_params": {"metadata": metadata}, "standard_logging_object": { "guardrail_information": guardrail_information, @@ -992,9 +986,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return tracer_to_use - def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]: + def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> dict | None: """Extract dynamic headers from kwargs if available.""" - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params" ) @@ -1007,9 +1001,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return dynamic_headers if dynamic_headers else None - def _get_dynamic_otel_config_from_kwargs(self, kwargs: dict) -> Optional[OpenTelemetryConfig]: + def _get_dynamic_otel_config_from_kwargs(self, kwargs: dict) -> OpenTelemetryConfig | None: """Extract a full dynamic exporter config from kwargs if available.""" - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params" ) @@ -1053,7 +1047,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def construct_dynamic_otel_headers( self, standard_callback_dynamic_params: StandardCallbackDynamicParams - ) -> Optional[dict]: + ) -> dict | None: """ Construct dynamic headers from standard callback dynamic params @@ -1066,7 +1060,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def construct_dynamic_otel_config( self, standard_callback_dynamic_params: StandardCallbackDynamicParams - ) -> Optional[OpenTelemetryConfig]: + ) -> OpenTelemetryConfig | None: """ Construct a full exporter config from standard callback dynamic params. @@ -1276,7 +1270,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs) - span_kwargs: Dict[str, Any] = { + span_kwargs: dict[str, Any] = { "name": self._get_span_name(kwargs), "start_time": self._to_ns(start_time), "context": context, @@ -1321,8 +1315,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _set_team_attributes_on_span( self, span: Span, - team_id: Optional[str], - team_alias: Optional[str], + team_id: str | None, + team_alias: str | None, ) -> None: """Stamp team_id / team_alias onto a span so every child span of a litellm_request trace carries them, not just the root span. @@ -1418,7 +1412,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): self.safe_set_attribute(span=span, key=PROVIDER_MODEL_ATTRIBUTE, value=provider_model) @staticmethod - def _team_metadata_json(value: Any, allowed_keys: List[str]) -> Optional[str]: + def _team_metadata_json(value: Any, allowed_keys: list[str]) -> str | None: """JSON-serialize only the allowlisted sub-keys of a team's metadata. Returns ``None`` when nothing is allowlisted or no allowlisted key is @@ -1450,7 +1444,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ) = _resolve_metric_attribute_filter(attributes) self._metric_attr_filter_resolved = True - def _filter_metric_attributes(self, attrs: Dict[str, Any]) -> Dict[str, Any]: + def _filter_metric_attributes(self, attrs: dict[str, Any]) -> dict[str, Any]: if not self._metric_attr_filter_resolved: self._ensure_metric_attribute_filter() if self._metric_attr_include is not None: @@ -1510,8 +1504,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): @staticmethod def _to_timestamp( - val: Optional[Union[datetime, float, str]], - ) -> Optional[float]: + val: datetime | float | str | None, + ) -> float | None: """Convert datetime/float/string to timestamp.""" if val is None: return None @@ -1555,7 +1549,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _record_time_per_output_token_metric( self, kwargs: dict, - response_obj: Optional[Any], + response_obj: Any | None, end_time: datetime, duration_s: float, common_attrs: dict, @@ -1619,7 +1613,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _record_response_duration_metric( self, kwargs: dict, - end_time: Union[datetime, float], + end_time: datetime | float, common_attrs: dict, ): """Record Total Generation Time (response duration) metric. @@ -1771,10 +1765,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): @staticmethod def _resolve_guardrail_context( - span: Optional[Any], - parent_span: Optional[Any], - fallback_ctx: Optional[Any], - ) -> Optional[Any]: + span: Any | None, + parent_span: Any | None, + fallback_ctx: Any | None, + ) -> Any | None: """ Return a valid OTEL context for guardrail child spans so they are never orphaned (Issue #5). Priority: @@ -1790,13 +1784,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return _trace.set_span_in_context(parent_span) return fallback_ctx - def _create_guardrail_span(self, kwargs: Optional[dict], context: Optional[Context]): + def _create_guardrail_span(self, kwargs: dict | None, context: Context | None): """ Creates a span for Guardrail, if any guardrail information is present in standard_logging_object """ # Create span for guardrail information kwargs = kwargs or {} - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: return @@ -1941,7 +1935,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if should_create_primary_span: # Span 1: Request sent to litellm SDK otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs) - span_kwargs: Dict[str, Any] = { + span_kwargs: dict[str, Any] = { "name": self._get_span_name(kwargs), "start_time": self._to_ns(start_time), "context": _parent_context, @@ -2003,7 +1997,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): span.record_exception(exception) # Get StandardLoggingPayload for structured error information - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: return @@ -2106,9 +2100,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ) except Exception as e: verbose_logger.error("OpenTelemetry: Error setting tools attributes: %s", str(e)) - pass - def cast_as_primitive_value_type(self, value) -> Union[str, bool, int, float]: + def cast_as_primitive_value_type(self, value) -> str | bool | int | float: """ Casts the value to a primitive OTEL type if it is not already a primitive type. @@ -2127,11 +2120,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): @staticmethod def _tool_calls_kv_pair( - tool_calls: List[ChatCompletionMessageToolCall], - ) -> Dict[str, Any]: + tool_calls: list[ChatCompletionMessageToolCall], + ) -> dict[str, Any]: from litellm.proxy._types import SpanAttributes - kv_pairs: Dict[str, Any] = {} + kv_pairs: dict[str, Any] = {} for idx, tool_call in enumerate(tool_calls): _function = tool_call.get("function") if not _function: @@ -2145,7 +2138,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return kv_pairs - def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): + def set_attributes(self, span: Span, kwargs, response_obj: Any | None): try: if self.callback_name == "langtrace": from litellm.integrations.langtrace import LangtraceAttributes @@ -2170,7 +2163,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): optional_params = kwargs.get("optional_params", {}) litellm_params = kwargs.get("litellm_params", {}) or {} - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") @@ -2181,7 +2174,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ############################################# metadata = standard_logging_payload["metadata"] for key, value in metadata.items(): - self.safe_set_attribute(span=span, key="metadata.{}".format(key), value=value) + self.safe_set_attribute(span=span, key=f"metadata.{key}", value=value) # get hidden params hidden_params = getattr(standard_logging_payload, "hidden_params", None) or ( @@ -2200,7 +2193,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): litellm_params=litellm_params, ) # Cost breakdown tracking - cost_breakdown: Optional[CostBreakdown] = standard_logging_payload.get("cost_breakdown") + cost_breakdown: CostBreakdown | None = standard_logging_payload.get("cost_breakdown") if cost_breakdown: for key, value in cost_breakdown.items(): if value is not None: @@ -2500,7 +2493,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry") verbose_logger.exception("OpenTelemetry logging error in set_attributes %s", str(e)) - def _cast_as_primitive_value_type(self, value) -> Union[str, bool, int, float]: + def _cast_as_primitive_value_type(self, value) -> str | bool | int | float: """ Casts the value to a primitive OTEL type if it is not already a primitive type. @@ -2524,7 +2517,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): primitive_value = self._cast_as_primitive_value_type(value) span.set_attribute(key, primitive_value) - def _transform_messages_to_otel_semantic_conventions(self, messages: Union[List[dict], str]) -> List[dict]: + def _transform_messages_to_otel_semantic_conventions(self, messages: list[dict] | str) -> list[dict]: """ Transforms LiteLLM/OpenAI style messages into OTEL GenAI 1.38 compliant format. OTEL expects a 'parts' array instead of a single 'content' string. @@ -2565,7 +2558,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return transformed - def _transform_choices_to_otel_semantic_conventions(self, choices: List[dict]) -> List[dict]: + def _transform_choices_to_otel_semantic_conventions(self, choices: list[dict]) -> list[dict]: """ Transforms choices into OTEL GenAI 1.38 compliant format for output.messages. """ @@ -2582,7 +2575,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return transformed @staticmethod - def _to_dict(obj) -> Optional[dict]: + def _to_dict(obj) -> dict | None: """Normalize an object to a plain dict. Handles three forms that appear in practice: @@ -2606,7 +2599,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return obj.model_dump() # type: ignore[union-attr] return None - def _transform_responses_api_output_to_otel(self, output: List) -> List[dict]: + def _transform_responses_api_output_to_otel(self, output: list) -> list[dict]: """ Transform Responses API output items into OTEL GenAI 1.38 format. @@ -2695,9 +2688,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ) except json.JSONDecodeError: verbose_logger.debug( - "litellm.integrations.opentelemetry.py::set_raw_request_attributes() - raw_response not json string - {}".format( - _raw_response - ) + f"litellm.integrations.opentelemetry.py::set_raw_request_attributes() - raw_response not json string - {_raw_response}" ) self.safe_set_attribute( @@ -2749,7 +2740,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return _parent_context - def _get_span_context(self, kwargs, default_span: Optional[Span] = None): + def _get_span_context(self, kwargs, default_span: Span | None = None): from opentelemetry import context, trace from opentelemetry.trace.propagation.tracecontext import ( TraceContextTextMapPropagator, @@ -2806,8 +2797,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _get_span_processor( self, - dynamic_headers: Optional[dict] = None, - config_override: Optional[OpenTelemetryConfig] = None, + dynamic_headers: dict | None = None, + config_override: OpenTelemetryConfig | None = None, ): from opentelemetry.sdk.trace.export import ( BatchSpanProcessor, @@ -3046,7 +3037,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): exporter = ConsoleMetricExporter() return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) - def _normalize_otel_endpoint(self, endpoint: Optional[str], signal_type: str) -> Optional[str]: + def _normalize_otel_endpoint(self, endpoint: str | None, signal_type: str) -> str | None: """ Normalize the endpoint URL for a specific OpenTelemetry signal type. @@ -3118,12 +3109,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): @staticmethod def _get_headers_dictionary( - headers: Optional[Union[str, dict]], - ) -> Dict[str, str]: + headers: str | dict | None, + ) -> dict[str, str]: """ Convert a string or dictionary of headers into a dictionary of headers. """ - _split_otel_headers: Dict[str, str] = {} + _split_otel_headers: dict[str, str] = {} if headers: if isinstance(headers, str): # when passed HEADERS="x-honeycomb-team=B85YgLm96******" @@ -3139,7 +3130,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): async def async_management_endpoint_success_hook( self, logging_payload: ManagementEndpointLoggingPayload, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ): from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -3197,7 +3188,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): async def async_management_endpoint_failure_hook( self, logging_payload: ManagementEndpointLoggingPayload, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ): from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -3266,7 +3257,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): self, start_time: datetime, headers: dict, - ) -> Optional[Span]: + ) -> Span | None: """ Create a span for the received proxy server request. """ @@ -3280,10 +3271,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def set_proxy_request_route_attributes( self, - span: Optional[Span], + span: Span | None, *, - url_path: Optional[str] = None, - http_route: Optional[str] = None, + url_path: str | None = None, + http_route: str | None = None, ) -> None: """ Set OTel-standard ``http.route`` / ``url.path`` on the proxy SERVER @@ -3297,7 +3288,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if http_route: self.safe_set_attribute(span=span, key=HTTP_ROUTE_ATTRIBUTE, value=http_route) - def set_response_status_code_attribute(self, span: Optional[Span], status_code: Optional[int]) -> None: + def set_response_status_code_attribute(self, span: Span | None, status_code: int | None) -> None: """ Set OTel-standard ``http.response.status_code`` (int) on the proxy SERVER span. The failure path sets this from the error code in @@ -3316,8 +3307,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def record_error_attributes_on_span( self, - span: Optional[Span], - exception: Optional[Exception], + span: Span | None, + exception: Exception | None, status_code: int, ) -> None: """Stamp structured ``error.*`` attributes on the SERVER span from the @@ -3336,7 +3327,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): kwargs={"standard_logging_object": {"error_information": error_information}}, ) - def set_preprocessing_duration_attribute(self, span: Optional[Span], container: Any) -> None: + def set_preprocessing_duration_attribute(self, span: Span | None, container: Any) -> None: """ Set ``litellm.preprocessing.duration_ms`` (proxy-receive -> first provider handoff) on the proxy SERVER span. ``litellm_received_at`` diff --git a/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py b/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py index f74da8231f3..71119f07afa 100644 --- a/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py +++ b/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py @@ -1,5 +1,5 @@ from abc import ABC -from typing import TYPE_CHECKING, Any, Dict, Union +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from opentelemetry.trace import Span @@ -7,7 +7,7 @@ if TYPE_CHECKING: class BaseLLMObsOTELAttributes(ABC): @staticmethod - def set_messages(span: "Span", kwargs: Dict[str, Any]): + def set_messages(span: "Span", kwargs: dict[str, Any]): pass @staticmethod @@ -15,7 +15,7 @@ class BaseLLMObsOTELAttributes(ABC): pass -def cast_as_primitive_value_type(value) -> Union[str, bool, int, float]: +def cast_as_primitive_value_type(value) -> str | bool | int | float: """ Converts a value to an OTEL-supported primitive for Arize/Phoenix observability. """ diff --git a/litellm/integrations/opentelemetry_utils/gen_ai_semconv.py b/litellm/integrations/opentelemetry_utils/gen_ai_semconv.py index 98d24f1f7cc..9a619bfc534 100644 --- a/litellm/integrations/opentelemetry_utils/gen_ai_semconv.py +++ b/litellm/integrations/opentelemetry_utils/gen_ai_semconv.py @@ -31,7 +31,7 @@ Events: from datetime import datetime from enum import Enum -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union +from typing import TYPE_CHECKING, Any, Union from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -77,7 +77,7 @@ _SEMCONV_CACHE_TOKEN_ATTRIBUTES = { _INFERENCE_DETAILS_EVENT_NAME = "gen_ai.client.inference.operation.details" -def parse_semconv_opt_in(raw: Optional[str]) -> Set[OTELSemconvCategory]: +def parse_semconv_opt_in(raw: str | None) -> set[OTELSemconvCategory]: """Parse the comma-separated OTEL_SEMCONV_STABILITY_OPT_IN value into the set of recognized categories. Unknown tokens are ignored per the spec.""" if not raw: @@ -121,13 +121,13 @@ class OTELGenAISemconvMixin: def _capture_in_event(self) -> bool: ... - def _transform_messages_to_otel_semantic_conventions(self, messages: Union[List[dict], str]) -> List[dict]: ... + def _transform_messages_to_otel_semantic_conventions(self, messages: list[dict] | str) -> list[dict]: ... - def _transform_choices_to_otel_semantic_conventions(self, choices: List[dict]) -> List[dict]: ... + def _transform_choices_to_otel_semantic_conventions(self, choices: list[dict]) -> list[dict]: ... def _to_ns(self, dt: datetime) -> int: ... - def _otel_log_types(self) -> Tuple[Any, Any]: ... + def _otel_log_types(self) -> tuple[Any, Any]: ... @property def _gen_ai_semconv_latest_experimental(self) -> bool: @@ -195,13 +195,13 @@ class OTELGenAISemconvMixin: if value: self.safe_set_attribute(span=span, key=semconv_key, value=value) - def _build_inference_details_attrs(self, kwargs: dict, response_obj: dict, provider: str) -> Dict[str, Any]: + def _build_inference_details_attrs(self, kwargs: dict, response_obj: dict, provider: str) -> dict[str, Any]: """Build the attribute payload for the inference-details event. Always includes provider/operation; input/output messages are added only when content capture is enabled and non-empty. Mixin-internal. """ - attrs: Dict[str, Any] = { + attrs: dict[str, Any] = { "event_name": _INFERENCE_DETAILS_EVENT_NAME, "gen_ai.provider.name": provider, "gen_ai.operation.name": self._gen_ai_operation_name(kwargs), diff --git a/litellm/integrations/opik/opik.py b/litellm/integrations/opik/opik.py index fd84ad56247..deb325286e9 100644 --- a/litellm/integrations/opik/opik.py +++ b/litellm/integrations/opik/opik.py @@ -5,7 +5,7 @@ Opik Logger that logs LLM events to an Opik server import asyncio import traceback from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -23,7 +23,7 @@ except Exception: opik_client = None -def _should_skip_event(kwargs: Dict[str, Any]) -> bool: +def _should_skip_event(kwargs: dict[str, Any]) -> bool: """Check if event should be skipped due to missing standard_logging_object.""" if kwargs.get("standard_logging_object") is None: verbose_logger.debug("OpikLogger skipping event; no standard_logging_object found") @@ -57,31 +57,31 @@ class OpikLogger(CustomBatchLogger): ) or "https://www.comet.com/opik/api" ) - opik_api_key: Optional[str] = utils.get_opik_config_variable( + opik_api_key: str | None = utils.get_opik_config_variable( "api_key", user_value=kwargs.get("api_key", None), default_value=None ) - opik_workspace: Optional[str] = utils.get_opik_config_variable( + opik_workspace: str | None = utils.get_opik_config_variable( "workspace", user_value=kwargs.get("workspace", None), default_value=None ) self.trace_url: str = f"{opik_base_url}/v1/private/traces/batch" self.span_url: str = f"{opik_base_url}/v1/private/spans/batch" - self.headers: Dict[str, str] = {} + self.headers: dict[str, str] = {} if opik_workspace: self.headers["Comet-Workspace"] = opik_workspace if opik_api_key: self.headers["authorization"] = opik_api_key - self.opik_workspace: Optional[str] = opik_workspace - self.opik_api_key: Optional[str] = opik_api_key + self.opik_workspace: str | None = opik_workspace + self.opik_api_key: str | None = opik_api_key try: asyncio.create_task(self.periodic_flush()) - self.flush_lock: Optional[asyncio.Lock] = asyncio.Lock() + self.flush_lock: asyncio.Lock | None = asyncio.Lock() except Exception as e: verbose_logger.exception( - f"OpikLogger - Asynchronous processing not initialized as we are not running in an async context {str(e)}" + f"OpikLogger - Asynchronous processing not initialized as we are not running in an async context {e!s}" ) self.flush_lock = None @@ -95,7 +95,7 @@ class OpikLogger(CustomBatchLogger): async def async_log_success_event( self, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], response_obj: Any, start_time: datetime, end_time: datetime, @@ -161,9 +161,9 @@ class OpikLogger(CustomBatchLogger): verbose_logger.debug("OpikLogger - Flushing batch") await self.flush_queue() except Exception as e: - verbose_logger.exception(f"OpikLogger failed to log success event - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"OpikLogger failed to log success event - {e!s}\n{traceback.format_exc()}") - def _sync_send(self, url: str, headers: Dict[str, str], batch: Dict[str, Any]) -> None: + def _sync_send(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None: try: response = self.sync_httpx_client.post( url=url, @@ -174,11 +174,11 @@ class OpikLogger(CustomBatchLogger): if response.status_code != 204: raise Exception(f"Response from opik API status_code: {response.status_code}, text: {response.text}") except Exception as e: - verbose_logger.exception(f"OpikLogger failed to send batch - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"OpikLogger failed to send batch - {e!s}\n{traceback.format_exc()}") def log_success_event( self, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], response_obj: Any, start_time: datetime, end_time: datetime, @@ -245,9 +245,9 @@ class OpikLogger(CustomBatchLogger): batch={"spans": [span_payload.__dict__]}, ) except Exception as e: - verbose_logger.exception(f"OpikLogger failed to log success event - {str(e)}\n{traceback.format_exc()}") + verbose_logger.exception(f"OpikLogger failed to log success event - {e!s}\n{traceback.format_exc()}") - async def _submit_batch(self, url: str, headers: Dict[str, str], batch: Dict[str, Any]) -> None: + async def _submit_batch(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None: try: response = await self.async_httpx_client.post( url=url, @@ -261,10 +261,10 @@ class OpikLogger(CustomBatchLogger): else: verbose_logger.info(f"OpikLogger - {len(self.log_queue)} Opik events submitted") except Exception as e: - verbose_logger.exception(f"OpikLogger failed to send batch - {str(e)}") + verbose_logger.exception(f"OpikLogger failed to send batch - {e!s}") - def _create_opik_headers(self) -> Dict[str, str]: - headers: Dict[str, str] = {} + def _create_opik_headers(self) -> dict[str, str]: + headers: dict[str, str] = {} if self.opik_workspace: headers["Comet-Workspace"] = self.opik_workspace diff --git a/litellm/integrations/opik/opik_payload_builder/api.py b/litellm/integrations/opik/opik_payload_builder/api.py index 6a5f9bfddc5..07ba2cd87e3 100644 --- a/litellm/integrations/opik/opik_payload_builder/api.py +++ b/litellm/integrations/opik/opik_payload_builder/api.py @@ -1,7 +1,7 @@ """Public API for Opik payload building.""" from datetime import datetime -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm.integrations.opik import utils @@ -9,12 +9,12 @@ from . import extractors, payload_builders, types def build_opik_payload( - kwargs: Dict[str, Any], - response_obj: Dict[str, Any], + kwargs: dict[str, Any], + response_obj: dict[str, Any], start_time: datetime, end_time: datetime, project_name: str, -) -> Tuple[Optional[types.TracePayload], types.SpanPayload]: +) -> tuple[types.TracePayload | None, types.SpanPayload]: """ Build Opik trace and span payloads from LiteLLM completion data. @@ -77,7 +77,7 @@ def build_opik_payload( output_data = standard_logging_object.get("response", {}) # Decide whether to create a new trace or attach to existing - trace_payload: Optional[types.TracePayload] = None + trace_payload: types.TracePayload | None = None if trace_id is None: trace_id = utils.create_uuid7() trace_payload = payload_builders.build_trace_payload( diff --git a/litellm/integrations/opik/opik_payload_builder/extractors.py b/litellm/integrations/opik/opik_payload_builder/extractors.py index 73058b2a524..f95bd110cb3 100644 --- a/litellm/integrations/opik/opik_payload_builder/extractors.py +++ b/litellm/integrations/opik/opik_payload_builder/extractors.py @@ -1,12 +1,12 @@ """Data extraction functions for Opik payload building.""" import json -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from litellm import _logging -def normalize_provider_name(provider: Optional[str]) -> Optional[str]: +def normalize_provider_name(provider: str | None) -> str | None: """ Normalize LiteLLM provider names to standardized string names. @@ -35,9 +35,9 @@ def normalize_provider_name(provider: Optional[str]) -> Optional[str]: def extract_opik_metadata( - litellm_metadata: Dict[str, Any], - standard_logging_metadata: Dict[str, Any], -) -> Dict[str, Any]: + litellm_metadata: dict[str, Any], + standard_logging_metadata: dict[str, Any], +) -> dict[str, Any]: """ Merge Opik metadata from three sources in increasing priority order: @@ -73,7 +73,7 @@ def extract_opik_metadata( def extract_span_identifiers( current_span_data: Any, -) -> Tuple[Optional[str], Optional[str]]: +) -> tuple[str | None, str | None]: """ Extract trace_id and parent_span_id from current_span_data. @@ -97,9 +97,9 @@ def extract_span_identifiers( def extract_tags( - opik_metadata: Dict[str, Any], - custom_llm_provider: Optional[str], -) -> List[str]: + opik_metadata: dict[str, Any], + custom_llm_provider: str | None, +) -> list[str]: """ Extract and build list of tags. @@ -120,10 +120,10 @@ def extract_tags( def apply_proxy_header_overrides( project_name: str, - tags: List[str], - thread_id: Optional[str], - proxy_headers: Dict[str, Any], -) -> Tuple[str, List[str], Optional[str]]: + tags: list[str], + thread_id: str | None, + proxy_headers: dict[str, Any], +) -> tuple[str, list[str], str | None]: """ Apply overrides from proxy request headers (opik_* prefix). @@ -158,11 +158,11 @@ def apply_proxy_header_overrides( def extract_and_build_metadata( - opik_metadata: Dict[str, Any], - standard_logging_metadata: Dict[str, Any], - standard_logging_object: Dict[str, Any], - litellm_kwargs: Dict[str, Any], -) -> Dict[str, Any]: + opik_metadata: dict[str, Any], + standard_logging_metadata: dict[str, Any], + standard_logging_object: dict[str, Any], + litellm_kwargs: dict[str, Any], +) -> dict[str, Any]: """ Build the complete metadata dictionary from all available sources. diff --git a/litellm/integrations/opik/opik_payload_builder/payload_builders.py b/litellm/integrations/opik/opik_payload_builder/payload_builders.py index 4d92650d2b8..517d5431b70 100644 --- a/litellm/integrations/opik/opik_payload_builder/payload_builders.py +++ b/litellm/integrations/opik/opik_payload_builder/payload_builders.py @@ -1,7 +1,7 @@ """Payload builders for Opik traces and spans.""" from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from litellm import _logging from litellm.integrations.opik import utils @@ -12,14 +12,14 @@ from . import types def build_trace_payload( project_name: str, trace_id: str, - response_obj: Dict[str, Any], + response_obj: dict[str, Any], start_time: datetime, end_time: datetime, input_data: Any, output_data: Any, - metadata: Dict[str, Any], - tags: List[str], - thread_id: Optional[str], + metadata: dict[str, Any], + tags: list[str], + thread_id: str | None, ) -> types.TracePayload: """Build a complete trace payload.""" trace_name = response_obj.get("object", "unknown type") @@ -41,17 +41,17 @@ def build_trace_payload( def build_span_payload( project_name: str, trace_id: str, - parent_span_id: Optional[str], - response_obj: Dict[str, Any], + parent_span_id: str | None, + response_obj: dict[str, Any], start_time: datetime, end_time: datetime, input_data: Any, output_data: Any, - metadata: Dict[str, Any], - tags: List[str], - usage: Dict[str, int], - provider: Optional[str] = None, - cost: Optional[float] = None, + metadata: dict[str, Any], + tags: list[str], + usage: dict[str, int], + provider: str | None = None, + cost: float | None = None, ) -> types.SpanPayload: """Build a complete span payload.""" span_id = utils.create_uuid7() diff --git a/litellm/integrations/opik/opik_payload_builder/types.py b/litellm/integrations/opik/opik_payload_builder/types.py index 070cb11489a..665a88bf0a5 100644 --- a/litellm/integrations/opik/opik_payload_builder/types.py +++ b/litellm/integrations/opik/opik_payload_builder/types.py @@ -1,7 +1,7 @@ """Type definitions for Opik payload building.""" from dataclasses import dataclass -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Any, Literal, Union @dataclass @@ -15,9 +15,9 @@ class TracePayload: end_time: str input: Any output: Any - metadata: Dict[str, Any] - tags: List[str] - thread_id: Optional[str] = None + metadata: dict[str, Any] + tags: list[str] + thread_id: str | None = None @dataclass @@ -34,13 +34,13 @@ class SpanPayload: end_time: str input: Any output: Any - metadata: Dict[str, Any] - tags: List[str] - usage: Dict[str, int] - parent_span_id: Optional[str] = None - provider: Optional[str] = None - total_cost: Optional[float] = None + metadata: dict[str, Any] + tags: list[str] + usage: dict[str, int] + parent_span_id: str | None = None + provider: str | None = None + total_cost: float | None = None PayloadItem = Union[TracePayload, SpanPayload] -TraceSpanPayloadTuple = Tuple[Optional[TracePayload], SpanPayload] +TraceSpanPayloadTuple = tuple[TracePayload | None, SpanPayload] diff --git a/litellm/integrations/opik/utils.py b/litellm/integrations/opik/utils.py index d4850d50778..c9220730a4d 100644 --- a/litellm/integrations/opik/utils.py +++ b/litellm/integrations/opik/utils.py @@ -2,7 +2,7 @@ import configparser import os import time import uuid -from typing import Any, Dict, Final, List, Optional, Tuple +from typing import Any, Final CONFIG_FILE_PATH_DEFAULT: Final[str] = "~/.opik.config" @@ -35,7 +35,7 @@ def create_uuid7() -> str: return str(uuid.UUID(bytes=bytes(uuid_bytes))) -def _read_opik_config_file() -> Dict[str, str]: +def _read_opik_config_file() -> dict[str, str]: config_path = os.path.expanduser(CONFIG_FILE_PATH_DEFAULT) config = configparser.ConfigParser() @@ -49,14 +49,12 @@ def _read_opik_config_file() -> Dict[str, str]: return {} -def _get_env_variable(key: str) -> Optional[str]: +def _get_env_variable(key: str) -> str | None: env_prefix = "opik_" return os.getenv((env_prefix + key).upper(), None) -def get_opik_config_variable( - key: str, user_value: Optional[str] = None, default_value: Optional[str] = None -) -> Optional[str]: +def get_opik_config_variable(key: str, user_value: str | None = None, default_value: str | None = None) -> str | None: """ Get the configuration value of a variable, order priority is: 1. user provided value @@ -95,14 +93,14 @@ def create_usage_object(usage): return usage_dict -def _remove_nulls(x: Dict[str, Any]) -> Dict[str, Any]: +def _remove_nulls(x: dict[str, Any]) -> dict[str, Any]: """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: list[dict[str, Any]], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: """ Separate traces and spans from payload. diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 5e167e006ff..8e11f55f46f 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -13,16 +13,16 @@ The ``LITELLM_OTEL_V2`` env var gates whether the factory in class (from :mod:`logger`). """ -from litellm.integrations.otel.model.config import ( - OTEL_V2_ENV, - OpenTelemetryV2Config, - is_otel_v2_enabled, -) from litellm.integrations.otel.model.baggage import ( BAGGAGE_PROMOTED_KEYS, DEFAULT_BAGGAGE_METADATA_KEYS, promoted_baggage, ) +from litellm.integrations.otel.model.config import ( + OTEL_V2_ENV, + OpenTelemetryV2Config, + is_otel_v2_enabled, +) from litellm.integrations.otel.model.metadata import ( RequestContext, RequestIdentity, diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py index 21a243de9b9..f5509b87f4f 100644 --- a/litellm/integrations/otel/emitter.py +++ b/litellm/integrations/otel/emitter.py @@ -126,7 +126,7 @@ class SpanEmitter: ) # Bounded LRU (ordered by insertion / most-recent touch). Storing keys # only — the value is unused — so it behaves like a capped set. - self._emitted: "OrderedDict[tuple[str, SpanRole], None]" = OrderedDict() + self._emitted: OrderedDict[tuple[str, SpanRole], None] = OrderedDict() # -- low-level helpers --------------------------------------------------- # diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 888d006e4f7..60291046361 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -158,7 +158,7 @@ class OpenTelemetryV2(CustomLogger): event_recorder=self._init_events(logger_provider), ) self._tenant_tracers = TenantTracerCache(self.config, callback_name, LITELLM_TRACER_NAME) - self._open_llm_calls: "OrderedDict[str, _LLMCallSpan]" = OrderedDict() + self._open_llm_calls: OrderedDict[str, _LLMCallSpan] = OrderedDict() self._init_otel_logger_on_litellm_proxy() def _init_metrics(self, meter_provider: Any | None) -> "GenAIMetricRecorder | None": @@ -669,9 +669,9 @@ class OpenTelemetryV2(CustomLogger): SDK dropped it, leaving the POST that actually failed unmarked.""" span = mcp_message_transport_span() or request_root_span() or user_api_key_dict.parent_otel_span if span is None or not is_recordable_span(span): - return None + return stamp_error(span, _span_error_from_exception(original_exception, traceback_str=traceback_str)) - return None + return def emit_guardrail_span(self, entry: "StandardLoggingGuardrailInformation") -> None: # Emitted by the guardrail-recording code the moment a guardrail finishes, diff --git a/litellm/integrations/otel/mappers/__init__.py b/litellm/integrations/otel/mappers/__init__.py index 2050677feb3..6ca549468df 100644 --- a/litellm/integrations/otel/mappers/__init__.py +++ b/litellm/integrations/otel/mappers/__init__.py @@ -55,9 +55,9 @@ def resolve_mappers(names: Iterable[str]) -> list[AttributeMapper]: __all__ = [ + "AttrValue", "AttributeMap", "AttributeMapper", - "AttrValue", "GenAIMapper", "LangfuseMapper", "LangtraceMapper", diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 270d9e33b86..ea7de8d54ef 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, List +from typing import Annotated, Any from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict @@ -150,7 +150,7 @@ class OpenTelemetryV2Config(BaseSettings): ), ) - mapper_names: Annotated[List[str], NoDecode] = Field( + mapper_names: Annotated[list[str], NoDecode] = Field( default_factory=lambda: ["genai"], description=( "Ordered attribute vocabularies to emit. ``genai`` is the " @@ -168,7 +168,7 @@ class OpenTelemetryV2Config(BaseSettings): ), ) - baggage_promoted_keys: Annotated[List[str], NoDecode] = Field( + baggage_promoted_keys: Annotated[list[str], NoDecode] = Field( default_factory=lambda: list(BAGGAGE_PROMOTED_KEYS), validation_alias=AliasChoices("baggage_promoted_keys", "LITELLM_OTEL_BAGGAGE_PROMOTED_KEYS"), description=( @@ -179,7 +179,7 @@ class OpenTelemetryV2Config(BaseSettings): "YAML list)." ), ) - baggage_metadata_keys: Annotated[List[str], NoDecode] = Field( + baggage_metadata_keys: Annotated[list[str], NoDecode] = Field( default_factory=lambda: list(DEFAULT_BAGGAGE_METADATA_KEYS), validation_alias=AliasChoices("baggage_metadata_keys", "LITELLM_OTEL_BAGGAGE_METADATA_KEYS"), description=( @@ -189,7 +189,7 @@ class OpenTelemetryV2Config(BaseSettings): "``callback_settings.otel.baggage_metadata_keys`` in config.yaml." ), ) - baggage_team_metadata_keys: Annotated[List[str], NoDecode] = Field( + baggage_team_metadata_keys: Annotated[list[str], NoDecode] = Field( default_factory=lambda: list(DEFAULT_BAGGAGE_TEAM_METADATA_KEYS), validation_alias=AliasChoices("baggage_team_metadata_keys", "LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS"), description=( diff --git a/litellm/integrations/otel/model/metadata.py b/litellm/integrations/otel/model/metadata.py index 8ac02e936a7..bacd3e472f1 100644 --- a/litellm/integrations/otel/model/metadata.py +++ b/litellm/integrations/otel/model/metadata.py @@ -66,7 +66,7 @@ class RequestIdentity: metadata: Mapping[str, str] = field(default_factory=dict) @classmethod - def from_payload(cls, payload: "StandardLoggingPayload") -> "RequestIdentity": + def from_payload(cls, payload: StandardLoggingPayload) -> RequestIdentity: """Parse caller identity out of a closed request's payload metadata. ``provider_model`` is resolved here too (see :func:`resolve_provider_model`) @@ -90,7 +90,7 @@ class RequestIdentity: ) @classmethod - def from_user_api_key_auth(cls, auth: object) -> "RequestIdentity": + def from_user_api_key_auth(cls, auth: object) -> RequestIdentity: """Identity from a ``UserAPIKeyAuth`` (duck-typed to keep this module free of a proxy import). @@ -145,7 +145,7 @@ class RequestContext: return self.identity.provider_model @classmethod - def from_standard_logging_payload(cls, payload: "StandardLoggingPayload") -> "RequestContext": + def from_standard_logging_payload(cls, payload: StandardLoggingPayload) -> RequestContext: raw_meta = cast(Mapping[str, object], payload.get("metadata") or {}) hidden = cast(Mapping[str, object], payload.get("hidden_params") or {}) raw_response = payload.get("response") @@ -191,7 +191,7 @@ class LLMCallEvent: # The ``StandardLoggingPayload`` carried on a success/failure callback; ``None`` # at ``pre_call``, or when the call closed before any payload materialized (so # there is nothing to stamp on the span). - payload: "StandardLoggingPayload | None" + payload: StandardLoggingPayload | None # The ``standard_callback_dynamic_params`` routing the call to a per-tenant # tracer (its own exporter/endpoint), or ``None`` when the call isn't scoped. dynamic_params: Any @@ -205,7 +205,7 @@ class LLMCallEvent: time_to_first_chunk_seconds: float | None @classmethod - def from_dict(cls, kwargs: Mapping[str, Any]) -> "LLMCallEvent": + def from_dict(cls, kwargs: Mapping[str, Any]) -> LLMCallEvent: raw_payload = kwargs.get("standard_logging_object") payload = cast("StandardLoggingPayload", raw_payload) if raw_payload else None operation = resolve_operation(as_str(kwargs.get("call_type"))) @@ -235,7 +235,7 @@ def time_to_first_chunk_seconds(kwargs: Mapping[str, Any]) -> float | None: return completion_start - api_call_start -def _call_id(payload: "StandardLoggingPayload | None", kwargs: Mapping[str, Any]) -> str | None: +def _call_id(payload: StandardLoggingPayload | None, kwargs: Mapping[str, Any]) -> str | None: """The call id from the payload (when closed) or the bare kwargs (at pre_call).""" if payload is not None: call_id = as_str(payload.get("litellm_call_id")) or as_str(payload.get("id")) @@ -255,7 +255,7 @@ def model_from_request_data(data: object) -> str | None: return None -def resolve_provider_model(payload: "StandardLoggingPayload") -> str | None: +def resolve_provider_model(payload: StandardLoggingPayload) -> str | None: """The model litellm dispatched to the provider, from the payload. Prefers the explicit ``hidden_params.litellm_model_name`` (set on call paths diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index 892da9080f5..b3eba835a03 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -31,8 +31,6 @@ from litellm.integrations.otel.model.utils import ( # :mod:`metadata`; re-exported here so existing ``model.payloads`` imports keep # resolving it. __all__ = [ - "RequestContext", - "RequestIdentity", "GuardrailSpanData", "LLMCallSpanData", "LLMCost", @@ -41,6 +39,8 @@ __all__ = [ "MCPListToolsSpanData", "MCPToolCallSpanData", "ProxyRequestSpanData", + "RequestContext", + "RequestIdentity", "ServerInfo", "ServiceSpanData", "SpanError", @@ -72,7 +72,7 @@ class LLMRequestParams: seed: int | None = None @classmethod - def from_model_parameters(cls, params: Mapping[str, object]) -> "LLMRequestParams": + def from_model_parameters(cls, params: Mapping[str, object]) -> LLMRequestParams: max_tokens = as_int(params.get("max_tokens")) if max_tokens is None: max_tokens = as_int(params.get("max_completion_tokens")) @@ -121,7 +121,7 @@ class LLMCost: margin_total_amount: float | None = None @classmethod - def from_breakdown(cls, breakdown: Mapping[str, object] | None) -> "LLMCost": + def from_breakdown(cls, breakdown: Mapping[str, object] | None) -> LLMCost: b = breakdown or {} return cls( input=as_float(b.get("input_cost")), @@ -196,7 +196,7 @@ class GuardrailSpanData: _ERROR_STATUSES: ClassVar[frozenset[str]] = frozenset({"guardrail_intervened", "guardrail_failed_to_respond"}) @classmethod - def from_logging_entry(cls, entry: "StandardLoggingGuardrailInformation") -> "GuardrailSpanData": + def from_logging_entry(cls, entry: StandardLoggingGuardrailInformation) -> GuardrailSpanData: """Build from one ``standard_logging_guardrail_information`` entry. Reads the canonical, provider-agnostic ``StandardLoggingGuardrailInformation`` @@ -247,9 +247,9 @@ class ServiceSpanData: @classmethod def from_payload( cls, - payload: "ServiceLoggerPayload", + payload: ServiceLoggerPayload, event_metadata: Mapping[str, object] | None = None, - ) -> "ServiceSpanData": + ) -> ServiceSpanData: # ``payload.service`` is a ``ServiceTypes(str, Enum)`` and ``error`` is # ``Optional[str]`` on the Pydantic model — no defensive reads needed. # ``event_metadata`` is sanitized: the legacy service decorators pass raw @@ -314,10 +314,10 @@ class LLMCallSpanData: @classmethod def from_standard_logging_payload( cls, - payload: "StandardLoggingPayload", + payload: StandardLoggingPayload, capture_content: bool = False, time_to_first_chunk_seconds: float | None = None, - ) -> "LLMCallSpanData": + ) -> LLMCallSpanData: params = cast(Mapping[str, object], payload.get("model_parameters") or {}) # The single parse of the request's metadata — the request-vs-provider # model split, the response model, api base, and identity all come from @@ -387,8 +387,8 @@ class MCPToolCallSpanData: @classmethod def from_standard_logging_payload( - cls, payload: "StandardLoggingPayload", capture_content: bool = False - ) -> "MCPToolCallSpanData": + cls, payload: StandardLoggingPayload, capture_content: bool = False + ) -> MCPToolCallSpanData: meta = _mcp_tool_call_metadata(cast(Mapping[str, object], payload)) return cls( operation=resolve_operation(as_str(payload.get("call_type"))), @@ -567,7 +567,7 @@ def _finish_reasons(choices: tuple[Mapping[str, object], ...]) -> tuple[str, ... return tuple(r for c in choices if (r := as_str(c.get("finish_reason")))) -def _parse_error(payload: "StandardLoggingPayload") -> SpanError | None: +def _parse_error(payload: StandardLoggingPayload) -> SpanError | None: """A ``SpanError`` for a failed request, or ``None`` on success.""" if payload.get("status") != "failure": return None diff --git a/litellm/integrations/otel/model/utils.py b/litellm/integrations/otel/model/utils.py index ab54a558a9a..fb35e9abf51 100644 --- a/litellm/integrations/otel/model/utils.py +++ b/litellm/integrations/otel/model/utils.py @@ -65,7 +65,7 @@ def as_str_tuple(value: object) -> tuple[str, ...] | None: return None -def to_ns(value: datetime | float | int | None) -> int | None: +def to_ns(value: datetime | float | None) -> int | None: """Coerce a datetime / epoch value to integer nanoseconds.""" if value is None: return None @@ -76,7 +76,7 @@ def to_ns(value: datetime | float | int | None) -> int | None: return None -def to_seconds(value: datetime | float | int | str | None) -> float | None: +def to_seconds(value: datetime | float | str | None) -> float | None: """Coerce a datetime / epoch / formatted-string value to epoch seconds.""" if value is None: return None diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index cd256852c13..266277af648 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -11,7 +11,7 @@ identical metrics. The attribute cardinality filter is reused from v1 by import from collections.abc import Mapping from dataclasses import dataclass from datetime import datetime -from typing import Any, Final, FrozenSet, Optional, TypeAlias +from typing import Any, Final, TypeAlias from opentelemetry.metrics import Histogram, Meter @@ -182,11 +182,11 @@ class GenAIMetricRecorder: survives. """ - def __init__(self, metrics: GenAIMetrics, callback_name: Optional[str] = None) -> None: + def __init__(self, metrics: GenAIMetrics, callback_name: str | None = None) -> None: self._metrics = metrics self._callback_name = callback_name - self._include: Optional[FrozenSet[str]] = None - self._exclude: Optional[FrozenSet[str]] = None + self._include: frozenset[str] | None = None + self._exclude: frozenset[str] | None = None self._filter_resolved = False def record( diff --git a/litellm/integrations/otel/plumbing/routing.py b/litellm/integrations/otel/plumbing/routing.py index 3c99e200bb4..f8972733e21 100644 --- a/litellm/integrations/otel/plumbing/routing.py +++ b/litellm/integrations/otel/plumbing/routing.py @@ -61,7 +61,7 @@ class TenantTracerCache: self._config = config self._callback_name = callback_name self._tracer_name = tracer_name - self._providers: "OrderedDict[tuple[tuple[str, str], ...], TracerProvider]" = OrderedDict() + self._providers: OrderedDict[tuple[tuple[str, str], ...], TracerProvider] = OrderedDict() def tracer_for(self, default: Tracer, dynamic_params: Any) -> Tracer: """Return the tracer for this request. diff --git a/litellm/integrations/otel/presets/__init__.py b/litellm/integrations/otel/presets/__init__.py index 31690d37b93..6448a30384c 100644 --- a/litellm/integrations/otel/presets/__init__.py +++ b/litellm/integrations/otel/presets/__init__.py @@ -62,12 +62,12 @@ def dynamic_otlp_headers( __all__ = [ - "PRESET_BY_CALLBACK", "DYNAMIC_HEADERS_BY_CALLBACK", + "PRESET_BY_CALLBACK", "Preset", - "dynamic_otlp_headers", "agentops_preset", "arize_preset", + "dynamic_otlp_headers", "langfuse_preset", "langtrace_preset", "levo_preset", diff --git a/litellm/integrations/otel/runtime.py b/litellm/integrations/otel/runtime.py index 13b36a5b1eb..1a9efb1e3b0 100644 --- a/litellm/integrations/otel/runtime.py +++ b/litellm/integrations/otel/runtime.py @@ -10,11 +10,11 @@ identity unconditionally. from collections.abc import Callable, Iterator from contextlib import contextmanager from functools import cache -from typing import Any, Optional +from typing import Any @cache -def _otel_runtime() -> "Optional[tuple[Callable[[str], Any], Callable[..., None]]]": +def _otel_runtime() -> "tuple[Callable[[str], Any], Callable[..., None]] | None": """Resolve the SDK-backed hooks once and cache the outcome, absence included. CPython never caches a failed import, so without this memoization every call diff --git a/litellm/integrations/posthog.py b/litellm/integrations/posthog.py index e519736e162..b61eeb8198f 100644 --- a/litellm/integrations/posthog.py +++ b/litellm/integrations/posthog.py @@ -12,16 +12,16 @@ For batching specific details see CustomBatchLogger class import asyncio import atexit import os -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.integrations.custom_batch_logger import CustomBatchLogger -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.integrations.posthog_mock_client import ( - should_use_posthog_mock, create_mock_posthog_client, + should_use_posthog_mock, ) +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, @@ -72,7 +72,7 @@ class PostHogLogger(CustomBatchLogger): super().__init__(**kwargs, flush_lock=None, batch_size=POSTHOG_MAX_BATCH_SIZE) except Exception as e: - verbose_logger.exception(f"PostHog: Got exception on init PostHog client {str(e)}") + verbose_logger.exception(f"PostHog: Got exception on init PostHog client {e!s}") raise e def log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -107,7 +107,7 @@ class PostHogLogger(CustomBatchLogger): verbose_logger.debug("PostHog: Sync event successfully sent") except Exception as e: - verbose_logger.exception(f"PostHog Sync Layer Error - {str(e)}") + verbose_logger.exception(f"PostHog Sync Layer Error - {e!s}") async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: @@ -115,8 +115,7 @@ class PostHogLogger(CustomBatchLogger): self._ensure_async_setup() # Lazy initialization await self._log_async_event(kwargs, response_obj, start_time, end_time) except Exception as e: - verbose_logger.exception(f"PostHog Layer Error - {str(e)}") - pass + verbose_logger.exception(f"PostHog Layer Error - {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: @@ -124,8 +123,7 @@ class PostHogLogger(CustomBatchLogger): self._ensure_async_setup() # Lazy initialization await self._log_async_event(kwargs, response_obj, start_time, end_time) except Exception as e: - verbose_logger.exception(f"PostHog Layer Error - {str(e)}") - pass + verbose_logger.exception(f"PostHog Layer Error - {e!s}") async def _log_async_event(self, kwargs, response_obj=None, start_time=0.0, end_time=0.0): # Note: response_obj, start_time, end_time not used - all data comes from kwargs @@ -139,7 +137,7 @@ class PostHogLogger(CustomBatchLogger): if len(self.log_queue) >= self.batch_size: await self.flush_queue() - def create_posthog_event_payload(self, kwargs: Dict[str, Any]) -> PostHogEventPayload: + def create_posthog_event_payload(self, kwargs: dict[str, Any]) -> PostHogEventPayload: """ Helper function to create a PostHog event payload for logging @@ -149,7 +147,7 @@ class PostHogLogger(CustomBatchLogger): Returns: PostHogEventPayload: defined in types.py """ - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: raise ValueError("standard_logging_object not found in kwargs") @@ -173,9 +171,9 @@ class PostHogLogger(CustomBatchLogger): def _create_posthog_properties( self, standard_logging_object: StandardLoggingPayload, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], event_name: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Create PostHog properties following LLM Analytics spec""" properties = {} @@ -220,7 +218,7 @@ class PostHogLogger(CustomBatchLogger): return properties - def _add_trace_properties(self, properties: Dict[str, Any], kwargs: Dict[str, Any]): + def _add_trace_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]): standard_logging_object = self._safe_get(kwargs, "standard_logging_object", {}) trace_id = self._safe_get(standard_logging_object, "trace_id", self._safe_uuid()) @@ -234,7 +232,7 @@ class PostHogLogger(CustomBatchLogger): if parent_id: properties["$ai_parent_id"] = parent_id - def _add_custom_metadata_properties(self, properties: Dict[str, Any], kwargs: Dict[str, Any]): + def _add_custom_metadata_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]): """Add custom metadata fields to PostHog properties""" metadata = self._extract_metadata(kwargs) if not isinstance(metadata, dict): @@ -269,7 +267,6 @@ class PostHogLogger(CustomBatchLogger): "deployment", "model_info", "api_base", - "caching_groups", "hidden_params", "parent_run_id", "parent_id", @@ -280,7 +277,7 @@ class PostHogLogger(CustomBatchLogger): if key not in litellm_internal_fields: properties[key] = value - def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: Dict[str, Any]) -> str: + def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: dict[str, Any]) -> str: metadata = self._extract_metadata(kwargs) user_id = self._safe_get(metadata, "user_id") if user_id: @@ -294,7 +291,7 @@ class PostHogLogger(CustomBatchLogger): return self._safe_uuid() - def _get_credentials_for_request(self, kwargs: Dict[str, Any]) -> Tuple[Optional[str], Optional[str]]: + def _get_credentials_for_request(self, kwargs: dict[str, Any]) -> tuple[str | None, str | None]: """ Get PostHog credentials for this request. @@ -307,7 +304,7 @@ class PostHogLogger(CustomBatchLogger): Returns: tuple[str, str]: (api_key, api_url) """ - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = kwargs.get( "standard_callback_dynamic_params", None ) @@ -337,7 +334,7 @@ class PostHogLogger(CustomBatchLogger): verbose_logger.debug("[POSTHOG MOCK] Mock mode enabled - API calls will be intercepted") # Group events by credentials for batch sending - batches_by_credentials: Dict[tuple[str, str], list] = {} + batches_by_credentials: dict[tuple[str, str], list] = {} for item in self.log_queue: key = (item["api_key"], item["api_url"]) if key not in batches_by_credentials: @@ -370,7 +367,7 @@ class PostHogLogger(CustomBatchLogger): else: verbose_logger.debug(f"PostHog: Batch of {len(self.log_queue)} events successfully sent") except Exception as e: - verbose_logger.exception(f"PostHog Error sending batch API - {str(e)}") + verbose_logger.exception(f"PostHog Error sending batch API - {e!s}") def _ensure_async_setup(self): if not self._async_initialized: @@ -380,17 +377,17 @@ class PostHogLogger(CustomBatchLogger): self._async_initialized = True verbose_logger.debug("PostHog: Async components initialized") except Exception as e: - verbose_logger.error(f"PostHog: Failed to initialize async components: {str(e)}") + verbose_logger.error(f"PostHog: Failed to initialize async components: {e!s}") raise - def _extract_metadata(self, kwargs: Dict[str, Any]) -> Dict[str, Any]: + def _extract_metadata(self, kwargs: dict[str, Any]) -> dict[str, Any]: litellm_params = kwargs.get("litellm_params", {}) or {} return litellm_params.get("metadata", {}) or {} def _safe_uuid(self) -> str: return str(uuid.uuid4()) - def _create_posthog_payload(self, events: list, api_key: str) -> Dict[str, Any]: + def _create_posthog_payload(self, events: list, api_key: str) -> dict[str, Any]: return {"api_key": api_key, "batch": events} def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any: @@ -415,7 +412,7 @@ class PostHogLogger(CustomBatchLogger): try: # Group events by credentials (same logic as async_send_batch) - batches_by_credentials: Dict[Tuple[str, str], list] = {} + batches_by_credentials: dict[tuple[str, str], list] = {} for item in self.log_queue: key = (item["api_key"], item["api_url"]) if key not in batches_by_credentials: @@ -448,4 +445,4 @@ class PostHogLogger(CustomBatchLogger): self.log_queue.clear() except Exception as e: - verbose_logger.error(f"PostHog: Error flushing events on exit: {str(e)}") + verbose_logger.error(f"PostHog: Error flushing events on exit: {e!s}") diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index f0251586c4f..c84a6c34f1f 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -12,12 +12,7 @@ from datetime import datetime, timedelta from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, - Union, cast, ) @@ -137,7 +132,7 @@ class PrometheusLogger(CustomLogger): _ADDITIVE_GUARDRAIL_MODES = frozenset((GuardrailEventHooks.pre_call.value, GuardrailEventHooks.post_call.value)) @staticmethod - def get_instance() -> Optional["PrometheusLogger"]: + def get_instance() -> PrometheusLogger | None: """Find the PrometheusLogger instance from litellm.callbacks, if registered.""" import litellm @@ -170,7 +165,7 @@ class PrometheusLogger(CustomLogger): # logger init time pins the label set for the lifetime of the # logger so toggling these flags only takes effect after a # restart, keeping init-time and runtime label sets in sync. - self._cached_metric_labels: Dict[str, List[str]] = {} + self._cached_metric_labels: dict[str, list[str]] = {} _custom_buckets = litellm.prometheus_latency_buckets self.latency_buckets = tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS @@ -688,10 +683,10 @@ class PrometheusLogger(CustomLogger): ) except Exception as e: - print_verbose(f"Got exception on init prometheus client {str(e)}") + print_verbose(f"Got exception on init prometheus client {e!s}") raise e - def _parse_prometheus_config(self) -> Dict[str, List[str]]: + def _parse_prometheus_config(self) -> dict[str, list[str]]: """Parse prometheus metrics configuration for label filtering and enabled metrics""" import litellm from litellm.types.integrations.prometheus import PrometheusMetricsConfig @@ -770,7 +765,7 @@ class PrometheusLogger(CustomLogger): ) return builtin_labels | _NON_ENUM_METRIC_LABELS | custom_metadata_labels | custom_tag_labels - def _validate_all_configurations(self, parsed_configs: List) -> ValidationResults: + def _validate_all_configurations(self, parsed_configs: list) -> ValidationResults: """Validate all metric configurations and return collected errors""" metric_errors = [] label_errors = [] @@ -791,7 +786,7 @@ class PrometheusLogger(CustomLogger): return ValidationResults(metric_errors=metric_errors, label_errors=label_errors) - def _validate_single_metric_name(self, metric_name: str) -> Optional[MetricValidationError]: + def _validate_single_metric_name(self, metric_name: str) -> MetricValidationError | None: """Validate a single metric name""" from typing import get_args @@ -802,7 +797,7 @@ class PrometheusLogger(CustomLogger): ) return None - def _validate_single_metric_labels(self, metric_name: str, labels: List[str]) -> Optional[LabelValidationError]: + def _validate_single_metric_labels(self, metric_name: str, labels: list[str]) -> LabelValidationError | None: """Validate labels for a single metric""" from typing import cast @@ -820,7 +815,7 @@ class PrometheusLogger(CustomLogger): ) return None - def _build_label_filters(self, parsed_configs: List) -> Dict[str, List[str]]: + def _build_label_filters(self, parsed_configs: list) -> dict[str, list[str]]: """Build label filters from validated configurations""" label_filters = {} @@ -833,7 +828,7 @@ class PrometheusLogger(CustomLogger): return label_filters - def _validate_configured_metric_labels(self, metric_name: str, labels: List[str]): + def _validate_configured_metric_labels(self, metric_name: str, labels: list[str]): """ Ensure that all the configured labels are valid for the metric @@ -929,7 +924,7 @@ class PrometheusLogger(CustomLogger): verbose_logger.error(label_error.message) def _pretty_print_invalid_labels_error( - self, metric_name: str, invalid_labels: List[str], valid_labels: List[str] + self, metric_name: str, invalid_labels: list[str], valid_labels: list[str] ) -> None: """Pretty print error message for invalid labels using rich""" try: @@ -1025,7 +1020,7 @@ class PrometheusLogger(CustomLogger): ) raise ValueError(error.message) - def _pretty_print_prometheus_config(self, label_filters: Dict[str, List[str]]) -> None: + def _pretty_print_prometheus_config(self, label_filters: dict[str, list[str]]) -> None: """Pretty print the processed prometheus configuration using rich""" try: from rich.console import Console @@ -1123,7 +1118,7 @@ class PrometheusLogger(CustomLogger): return factory - def get_labels_for_metric(self, metric_name: DEFINED_PROMETHEUS_METRICS) -> List[str]: + def get_labels_for_metric(self, metric_name: DEFINED_PROMETHEUS_METRICS) -> list[str]: """ Get the labels for a metric, filtered if configured. @@ -1192,7 +1187,7 @@ class PrometheusLogger(CustomLogger): self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ) -> None: """Record litellm_overhead_with_guardrails_latency_metric (seconds): SDK overhead + pre/post-call guardrail time. Recorded outside the SDK-overhead gate so @@ -1218,7 +1213,7 @@ class PrometheusLogger(CustomLogger): self, metric: Any, metric_name: DEFINED_PROMETHEUS_METRICS, - labels: Dict[str, Optional[str]], + labels: dict[str, str | None], ) -> None: """ Cap the cardinality of metrics that include the ``end_user`` label. @@ -1253,7 +1248,7 @@ class PrometheusLogger(CustomLogger): counter: Any, metric_name: DEFINED_PROMETHEUS_METRICS, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, amount: float = 1.0, ) -> None: _labels = prometheus_label_factory( @@ -1272,7 +1267,7 @@ class PrometheusLogger(CustomLogger): ) # unpack kwargs - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_payload is None or not isinstance(standard_logging_payload, dict): raise ValueError(f"standard_logging_object is required, got={standard_logging_payload}") @@ -1466,15 +1461,15 @@ class PrometheusLogger(CustomLogger): def _increment_token_metrics( self, standard_logging_payload: StandardLoggingPayload, - end_user_id: Optional[str], - user_api_key: Optional[str], - user_api_key_alias: Optional[str], - model: Optional[str], - user_api_team: Optional[str], - user_api_team_alias: Optional[str], - user_id: Optional[str], + end_user_id: str | None, + user_api_key: str | None, + user_api_key_alias: str | None, + model: str | None, + user_api_team: str | None, + user_api_team_alias: str | None, + user_id: str | None, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ): verbose_logger.debug("prometheus Logging - Enters token metrics function") # token metrics @@ -1520,7 +1515,7 @@ class PrometheusLogger(CustomLogger): self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ) -> None: """ Increment per-token-type counters from the Usage object that providers @@ -1542,7 +1537,7 @@ class PrometheusLogger(CustomLogger): cache_creation_detail_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details) - detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [ + detail_metrics: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [ ( self.litellm_input_cached_tokens_metric, "litellm_input_cached_tokens_metric", @@ -1643,7 +1638,7 @@ class PrometheusLogger(CustomLogger): self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ): """ Increment cache-related Prometheus metrics based on cache hit/miss status. @@ -1796,14 +1791,14 @@ class PrometheusLogger(CustomLogger): async def _increment_remaining_budget_metrics( self, - user_api_team: Optional[str], - user_api_team_alias: Optional[str], - user_api_key: Optional[str], - user_api_key_alias: Optional[str], + user_api_team: str | None, + user_api_team_alias: str | None, + user_api_key: str | None, + user_api_key_alias: str | None, litellm_params: dict, response_cost: float, - user_id: Optional[str] = None, - user_api_key_org_id: Optional[str] = None, + user_id: str | None = None, + user_api_key_org_id: str | None = None, ): if ( isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric) @@ -1876,16 +1871,16 @@ class PrometheusLogger(CustomLogger): def _increment_top_level_request_and_spend_metrics( self, - end_user_id: Optional[str], - user_api_key: Optional[str], - user_api_key_alias: Optional[str], - model: Optional[str], - user_api_team: Optional[str], - user_api_team_alias: Optional[str], - user_id: Optional[str], + end_user_id: str | None, + user_api_key: str | None, + user_api_key_alias: str | None, + model: str | None, + user_api_team: str | None, + user_api_team_alias: str | None, + user_id: str | None, response_cost: float, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ): PrometheusLogger._inc_labeled_counter( self, @@ -1934,11 +1929,11 @@ class PrometheusLogger(CustomLogger): def _set_virtual_key_rate_limit_metrics( self, - user_api_key: Optional[str], - user_api_key_alias: Optional[str], + user_api_key: str | None, + user_api_key_alias: str | None, kwargs: dict, metadata: dict, - model_id: Optional[str] = None, + model_id: str | None = None, ): from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, @@ -1995,17 +1990,17 @@ class PrometheusLogger(CustomLogger): def _set_latency_metrics( self, kwargs: dict, - model: Optional[str], - user_api_key: Optional[str], - user_api_key_alias: Optional[str], - user_api_team: Optional[str], - user_api_team_alias: Optional[str], + model: str | None, + user_api_key: str | None, + user_api_key_alias: str | None, + user_api_team: str | None, + user_api_team_alias: str | None, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ): # latency metrics end_time: datetime = kwargs.get("end_time") or datetime.now() - start_time: Optional[datetime] = kwargs.get("start_time") + start_time: datetime | None = kwargs.get("start_time") api_call_start_time = kwargs.get("api_call_start_time", None) completion_start_time = kwargs.get("completion_start_time", None) time_to_first_token_seconds = self._safe_duration_seconds( @@ -2137,16 +2132,14 @@ class PrometheusLogger(CustomLogger): response_cost=0, ) except Exception as e: - verbose_logger.exception("prometheus Layer Error(): Exception occured - {}".format(str(e))) - pass - pass + verbose_logger.exception(f"prometheus Layer Error(): Exception occured - {e!s}") def _extract_status_code( self, - kwargs: Optional[dict] = None, - enum_values: Optional[Any] = None, - exception: Optional[Exception] = None, - ) -> Optional[int]: + kwargs: dict | None = None, + enum_values: Any | None = None, + exception: Exception | None = None, + ) -> int | None: """ Extract HTTP status code from various input formats for validation. @@ -2196,8 +2189,8 @@ class PrometheusLogger(CustomLogger): def _is_invalid_api_key_request( self, - status_code: Optional[int], - exception: Optional[Exception] = None, + status_code: int | None, + exception: Exception | None = None, ) -> bool: """ Determine if a request has an invalid API key based on status code and exception. @@ -2235,11 +2228,11 @@ class PrometheusLogger(CustomLogger): def _should_skip_metrics_for_invalid_key( self, - kwargs: Optional[dict] = None, - user_api_key_dict: Optional[Any] = None, - enum_values: Optional[Any] = None, - standard_logging_payload: Optional[Union[dict, StandardLoggingPayload]] = None, - exception: Optional[Exception] = None, + kwargs: dict | None = None, + user_api_key_dict: Any | None = None, + enum_values: Any | None = None, + standard_logging_payload: dict | StandardLoggingPayload | None = None, + exception: Exception | None = None, ) -> bool: """ Determine if Prometheus metrics should be skipped for invalid API key requests. @@ -2277,7 +2270,7 @@ class PrometheusLogger(CustomLogger): return False @staticmethod - def _extract_api_provider_from_request_data(request_data: dict) -> Optional[str]: + def _extract_api_provider_from_request_data(request_data: dict) -> str | None: """ Best-effort provider for the client-side failure path. @@ -2318,7 +2311,7 @@ class PrometheusLogger(CustomLogger): request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, + traceback_str: str | None = None, ): """ Track client side failures @@ -2390,8 +2383,7 @@ class PrometheusLogger(CustomLogger): ) except Exception as e: - verbose_logger.exception("prometheus Layer Error(): Exception occured - {}".format(str(e))) - pass + verbose_logger.exception(f"prometheus Layer Error(): Exception occured - {e!s}") async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): """ @@ -2401,7 +2393,6 @@ class PrometheusLogger(CustomLogger): double-counting. It is incremented in async_log_success_event which fires for all successful requests (both streaming and non-streaming). """ - pass def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any: """Get value from dict or Pydantic model.""" @@ -2411,7 +2402,7 @@ class PrometheusLogger(CustomLogger): return obj.get(key, default) return getattr(obj, key, default) - def _extract_deployment_failure_label_values(self, request_kwargs: dict) -> Dict[str, Optional[str]]: + def _extract_deployment_failure_label_values(self, request_kwargs: dict) -> dict[str, str | None]: """ Extract label values for deployment failure metrics from all available sources in request_kwargs. Falls back to litellm_params metadata and @@ -2436,7 +2427,7 @@ class PrometheusLogger(CustomLogger): # Extract user_api_key_auth if present (proxy injects this, skipped in merge) user_api_key_auth = _litellm_params_metadata.get("user_api_key_auth") - def _get_api_key_alias() -> Optional[str]: + def _get_api_key_alias() -> str | None: val = _metadata.get("user_api_key_alias") if val is not None: return val @@ -2447,7 +2438,7 @@ class PrometheusLogger(CustomLogger): return getattr(user_api_key_auth, "key_alias", None) return None - def _get_team_id() -> Optional[str]: + def _get_team_id() -> str | None: val = _metadata.get("user_api_key_team_id") if val is not None: return val @@ -2458,7 +2449,7 @@ class PrometheusLogger(CustomLogger): return getattr(user_api_key_auth, "team_id", None) return None - def _get_team_alias() -> Optional[str]: + def _get_team_alias() -> str | None: val = _metadata.get("user_api_key_team_alias") if val is not None: return val @@ -2469,7 +2460,7 @@ class PrometheusLogger(CustomLogger): return getattr(user_api_key_auth, "team_alias", None) return None - def _get_hashed_api_key() -> Optional[str]: + def _get_hashed_api_key() -> str | None: val = _metadata.get("user_api_key_hash") if val is not None: return val @@ -2616,20 +2607,17 @@ class PrometheusLogger(CustomLogger): label_context=_deployment_label_ctx, ) - pass except Exception as e: - verbose_logger.debug( - "Prometheus Error: set_llm_deployment_failure_metrics. Exception occured - {}".format(str(e)) - ) + verbose_logger.debug(f"Prometheus Error: set_llm_deployment_failure_metrics. Exception occured - {e!s}") def _set_deployment_tpm_rpm_limit_metrics( self, model_info: dict, litellm_params: dict, - litellm_model_name: Optional[str], - model_id: Optional[str], - api_base: Optional[str], - llm_provider: Optional[str], + litellm_model_name: str | None, + model_id: str | None, + api_base: str | None, + llm_provider: str | None, ): """ Set the deployment TPM and RPM limits metrics @@ -2665,7 +2653,7 @@ class PrometheusLogger(CustomLogger): self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ) -> None: """ Populate ``litellm_remaining_tokens_metric`` / @@ -2735,7 +2723,7 @@ class PrometheusLogger(CustomLogger): self.litellm_remaining_requests_metric.labels(**_labels).set(remaining_requests) except Exception as e: verbose_logger.exception( - "Prometheus Error: _async_set_router_remaining_metrics. Exception occured - {}".format(str(e)) + f"Prometheus Error: _async_set_router_remaining_metrics. Exception occured - {e!s}" ) def set_llm_deployment_success_metrics( @@ -2745,11 +2733,11 @@ class PrometheusLogger(CustomLogger): end_time, enum_values: UserAPIKeyLabelValues, output_tokens: float = 1.0, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ): try: verbose_logger.debug("setting remaining tokens requests metric") - standard_logging_payload: Optional[StandardLoggingPayload] = request_kwargs.get("standard_logging_object") + standard_logging_payload: StandardLoggingPayload | None = request_kwargs.get("standard_logging_object") if standard_logging_payload is None: return @@ -2780,8 +2768,8 @@ class PrometheusLogger(CustomLogger): llm_provider=llm_provider, ) - remaining_requests: Optional[int] = None - remaining_tokens: Optional[int] = None + remaining_requests: int | None = None + remaining_tokens: int | None = None if additional_headers := standard_logging_payload["hidden_params"]["additional_headers"]: # OpenAI / OpenAI Compatible headers remaining_requests = additional_headers.get("x_ratelimit_remaining_requests", None) @@ -2853,7 +2841,7 @@ class PrometheusLogger(CustomLogger): # Track deployment Latency response_ms: timedelta = end_time - start_time - time_to_first_token_response_time: Optional[timedelta] = None + time_to_first_token_response_time: timedelta | None = None if request_kwargs.get("stream", None) is not None and request_kwargs["stream"] is True: # only log ttft for streaming request @@ -2879,9 +2867,7 @@ class PrometheusLogger(CustomLogger): self.litellm_deployment_latency_per_output_token.labels(**_labels).observe(latency_per_token) except Exception as e: - verbose_logger.exception( - "Prometheus Error: set_llm_deployment_success_metrics. Exception occured - {}".format(str(e)) - ) + verbose_logger.exception(f"Prometheus Error: set_llm_deployment_success_metrics. Exception occured - {e!s}") return def _record_guardrail_metrics( @@ -2889,7 +2875,7 @@ class PrometheusLogger(CustomLogger): guardrail_name: str, latency_seconds: float, status: str, - error_type: Optional[str], + error_type: str | None, hook_type: str, ): """ @@ -2926,7 +2912,7 @@ class PrometheusLogger(CustomLogger): hook_type=hook_type, ).inc() except Exception as e: - verbose_logger.debug(f"Error recording guardrail metrics: {str(e)}") + verbose_logger.debug(f"Error recording guardrail metrics: {e!s}") ######################################## # Managed Batch Metric Recording Methods @@ -2934,11 +2920,11 @@ class PrometheusLogger(CustomLogger): def record_managed_batch_created( self, - model: Optional[str], - api_provider: Optional[str], - user: Optional[str], - user_email: Optional[str], - api_key_alias: Optional[str], + model: str | None, + api_provider: str | None, + user: str | None, + user_email: str | None, + api_key_alias: str | None, ): try: self.litellm_managed_batch_created_total.labels( @@ -2956,9 +2942,9 @@ class PrometheusLogger(CustomLogger): size_bytes: int, purpose: str, file_type: str, - model: Optional[str] = None, - api_provider: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + api_provider: str | None = None, + user: str | None = None, ): """Record the size of a managed file. Uses a gauge (last-seen value per label combination).""" try: @@ -2975,8 +2961,8 @@ class PrometheusLogger(CustomLogger): def record_managed_batch_duration( self, duration_seconds: float, - model: Optional[str] = None, - api_provider: Optional[str] = None, + model: str | None = None, + api_provider: str | None = None, ): try: self.litellm_managed_batch_duration_seconds.labels( @@ -2988,11 +2974,11 @@ class PrometheusLogger(CustomLogger): def record_managed_file_created( self, - model: Optional[str], - api_provider: Optional[str], - user: Optional[str], - user_email: Optional[str], - api_key_alias: Optional[str], + model: str | None, + api_provider: str | None, + user: str | None, + user_email: str | None, + api_key_alias: str | None, ): try: self.litellm_managed_file_created_total.labels( @@ -3015,7 +3001,7 @@ class PrometheusLogger(CustomLogger): def record_check_batch_cost_run( self, jobs_polled: int, - processed_models: Optional[List[Tuple[Optional[str], Optional[str]]]] = None, + processed_models: list[tuple[str | None, str | None]] | None = None, ): """ Record CheckBatchCost polling metrics. @@ -3088,8 +3074,8 @@ class PrometheusLogger(CustomLogger): @staticmethod def _extract_rate_limit_labels( - exception: Optional[Exception], - ) -> Tuple[Optional[str], Optional[str]]: + exception: Exception | None, + ) -> tuple[str | None, str | None]: """ Pull the unified ``category`` / ``rate_limit_type`` fields off any exception that declares them (``litellm.RateLimitError`` and bare- @@ -3129,7 +3115,7 @@ class PrometheusLogger(CustomLogger): metadata=_metadata ) _new_model = kwargs.get("model") - _tags = cast(List[str], kwargs.get("tags") or []) + _tags = cast(list[str], kwargs.get("tags") or []) enum_values = UserAPIKeyLabelValues( requested_model=original_model_group, @@ -3167,7 +3153,7 @@ class PrometheusLogger(CustomLogger): _new_model = kwargs.get("model") _metadata_key = get_metadata_variable_name_from_kwargs(kwargs) _metadata = kwargs.get(_metadata_key) or {} - _tags = cast(List[str], kwargs.get("tags") or []) + _tags = cast(list[str], kwargs.get("tags") or []) standard_metadata: StandardLoggingMetadata = StandardLoggingPayloadSetup.get_standard_logging_metadata( metadata=_metadata ) @@ -3196,8 +3182,8 @@ class PrometheusLogger(CustomLogger): self, state: int, litellm_model_name: str, - model_id: Optional[str], - api_base: Optional[str], + model_id: str | None, + api_base: str | None, api_provider: str, ): """ @@ -3227,8 +3213,8 @@ class PrometheusLogger(CustomLogger): def set_deployment_partial_outage( self, litellm_model_name: str, - model_id: Optional[str], - api_base: Optional[str], + model_id: str | None, + api_base: str | None, api_provider: str, ): self.set_litellm_deployment_state(1, litellm_model_name, model_id, api_base, api_provider) @@ -3236,8 +3222,8 @@ class PrometheusLogger(CustomLogger): def set_deployment_complete_outage( self, litellm_model_name: str, - model_id: Optional[str], - api_base: Optional[str], + model_id: str | None, + api_base: str | None, api_provider: str, ): self.set_litellm_deployment_state(2, litellm_model_name, model_id, api_base, api_provider) @@ -3281,7 +3267,7 @@ class PrometheusLogger(CustomLogger): ) ) - def _safe_get_remaining_budget(self, max_budget: Optional[float], spend: Optional[float]) -> float: + def _safe_get_remaining_budget(self, max_budget: float | None, spend: float | None) -> float: if max_budget is None: return float("inf") @@ -3292,8 +3278,8 @@ class PrometheusLogger(CustomLogger): async def _initialize_budget_metrics( self, - data_fetch_function: Callable[..., Awaitable[Tuple[List[Any], Optional[int]]]], - set_metrics_function: Callable[[List[Any]], Awaitable[None]], + data_fetch_function: Callable[..., Awaitable[tuple[list[Any], int | None]]], + set_metrics_function: Callable[[list[Any]], Awaitable[None]], data_type: Literal["teams", "keys", "users", "orgs"], ): """ @@ -3329,7 +3315,7 @@ class PrometheusLogger(CustomLogger): await set_metrics_function(data) except Exception as e: - verbose_logger.exception(f"Error initializing {data_type} budget metrics: {str(e)}") + verbose_logger.exception(f"Error initializing {data_type} budget metrics: {e!s}") async def _initialize_team_budget_metrics(self): """ @@ -3344,7 +3330,7 @@ class PrometheusLogger(CustomLogger): verbose_logger.debug("Prometheus: skipping team metrics initialization, DB not initialized") return - async def fetch_teams(page_size: int, page: int) -> Tuple[List[LiteLLM_TeamTable], Optional[int]]: + async def fetch_teams(page_size: int, page: int) -> tuple[list[LiteLLM_TeamTable], int | None]: teams, total_count = await get_paginated_teams(prisma_client=prisma_client, page_size=page_size, page=page) if total_count is None: total_count = len(teams) @@ -3372,9 +3358,9 @@ class PrometheusLogger(CustomLogger): async def fetch_keys( page_size: int, page: int - ) -> Tuple[ - List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], - Optional[int], + ) -> tuple[ + list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken], + int | None, ]: key_list_response = await _list_key_helper( prisma_client=prisma_client, @@ -3410,7 +3396,7 @@ class PrometheusLogger(CustomLogger): verbose_logger.debug("Prometheus: skipping user metrics initialization, DB not initialized") return - async def fetch_users(page_size: int, page: int) -> Tuple[List[LiteLLM_UserTable], Optional[int]]: + async def fetch_users(page_size: int, page: int) -> tuple[list[LiteLLM_UserTable], int | None]: skip = (page - 1) * page_size users = await UserRepository(prisma_client).table.find_many( skip=skip, @@ -3436,7 +3422,7 @@ class PrometheusLogger(CustomLogger): verbose_logger.debug("Prometheus: skipping org metrics initialization, DB not initialized") return - async def fetch_orgs(page_size: int, page: int) -> Tuple[list, Optional[int]]: + async def fetch_orgs(page_size: int, page: int) -> tuple[list, int | None]: skip = (page - 1) * page_size orgs = await OrganizationRepository(prisma_client).table.find_many( skip=skip, @@ -3520,20 +3506,20 @@ class PrometheusLogger(CustomLogger): self.litellm_teams_count_metric.set(total_teams) verbose_logger.debug(f"Prometheus: set litellm_teams_count to {total_teams}") except Exception as e: - verbose_logger.exception(f"Error initializing user/team count metrics: {str(e)}") + verbose_logger.exception(f"Error initializing user/team count metrics: {e!s}") - async def _set_key_list_budget_metrics(self, keys: List[Union[str, UserAPIKeyAuth]]): + async def _set_key_list_budget_metrics(self, keys: list[str | UserAPIKeyAuth]): """Helper function to set budget metrics for a list of keys""" for key in keys: if isinstance(key, UserAPIKeyAuth): self._set_key_budget_metrics(key) - async def _set_team_list_budget_metrics(self, teams: List[LiteLLM_TeamTable]): + async def _set_team_list_budget_metrics(self, teams: list[LiteLLM_TeamTable]): """Helper function to set budget metrics for a list of teams""" for team in teams: self._set_team_budget_metrics(team) - async def _set_user_list_budget_metrics(self, users: List[LiteLLM_UserTable]): + async def _set_user_list_budget_metrics(self, users: list[LiteLLM_UserTable]): """Helper function to set budget metrics for a list of users""" for user in users: self._set_user_budget_metrics(user) @@ -3552,10 +3538,10 @@ class PrometheusLogger(CustomLogger): async def _set_team_budget_metrics_after_api_request( self, - user_api_team: Optional[str], - user_api_team_alias: Optional[str], - team_spend: Optional[float], - team_max_budget: Optional[float], + user_api_team: str | None, + user_api_team_alias: str | None, + team_spend: float | None, + team_max_budget: float | None, response_cost: float, ): """ @@ -3583,8 +3569,8 @@ class PrometheusLogger(CustomLogger): self, team_id: str, team_alias: str, - spend: Optional[float], - max_budget: Optional[float], + spend: float | None, + max_budget: float | None, response_cost: float, ) -> LiteLLM_TeamTable: """ @@ -3611,7 +3597,7 @@ class PrometheusLogger(CustomLogger): user_api_key_cache=user_api_key_cache, ) except Exception as e: - verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting team info: {str(e)}") + verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting team info: {e!s}") return team_object if team_info: @@ -3680,7 +3666,7 @@ class PrometheusLogger(CustomLogger): async def _set_org_budget_metrics_after_api_request( self, - org_id: Optional[str], + org_id: str | None, response_cost: float, ): """ @@ -3709,7 +3695,7 @@ class PrometheusLogger(CustomLogger): include_budget_table=True, ) except Exception as e: - verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting org info: {str(e)}") + verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting org info: {e!s}") return if org_info is None: @@ -3734,8 +3720,8 @@ class PrometheusLogger(CustomLogger): org_id: str, org_alias: str, spend: float, - max_budget: Optional[float], - budget_reset_at: Optional[datetime], + max_budget: float | None, + budget_reset_at: datetime | None, ): """ Set org budget metrics for a single org @@ -3815,11 +3801,11 @@ class PrometheusLogger(CustomLogger): async def _set_api_key_budget_metrics_after_api_request( self, - user_api_key: Optional[str], - user_api_key_alias: Optional[str], + user_api_key: str | None, + user_api_key_alias: str | None, response_cost: float, - key_max_budget: Optional[float], - key_spend: Optional[float], + key_max_budget: float | None, + key_spend: float | None, ): if isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric): return @@ -3838,8 +3824,8 @@ class PrometheusLogger(CustomLogger): self, user_api_key: str, user_api_key_alias: str, - key_max_budget: Optional[float], - key_spend: Optional[float], + key_max_budget: float | None, + key_spend: float | None, response_cost: float, ) -> UserAPIKeyAuth: """ @@ -3866,15 +3852,15 @@ class PrometheusLogger(CustomLogger): if key_object: user_api_key_dict.budget_reset_at = key_object.budget_reset_at except Exception as e: - verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting key info: {str(e)}") + verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting key info: {e!s}") return user_api_key_dict async def _set_user_budget_metrics_after_api_request( self, - user_id: Optional[str], - user_spend: Optional[float], - user_max_budget: Optional[float], + user_id: str | None, + user_spend: float | None, + user_max_budget: float | None, response_cost: float, ): """ @@ -3900,8 +3886,8 @@ class PrometheusLogger(CustomLogger): async def _assemble_user_object( self, user_id: str, - spend: Optional[float], - max_budget: Optional[float], + spend: float | None, + max_budget: float | None, response_cost: float, ) -> LiteLLM_UserTable: """ @@ -3931,7 +3917,7 @@ class PrometheusLogger(CustomLogger): check_db_only=False, ) except Exception as e: - verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting user info: {str(e)}") + verbose_logger.debug(f"[Non-Blocking] Prometheus: Error getting user info: {e!s}") return user_object if user_info: @@ -4001,7 +3987,7 @@ class PrometheusLogger(CustomLogger): self, start_time: Any, end_time: Any, - ) -> Optional[float]: + ) -> float | None: """ Compute the duration in seconds between two objects. @@ -4021,7 +4007,7 @@ class PrometheusLogger(CustomLogger): """ from litellm.constants import PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES - prometheus_loggers: List[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( + prometheus_loggers: list[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( callback_type=PrometheusLogger ) # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them @@ -4071,10 +4057,10 @@ class PrometheusLogger(CustomLogger): def _prometheus_labels_from_context( - supported_enum_labels: List[str], + supported_enum_labels: list[str], ctx: PrometheusLabelFactoryContext, -) -> Dict[str, Optional[str]]: - filtered_labels: Dict[str, Optional[str]] = { +) -> dict[str, str | None]: + filtered_labels: dict[str, str | None] = { label: ctx._sanitized_enum[label] for label in supported_enum_labels if label in ctx._sanitized_enum } @@ -4097,11 +4083,11 @@ def _prometheus_labels_from_context( def prometheus_label_factory( - supported_enum_labels: List[str], + supported_enum_labels: list[str], enum_values: UserAPIKeyLabelValues, - tag: Optional[str] = None, + tag: str | None = None, *, - label_context: Optional[PrometheusLabelFactoryContext] = None, + label_context: PrometheusLabelFactoryContext | None = None, ) -> dict: """ Returns a dictionary of label + values for prometheus. @@ -4156,7 +4142,7 @@ def prometheus_label_factory( return filtered_labels -def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]: +def get_custom_labels_from_metadata(metadata: dict) -> dict[str, str]: """ Get custom labels from metadata """ @@ -4164,7 +4150,7 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]: if keys is None or len(keys) == 0: return {} - result: Dict[str, str] = {} + result: dict[str, str] = {} for key in keys: # Split the dot notation key into parts @@ -4227,8 +4213,8 @@ def get_service_tier_from_standard_logging_payload( def _get_combined_custom_metadata_from_standard_logging_payload( - standard_logging_payload: Optional[dict], -) -> Dict[str, Any]: + standard_logging_payload: dict | None, +) -> dict[str, Any]: """ Combine the metadata sources that can supply custom Prometheus labels. @@ -4286,7 +4272,7 @@ def _tag_matches_wildcard_configured_pattern(tags: Sequence[str], configured_tag return any(re.match(pattern=regex_pattern, string=tag) for tag in tags) -def get_custom_labels_from_tags(tags: Sequence[str]) -> Dict[str, str]: +def get_custom_labels_from_tags(tags: Sequence[str]) -> dict[str, str]: """ Get custom labels from tags based on admin configuration. @@ -4314,7 +4300,7 @@ def get_custom_labels_from_tags(tags: Sequence[str]) -> Dict[str, str]: if configured_tags is None or len(configured_tags) == 0: return {} - result: Dict[str, str] = {} + result: dict[str, str] = {} for configured_tag in configured_tags: label_name = _sanitize_prometheus_label_name(f"tag_{configured_tag}") diff --git a/litellm/integrations/prometheus_helpers/__init__.py b/litellm/integrations/prometheus_helpers/__init__.py index 7de072ecd03..cad83a013c6 100644 --- a/litellm/integrations/prometheus_helpers/__init__.py +++ b/litellm/integrations/prometheus_helpers/__init__.py @@ -6,7 +6,7 @@ Helpers for the Prometheus integration (extracted to keep ``prometheus.py`` smal from __future__ import annotations -from typing import Any, Dict, Optional, cast +from typing import Any, cast from litellm.types.integrations.prometheus import ( UserAPIKeyLabelValues, @@ -38,11 +38,11 @@ class PrometheusLabelFactoryContext: """ __slots__ = ( - "enum_values", - "_sanitized_enum", "_custom_by_sanitized_key", - "_tag_labels", "_resolved_end_user", + "_sanitized_enum", + "_tag_labels", + "enum_values", ) _END_USER_NOT_COMPUTED = object() @@ -50,15 +50,15 @@ class PrometheusLabelFactoryContext: def __init__(self, enum_values: UserAPIKeyLabelValues) -> None: self.enum_values = enum_values enum_dict = enum_values.model_dump() - self._sanitized_enum: Dict[str, Optional[str]] = { + self._sanitized_enum: dict[str, str | None] = { k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items() } - self._custom_by_sanitized_key: Dict[str, Optional[str]] = {} + self._custom_by_sanitized_key: dict[str, str | None] = {} if enum_values.custom_metadata_labels is not None: for key, value in enum_values.custom_metadata_labels.items(): sk = _sanitize_prometheus_label_name(key) self._custom_by_sanitized_key[sk] = _sanitize_prometheus_label_value(value) - self._tag_labels: Dict[str, Optional[str]] = {} + self._tag_labels: dict[str, str | None] = {} if enum_values.tags is not None: # Late import avoids circular import: ``prometheus`` imports this module. from litellm.integrations.prometheus import get_custom_labels_from_tags @@ -68,11 +68,11 @@ class PrometheusLabelFactoryContext: # Use a dedicated sentinel so `None` can be cached as a computed result. self._resolved_end_user: Any = self._END_USER_NOT_COMPUTED - def get_resolved_end_user(self) -> Optional[str]: + def get_resolved_end_user(self) -> str | None: if self._resolved_end_user is self._END_USER_NOT_COMPUTED: fn = _get_cached_end_user_id_for_cost_tracking() self._resolved_end_user = fn( litellm_params={"user_api_key_end_user_id": self.enum_values.end_user}, service_type="prometheus", ) - return cast(Optional[str], self._resolved_end_user) + return cast(str | None, self._resolved_end_user) diff --git a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py index 61b4d5ab96e..89f6c9610b6 100644 --- a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py +++ b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py @@ -3,7 +3,7 @@ from __future__ import annotations import time from collections import OrderedDict from threading import RLock -from typing import Any, Dict, Optional +from typing import Any class BoundedPrometheusSeriesTracker: @@ -15,18 +15,18 @@ class BoundedPrometheusSeriesTracker: """ def __init__(self) -> None: - self._series: Dict[str, OrderedDict[tuple[Optional[str], ...], float]] = {} - self._last_ttl_cleanup: Dict[str, float] = {} + self._series: dict[str, OrderedDict[tuple[str | None, ...], float]] = {} + self._last_ttl_cleanup: dict[str, float] = {} self.lock = RLock() def track_series( self, metric: Any, metric_name: str, - label_values: tuple[Optional[str], ...], - max_series: Optional[int], - ttl_seconds: Optional[float], - cleanup_interval_seconds: Optional[float], + label_values: tuple[str | None, ...], + max_series: int | None, + ttl_seconds: float | None, + cleanup_interval_seconds: float | None, ) -> None: if max_series is None and ttl_seconds is None: return @@ -64,7 +64,7 @@ class BoundedPrometheusSeriesTracker: self, metric_name: str, now: float, - cleanup_interval_seconds: Optional[float], + cleanup_interval_seconds: float | None, ) -> bool: if cleanup_interval_seconds is None or cleanup_interval_seconds <= 0: self._last_ttl_cleanup[metric_name] = now @@ -79,14 +79,14 @@ class BoundedPrometheusSeriesTracker: def _remove_metric_series( self, metric: Any, - series: OrderedDict[tuple[Optional[str], ...], float], - label_values: tuple[Optional[str], ...], + series: OrderedDict[tuple[str | None, ...], float], + label_values: tuple[str | None, ...], ) -> None: if self._remove_metric_child(metric, label_values): series.pop(label_values, None) @staticmethod - def _remove_metric_child(metric: Any, label_values: tuple[Optional[str], ...]) -> bool: + def _remove_metric_child(metric: Any, label_values: tuple[str | None, ...]) -> bool: """ Remove the Prometheus child for ``label_values`` and report whether the tracker should commit the matching state change. diff --git a/litellm/integrations/prometheus_helpers/prometheus_api.py b/litellm/integrations/prometheus_helpers/prometheus_api.py index 038788f0522..d51b03b9b1b 100644 --- a/litellm/integrations/prometheus_helpers/prometheus_api.py +++ b/litellm/integrations/prometheus_helpers/prometheus_api.py @@ -5,7 +5,6 @@ Helper functions to query prometheus API import json import time from datetime import datetime, timedelta -from typing import Optional from litellm import get_secret from litellm._logging import verbose_logger @@ -14,8 +13,8 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) -PROMETHEUS_URL: Optional[str] = get_secret("PROMETHEUS_URL") # type: ignore -PROMETHEUS_SELECTED_INSTANCE: Optional[str] = get_secret("PROMETHEUS_SELECTED_INSTANCE") # type: ignore +PROMETHEUS_URL: str | None = get_secret("PROMETHEUS_URL") # type: ignore +PROMETHEUS_SELECTED_INSTANCE: str | None = get_secret("PROMETHEUS_SELECTED_INSTANCE") # type: ignore async_http_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) @@ -96,7 +95,7 @@ def _quote_promql_string_literal(value: str) -> str: return json.dumps(value, ensure_ascii=False) -async def get_daily_spend_from_prometheus(api_key: Optional[str]): +async def get_daily_spend_from_prometheus(api_key: str | None): """ Expected Response Format: [ diff --git a/litellm/integrations/prometheus_services.py b/litellm/integrations/prometheus_services.py index db005aaffc5..f07606a3192 100644 --- a/litellm/integrations/prometheus_services.py +++ b/litellm/integrations/prometheus_services.py @@ -3,8 +3,6 @@ # On success + failure, log events to Prometheus for litellm / adjacent services (litellm, redis, postgres, llm api providers) -from typing import Dict, List, Optional, Union - import litellm from litellm._logging import print_verbose, verbose_logger from litellm.types.integrations.prometheus import LATENCY_BUCKETS @@ -44,10 +42,10 @@ class PrometheusServicesLogger: verbose_logger.debug("in init prometheus services metrics") - self.payload_to_prometheus_map: Dict[str, List[Union[Histogram, Counter, Gauge, Collector]]] = {} + self.payload_to_prometheus_map: dict[str, list[Histogram | Counter | Gauge | Collector]] = {} for service in ServiceTypes: - service_metrics: List[Union[Histogram, Counter, Gauge, Collector]] = [] + service_metrics: list[Histogram | Counter | Gauge | Collector] = [] metrics_to_initialize = self._get_service_metrics_initialize(service) @@ -84,10 +82,10 @@ class PrometheusServicesLogger: self.mock_testing_failure_calls = 0 except Exception as e: - print_verbose(f"Got exception on init prometheus client {str(e)}") + print_verbose(f"Got exception on init prometheus client {e!s}") raise e - def _get_service_metrics_initialize(self, service: ServiceTypes) -> List[ServiceMetrics]: + def _get_service_metrics_initialize(self, service: ServiceTypes) -> list[ServiceMetrics]: DEFAULT_METRICS = [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM] if service not in DEFAULT_SERVICE_CONFIGS: return DEFAULT_METRICS @@ -116,37 +114,37 @@ class PrometheusServicesLogger: return self.REGISTRY._names_to_collectors.get(metric_name) def create_histogram(self, service: str, type_of_request: str): - metric_name = "litellm_{}_{}".format(service, type_of_request) + metric_name = f"litellm_{service}_{type_of_request}" is_registered = self.is_metric_registered(metric_name) if is_registered: return self._get_metric(metric_name) return self.Histogram( metric_name, - "Latency for {} service".format(service), + f"Latency for {service} service", labelnames=[service], buckets=self.latency_buckets, ) def create_gauge(self, service: str, type_of_request: str): - metric_name = "litellm_{}_{}".format(service, type_of_request) + metric_name = f"litellm_{service}_{type_of_request}" is_registered = self.is_metric_registered(metric_name) if is_registered: return self._get_metric(metric_name) - return self.Gauge(metric_name, "Gauge for {} service".format(service), labelnames=[service]) + return self.Gauge(metric_name, f"Gauge for {service} service", labelnames=[service]) def create_counter( self, service: str, type_of_request: str, - additional_labels: Optional[List[str]] = None, + additional_labels: list[str] | None = None, ): - metric_name = "litellm_{}_{}".format(service, type_of_request) + metric_name = f"litellm_{service}_{type_of_request}" is_registered = self.is_metric_registered(metric_name) if is_registered: return self._get_metric(metric_name) return self.Counter( metric_name, - "Total {} for {} service".format(type_of_request, service), + f"Total {type_of_request} for {service} service", labelnames=[service] + (additional_labels or []), ) @@ -174,7 +172,7 @@ class PrometheusServicesLogger: counter, labels: str, amount: float, - additional_labels: Optional[List[str]] = [], + additional_labels: list[str] | None = [], ): assert isinstance(counter, self.Counter) @@ -250,7 +248,7 @@ class PrometheusServicesLogger: async def async_service_failure_hook( self, payload: ServiceLoggerPayload, - error: Union[str, Exception], + error: str | Exception, ): if self.mock_testing: self.mock_testing_failure_calls += 1 diff --git a/litellm/integrations/prompt_layer.py b/litellm/integrations/prompt_layer.py index 52209b2953f..9402055f414 100644 --- a/litellm/integrations/prompt_layer.py +++ b/litellm/integrations/prompt_layer.py @@ -80,4 +80,3 @@ class PromptLayerLogger: except Exception: print_verbose(f"error: Prompt Layer Error - {traceback.format_exc()}") - pass diff --git a/litellm/integrations/prompt_management_base.py b/litellm/integrations/prompt_management_base.py index a85ef4a3137..341ceff36db 100644 --- a/litellm/integrations/prompt_management_base.py +++ b/litellm/integrations/prompt_management_base.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from typing_extensions import TypedDict @@ -12,11 +12,11 @@ if TYPE_CHECKING: class PromptManagementClient(TypedDict): - prompt_id: Optional[str] - prompt_template: List[AllMessageValues] - prompt_template_model: Optional[str] - prompt_template_optional_params: Optional[Dict[str, Any]] - completed_messages: Optional[List[AllMessageValues]] + prompt_id: str | None + prompt_template: list[AllMessageValues] + prompt_template_model: str | None + prompt_template_optional_params: dict[str, Any] | None + completed_messages: list[AllMessageValues] | None class PromptManagementBase(ABC): @@ -28,8 +28,8 @@ class PromptManagementBase(ABC): @abstractmethod def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: pass @@ -37,43 +37,43 @@ class PromptManagementBase(ABC): @abstractmethod def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: pass @abstractmethod async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: pass def merge_messages( self, - prompt_template: List[AllMessageValues], - client_messages: List[AllMessageValues], - ) -> List[AllMessageValues]: + prompt_template: list[AllMessageValues], + client_messages: list[AllMessageValues], + ) -> list[AllMessageValues]: return prompt_template + client_messages def compile_prompt( self, prompt_id: str, - prompt_variables: Optional[dict], - client_messages: List[AllMessageValues], + prompt_variables: dict | None, + client_messages: list[AllMessageValues], dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - prompt_spec: Optional[PromptSpec] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + prompt_spec: PromptSpec | None = None, ) -> PromptManagementClient: compiled_prompt_client = self._compile_prompt_helper( prompt_id=prompt_id, @@ -94,13 +94,13 @@ class PromptManagementBase(ABC): async def async_compile_prompt( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], - client_messages: List[AllMessageValues], + prompt_id: str | None, + prompt_variables: dict | None, + client_messages: list[AllMessageValues], dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: compiled_prompt_client = await self.async_compile_prompt_helper( prompt_id=prompt_id, @@ -123,16 +123,16 @@ class PromptManagementBase(ABC): if prompt_management_client["prompt_template_model"] is not None: return prompt_management_client["prompt_template_model"] else: - return model.replace("{}/".format(self.integration_name), "") + return model.replace(f"{self.integration_name}/", "") def post_compile_prompt_processing( self, prompt_template: PromptManagementClient, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, model: str, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, ): completed_messages = prompt_template["completed_messages"] or messages @@ -153,17 +153,17 @@ class PromptManagementBase(ABC): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: if prompt_id is None: raise ValueError("prompt_id is required for Prompt Management Base class") if not self.should_run_prompt_management( @@ -194,19 +194,19 @@ class PromptManagementBase(ABC): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: "LiteLLMLoggingObj", - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: if not self.should_run_prompt_management( prompt_id=prompt_id, prompt_spec=prompt_spec, diff --git a/litellm/integrations/rubrik.py b/litellm/integrations/rubrik.py index 11809ee6361..2e49da45ce9 100644 --- a/litellm/integrations/rubrik.py +++ b/litellm/integrations/rubrik.py @@ -7,9 +7,10 @@ import time import urllib.parse import uuid from collections import Counter -from typing import TYPE_CHECKING, Any, List, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Optional import httpx + from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.integrations.custom_guardrail import ( @@ -53,7 +54,7 @@ class _MalformedToolBlockingResponseError(Exception): class RubrikLogger(CustomGuardrail, CustomBatchLogger): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] def __init__( @@ -134,9 +135,9 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger): # Periodic flush is started lazily on the first log event so that # low-traffic deployments still get their batches drained even when the # logger is instantiated outside a running event loop (sync init). - self._flush_task: Optional[asyncio.Task[Any]] = self._start_periodic_flush_task() + self._flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task() - def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]: + def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None: """Start the periodic flush task only when an event loop is already running.""" try: loop = asyncio.get_running_loop() @@ -519,7 +520,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger): def _extract_blocked_tools( service_response: dict[str, Any], all_tool_calls: list[ChatCompletionMessageToolCall], - ) -> Optional[str]: + ) -> str | None: """Return the blocking explanation if any tool calls were blocked. Compares the service response (which contains only allowed tools) against diff --git a/litellm/integrations/s3.py b/litellm/integrations/s3.py index 07bd957b5a3..51de43e302c 100644 --- a/litellm/integrations/s3.py +++ b/litellm/integrations/s3.py @@ -2,7 +2,7 @@ # On success + failure, log events to Supabase from datetime import datetime -from typing import Optional, cast +from typing import cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -78,7 +78,7 @@ class S3Logger: **kwargs, ) except Exception as e: - print_verbose(f"Got exception on init s3 client {str(e)}") + print_verbose(f"Got exception on init s3 client {e!s}") raise e async def _async_log_event(self, kwargs, response_obj, start_time, end_time, print_verbose): @@ -111,8 +111,8 @@ class S3Logger: clean_metadata[key] = value # Ensure everything in the payload is converted to str - payload: Optional[StandardLoggingPayload] = cast( - Optional[StandardLoggingPayload], + payload: StandardLoggingPayload | None = cast( + StandardLoggingPayload | None, kwargs.get("standard_logging_object", None), ) @@ -127,7 +127,7 @@ class S3Logger: s3_file_name = litellm.utils.get_logging_id(start_time, payload) or "" s3_object_key = get_s3_object_key( - cast(Optional[str], self.s3_path) or "", + cast(str | None, self.s3_path) or "", team_alias_prefix, start_time, s3_file_name, @@ -163,13 +163,12 @@ class S3Logger: **sse_params, ) - print_verbose(f"Response from s3:{str(response)}") + print_verbose(f"Response from s3:{response!s}") print_verbose(f"s3 Layer Logging - final response object: {response_obj}") return response except Exception as e: - verbose_logger.exception(f"s3 Layer Error - {str(e)}") - pass + verbose_logger.exception(f"s3 Layer Error - {e!s}") def _validated_sse_value(name: str, value: str | None) -> str | None: diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index 7fa78f39460..8c6cadd5356 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -10,7 +10,7 @@ import asyncio import time from collections.abc import Mapping from datetime import datetime -from typing import List, Optional, cast +from typing import cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -33,31 +33,31 @@ from .custom_batch_logger import CustomBatchLogger class S3Logger(CustomBatchLogger, BaseAWSLLM): def __init__( self, - s3_bucket_name: Optional[str] = None, - s3_path: Optional[str] = None, - s3_region_name: Optional[str] = None, - s3_api_version: Optional[str] = None, + s3_bucket_name: str | None = None, + s3_path: str | None = None, + s3_region_name: str | None = None, + s3_api_version: str | None = None, s3_use_ssl: bool = True, - s3_verify: Optional[bool] = None, - s3_endpoint_url: Optional[str] = None, - s3_aws_access_key_id: Optional[str] = None, - s3_aws_secret_access_key: Optional[str] = None, - s3_aws_session_token: Optional[str] = None, - s3_aws_session_name: Optional[str] = None, - s3_aws_profile_name: Optional[str] = None, - s3_aws_role_name: Optional[str] = None, - s3_aws_web_identity_token: Optional[str] = None, - s3_aws_sts_endpoint: Optional[str] = None, - s3_flush_interval: Optional[int] = DEFAULT_S3_FLUSH_INTERVAL_SECONDS, - s3_batch_size: Optional[int] = DEFAULT_S3_BATCH_SIZE, + s3_verify: bool | None = None, + s3_endpoint_url: str | None = None, + s3_aws_access_key_id: str | None = None, + s3_aws_secret_access_key: str | None = None, + s3_aws_session_token: str | None = None, + s3_aws_session_name: str | None = None, + s3_aws_profile_name: str | None = None, + s3_aws_role_name: str | None = None, + s3_aws_web_identity_token: str | None = None, + s3_aws_sts_endpoint: str | None = None, + s3_flush_interval: int | None = DEFAULT_S3_FLUSH_INTERVAL_SECONDS, + s3_batch_size: int | None = DEFAULT_S3_BATCH_SIZE, s3_config=None, s3_use_team_prefix: bool = False, s3_strip_base64_files: bool = False, s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, - s3_server_side_encryption: Optional[str] = None, + s3_server_side_encryption: str | None = None, s3_sse_kms_key_id: str | None = None, - s3_callback_params_override: Optional[dict] = None, + s3_callback_params_override: dict | None = None, **kwargs, ): try: @@ -119,40 +119,40 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): flush_interval=s3_flush_interval, batch_size=s3_batch_size, ) - self.log_queue: List[s3BatchLoggingElement] = [] + self.log_queue: list[s3BatchLoggingElement] = [] # Call BaseAWSLLM's __init__ BaseAWSLLM.__init__(self) except Exception as e: - print_verbose(f"Got exception on init s3 client {str(e)}") + print_verbose(f"Got exception on init s3 client {e!s}") raise e def _init_s3_params( self, - s3_bucket_name: Optional[str] = None, - s3_region_name: Optional[str] = None, - s3_api_version: Optional[str] = None, + s3_bucket_name: str | None = None, + s3_region_name: str | None = None, + s3_api_version: str | None = None, s3_use_ssl: bool = True, - s3_verify: Optional[bool] = None, - s3_endpoint_url: Optional[str] = None, - s3_aws_access_key_id: Optional[str] = None, - s3_aws_secret_access_key: Optional[str] = None, - s3_aws_session_token: Optional[str] = None, - s3_aws_session_name: Optional[str] = None, - s3_aws_profile_name: Optional[str] = None, - s3_aws_role_name: Optional[str] = None, - s3_aws_web_identity_token: Optional[str] = None, - s3_aws_sts_endpoint: Optional[str] = None, + s3_verify: bool | None = None, + s3_endpoint_url: str | None = None, + s3_aws_access_key_id: str | None = None, + s3_aws_secret_access_key: str | None = None, + s3_aws_session_token: str | None = None, + s3_aws_session_name: str | None = None, + s3_aws_profile_name: str | None = None, + s3_aws_role_name: str | None = None, + s3_aws_web_identity_token: str | None = None, + s3_aws_sts_endpoint: str | None = None, s3_config=None, - s3_path: Optional[str] = None, + s3_path: str | None = None, s3_use_team_prefix: bool = False, s3_strip_base64_files: bool = False, s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, - s3_server_side_encryption: Optional[str] = None, + s3_server_side_encryption: str | None = None, s3_sse_kms_key_id: str | None = None, - params_source: Optional[dict] = None, + params_source: dict | None = None, ): """ Initialize the s3 params for this logging callback. Reads from @@ -206,8 +206,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id, ) - return - def _sse_headers(self) -> Mapping[str, str]: candidates = { "x-amz-server-side-encryption": self.s3_server_side_encryption, @@ -230,7 +228,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): start_time=start_time, end_time=end_time, ) - pass async def async_log_audit_log_event(self, audit_log: StandardAuditLogPayload) -> None: """Batch audit logs and upload to S3 under audit_logs/ prefix.""" @@ -240,7 +237,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): now = datetime.now(timezone.utc) audit_log_id = audit_log.get("id", "unknown") - s3_path = cast(Optional[str], self.s3_path) or "" + s3_path = cast(str | None, self.s3_path) or "" s3_path = s3_path.rstrip("/") + "/" if s3_path else "" s3_object_key = ( @@ -287,7 +284,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): self.batch_size, ) except Exception as e: - verbose_logger.exception(f"s3 Layer Error - {str(e)}") + verbose_logger.exception(f"s3 Layer Error - {e!s}") self.handle_callback_failure(callback_name="S3Logger") async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement): @@ -386,7 +383,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): response.raise_for_status() break except Exception as e: - verbose_logger.exception(f"Error uploading to s3: {str(e)}") + verbose_logger.exception(f"Error uploading to s3: {e!s}") self.handle_callback_failure(callback_name="S3Logger") async def async_send_batch(self): @@ -413,8 +410,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): def create_s3_batch_logging_element( self, start_time: datetime, - standard_logging_payload: Optional[StandardLoggingPayload], - ) -> Optional[s3BatchLoggingElement]: + standard_logging_payload: StandardLoggingPayload | None, + ) -> s3BatchLoggingElement | None: """ Helper function to create an s3BatchLoggingElement. @@ -453,7 +450,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): f"Creating s3 file with prefix_components={prefix_components},prefix_path={prefix_path} and {s3_file_name}" ) s3_object_key = get_s3_object_key( - s3_path=cast(Optional[str], self.s3_path) or "", + s3_path=cast(str | None, self.s3_path) or "", prefix=prefix_path, start_time=start_time, s3_file_name=s3_file_name, @@ -560,10 +557,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): response.raise_for_status() break except Exception as e: - verbose_logger.exception(f"Error uploading to s3: {str(e)}") + verbose_logger.exception(f"Error uploading to s3: {e!s}") self.handle_callback_failure(callback_name="S3Logger") - async def _download_object_from_s3(self, s3_object_key: str) -> Optional[dict]: + async def _download_object_from_s3(self, s3_object_key: str) -> dict | None: """ Download and parse JSON object from S3. @@ -645,13 +642,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): return response.json() except Exception as e: - verbose_logger.exception(f"Error downloading from S3: {str(e)}") + verbose_logger.exception(f"Error downloading from S3: {e!s}") return None async def get_proxy_server_request_from_cold_storage_with_object_key( self, object_key: str, - ) -> Optional[dict]: + ) -> dict | None: """ Get the proxy server request from cold storage @@ -669,5 +666,5 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): downloaded_object = await self._download_object_from_s3(object_key) return downloaded_object except Exception as e: - verbose_logger.exception(f"Error retrieving object {object_key} from cold storage: {str(e)}") + verbose_logger.exception(f"Error retrieving object {object_key} from cold storage: {e!s}") return None diff --git a/litellm/integrations/sqs.py b/litellm/integrations/sqs.py index 8c0b06df888..18717790207 100644 --- a/litellm/integrations/sqs.py +++ b/litellm/integrations/sqs.py @@ -11,7 +11,6 @@ import base64 import json import re import traceback -from typing import List, Optional import litellm from litellm._logging import print_verbose, verbose_logger @@ -27,10 +26,10 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus from litellm.types.utils import StandardLoggingPayload from .custom_batch_logger import CustomBatchLogger -from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus _BASE64_INLINE_PATTERN = re.compile( r"data:(?:application|image|audio|video)/[a-zA-Z0-9.+-]+;base64,[A-Za-z0-9+/=\s]+", @@ -44,28 +43,28 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): def __init__( self, # --- Standard SQS params --- - sqs_queue_url: Optional[str] = None, - sqs_region_name: Optional[str] = None, - sqs_api_version: Optional[str] = None, + sqs_queue_url: str | None = None, + sqs_region_name: str | None = None, + sqs_api_version: str | None = None, sqs_use_ssl: bool = True, - sqs_verify: Optional[bool] = None, - sqs_endpoint_url: Optional[str] = None, - sqs_aws_access_key_id: Optional[str] = None, - sqs_aws_secret_access_key: Optional[str] = None, - sqs_aws_session_token: Optional[str] = None, - sqs_aws_session_name: Optional[str] = None, - sqs_aws_profile_name: Optional[str] = None, - sqs_aws_role_name: Optional[str] = None, - sqs_aws_web_identity_token: Optional[str] = None, - sqs_aws_sts_endpoint: Optional[str] = None, - sqs_flush_interval: Optional[int] = DEFAULT_SQS_FLUSH_INTERVAL_SECONDS, - sqs_batch_size: Optional[int] = DEFAULT_SQS_BATCH_SIZE, + sqs_verify: bool | None = None, + sqs_endpoint_url: str | None = None, + sqs_aws_access_key_id: str | None = None, + sqs_aws_secret_access_key: str | None = None, + sqs_aws_session_token: str | None = None, + sqs_aws_session_name: str | None = None, + sqs_aws_profile_name: str | None = None, + sqs_aws_role_name: str | None = None, + sqs_aws_web_identity_token: str | None = None, + sqs_aws_sts_endpoint: str | None = None, + sqs_flush_interval: int | None = DEFAULT_SQS_FLUSH_INTERVAL_SECONDS, + sqs_batch_size: int | None = DEFAULT_SQS_BATCH_SIZE, sqs_config=None, sqs_strip_base64_files: bool = False, # --- 🔐 Application-level encryption params --- sqs_aws_use_application_level_encryption: bool = False, - sqs_app_encryption_key_b64: Optional[str] = None, - sqs_app_encryption_aad: Optional[str] = None, + sqs_app_encryption_key_b64: str | None = None, + sqs_app_encryption_aad: str | None = None, **kwargs, ) -> None: try: @@ -110,33 +109,33 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): batch_size=sqs_batch_size, ) - self.log_queue: List[StandardLoggingPayload] = [] + self.log_queue: list[StandardLoggingPayload] = [] BaseAWSLLM.__init__(self) except Exception as e: - print_verbose(f"Got exception on init sqs client {str(e)}") + print_verbose(f"Got exception on init sqs client {e!s}") raise e def _init_sqs_params( self, - sqs_queue_url: Optional[str] = None, - sqs_region_name: Optional[str] = None, - sqs_api_version: Optional[str] = None, + sqs_queue_url: str | None = None, + sqs_region_name: str | None = None, + sqs_api_version: str | None = None, sqs_use_ssl: bool = True, - sqs_verify: Optional[bool] = None, - sqs_endpoint_url: Optional[str] = None, - sqs_aws_access_key_id: Optional[str] = None, - sqs_aws_secret_access_key: Optional[str] = None, - sqs_aws_session_token: Optional[str] = None, - sqs_aws_session_name: Optional[str] = None, - sqs_aws_profile_name: Optional[str] = None, - sqs_aws_role_name: Optional[str] = None, - sqs_aws_web_identity_token: Optional[str] = None, - sqs_aws_sts_endpoint: Optional[str] = None, + sqs_verify: bool | None = None, + sqs_endpoint_url: str | None = None, + sqs_aws_access_key_id: str | None = None, + sqs_aws_secret_access_key: str | None = None, + sqs_aws_session_token: str | None = None, + sqs_aws_session_name: str | None = None, + sqs_aws_profile_name: str | None = None, + sqs_aws_role_name: str | None = None, + sqs_aws_web_identity_token: str | None = None, + sqs_aws_sts_endpoint: str | None = None, sqs_strip_base64_files: bool = False, sqs_aws_use_application_level_encryption: bool = False, - sqs_app_encryption_key_b64: Optional[str] = None, - sqs_app_encryption_aad: Optional[str] = None, + sqs_app_encryption_key_b64: str | None = None, + sqs_app_encryption_aad: str | None = None, sqs_config=None, ) -> None: litellm.aws_sqs_callback_params = litellm.aws_sqs_callback_params or {} @@ -189,7 +188,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): self.sqs_app_encryption_aad = ( litellm.aws_sqs_callback_params.get("sqs_app_encryption_aad") or sqs_app_encryption_aad ) - self.app_crypto: Optional["AppCrypto"] = None + self.app_crypto: AppCrypto | None = None if self.sqs_aws_use_application_level_encryption: from litellm.litellm_core_utils.app_crypto import AppCrypto @@ -216,7 +215,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): self.batch_size, ) except Exception as e: - verbose_logger.exception(f"sqs Layer Error - {str(e)}") + verbose_logger.exception(f"sqs Layer Error - {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: @@ -234,8 +233,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): ) except Exception as e: - verbose_logger.exception(f"Datadog Layer Error - {str(e)}\n{traceback.format_exc()}") - pass + verbose_logger.exception(f"Datadog Layer Error - {e!s}\n{traceback.format_exc()}") async def async_send_batch(self) -> None: verbose_logger.debug(f"sqs logger - sending batch of {len(self.log_queue)}") @@ -307,7 +305,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM): ) response.raise_for_status() except Exception as e: - verbose_logger.exception(f"Error sending to SQS: {str(e)}") + verbose_logger.exception(f"Error sending to SQS: {e!s}") async def async_health_check(self) -> IntegrationHealthCheckStatus: """ diff --git a/litellm/integrations/supabase.py b/litellm/integrations/supabase.py index 18cf4f9549c..de37e32613c 100644 --- a/litellm/integrations/supabase.py +++ b/litellm/integrations/supabase.py @@ -45,7 +45,6 @@ class Supabase: print_verbose(f"data: {data}") except Exception: print_verbose(f"Supabase Logging Error - {traceback.format_exc()}") - pass def log_event( self, @@ -103,4 +102,3 @@ class Supabase: except Exception: print_verbose(f"Supabase Logging Error - {traceback.format_exc()}") - pass diff --git a/litellm/integrations/vantage/vantage_logger.py b/litellm/integrations/vantage/vantage_logger.py index be8907f07ff..6ce0af7795d 100644 --- a/litellm/integrations/vantage/vantage_logger.py +++ b/litellm/integrations/vantage/vantage_logger.py @@ -7,7 +7,7 @@ so users can simply set ``success_callback: ["vantage"]`` in their proxy config. from __future__ import annotations import os -from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_logger @@ -36,11 +36,11 @@ class VantageLogger(FocusLogger): def __init__( self, *, - api_key: Optional[str] = None, - integration_token: Optional[str] = None, - base_url: Optional[str] = None, - frequency: Optional[str] = None, - interval_seconds: Optional[int] = None, + api_key: str | None = None, + integration_token: str | None = None, + base_url: str | None = None, + frequency: str | None = None, + interval_seconds: int | None = None, **kwargs: Any, ) -> None: resolved_api_key = api_key or os.getenv("VANTAGE_API_KEY") @@ -49,7 +49,7 @@ class VantageLogger(FocusLogger): resolved_frequency = (frequency or os.getenv("VANTAGE_EXPORT_FREQUENCY") or "hourly").lower() raw_interval = interval_seconds or os.getenv("VANTAGE_EXPORT_INTERVAL_SECONDS") - resolved_interval: Optional[int] = None + resolved_interval: int | None = None if raw_interval is not None: try: resolved_interval = int(raw_interval) @@ -59,7 +59,7 @@ class VantageLogger(FocusLogger): raw_interval, ) - destination_config: Dict[str, Any] = {} + destination_config: dict[str, Any] = {} if resolved_api_key: destination_config["api_key"] = resolved_api_key if resolved_token: @@ -114,7 +114,7 @@ class VantageLogger(FocusLogger): scheduler: AsyncIOScheduler, ) -> None: """Register the Vantage export job with the provided scheduler.""" - vantage_loggers: List[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( + vantage_loggers: list[CustomLogger] = litellm.logging_callback_manager.get_custom_loggers_for_type( callback_type=VantageLogger ) if not vantage_loggers: 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 0ba6da78b27..73c48f72d34 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 @@ -5,7 +5,7 @@ This hook is called before making an LLM request when a vector store is configur It searches the vector store for relevant context and appends it to the messages. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import litellm import litellm.vector_stores @@ -44,19 +44,19 @@ class VectorStorePreCallHook(CustomLogger): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: Optional[PromptSpec] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Perform vector store search and append results as context to messages. @@ -88,7 +88,7 @@ class VectorStorePreCallHook(CustomLogger): pass # Use database fallback to ensure synchronization across instances - vector_stores_to_run: List[ + vector_stores_to_run: list[ LiteLLM_ManagedVectorStore ] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback( non_default_params=non_default_params, @@ -106,8 +106,8 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug("No query found in messages for vector store search") return model, messages, non_default_params - modified_messages: List[AllMessageValues] = messages.copy() - all_search_results: List[VectorStoreSearchResponse] = [] + modified_messages: list[AllMessageValues] = messages.copy() + all_search_results: list[VectorStoreSearchResponse] = [] for vector_store_to_run in vector_stores_to_run: # Get vector store id from the vector store config @@ -146,11 +146,11 @@ class VectorStorePreCallHook(CustomLogger): return model, modified_messages, non_default_params except Exception as e: - verbose_logger.exception(f"Error in VectorStorePreCallHook: {str(e)}") + verbose_logger.exception(f"Error in VectorStorePreCallHook: {e!s}") # Return original parameters on error return model, messages, non_default_params - def _extract_query_from_messages(self, messages: List[AllMessageValues]) -> Optional[str]: + def _extract_query_from_messages(self, messages: list[AllMessageValues]) -> str | None: """ Extract the query from the last user message. @@ -181,9 +181,9 @@ class VectorStorePreCallHook(CustomLogger): def _append_search_results_to_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], search_response: VectorStoreSearchResponse, - ) -> List[AllMessageValues]: + ) -> list[AllMessageValues]: """ Append search results as context to the messages. @@ -194,17 +194,17 @@ class VectorStorePreCallHook(CustomLogger): Returns: Modified list of messages with context appended """ - search_response_data: Optional[List[VectorStoreSearchResult]] = search_response.get("data") + search_response_data: list[VectorStoreSearchResult] | None = search_response.get("data") if not search_response_data: return messages context_content = self.CONTENT_PREFIX_STRING for result in search_response_data: - result_content: Optional[List[VectorStoreResultContent]] = result.get("content") + result_content: list[VectorStoreResultContent] | None = result.get("content") if result_content: for content_item in result_content: - content_text: Optional[str] = content_item.get("text") + content_text: str | None = content_item.get("text") if content_text: context_content += content_text + "\n\n" @@ -226,8 +226,8 @@ class VectorStorePreCallHook(CustomLogger): self, request_data: dict, response: Any, - call_type: Optional[Any], - ) -> Optional[Any]: + call_type: Any | None, + ) -> Any | None: """ Add search results to the response after successful LLM call. @@ -246,7 +246,7 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug(f"model_call_details keys: {list(litellm_logging_obj.model_call_details.keys())}") # Get search results from model_call_details (already in OpenAI format) - search_results: Optional[List[VectorStoreSearchResponse]] = litellm_logging_obj.model_call_details.get( + search_results: list[VectorStoreSearchResponse] | None = litellm_logging_obj.model_call_details.get( "search_results" ) @@ -275,7 +275,7 @@ class VectorStorePreCallHook(CustomLogger): return response except Exception as e: - verbose_logger.exception(f"Error adding search results to response: {str(e)}") + verbose_logger.exception(f"Error adding search results to response: {e!s}") # Don't fail the request if search results fail to be added return None @@ -283,8 +283,8 @@ class VectorStorePreCallHook(CustomLogger): self, request_data: dict, response_chunk: Any, - call_type: Optional[Any], - ) -> Optional[Any]: + call_type: Any | None, + ) -> Any | None: """ Add search results to the final streaming chunk. @@ -295,7 +295,7 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug("VectorStorePreCallHook.async_post_call_streaming_deployment_hook called") # Get search results from model_call_details (already in OpenAI format) - search_results: Optional[List[VectorStoreSearchResponse]] = request_data.get("search_results") + search_results: list[VectorStoreSearchResponse] | None = request_data.get("search_results") verbose_logger.debug(f"Search results found for streaming chunk: {search_results is not None}") @@ -322,6 +322,6 @@ class VectorStorePreCallHook(CustomLogger): return response_chunk except Exception as e: - verbose_logger.exception(f"Error adding search results to streaming chunk: {str(e)}") + verbose_logger.exception(f"Error adding search results to streaming chunk: {e!s}") # Don't fail the request if search results fail to be added return response_chunk diff --git a/litellm/integrations/weave/weave_otel.py b/litellm/integrations/weave/weave_otel.py index c43afe7b6ca..321dda2983d 100644 --- a/litellm/integrations/weave/weave_otel.py +++ b/litellm/integrations/weave/weave_otel.py @@ -3,7 +3,7 @@ from __future__ import annotations import base64 import json import os -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from opentelemetry.trace import Status, StatusCode from typing_extensions import override @@ -43,7 +43,7 @@ class WeaveLLMObsOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @override - def set_messages(span: "Span", kwargs: dict[str, Any]): + def set_messages(span: Span, kwargs: dict[str, Any]): """Set input messages as span attributes using OpenInference conventions.""" messages = kwargs.get("messages") or [] @@ -203,8 +203,8 @@ class WeaveOtelLogger(OpenTelemetry): def __init__( self, - config: Optional[OpenTelemetryConfig] = None, - callback_name: Optional[str] = "weave_otel", + config: OpenTelemetryConfig | None = None, + callback_name: str | None = "weave_otel", **kwargs, ): """ @@ -233,7 +233,6 @@ class WeaveOtelLogger(OpenTelemetry): already contains all the necessary attributes, so the child span is redundant. """ - pass def _start_primary_span( self, diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 531caf273f1..54278afafc4 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -10,7 +10,7 @@ import asyncio import math import uuid from collections.abc import AsyncIterator, Mapping -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_logger @@ -80,8 +80,8 @@ class WebSearchInterceptionLogger(CustomLogger): def __init__( self, - enabled_providers: Optional[List[Union[LlmProviders, str]]] = None, - search_tool_name: Optional[str] = None, + enabled_providers: list[LlmProviders | str] | None = None, + search_tool_name: str | None = None, ): """ Args: @@ -104,9 +104,9 @@ class WebSearchInterceptionLogger(CustomLogger): async def try_short_circuit_search( self, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - custom_llm_provider: Optional[str], + messages: list[dict], + tools: list[dict] | None, + custom_llm_provider: str | None, kwargs: Mapping[str, object] | None = None, ) -> dict[str, object] | None: """ @@ -171,7 +171,7 @@ class WebSearchInterceptionLogger(CustomLogger): get_last_user_message, ) - query = get_last_user_message(cast(List[AllMessageValues], messages)) + query = get_last_user_message(cast(list[AllMessageValues], messages)) if not query: return None @@ -224,7 +224,7 @@ class WebSearchInterceptionLogger(CustomLogger): content.append({"type": "text", "text": search_result_text}) response: dict[str, object] = { - "id": f"msg_{str(uuid.uuid4())}", + "id": f"msg_{uuid.uuid4()!s}", "type": "message", "role": "assistant", "model": model, @@ -241,9 +241,7 @@ class WebSearchInterceptionLogger(CustomLogger): ) return response - async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: """ Pre-call hook to convert native Anthropic web_search tools to regular tools. @@ -359,7 +357,7 @@ class WebSearchInterceptionLogger(CustomLogger): search_tool_name = config.get("search_tool_name", None) # Convert string provider names to LlmProviders enum values - enabled_providers: Optional[List[Union[LlmProviders, str]]] = None + enabled_providers: list[LlmProviders | str] | None = None if enabled_providers_str is not None: enabled_providers = [] for provider in enabled_providers_str: @@ -377,7 +375,7 @@ class WebSearchInterceptionLogger(CustomLogger): ) @staticmethod - def _tool_name(tool: dict[str, Any]) -> Optional[str]: + def _tool_name(tool: dict[str, Any]) -> str | None: """Effective tool name, handling OpenAI ``function`` wrapper shape.""" fn = tool.get("function") if tool.get("type") == "function" and isinstance(fn, dict): @@ -402,7 +400,7 @@ class WebSearchInterceptionLogger(CustomLogger): return tool_choice return {**tool_choice, "name": LITELLM_WEB_SEARCH_TOOL_NAME} - async def async_pre_request_hook(self, model: str, messages: List[Dict], kwargs: Dict) -> Optional[Dict]: + async def async_pre_request_hook(self, model: str, messages: list[dict], kwargs: dict) -> dict | None: """ Pre-request hook to convert native web search tools to LiteLLM standard. @@ -485,12 +483,12 @@ class WebSearchInterceptionLogger(CustomLogger): self, response: object, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Tuple[bool, Dict]: + kwargs: dict, + ) -> tuple[bool, dict]: if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE: return await self.async_should_run_chat_completion_agentic_loop( response=response, @@ -550,7 +548,7 @@ class WebSearchInterceptionLogger(CustomLogger): # When extended thinking is enabled, the model response includes # thinking/redacted_thinking blocks that must be preserved and # prepended to the follow-up assistant message. - thinking_blocks: List[Dict] = [] + thinking_blocks: list[dict] = [] if isinstance(response, dict): content = response.get("content", []) else: @@ -568,7 +566,7 @@ class WebSearchInterceptionLogger(CustomLogger): else: # Convert object to dict using getattr, matching the # pattern in _detect_from_non_streaming_response - thinking_block_dict: Dict = {"type": block_type} + thinking_block_dict: dict = {"type": block_type} if block_type == "thinking": thinking_block_dict["thinking"] = getattr(block, "thinking", "") thinking_block_dict["signature"] = getattr(block, "signature", "") @@ -595,12 +593,12 @@ class WebSearchInterceptionLogger(CustomLogger): self, response: object, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Tuple[bool, Dict]: + kwargs: dict, + ) -> tuple[bool, dict]: """ Check if WebSearch tool interception is needed for Chat Completions API. @@ -699,15 +697,15 @@ class WebSearchInterceptionLogger(CustomLogger): async def async_run_agentic_loop( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: object, anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None", - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> "AnthropicMessagesResponse | AsyncIterator[object]": """ Execute agentic loop with WebSearch execution for Anthropic Messages API. @@ -733,15 +731,15 @@ class WebSearchInterceptionLogger(CustomLogger): async def async_build_agentic_loop_plan( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: object, anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None", - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> AgenticLoopPlan: if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE: return await self.async_build_chat_completion_agentic_loop_plan( @@ -804,7 +802,7 @@ class WebSearchInterceptionLogger(CustomLogger): self, response: object, plan: AgenticLoopPlan, - kwargs: Dict, + kwargs: dict, ) -> object: """ Inject Anthropic-native ``web_search_tool_result`` blocks into the @@ -823,8 +821,8 @@ class WebSearchInterceptionLogger(CustomLogger): @staticmethod def _build_native_result_blocks( - tool_calls: List[Dict], - structured_results: List[Optional[SearchResponse]], + tool_calls: list[dict], + structured_results: list[SearchResponse | None], ) -> list[dict[str, object]]: """Build one ``web_search_tool_result`` block per tool_call.""" blocks: list[dict[str, object]] = [] @@ -861,14 +859,14 @@ class WebSearchInterceptionLogger(CustomLogger): async def async_run_chat_completion_agentic_loop( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: object, - optional_params: Dict, + optional_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> "ModelResponse | CustomStreamWrapper": """ Execute agentic loop with WebSearch execution for Chat Completions API. @@ -896,14 +894,14 @@ class WebSearchInterceptionLogger(CustomLogger): async def async_build_chat_completion_agentic_loop_plan( self, - tools: Dict, + tools: dict, model: str, - messages: List[Dict], + messages: list[dict], response: object, - optional_params: Dict, + optional_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> AgenticLoopPlan: tool_calls = tools["tool_calls"] response_format = tools.get("response_format", "openai") @@ -949,7 +947,7 @@ class WebSearchInterceptionLogger(CustomLogger): async def _build_responses_request_patch( self, model: str, - messages: Union[str, list[dict]], + messages: str | list[dict], tool_calls: list[dict], optional_params: dict, kwargs: dict, @@ -1030,7 +1028,7 @@ class WebSearchInterceptionLogger(CustomLogger): ) @staticmethod - def _normalize_responses_input(messages: Union[str, list[dict]]) -> list[dict]: + def _normalize_responses_input(messages: str | list[dict]) -> list[dict]: if isinstance(messages, str): return [{"role": "user", "content": messages}] if isinstance(messages, list): @@ -1040,8 +1038,8 @@ class WebSearchInterceptionLogger(CustomLogger): @staticmethod def _extract_search_text(result: object) -> str: if isinstance(result, Exception): - verbose_logger.error(f"WebSearchInterception: Responses search failed with error: {str(result)}") - return f"Search failed: {str(result)}" + verbose_logger.error(f"WebSearchInterception: Responses search failed with error: {result!s}") + return f"Search failed: {result!s}" if isinstance(result, tuple) and len(result) == 2: text_value, _ = result return text_value if isinstance(text_value, str) else str(text_value) @@ -1050,8 +1048,8 @@ class WebSearchInterceptionLogger(CustomLogger): @staticmethod def _resolve_max_tokens( - optional_params: Dict, - kwargs: Dict, + optional_params: dict, + kwargs: dict, ) -> int: """Extract max_tokens and validate against thinking.budget_tokens. @@ -1084,7 +1082,7 @@ class WebSearchInterceptionLogger(CustomLogger): return max_tokens @staticmethod - def _prepare_followup_kwargs(kwargs: Dict) -> Dict: + def _prepare_followup_kwargs(kwargs: dict) -> dict: """Build kwargs for the follow-up call, excluding internal keys. ``litellm_logging_obj`` MUST be excluded so the follow-up call creates @@ -1102,13 +1100,13 @@ class WebSearchInterceptionLogger(CustomLogger): async def _execute_agentic_loop( self, model: str, - messages: List[Dict], - tool_calls: List[Dict], - thinking_blocks: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + tool_calls: list[dict], + thinking_blocks: list[dict], + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, ) -> "AnthropicMessagesResponse | AsyncIterator[object]": """Legacy path: execute search + build patch + run follow-up call.""" request_patch, structured_results = await self._build_anthropic_request_patch( @@ -1127,7 +1125,7 @@ class WebSearchInterceptionLogger(CustomLogger): optional_params.update(request_patch.optional_params) max_tokens = request_patch.max_tokens if max_tokens is None: - max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None)) + max_tokens = cast(int | None, optional_params.pop("max_tokens", None)) else: optional_params.pop("max_tokens", None) if max_tokens is None: @@ -1156,13 +1154,13 @@ class WebSearchInterceptionLogger(CustomLogger): async def _build_anthropic_request_patch( self, model: str, - messages: List[Dict], - tool_calls: List[Dict], - thinking_blocks: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + tool_calls: list[dict], + thinking_blocks: list[dict], + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj | None", - kwargs: Dict, - ) -> Tuple[AgenticLoopRequestPatch, List[Optional[SearchResponse]]]: + kwargs: dict, + ) -> tuple[AgenticLoopRequestPatch, list[SearchResponse | None]]: """ Execute litellm.search() and build follow-up request patch. @@ -1192,12 +1190,12 @@ class WebSearchInterceptionLogger(CustomLogger): # Split the gathered (text, structured) tuples into two parallel lists. # The text list feeds the follow-up model call; the structured list # is returned to the caller for native-block emission. - final_search_results: List[str] = [] - structured_results: List[Optional[SearchResponse]] = [] + final_search_results: list[str] = [] + structured_results: list[SearchResponse | None] = [] for i, result in enumerate(search_results): if isinstance(result, Exception): - verbose_logger.error(f"WebSearchInterception: Search {i} failed with error: {str(result)}") - final_search_results.append(f"Search failed: {str(result)}") + verbose_logger.error(f"WebSearchInterception: Search {i} failed with error: {result!s}") + final_search_results.append(f"Search failed: {result!s}") structured_results.append(None) elif isinstance(result, tuple) and len(result) == 2: text_value, structured_value = result @@ -1217,7 +1215,7 @@ class WebSearchInterceptionLogger(CustomLogger): thinking_blocks=thinking_blocks, ) - follow_up_messages = messages + [assistant_message, cast(Dict, user_message)] + follow_up_messages = messages + [assistant_message, cast(dict, user_message)] # Correlation context for structured logging _call_id = getattr(logging_obj, "litellm_call_id", None) or kwargs.get("litellm_call_id", "unknown") @@ -1254,7 +1252,7 @@ class WebSearchInterceptionLogger(CustomLogger): async def _execute_search( self, query: str, kwargs: Mapping[str, object] | None = None - ) -> Tuple[str, Optional[SearchResponse]]: + ) -> tuple[str, SearchResponse | None]: """ Execute a single web search using router's search tools. @@ -1277,7 +1275,7 @@ class WebSearchInterceptionLogger(CustomLogger): llm_router = None search_tool = self._select_search_tool_from_router(llm_router=llm_router) - search_provider: Optional[str] = None + search_provider: str | None = None search_litellm_params: dict[str, Any] = {} if search_tool is not None: await self._authorize_search_tool(search_tool=search_tool, kwargs=kwargs) @@ -1310,7 +1308,7 @@ class WebSearchInterceptionLogger(CustomLogger): ) return search_result_text, result except Exception as e: - verbose_logger.error(f"WebSearchInterception: Search failed for '{query}': {str(e)}") + verbose_logger.error(f"WebSearchInterception: Search failed for '{query}': {e!s}") raise async def _authorize_search_tool( @@ -1378,7 +1376,7 @@ class WebSearchInterceptionLogger(CustomLogger): return None - def _select_search_tool_from_router(self, llm_router: object) -> Optional[dict[str, Any]]: + def _select_search_tool_from_router(self, llm_router: object) -> dict[str, Any] | None: if llm_router is None or not hasattr(llm_router, "search_tools"): return None search_tools = list(getattr(llm_router, "search_tools") or []) @@ -1388,7 +1386,7 @@ class WebSearchInterceptionLogger(CustomLogger): self, search_tools: list[dict[str, Any]], source: str, - ) -> Optional[dict[str, Any]]: + ) -> dict[str, Any] | None: if self.search_tool_name: matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name] if matching_tools: @@ -1417,12 +1415,12 @@ class WebSearchInterceptionLogger(CustomLogger): async def _execute_chat_completion_agentic_loop( self, model: str, - messages: List[Dict], - tool_calls: List[Dict], - optional_params: Dict, + messages: list[dict], + tool_calls: list[dict], + optional_params: dict, logging_obj: "LiteLLMLoggingObj | None", stream: bool, - kwargs: Dict, + kwargs: dict, response_format: str = "openai", ) -> "ModelResponse | CustomStreamWrapper": """Legacy path: execute search + build patch + run follow-up call.""" @@ -1449,10 +1447,10 @@ class WebSearchInterceptionLogger(CustomLogger): async def _build_chat_completion_request_patch( self, model: str, - messages: List[Dict], - tool_calls: List[Dict], - optional_params: Dict, - kwargs: Dict, + messages: list[dict], + tool_calls: list[dict], + optional_params: dict, + kwargs: dict, response_format: str = "openai", ) -> AgenticLoopRequestPatch: """Execute litellm.search() and build chat-completion rerun patch.""" @@ -1485,11 +1483,11 @@ class WebSearchInterceptionLogger(CustomLogger): # Chat-completion path only needs text — OpenAI tool_result format # has no equivalent of Anthropic's web_search_tool_result block. - final_search_results: List[str] = [] + final_search_results: list[str] = [] for i, result in enumerate(search_results): if isinstance(result, Exception): - verbose_logger.error(f"WebSearchInterception: Search {i} failed with error: {str(result)}") - final_search_results.append(f"Search failed: {str(result)}") + verbose_logger.error(f"WebSearchInterception: Search {i} failed with error: {result!s}") + final_search_results.append(f"Search failed: {result!s}") elif isinstance(result, tuple) and len(result) == 2: text_value, _ = result final_search_results.append(cast(str, text_value) if isinstance(text_value, str) else str(text_value)) @@ -1510,12 +1508,12 @@ class WebSearchInterceptionLogger(CustomLogger): # Make follow-up request with search results # For OpenAI format, tool_messages_or_user is a list of tool messages if response_format == "openai": - follow_up_messages = messages + [assistant_message] + cast(List[Dict], tool_messages_or_user) + follow_up_messages = messages + [assistant_message] + cast(list[dict], tool_messages_or_user) else: # For Anthropic format (shouldn't happen in this method, but handle it) follow_up_messages = messages + [ assistant_message, - cast(Dict, tool_messages_or_user), + cast(dict, tool_messages_or_user), ] verbose_logger.debug("WebSearchInterception: Making follow-up chat completion request with search results") @@ -1573,14 +1571,14 @@ class WebSearchInterceptionLogger(CustomLogger): async def _create_empty_search_result( self, - ) -> Tuple[str, Optional[SearchResponse]]: + ) -> tuple[str, SearchResponse | None]: """Create an empty search result for tool calls without queries""" return "No search query provided", None @staticmethod def initialize_from_proxy_config( - litellm_settings: Dict[str, Any], - callback_specific_params: Dict[str, Any], + litellm_settings: dict[str, Any], + callback_specific_params: dict[str, Any], ) -> "WebSearchInterceptionLogger": """ Static method to initialize WebSearchInterceptionLogger from proxy config. diff --git a/litellm/integrations/websearch_interception/tools.py b/litellm/integrations/websearch_interception/tools.py index 14c8aea0908..ad0c3d5688f 100644 --- a/litellm/integrations/websearch_interception/tools.py +++ b/litellm/integrations/websearch_interception/tools.py @@ -6,12 +6,12 @@ Native provider tools (like Anthropic's web_search_20250305) are converted to this format for consistent interception and execution. """ -from typing import Any, Dict +from typing import Any from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME -def get_litellm_web_search_tool() -> Dict[str, Any]: +def get_litellm_web_search_tool() -> dict[str, Any]: """ Get the standard LiteLLM web search tool definition. @@ -49,7 +49,7 @@ def get_litellm_web_search_tool() -> Dict[str, Any]: } -def get_litellm_web_search_tool_openai() -> Dict[str, Any]: +def get_litellm_web_search_tool_openai() -> dict[str, Any]: """ Get the standard LiteLLM web search tool definition in OpenAI format. @@ -151,7 +151,7 @@ def is_web_search_tool_responses(tool: dict[str, Any]) -> bool: return tool_type == "web_search" or tool_type.startswith("web_search_") -def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool: +def is_web_search_tool_chat_completion(tool: dict[str, Any]) -> bool: """ Check if a tool is a web search tool for Chat Completions API (strict check). @@ -195,7 +195,7 @@ def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool: return False -def is_anthropic_native_web_search_tool(tool: Dict[str, Any]) -> bool: +def is_anthropic_native_web_search_tool(tool: dict[str, Any]) -> bool: """ Check if a tool is an Anthropic-native ``web_search_*`` tool. @@ -216,7 +216,7 @@ def is_anthropic_native_web_search_tool(tool: Dict[str, Any]) -> bool: return tool_type.startswith("web_search_") and tool_type != "function" -def is_web_search_tool(tool: Dict[str, Any]) -> bool: +def is_web_search_tool(tool: dict[str, Any]) -> bool: """ Check if a tool is a web search tool (native or LiteLLM standard). diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index 282d75d3d4d..9dd0c155142 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -5,7 +5,7 @@ Transforms between Anthropic/OpenAI tool_use format and LiteLLM search format. """ import json -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any from litellm._logging import verbose_logger from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME @@ -27,7 +27,7 @@ class WebSearchTransformation: response: Any, stream: bool, response_format: str = "anthropic", - ) -> Tuple[bool, List[Dict]]: + ) -> tuple[bool, list[dict]]: """ Transform model response to extract WebSearch tool calls. @@ -129,7 +129,7 @@ class WebSearchTransformation: @staticmethod def _detect_from_non_streaming_response( response: Any, - ) -> Tuple[bool, List[Dict]]: + ) -> tuple[bool, list[dict]]: """Parse non-streaming response for WebSearch tool_use""" # Handle both dict and object responses @@ -185,7 +185,7 @@ class WebSearchTransformation: @staticmethod def _detect_from_openai_response( response: Any, - ) -> Tuple[bool, List[Dict]]: + ) -> tuple[bool, list[dict]]: """Parse OpenAI-style response for WebSearch tool_calls""" # Handle both dict and ModelResponse objects @@ -279,11 +279,11 @@ class WebSearchTransformation: @staticmethod def transform_response( - tool_calls: List[Dict], - search_results: List[str], + tool_calls: list[dict], + search_results: list[str], response_format: str = "anthropic", - thinking_blocks: Optional[List[Dict]] = None, - ) -> Tuple[Dict, Union[Dict, List[Dict]]]: + thinking_blocks: list[dict] | None = None, + ) -> tuple[dict, dict | list[dict]]: """ Transform LiteLLM search results to Anthropic/OpenAI tool_result format. @@ -313,13 +313,13 @@ class WebSearchTransformation: @staticmethod def _transform_response_anthropic( - tool_calls: List[Dict], - search_results: List[str], - thinking_blocks: Optional[List[Dict]] = None, - ) -> Tuple[Dict, Dict]: + tool_calls: list[dict], + search_results: list[str], + thinking_blocks: list[dict] | None = None, + ) -> tuple[dict, dict]: """Transform to Anthropic format (single user message with tool_result blocks)""" # Build assistant message content - assistant_content: List[Dict] = [] + assistant_content: list[dict] = [] # Prepend thinking blocks if present. # When extended thinking is enabled, Anthropic requires the assistant @@ -363,9 +363,9 @@ class WebSearchTransformation: @staticmethod def _transform_response_openai( - tool_calls: List[Dict], - search_results: List[str], - ) -> Tuple[Dict, List[Dict]]: + tool_calls: list[dict], + search_results: list[str], + ) -> tuple[dict, list[dict]]: """Transform to OpenAI format (assistant with tool_calls, separate tool messages)""" # Build assistant message with tool_calls assistant_message = { @@ -398,8 +398,8 @@ class WebSearchTransformation: @staticmethod def build_web_search_tool_result_block( tool_use_id: str, - search_response: Optional[SearchResponse], - ) -> Dict[str, Any]: + search_response: SearchResponse | None, + ) -> dict[str, Any]: """ Build an Anthropic-native ``web_search_tool_result`` content block. @@ -424,7 +424,7 @@ class WebSearchTransformation: emitted with an empty result list (signals "search ran, no results" rather than "search did not run"). """ - items: List[Dict[str, Any]] = [] + items: list[dict[str, Any]] = [] if search_response is not None: results = getattr(search_response, "results", None) or [] for r in results: diff --git a/litellm/integrations/weights_biases.py b/litellm/integrations/weights_biases.py index d2c93bc1bf4..0fe2a70ab66 100644 --- a/litellm/integrations/weights_biases.py +++ b/litellm/integrations/weights_biases.py @@ -3,7 +3,7 @@ try: import io import logging import sys - from typing import Any, Dict, List, Optional, TypeVar + from typing import Any, TypeVar from wandb.sdk.data_types import trace_tree @@ -25,15 +25,15 @@ try: def __getitem__(self, key: K) -> V: ... - def get(self, key: K, default: Optional[V] = None) -> Optional[V]: ... # pragma: no cover + def get(self, key: K, default: V | None = None) -> V | None: ... # pragma: no cover class OpenAIRequestResponseResolver: def __call__( self, - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, time_elapsed: float, - ) -> Optional[trace_tree.WBTraceTree]: + ) -> trace_tree.WBTraceTree | None: try: if response["object"] == "edit": return self._resolve_edit(request, response, time_elapsed) @@ -49,9 +49,9 @@ try: @staticmethod def results_to_trace_tree( - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, - results: List[trace_tree.Result], + results: list[trace_tree.Result], time_elapsed: float, ) -> trace_tree.WBTraceTree: """Converts the request, response, and results into a trace tree. @@ -79,7 +79,7 @@ try: def _resolve_edit( self, - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, time_elapsed: float, ) -> trace_tree.WBTraceTree: @@ -97,7 +97,7 @@ try: def _resolve_completion( self, - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, time_elapsed: float, ) -> trace_tree.WBTraceTree: @@ -115,7 +115,7 @@ try: def _resolve_chat_completion( self, - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, time_elapsed: float, ) -> trace_tree.WBTraceTree: @@ -140,10 +140,10 @@ try: def _request_response_result_to_trace( self, - request: Dict[str, Any], + request: dict[str, Any], response: OpenAIResponse, request_str: str, - choices: List[str], + choices: list[str], time_elapsed: float, ) -> trace_tree.WBTraceTree: """Resolves the request and response objects for `openai.Completion`.""" @@ -196,4 +196,3 @@ class WeightsBiasesLogger: print_verbose(f"W&B Logging Logging - final response object: {response_obj}") except Exception: print_verbose(f"W&B Logging Layer Error - {traceback.format_exc()}") - pass diff --git a/litellm/interactions/agents/__init__.py b/litellm/interactions/agents/__init__.py index 711a54fdcbb..21fd2dc0742 100644 --- a/litellm/interactions/agents/__init__.py +++ b/litellm/interactions/agents/__init__.py @@ -26,14 +26,14 @@ from litellm.interactions.agents.main import ( ) __all__ = [ - "create", "acreate", - "list", - "alist", - "get", - "aget", - "delete", "adelete", - "list_versions", + "aget", + "alist", "alist_versions", + "create", + "delete", + "get", + "list", + "list_versions", ] diff --git a/litellm/interactions/agents/http_handler.py b/litellm/interactions/agents/http_handler.py index 6359730519c..03ce26f4711 100644 --- a/litellm/interactions/agents/http_handler.py +++ b/litellm/interactions/agents/http_handler.py @@ -7,7 +7,7 @@ duplicated. BaseAgentsAPIConfig stays as pure transform code. """ from collections.abc import Coroutine -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -38,12 +38,12 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + ) -> AgentCreateResponse | Coroutine[Any, Any, AgentCreateResponse]: if _is_async: return self.async_create_agent( agents_api_config=agents_api_config, @@ -93,10 +93,10 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AgentCreateResponse: async_httpx_client = self._async_client(litellm_params, client) headers = agents_api_config.validate_environment( @@ -141,11 +141,11 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): agents_api_config: BaseAgentsAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: + ) -> AgentListResponse | Coroutine[Any, Any, AgentListResponse]: if _is_async: return self.async_list_agents( agents_api_config=agents_api_config, @@ -181,9 +181,9 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): agents_api_config: BaseAgentsAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AgentListResponse: async_httpx_client = self._async_client(litellm_params, client) headers = agents_api_config.validate_environment( @@ -216,11 +216,11 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + ) -> AgentCreateResponse | Coroutine[Any, Any, AgentCreateResponse]: if _is_async: return self.async_get_agent( agents_api_config=agents_api_config, @@ -259,9 +259,9 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AgentCreateResponse: async_httpx_client = self._async_client(litellm_params, client) headers = agents_api_config.validate_environment( @@ -295,11 +295,11 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: + ) -> AgentDeleteResult | Coroutine[Any, Any, AgentDeleteResult]: if _is_async: return self.async_delete_agent( agents_api_config=agents_api_config, @@ -338,9 +338,9 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AgentDeleteResult: async_httpx_client = self._async_client(litellm_params, client) headers = agents_api_config.validate_environment( @@ -374,11 +374,11 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: + ) -> AgentVersionsResponse | Coroutine[Any, Any, AgentVersionsResponse]: if _is_async: return self.async_list_agent_versions( agents_api_config=agents_api_config, @@ -417,9 +417,9 @@ class AgentsHTTPHandler(InteractionsHTTPHandler): name: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AgentVersionsResponse: async_httpx_client = self._async_client(litellm_params, client) headers = agents_api_config.validate_environment( diff --git a/litellm/interactions/agents/main.py b/litellm/interactions/agents/main.py index 9f54b929bdc..dfd3374c53a 100644 --- a/litellm/interactions/agents/main.py +++ b/litellm/interactions/agents/main.py @@ -32,7 +32,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -71,14 +71,14 @@ def _get_agents_api_config(custom_llm_provider: str): def _make_logging_obj( - kwargs: Dict[str, Any], + kwargs: dict[str, Any], model: str, custom_llm_provider: str, call_type: str, - optional_params: Dict[str, Any], + optional_params: dict[str, Any], ) -> LiteLLMLoggingObj: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model, @@ -97,13 +97,13 @@ def _make_logging_obj( @client async def acreate( name: str, - base_agent: Optional[str] = None, - instructions: Optional[str] = None, - base_environment: Optional[InteractionEnvironment] = None, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + base_agent: str | None = None, + instructions: str | None = None, + base_environment: InteractionEnvironment | None = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> AgentCreateResponse: """Async: Create a managed agent on the provider side.""" @@ -141,15 +141,15 @@ async def acreate( @client def create( name: str, - base_agent: Optional[str] = None, - instructions: Optional[str] = None, - base_environment: Optional[InteractionEnvironment] = None, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + base_agent: str | None = None, + instructions: str | None = None, + base_environment: InteractionEnvironment | None = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: +) -> AgentCreateResponse | Coroutine[Any, Any, AgentCreateResponse]: """ Sync: Create a managed agent on the provider side. @@ -206,9 +206,9 @@ def create( @client async def alist( - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> AgentListResponse: """Async: List all agents on the provider side.""" @@ -240,11 +240,11 @@ async def alist( @client def list( - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: +) -> AgentListResponse | Coroutine[Any, Any, AgentListResponse]: """Sync: List all agents on the provider side.""" local_vars = locals() custom_llm_provider = custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" @@ -280,9 +280,9 @@ def list( @client async def aget( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> AgentCreateResponse: """Async: Get a specific agent by name.""" @@ -316,11 +316,11 @@ async def aget( @client def get( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: +) -> AgentCreateResponse | Coroutine[Any, Any, AgentCreateResponse]: """Sync: Get a specific agent by name.""" local_vars = locals() custom_llm_provider = custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" @@ -357,9 +357,9 @@ def get( @client async def adelete( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> AgentDeleteResult: """Async: Delete a specific agent by name.""" @@ -393,11 +393,11 @@ async def adelete( @client def delete( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: +) -> AgentDeleteResult | Coroutine[Any, Any, AgentDeleteResult]: """Sync: Delete a specific agent by name.""" local_vars = locals() custom_llm_provider = custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" @@ -434,9 +434,9 @@ def delete( @client async def alist_versions( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> AgentVersionsResponse: """Async: List versions of a specific agent.""" @@ -470,11 +470,11 @@ async def alist_versions( @client def list_versions( name: str, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: +) -> AgentVersionsResponse | Coroutine[Any, Any, AgentVersionsResponse]: """Sync: List versions of a specific agent.""" local_vars = locals() custom_llm_provider = custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" diff --git a/litellm/interactions/agents/utils.py b/litellm/interactions/agents/utils.py index a195dca1d25..56f38d9a621 100644 --- a/litellm/interactions/agents/utils.py +++ b/litellm/interactions/agents/utils.py @@ -3,16 +3,15 @@ Utility functions for the Agents API SDK. """ from collections.abc import Mapping -from typing import Dict, Optional from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig def merge_agent_headers( *, - dynamic_headers: Optional[Mapping[str, str]] = None, - static_headers: Optional[Mapping[str, str]] = None, -) -> Optional[Dict[str, str]]: + dynamic_headers: Mapping[str, str] | None = None, + static_headers: Mapping[str, str] | None = None, +) -> dict[str, str] | None: """Merge outbound HTTP headers for A2A agent calls. Merge rules: @@ -24,7 +23,7 @@ def merge_agent_headers( If both contain the same header (case-insensitively), ``static_headers`` wins. """ - merged: Dict[str, str] = {} + merged: dict[str, str] = {} if dynamic_headers: merged.update({str(k): str(v) for k, v in dynamic_headers.items()}) @@ -38,8 +37,8 @@ def merge_agent_headers( def get_provider_agents_api_config( - custom_llm_provider: Optional[str], -) -> Optional[BaseAgentsAPIConfig]: + custom_llm_provider: str | None, +) -> BaseAgentsAPIConfig | None: """ Return a provider-specific BaseAgentsAPIConfig if the provider has a native agent-creation API, or None otherwise. diff --git a/litellm/interactions/http_handler.py b/litellm/interactions/http_handler.py index a04c2ca3807..ab5b5f6e9d9 100644 --- a/litellm/interactions/http_handler.py +++ b/litellm/interactions/http_handler.py @@ -7,9 +7,6 @@ This module handles the HTTP communication for the Google Interactions API. from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( Any, - Dict, - Optional, - Union, ) import httpx @@ -60,14 +57,14 @@ class _BaseHTTPHandler: def _sync_client( self, litellm_params: GenericLiteLLMParams, - client: Optional[HTTPHandler], + client: HTTPHandler | None, ) -> HTTPHandler: return client or _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) def _async_client( self, litellm_params: GenericLiteLLMParams, - client: Optional[AsyncHTTPHandler], + client: AsyncHTTPHandler | None, ) -> AsyncHTTPHandler: # GenericLiteLLMParams.get uses getattr; an unset field is None, not the default. custom_llm_provider = litellm_params.get("custom_llm_provider") or "gemini" @@ -98,24 +95,20 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - model: Optional[str] = None, - agent: Optional[str] = None, - input: Optional[InteractionInput] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + model: str | None = None, + agent: str | None = None, + input: InteractionInput | None = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - stream: Optional[bool] = None, - ) -> Union[ - InteractionsAPIResponse, - Iterator[InteractionsAPIStreamingResponse], - Coroutine[ - Any, - Any, - Union[InteractionsAPIResponse, AsyncIterator[InteractionsAPIStreamingResponse]], - ], - ]: + stream: bool | None = None, + ) -> ( + InteractionsAPIResponse + | Iterator[InteractionsAPIStreamingResponse] + | Coroutine[Any, Any, InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]] + ): """ Create a new interaction (synchronous or async based on _is_async flag). @@ -217,15 +210,15 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - model: Optional[str] = None, - agent: Optional[str] = None, - input: Optional[InteractionInput] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, - stream: Optional[bool] = None, - ) -> Union[InteractionsAPIResponse, AsyncIterator[InteractionsAPIStreamingResponse]]: + model: str | None = None, + agent: str | None = None, + input: InteractionInput | None = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, + stream: bool | None = None, + ) -> InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]: """ Create a new interaction (async version). """ @@ -308,7 +301,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): def _create_sync_streaming_iterator( self, response: httpx.Response, - model: Optional[str], + model: str | None, logging_obj: LiteLLMLoggingObj, interactions_api_config: BaseInteractionsAPIConfig, ) -> SyncInteractionsAPIStreamingIterator: @@ -327,7 +320,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): def _create_async_streaming_iterator( self, response: httpx.Response, - model: Optional[str], + model: str | None, logging_obj: LiteLLMLoggingObj, interactions_api_config: BaseInteractionsAPIConfig, ) -> InteractionsAPIStreamingIterator: @@ -354,11 +347,11 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[InteractionsAPIResponse, Coroutine[Any, Any, InteractionsAPIResponse]]: + ) -> InteractionsAPIResponse | Coroutine[Any, Any, InteractionsAPIResponse]: """Get an interaction by ID.""" if _is_async: return self.async_get_interaction( @@ -416,9 +409,9 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> InteractionsAPIResponse: """Get an interaction by ID (async version).""" if client is None: @@ -473,11 +466,11 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[DeleteInteractionResult, Coroutine[Any, Any, DeleteInteractionResult]]: + ) -> DeleteInteractionResult | Coroutine[Any, Any, DeleteInteractionResult]: """Delete an interaction by ID.""" if _is_async: return self.async_delete_interaction( @@ -536,9 +529,9 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> DeleteInteractionResult: """Delete an interaction by ID (async version).""" if client is None: @@ -594,11 +587,11 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, _is_async: bool = False, - ) -> Union[CancelInteractionResult, Coroutine[Any, Any, CancelInteractionResult]]: + ) -> CancelInteractionResult | Coroutine[Any, Any, CancelInteractionResult]: """Cancel an interaction by ID.""" if _is_async: return self.async_cancel_interaction( @@ -657,9 +650,9 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> CancelInteractionResult: """Cancel an interaction by ID (async version).""" if client is None: diff --git a/litellm/interactions/litellm_responses_transformation/__init__.py b/litellm/interactions/litellm_responses_transformation/__init__.py index 6f6b32503d2..e1932e002b2 100644 --- a/litellm/interactions/litellm_responses_transformation/__init__.py +++ b/litellm/interactions/litellm_responses_transformation/__init__.py @@ -10,6 +10,6 @@ from litellm.interactions.litellm_responses_transformation.transformation import ) __all__ = [ - "LiteLLMResponsesInteractionsHandler", "LiteLLMResponsesInteractionsConfig", # Transformation config class (not BaseInteractionsAPIConfig) + "LiteLLMResponsesInteractionsHandler", ] diff --git a/litellm/interactions/litellm_responses_transformation/handler.py b/litellm/interactions/litellm_responses_transformation/handler.py index 28b485cb3f5..04409363d5f 100644 --- a/litellm/interactions/litellm_responses_transformation/handler.py +++ b/litellm/interactions/litellm_responses_transformation/handler.py @@ -5,9 +5,6 @@ Handler for transforming interactions API requests to litellm.responses requests from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( Any, - Dict, - Optional, - Union, cast, ) @@ -34,24 +31,17 @@ class LiteLLMResponsesInteractionsHandler: def interactions_api_handler( self, model: str, - input: Optional[InteractionInput], + input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, _is_async: bool = False, - stream: Optional[bool] = None, + stream: bool | None = None, **kwargs, - ) -> Union[ - InteractionsAPIResponse, - Iterator[InteractionsAPIStreamingResponse], - Coroutine[ - Any, - Any, - Union[ - InteractionsAPIResponse, - AsyncIterator[InteractionsAPIStreamingResponse], - ], - ], - ]: + ) -> ( + InteractionsAPIResponse + | Iterator[InteractionsAPIStreamingResponse] + | Coroutine[Any, Any, InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]] + ): """ Handle Interactions API request by calling litellm.responses(). @@ -116,12 +106,12 @@ class LiteLLMResponsesInteractionsHandler: async def async_interactions_api_handler( self, - responses_request: Dict[str, Any], + responses_request: dict[str, Any], model: str, - input: Optional[InteractionInput], + input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, **kwargs, - ) -> Union[InteractionsAPIResponse, AsyncIterator[InteractionsAPIStreamingResponse]]: + ) -> InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]: """Async handler for interactions API requests.""" # Call litellm.aresponses() # Note: litellm.aresponses() returns Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator] diff --git a/litellm/interactions/litellm_responses_transformation/streaming_iterator.py b/litellm/interactions/litellm_responses_transformation/streaming_iterator.py index 670f18c598c..b13f6661872 100644 --- a/litellm/interactions/litellm_responses_transformation/streaming_iterator.py +++ b/litellm/interactions/litellm_responses_transformation/streaming_iterator.py @@ -6,10 +6,6 @@ from collections import deque from collections.abc import AsyncIterator, Iterator from typing import ( Any, - Deque, - Dict, - List, - Optional, cast, ) @@ -51,10 +47,10 @@ class LiteLLMResponsesInteractionsStreamingIterator: self, model: str, litellm_custom_stream_wrapper: BaseResponsesAPIStreamingIterator, - request_input: Optional[InteractionInput], + request_input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, - custom_llm_provider: Optional[str] = None, - litellm_metadata: Optional[Dict[str, Any]] = None, + custom_llm_provider: str | None = None, + litellm_metadata: dict[str, Any] | None = None, ): import litellm @@ -78,14 +74,14 @@ class LiteLLMResponsesInteractionsStreamingIterator: # produces interaction.created + step.start + step.delta), and the # terminal sequence on stream end may also span multiple events # (step.stop + interaction.completed). - self._pending_events: Deque[InteractionsAPIStreamingResponse] = deque() + self._pending_events: deque[InteractionsAPIStreamingResponse] = deque() # Tracks whether we've already emitted a terminal completion event so # the StopIteration fallback path doesn't double-emit. self._sent_completion_event = False # ID resolved from the first upstream chunk (item_id on a text delta or # response.id on response.created). Persisted so the EOF terminal # events stay correlated with the start events delivered earlier. - self._interaction_id: Optional[str] = None + self._interaction_id: str | None = None # ------------------------------------------------------------------ # Event builders @@ -129,7 +125,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: delta={"type": "text", "text": delta_text}, ) - def _build_content_stop_event(self, interaction_id: Optional[str]) -> InteractionsAPIStreamingResponse: + def _build_content_stop_event(self, interaction_id: str | None) -> InteractionsAPIStreamingResponse: if self._use_legacy: return InteractionsAPIStreamingResponse( event_type="content.stop", @@ -172,7 +168,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: def _events_for_chunk( self, responses_chunk: ResponsesAPIStreamingResponse - ) -> List[InteractionsAPIStreamingResponse]: + ) -> list[InteractionsAPIStreamingResponse]: """ Translate a single upstream Responses API chunk into the list of Interactions API events it should produce. @@ -192,7 +188,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: if self._interaction_id is None: self._interaction_id = interaction_id - events: List[InteractionsAPIStreamingResponse] = [] + events: list[InteractionsAPIStreamingResponse] = [] if not self.sent_interaction_start: self.sent_interaction_start = True events.append(self._build_interaction_start_event(interaction_id)) @@ -226,7 +222,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: response = responses_chunk.response response_id = self._interaction_id or getattr(response, "id", None) or f"interaction_{id(self)}" - terminal: List[InteractionsAPIStreamingResponse] = [] + terminal: list[InteractionsAPIStreamingResponse] = [] if self.sent_content_start: terminal.append(self._build_content_stop_event(response_id)) terminal.append(self._build_completion_event(response_id)) @@ -237,7 +233,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: def _build_terminal_events_on_eof( self, - ) -> List[InteractionsAPIStreamingResponse]: + ) -> list[InteractionsAPIStreamingResponse]: """ Build the events to flush when the upstream stream ends without a ResponseCompletedEvent. Ensures consumers always observe a terminal @@ -247,7 +243,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: return [] fallback_id = self._interaction_id or f"interaction_{id(self)}" - terminal: List[InteractionsAPIStreamingResponse] = [] + terminal: list[InteractionsAPIStreamingResponse] = [] if self.sent_content_start: terminal.append(self._build_content_stop_event(fallback_id)) if self.sent_interaction_start or self.collected_text: @@ -319,7 +315,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: def _transform_responses_chunk_to_interactions_chunk( self, responses_chunk: ResponsesAPIStreamingResponse, - ) -> Optional[InteractionsAPIStreamingResponse]: + ) -> InteractionsAPIStreamingResponse | None: """ Compatibility shim: returns the *first* event produced for this chunk and queues any remaining events on ``self._pending_events`` so they diff --git a/litellm/interactions/litellm_responses_transformation/transformation.py b/litellm/interactions/litellm_responses_transformation/transformation.py index a2d8ebc5d4c..b4849190eaa 100644 --- a/litellm/interactions/litellm_responses_transformation/transformation.py +++ b/litellm/interactions/litellm_responses_transformation/transformation.py @@ -6,7 +6,7 @@ This module handles transforming between: - Responses API format (OpenAI's format with input[], instructions, etc.) """ -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast from litellm.types.interactions import ( InteractionInput, @@ -26,10 +26,10 @@ class LiteLLMResponsesInteractionsConfig: @staticmethod def transform_interactions_request_to_responses_request( model: str, - input: Optional[InteractionInput], + input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform an Interactions API request to a Responses API request. @@ -39,7 +39,7 @@ class LiteLLMResponsesInteractionsConfig: - tools -> tools (similar format) - generation_config -> temperature, top_p, etc. """ - responses_request: Dict[str, Any] = { + responses_request: dict[str, Any] = { "model": model, } @@ -127,7 +127,7 @@ class LiteLLMResponsesInteractionsConfig: # Ensure content is a list for _transform_content_array # Cast to List[Any] to handle various content types if isinstance(content, list): - content_list: List[Any] = list(content) + content_list: list[Any] = list(content) elif content is not None: content_list = [content] else: @@ -162,13 +162,13 @@ class LiteLLMResponsesInteractionsConfig: return cast(ResponseInputParam, str(input)) @staticmethod - def _transform_content_array(content: List[Any]) -> List[Dict[str, Any]]: + def _transform_content_array(content: list[Any]) -> list[dict[str, Any]]: """Transform Interactions API content array to Responses API format.""" if not isinstance(content, list): # Single content item - wrap in array content = [content] - transformed: List[Dict[str, Any]] = [] + transformed: list[dict[str, Any]] = [] for item in content: if isinstance(item, dict): # Already in dict format, pass through @@ -201,7 +201,7 @@ class LiteLLMResponsesInteractionsConfig: @staticmethod def transform_responses_response_to_interactions_response( responses_response: ResponsesAPIResponse, - model: Optional[str] = None, + model: str | None = None, ) -> InteractionsAPIResponse: """ Transform a Responses API response to an Interactions API response. @@ -213,15 +213,15 @@ class LiteLLMResponsesInteractionsConfig: - Extract usage """ # Extract text from outputs and build both `outputs` (legacy) and `steps` (new schema). - outputs: List[Dict[str, Any]] = [] - steps: List[Dict[str, Any]] = [] + outputs: list[dict[str, Any]] = [] + steps: list[dict[str, Any]] = [] if hasattr(responses_response, "output") and responses_response.output: for output_item in responses_response.output: # Use getattr with None default to safely access content content = getattr(output_item, "content", None) if content is not None: content_items = content if isinstance(content, list) else [content] - model_output_contents: List[Dict[str, Any]] = [] + model_output_contents: list[dict[str, Any]] = [] for content_item in content_items: # Check if content_item has text attribute text = getattr(content_item, "text", None) @@ -263,7 +263,7 @@ class LiteLLMResponsesInteractionsConfig: # Build interactions response — populate both `outputs` (legacy schema) and # `steps` (new schema) so callers work regardless of which schema they expect. - interactions_response_dict: Dict[str, Any] = { + interactions_response_dict: dict[str, Any] = { "id": getattr(responses_response, "id", ""), "object": "interaction", "status": interactions_status, diff --git a/litellm/interactions/main.py b/litellm/interactions/main.py index 481b1052790..44af5ffde93 100644 --- a/litellm/interactions/main.py +++ b/litellm/interactions/main.py @@ -35,7 +35,7 @@ import asyncio import contextvars from collections.abc import AsyncIterator, Coroutine, Iterator from functools import partial -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -66,38 +66,38 @@ from litellm.utils import client @client async def acreate( # Model or Agent (one required per OpenAPI spec) - model: Optional[str] = None, - agent: Optional[str] = None, + model: str | None = None, + agent: str | None = None, # Input (required) - input: Optional[InteractionInput] = None, + input: InteractionInput | None = None, # Tools (for model interactions) - tools: Optional[List[InteractionTool]] = None, + tools: list[InteractionTool] | None = None, # System instruction - system_instruction: Optional[str] = None, + system_instruction: str | None = None, # Generation config - generation_config: Optional[Dict[str, Any]] = None, + generation_config: dict[str, Any] | None = None, # Streaming - stream: Optional[bool] = None, + stream: bool | None = None, # Storage - store: Optional[bool] = None, + store: bool | None = None, # Background execution - background: Optional[bool] = None, + background: bool | None = None, # Agent execution environment ("remote", env id, or remote config object) - environment: Optional[InteractionEnvironment] = None, + environment: InteractionEnvironment | None = None, # Response format - response_modalities: Optional[List[str]] = None, - response_format: Optional[Dict[str, Any]] = None, - response_mime_type: Optional[str] = None, + response_modalities: list[str] | None = None, + response_format: dict[str, Any] | None = None, + response_mime_type: str | None = None, # Continuation - previous_interaction_id: Optional[str] = None, + previous_interaction_id: str | None = None, # Extra params - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM params - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[InteractionsAPIResponse, AsyncIterator[InteractionsAPIStreamingResponse]]: +) -> InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]: """ Async: Create a new interaction using Google's Interactions API. @@ -185,46 +185,42 @@ async def acreate( @client def create( # Model or Agent (one required per OpenAPI spec) - model: Optional[str] = None, - agent: Optional[str] = None, + model: str | None = None, + agent: str | None = None, # Input (required) - input: Optional[InteractionInput] = None, + input: InteractionInput | None = None, # Tools (for model interactions) - tools: Optional[List[InteractionTool]] = None, + tools: list[InteractionTool] | None = None, # System instruction - system_instruction: Optional[str] = None, + system_instruction: str | None = None, # Generation config - generation_config: Optional[Dict[str, Any]] = None, + generation_config: dict[str, Any] | None = None, # Streaming - stream: Optional[bool] = None, + stream: bool | None = None, # Storage - store: Optional[bool] = None, + store: bool | None = None, # Background execution - background: Optional[bool] = None, + background: bool | None = None, # Agent execution environment ("remote", env id, or remote config object) - environment: Optional[InteractionEnvironment] = None, + environment: InteractionEnvironment | None = None, # Response format - response_modalities: Optional[List[str]] = None, - response_format: Optional[Dict[str, Any]] = None, - response_mime_type: Optional[str] = None, + response_modalities: list[str] | None = None, + response_format: dict[str, Any] | None = None, + response_mime_type: str | None = None, # Continuation - previous_interaction_id: Optional[str] = None, + previous_interaction_id: str | None = None, # Extra params - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM params - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ - InteractionsAPIResponse, - Iterator[InteractionsAPIStreamingResponse], - Coroutine[ - Any, - Any, - Union[InteractionsAPIResponse, AsyncIterator[InteractionsAPIStreamingResponse]], - ], -]: +) -> ( + InteractionsAPIResponse + | Iterator[InteractionsAPIStreamingResponse] + | Coroutine[Any, Any, InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]] +): """ Sync: Create a new interaction using Google's Interactions API. @@ -260,7 +256,7 @@ def create( try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acreate_interaction", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -353,9 +349,9 @@ def create( @client async def aget( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> InteractionsAPIResponse: """Async: Get an interaction by its ID.""" @@ -396,18 +392,18 @@ async def aget( @client def get( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[InteractionsAPIResponse, Coroutine[Any, Any, InteractionsAPIResponse]]: +) -> InteractionsAPIResponse | Coroutine[Any, Any, InteractionsAPIResponse]: """Sync: Get an interaction by its ID.""" local_vars = locals() custom_llm_provider = custom_llm_provider or "gemini" try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aget_interaction", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -455,9 +451,9 @@ def get( @client async def adelete( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> DeleteInteractionResult: """Async: Delete an interaction by its ID.""" @@ -498,18 +494,18 @@ async def adelete( @client def delete( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[DeleteInteractionResult, Coroutine[Any, Any, DeleteInteractionResult]]: +) -> DeleteInteractionResult | Coroutine[Any, Any, DeleteInteractionResult]: """Sync: Delete an interaction by its ID.""" local_vars = locals() custom_llm_provider = custom_llm_provider or "gemini" try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete_interaction", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -557,9 +553,9 @@ def delete( @client async def acancel( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> CancelInteractionResult: """Async: Cancel an interaction by its ID.""" @@ -600,18 +596,18 @@ async def acancel( @client def cancel( interaction_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[CancelInteractionResult, Coroutine[Any, Any, CancelInteractionResult]]: +) -> CancelInteractionResult | Coroutine[Any, Any, CancelInteractionResult]: """Sync: Cancel an interaction by its ID.""" local_vars = locals() custom_llm_provider = custom_llm_provider or "gemini" try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acancel_interaction", False) is True litellm_params = GenericLiteLLMParams(**kwargs) diff --git a/litellm/interactions/streaming_iterator.py b/litellm/interactions/streaming_iterator.py index 0d9d1b4579c..f8ac4f25ebc 100644 --- a/litellm/interactions/streaming_iterator.py +++ b/litellm/interactions/streaming_iterator.py @@ -8,7 +8,7 @@ from the Google Interactions API, similar to the responses API streaming iterato import asyncio import json from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any import httpx @@ -36,18 +36,18 @@ class BaseInteractionsAPIStreamingIterator: def __init__( self, response: httpx.Response, - model: Optional[str], + model: str | None, interactions_api_config: BaseInteractionsAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, ): self.response = response self.model = model self.logging_obj = logging_obj self.finished = False self.interactions_api_config = interactions_api_config - self.completed_response: Optional[InteractionsAPIStreamingResponse] = None + self.completed_response: InteractionsAPIStreamingResponse | None = None self.start_time = datetime.now() # set request kwargs @@ -59,14 +59,14 @@ class BaseInteractionsAPIStreamingIterator: model=model or "", optional_params=self.logging_obj.model_call_details.get("litellm_params", {}), ) - _model_info: Dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {} + _model_info: dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {} self._hidden_params = { "model_id": _model_info.get("id", None), "api_base": _api_base, } self._hidden_params["additional_headers"] = process_response_headers(self.response.headers or {}) - def _process_chunk(self, chunk: str) -> Optional[InteractionsAPIStreamingResponse]: + def _process_chunk(self, chunk: str) -> InteractionsAPIStreamingResponse | None: """Process a single chunk of data from the stream.""" if not chunk: return None @@ -114,7 +114,6 @@ class BaseInteractionsAPIStreamingIterator: def _handle_logging_completed_response(self): """Base implementation - should be overridden by subclasses.""" - pass class InteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator): @@ -125,11 +124,11 @@ class InteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator): def __init__( self, response: httpx.Response, - model: Optional[str], + model: str | None, interactions_api_config: BaseInteractionsAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, ): super().__init__( response=response, @@ -192,11 +191,11 @@ class SyncInteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator) def __init__( self, response: httpx.Response, - model: Optional[str], + model: str | None, interactions_api_config: BaseInteractionsAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, ): super().__init__( response=response, diff --git a/litellm/interactions/utils.py b/litellm/interactions/utils.py index 3dffaa538ba..135ec6f11dc 100644 --- a/litellm/interactions/utils.py +++ b/litellm/interactions/utils.py @@ -2,7 +2,7 @@ Utility functions for Interactions API. """ -from typing import Any, Dict, Optional, cast +from typing import Any, cast from litellm.llms.base_llm.interactions.transformation import BaseInteractionsAPIConfig from litellm.types.interactions import InteractionsAPIOptionalRequestParams @@ -26,8 +26,8 @@ INTERACTIONS_API_OPTIONAL_PARAMS = { def get_provider_interactions_api_config( provider: str, - model: Optional[str] = None, -) -> Optional[BaseInteractionsAPIConfig]: + model: str | None = None, +) -> BaseInteractionsAPIConfig | None: """ Get the interactions API config for the given provider. @@ -55,7 +55,7 @@ class InteractionsAPIRequestUtils: @staticmethod def get_requested_interactions_api_optional_params( - params: Dict[str, Any], + params: dict[str, Any], ) -> InteractionsAPIOptionalRequestParams: """ Filter parameters to only include valid optional params per OpenAPI spec. diff --git a/litellm/litellm_core_utils/api_route_to_call_types.py b/litellm/litellm_core_utils/api_route_to_call_types.py index 2ae9986ce94..ec5ca46399c 100644 --- a/litellm/litellm_core_utils/api_route_to_call_types.py +++ b/litellm/litellm_core_utils/api_route_to_call_types.py @@ -8,8 +8,6 @@ Route patterns may contain placeholders like {agent_id}, {model}, {batch_id}; th match a single path segment when resolving call types for a concrete path. """ -from typing import List, Optional - from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes @@ -30,7 +28,7 @@ def _route_matches_pattern(route: str, pattern: str) -> bool: return True -def get_call_types_for_route(route: str) -> Optional[List[CallTypes]]: +def get_call_types_for_route(route: str) -> list[CallTypes] | None: """ Get the list of CallTypes for a given API route. diff --git a/litellm/litellm_core_utils/app_crypto.py b/litellm/litellm_core_utils/app_crypto.py index e47962d6a36..862b53eaf86 100644 --- a/litellm/litellm_core_utils/app_crypto.py +++ b/litellm/litellm_core_utils/app_crypto.py @@ -1,7 +1,6 @@ import base64 import json import os -from typing import Optional from cryptography.hazmat.primitives.ciphers.aead import AESGCM @@ -12,7 +11,7 @@ class AppCrypto: raise ValueError("Master key must be 32 bytes for AES-256-GCM") self.key = master_key - def encrypt_json(self, data: dict, aad: Optional[bytes] = None) -> dict: + def encrypt_json(self, data: dict, aad: bytes | None = None) -> dict: aes = AESGCM(self.key) nonce = os.urandom(12) plaintext = json.dumps(data).encode("utf-8") @@ -24,7 +23,7 @@ class AppCrypto: "tag": base64.b64encode(tag).decode(), } - def decrypt_json(self, enc: dict, aad: Optional[bytes] = None) -> dict: + def decrypt_json(self, enc: dict, aad: bytes | None = None) -> dict: aes = AESGCM(self.key) nonce = base64.b64decode(enc["nonce"]) ct = base64.b64decode(enc["ciphertext"]) diff --git a/litellm/litellm_core_utils/asyncify.py b/litellm/litellm_core_utils/asyncify.py index 0b174c8621c..bdfd6d0cd3b 100644 --- a/litellm/litellm_core_utils/asyncify.py +++ b/litellm/litellm_core_utils/asyncify.py @@ -1,7 +1,6 @@ import asyncio import functools from collections.abc import Awaitable, Callable -from typing import Optional import anyio import anyio.to_thread @@ -23,7 +22,7 @@ def asyncify( function: Callable[T_ParamSpec, T_Retval], *, cancellable: bool = False, - limiter: Optional[anyio.CapacityLimiter] = None, + limiter: anyio.CapacityLimiter | None = None, ) -> Callable[T_ParamSpec, Awaitable[T_Retval]]: """ Take a blocking function and create an async one that receives the same diff --git a/litellm/litellm_core_utils/audio_utils/utils.py b/litellm/litellm_core_utils/audio_utils/utils.py index e5007ceec34..a78d2ca9332 100644 --- a/litellm/litellm_core_utils/audio_utils/utils.py +++ b/litellm/litellm_core_utils/audio_utils/utils.py @@ -5,7 +5,6 @@ Utils used for litellm.transcription() and litellm.atranscription() import hashlib import os from dataclasses import dataclass -from typing import Optional from litellm.types.files import get_file_mime_type_from_extension from litellm.types.utils import FileTypes @@ -174,8 +173,8 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: Compute SHA-256 hash of audio file content for cache keys. Falls back to filename hash if content extraction fails. """ - file_content: Optional[bytes] = None - fallback_filename: Optional[str] = None + file_content: bytes | None = None + fallback_filename: str | None = None if isinstance(file_obj, tuple): if len(file_obj) < 2: @@ -203,7 +202,7 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: file_content = f.read() if fallback_filename is None: fallback_filename = str(file_content_obj) - except (OSError, IOError): + except OSError: fallback_filename = str(file_content_obj) file_content = None elif hasattr(file_content_obj, "read"): @@ -214,7 +213,7 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: file_content = file_content_obj.read() # type: ignore if current_position is not None and hasattr(file_content_obj, "seek"): file_content_obj.seek(current_position) # type: ignore - except (OSError, IOError, AttributeError): + except (OSError, AttributeError): file_content = None else: file_content = None @@ -248,7 +247,7 @@ def get_audio_file_for_health_check() -> FileTypes: return open(file_path, "rb") -def calculate_request_duration(file: FileTypes) -> Optional[float]: +def calculate_request_duration(file: FileTypes) -> float | None: """ Calculate audio duration from file content. @@ -268,7 +267,7 @@ def calculate_request_duration(file: FileTypes) -> Optional[float]: import io # Handle different file input types - file_content: Optional[bytes] = None + file_content: bytes | None = None if isinstance(file, (bytes, bytearray)): # Raw bytes diff --git a/litellm/litellm_core_utils/cached_imports.py b/litellm/litellm_core_utils/cached_imports.py index b897f61757a..2f600b91935 100644 --- a/litellm/litellm_core_utils/cached_imports.py +++ b/litellm/litellm_core_utils/cached_imports.py @@ -6,7 +6,7 @@ inside functions that are critical to performance. """ from collections.abc import Callable -from typing import TYPE_CHECKING, Optional, Type +from typing import TYPE_CHECKING, Optional # Type annotations for cached imports if TYPE_CHECKING: @@ -14,12 +14,12 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging # Global cache variables -_LiteLLMLogging: Optional[Type["Logging"]] = None +_LiteLLMLogging: type["Logging"] | None = None _coroutine_checker: Optional["CoroutineChecker"] = None -_set_callbacks: Optional[Callable] = None +_set_callbacks: Callable | None = None -def get_litellm_logging_class() -> Type["Logging"]: +def get_litellm_logging_class() -> type["Logging"]: """Get the cached LiteLLM Logging class, initializing if needed.""" global _LiteLLMLogging if _LiteLLMLogging is not None: diff --git a/litellm/litellm_core_utils/cli_token_utils.py b/litellm/litellm_core_utils/cli_token_utils.py index e730f60bc3b..71324dcd705 100644 --- a/litellm/litellm_core_utils/cli_token_utils.py +++ b/litellm/litellm_core_utils/cli_token_utils.py @@ -9,7 +9,6 @@ import json import os import time from pathlib import Path -from typing import Optional def get_cli_token_file_path() -> str: @@ -19,7 +18,7 @@ def get_cli_token_file_path() -> str: return str(config_dir / "token.json") -def load_cli_token() -> Optional[dict]: +def load_cli_token() -> dict | None: """Load CLI token data from file""" token_file = get_cli_token_file_path() if not os.path.exists(token_file): @@ -28,13 +27,13 @@ def load_cli_token() -> Optional[dict]: try: with open(token_file, "r") as f: return json.load(f) - except (json.JSONDecodeError, IOError): + except (OSError, json.JSONDecodeError): return None def get_litellm_gateway_api_key( - expected_base_url: Optional[str] = None, -) -> Optional[str]: + expected_base_url: str | None = None, +) -> str | None: """ Get the stored CLI API key for use with LiteLLM SDK. diff --git a/litellm/litellm_core_utils/cloud_storage_security.py b/litellm/litellm_core_utils/cloud_storage_security.py index 077af32e3d0..e106f20dee0 100644 --- a/litellm/litellm_core_utils/cloud_storage_security.py +++ b/litellm/litellm_core_utils/cloud_storage_security.py @@ -2,7 +2,7 @@ import posixpath import re from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import Any, Optional, Tuple, cast +from typing import Any, cast from urllib.parse import quote, unquote from litellm._uuid import uuid @@ -34,7 +34,7 @@ def is_managed_cloud_storage_uri(file_id: str) -> bool: _SAFE_OBJECT_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") -def sanitize_cloud_object_component(value: Optional[str], fallback: str = "file") -> str: +def sanitize_cloud_object_component(value: str | None, fallback: str = "file") -> str: if not isinstance(value, str): return fallback @@ -50,7 +50,7 @@ def sanitize_cloud_object_component(value: Optional[str], fallback: str = "file" return component[:255] -def sanitize_cloud_object_path(value: Optional[str], fallback: str = "file") -> str: +def sanitize_cloud_object_path(value: str | None, fallback: str = "file") -> str: if not isinstance(value, str): return fallback @@ -65,7 +65,7 @@ def sanitize_cloud_object_path(value: Optional[str], fallback: str = "file") -> return "/".join(segments) -def build_managed_cloud_object_name(prefix: str, filename: Optional[str], fallback_filename: str = "file") -> str: +def build_managed_cloud_object_name(prefix: str, filename: str | None, fallback_filename: str = "file") -> str: safe_filename = sanitize_cloud_object_component(filename, fallback=fallback_filename) return f"{prefix}{uuid.uuid4().hex}-{safe_filename}" @@ -84,7 +84,7 @@ def _validate_cloud_object_path(object_name: str) -> None: raise ValueError("Cloud storage object name contains an invalid path segment") -def split_configured_cloud_bucket_name(bucket_name: str) -> Tuple[str, str]: +def split_configured_cloud_bucket_name(bucket_name: str) -> tuple[str, str]: if not isinstance(bucket_name, str) or not bucket_name.strip(): raise ValueError("Cloud storage bucket name is required") @@ -116,7 +116,7 @@ def encode_s3_object_key_for_url(object_key: str) -> str: def should_allow_legacy_cloud_file_ids( - litellm_params: Optional[Mapping[str, Any]] = None, + litellm_params: Mapping[str, Any] | None = None, ) -> bool: value = None if isinstance(litellm_params, Mapping): @@ -137,7 +137,7 @@ def validate_managed_cloud_file_id( configured_bucket_name: str, allowed_object_prefixes: Sequence[str], allow_legacy_cloud_file_ids: bool = False, -) -> Tuple[str, str]: +) -> tuple[str, str]: decoded_file_id = unquote(file_id) if not decoded_file_id.startswith(scheme): raise ValueError(f"file_id must be a {scheme} URI") diff --git a/litellm/litellm_core_utils/completion_timeout.py b/litellm/litellm_core_utils/completion_timeout.py index bf6e6315c66..9f08fcc6bc4 100644 --- a/litellm/litellm_core_utils/completion_timeout.py +++ b/litellm/litellm_core_utils/completion_timeout.py @@ -3,7 +3,6 @@ from __future__ import annotations from collections.abc import Callable -from typing import Optional, Union import httpx @@ -15,7 +14,7 @@ class CompletionTimeout: @staticmethod def _fallback_when_no_explicit_timeout( - global_timeout: Optional[Union[float, str]], + global_timeout: float | str | None, ) -> float: """ Used when ``model_timeout`` and kwargs timeouts are all unset. @@ -31,13 +30,13 @@ class CompletionTimeout: @staticmethod def resolve( - model_timeout: Optional[Union[float, str, httpx.Timeout]], + model_timeout: float | str | httpx.Timeout | None, kwargs: dict, custom_llm_provider: str, *, - global_timeout: Optional[Union[float, str]], + global_timeout: float | str | None, supports_httpx_timeout: Callable[[str], bool], - ) -> Union[float, httpx.Timeout]: + ) -> float | httpx.Timeout: """ Resolution order (first non-None wins): @@ -49,7 +48,7 @@ class CompletionTimeout: Coerce :class:`httpx.Timeout` when the provider does not support it. """ - resolved: Union[float, str, httpx.Timeout] + resolved: float | str | httpx.Timeout if model_timeout is not None: resolved = model_timeout elif kwargs.get("timeout") is not None: diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 60fca1ae037..f1f0f73889d 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -2,7 +2,7 @@ ## Helper utilities import copy from collections.abc import Iterable -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Literal, Union import httpx @@ -19,7 +19,7 @@ else: Span = Any -def safe_divide_seconds(seconds: float, denominator: float, default: Optional[float] = None) -> Optional[float]: +def safe_divide_seconds(seconds: float, denominator: float, default: float | None = None) -> float | None: """ Safely divide seconds by denominator, handling zero division. @@ -38,10 +38,10 @@ def safe_divide_seconds(seconds: float, denominator: float, default: Optional[fl def safe_divide( - numerator: Union[int, float], - denominator: Union[int, float], - default: Union[int, float] = 0, -) -> Union[int, float]: + numerator: float, + denominator: float, + default: float = 0, +) -> int | float: """ Safely divide two numbers, returning a default value if denominator is zero. @@ -143,7 +143,7 @@ def map_finish_reason(finish_reason: str) -> OpenAIChatCompletionFinishReason: def remove_index_from_tool_calls( - messages: Optional[List[AllMessageValues]], + messages: list[AllMessageValues] | None, ): if messages is not None: for message in messages: @@ -153,10 +153,8 @@ def remove_index_from_tool_calls( if isinstance(tool_call, dict) and "index" in tool_call: # Type guard to ensure it's a dict tool_call.pop("index", None) - return - -def remove_items_at_indices(items: Optional[List[Any]], indices: Iterable[int]) -> None: +def remove_items_at_indices(items: list[Any] | None, indices: Iterable[int]) -> None: """Remove items from a list in-place by index""" if items is None: return @@ -237,7 +235,7 @@ def get_litellm_metadata_from_kwargs(kwargs: dict): def reconstruct_model_name( model_name: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, metadata: dict, ) -> str: """Reconstruct full model name with provider prefix for logging.""" @@ -256,8 +254,8 @@ def reconstruct_model_name( # Helper functions used for OTEL logging def _get_parent_otel_span_from_kwargs( - kwargs: Optional[dict] = None, -) -> Union[Span, None]: + kwargs: dict | None = None, +) -> Span | None: try: if kwargs is None: return None @@ -280,7 +278,7 @@ def _get_parent_otel_span_from_kwargs( def process_response_headers( - response_headers: Union[httpx.Headers, dict], + response_headers: httpx.Headers | dict, preserve_litellm_internal_headers: bool = False, ) -> dict: """ @@ -357,7 +355,7 @@ def safe_deep_copy(data): if litellm.safe_memory_mode is True: return data - litellm_parent_otel_span: Optional[Any] = None + litellm_parent_otel_span: Any | None = None # Step 1: Remove the litellm_parent_otel_span litellm_parent_otel_span = None if isinstance(data, dict): @@ -455,7 +453,7 @@ def filter_exceptions_from_params(data: Any, max_depth: int = 20) -> Any: return data -def filter_internal_params(data: dict, additional_internal_params: Optional[set] = None) -> dict: +def filter_internal_params(data: dict, additional_internal_params: set | None = None) -> dict: """ Filter out LiteLLM internal parameters that shouldn't be sent to provider APIs. @@ -488,8 +486,8 @@ def filter_internal_params(data: dict, additional_internal_params: Optional[set] def redact_nested_match_and_regex_keys( - payload: Union[dict, List[Any], str, None], -) -> Union[dict, List[Any], str, None]: + payload: dict | list[Any] | str | None, +) -> dict | list[Any] | str | None: """ Deep-copy `payload` and replace every `match` / `regex` string field with "[REDACTED]" anywhere in nested dict/list structures. @@ -499,14 +497,14 @@ def redact_nested_match_and_regex_keys( if payload is None or isinstance(payload, str): return payload try: - redacted: Union[dict, List[Any], str, None] = copy.deepcopy(payload) + redacted: dict | list[Any] | str | None = copy.deepcopy(payload) except Exception: return payload # Iterative traversal; `seen` guards against cyclic refs preserved by deepcopy. try: seen: set = set() - stack: List[Any] = [redacted] + stack: list[Any] = [redacted] while stack: node = stack.pop() node_id = id(node) diff --git a/litellm/litellm_core_utils/coroutine_checker.py b/litellm/litellm_core_utils/coroutine_checker.py index bf065e5a153..e1ab495ca43 100644 --- a/litellm/litellm_core_utils/coroutine_checker.py +++ b/litellm/litellm_core_utils/coroutine_checker.py @@ -3,6 +3,7 @@ import inspect from typing import Any from weakref import WeakKeyDictionary + from litellm.constants import ( COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY, ) diff --git a/litellm/litellm_core_utils/credential_accessor.py b/litellm/litellm_core_utils/credential_accessor.py index 45e1ea2c498..fa4da59579a 100644 --- a/litellm/litellm_core_utils/credential_accessor.py +++ b/litellm/litellm_core_utils/credential_accessor.py @@ -1,7 +1,5 @@ """Utils for accessing credentials.""" -from typing import List - import litellm from litellm.types.utils import CredentialItem @@ -19,7 +17,7 @@ class CredentialAccessor: return {} @staticmethod - def upsert_credentials(credentials: List[CredentialItem]): + def upsert_credentials(credentials: list[CredentialItem]): """Add a credential to the list of credentials.""" credential_names = [cred.credential_name for cred in litellm.credential_list] diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index a7fae104c92..53ffab07c13 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -8,8 +8,6 @@ Example: "prometheus" -> PrometheusLogger """ -from typing import Union - from litellm import _custom_logger_compatible_callbacks_literal from litellm.integrations.agentops import AgentOps from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook @@ -25,8 +23,6 @@ from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger from litellm.integrations.deepeval import DeepEvalLogger from litellm.integrations.dotprompt import DotpromptManager from litellm.integrations.focus.focus_logger import FocusLogger -from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import MavvrikFocusLogger -from litellm.integrations.vantage.vantage_logger import VantageLogger from litellm.integrations.galileo import GalileoObserve from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger from litellm.integrations.gcs_pubsub.pub_sub import GcsPubSubLogger @@ -39,6 +35,7 @@ from litellm.integrations.langfuse.langfuse_prompt_management import ( from litellm.integrations.langsmith import LangsmithLogger from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver from litellm.integrations.literal_ai import LiteralAILogger +from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import MavvrikFocusLogger from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.newrelic import NewRelicLogger from litellm.integrations.openmeter import OpenMeterLogger @@ -48,6 +45,7 @@ from litellm.integrations.posthog import PostHogLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.integrations.s3_v2 import S3Logger from litellm.integrations.sqs import SQSLogger +from litellm.integrations.vantage.vantage_logger import VantageLogger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( VectorStorePreCallHook, ) @@ -140,7 +138,7 @@ class CustomLoggerRegistry: pass # enterprise not installed @classmethod - def get_callback_str_from_class_type(cls, class_type: type) -> Union[str, None]: + def get_callback_str_from_class_type(cls, class_type: type) -> str | None: """ Get the callback string from the class type. diff --git a/litellm/litellm_core_utils/dd_tracing.py b/litellm/litellm_core_utils/dd_tracing.py index 3a1bd72e1a5..aa5e23d3868 100644 --- a/litellm/litellm_core_utils/dd_tracing.py +++ b/litellm/litellm_core_utils/dd_tracing.py @@ -5,7 +5,7 @@ If the ddtrace package is not installed, the tracer will be a no-op. """ from contextlib import contextmanager -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any from litellm.secret_managers.main import get_secret_bool @@ -64,7 +64,7 @@ def _should_use_dd_profiler(): # Initialize tracer should_use_dd_tracer = _should_use_dd_tracer() -tracer: Union[NullTracer, DD_TRACER] = NullTracer() +tracer: NullTracer | DD_TRACER = NullTracer() # We need to ensure tracer is never None and always has the required methods if should_use_dd_tracer: try: @@ -78,7 +78,7 @@ else: tracer = NullTracer() -def get_active_span() -> Optional[Any]: +def get_active_span() -> Any | None: """ Return the active Datadog span, checking current span first and then root span. """ diff --git a/litellm/litellm_core_utils/default_encoding.py b/litellm/litellm_core_utils/default_encoding.py index 38aacb47f04..5a763ffc703 100644 --- a/litellm/litellm_core_utils/default_encoding.py +++ b/litellm/litellm_core_utils/default_encoding.py @@ -28,9 +28,10 @@ os.environ["TIKTOKEN_CACHE_DIR"] = ( cache_dir # use local copy of tiktoken b/c of - https://github.com/BerriAI/litellm/issues/1071 ) -import tiktoken -import time import random +import time + +import tiktoken # Retry logic to handle race conditions when multiple processes try to create # the tiktoken cache file simultaneously (common in parallel test execution on Windows) diff --git a/litellm/litellm_core_utils/dot_notation_indexing.py b/litellm/litellm_core_utils/dot_notation_indexing.py index 85abbdddffc..d6fd2cffaf6 100644 --- a/litellm/litellm_core_utils/dot_notation_indexing.py +++ b/litellm/litellm_core_utils/dot_notation_indexing.py @@ -23,12 +23,12 @@ Used by JWT Auth to get the user role from the token, and by additional_drop_params to remove nested fields from optional parameters. """ -from typing import Any, Dict, List, Optional, TypeVar, Union +from typing import Any, TypeVar T = TypeVar("T") -def get_nested_value(data: Dict[str, Any], key_path: str, default: Optional[T] = None) -> Optional[T]: +def get_nested_value(data: dict[str, Any], key_path: str, default: T | None = None) -> T | None: """ Retrieves a value from a nested dictionary using dot notation. @@ -107,7 +107,7 @@ def _parse_path_segments(path: str) -> list: def _delete_nested_value_custom( - data: Union[Dict[str, Any], List[Any]], + data: dict[str, Any] | list[Any], segments: list, segment_index: int = 0, ) -> None: @@ -178,11 +178,11 @@ def _delete_nested_value_custom( def delete_nested_value( - data: Dict[str, Any], + data: dict[str, Any], path: str, depth: int = 0, max_depth: int = 20, -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Delete a field from nested data using JSONPath notation. diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index b78a314dc45..5cc75b7dac2 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -9,7 +9,7 @@ duration_in_seconds is used in diff parts of the code base, example import re import time as time_module from datetime import datetime, time, timedelta, timezone, tzinfo -from typing import Final, Optional, Tuple +from typing import Final from zoneinfo import ZoneInfo from litellm._logging import verbose_logger @@ -26,7 +26,7 @@ def _normalize_duration(duration: str) -> str: return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration) -def _extract_from_regex(duration: str) -> Tuple[int, str]: +def _extract_from_regex(duration: str) -> tuple[int, str]: match = re.match(r"(\d+)(mo|[smhdw]?)", duration) if not match: @@ -86,8 +86,7 @@ def duration_in_seconds(duration: str) -> int: target_day = current_time.day last_day_of_target_month = get_last_day_of_month(target_year, target_month) - if target_day > last_day_of_target_month: - target_day = last_day_of_target_month + target_day = min(target_day, last_day_of_target_month) next_month = datetime( year=target_year, @@ -167,7 +166,7 @@ def get_next_standardized_reset_time( return base_midnight + timedelta(days=1) -def _setup_timezone(current_time: datetime, timezone_str: str = "UTC") -> Tuple[datetime, tzinfo]: +def _setup_timezone(current_time: datetime, timezone_str: str = "UTC") -> tuple[datetime, tzinfo]: """Set up timezone and normalize current time to that timezone.""" try: if timezone_str is None: @@ -190,7 +189,7 @@ def _setup_timezone(current_time: datetime, timezone_str: str = "UTC") -> Tuple[ return current_time, tz -def _parse_duration(duration: str) -> Tuple[Optional[int], Optional[str]]: +def _parse_duration(duration: str) -> tuple[int | None, str | None]: """Parse the duration string into value and unit.""" match = re.match(r"(\d+)([a-z]+)", duration) if not match: diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index fdab3d5b9d4..47b7aa6d568 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -1,7 +1,7 @@ import json import re import traceback -from typing import Any, Optional, Protocol, cast +from typing import Any, Protocol, cast import httpx @@ -121,7 +121,7 @@ class ExceptionCheckers: return False -def get_error_message(error_obj) -> Optional[str]: +def get_error_message(error_obj) -> str | None: """ OpenAI Returns Error message that is nested, this extract the message @@ -177,13 +177,13 @@ def _get_body_error_code(error_str: str) -> int | None: return None -def _get_response_headers(original_exception: Exception) -> Optional[httpx.Headers]: +def _get_response_headers(original_exception: Exception) -> httpx.Headers | None: """ Extract and return the response headers from an exception, if present. Used for accurate retry logic. """ - _response_headers: Optional[httpx.Headers] = None + _response_headers: httpx.Headers | None = None try: _response_headers = getattr(original_exception, "headers", None) error_response = getattr(original_exception, "response", None) @@ -198,7 +198,7 @@ def _get_response_headers(original_exception: Exception) -> Optional[httpx.Heade def extract_and_raise_litellm_exception( - response: Optional[Any], + response: Any | None, error_str: str, model: str, custom_llm_provider: str, @@ -273,7 +273,7 @@ def _map_openai_exception( message = message.replace("OPENAI", custom_llm_provider.upper()) message = message.replace( "openai.OpenAIError", - "{}.{}Error".format(custom_llm_provider, custom_llm_provider), + f"{custom_llm_provider}.{custom_llm_provider}Error", ) if custom_llm_provider == "openai": exception_provider = "OpenAI" + "Exception" @@ -507,31 +507,31 @@ def _map_anthropic_exception( or ExceptionCheckers.is_error_str_context_window_exceeded(error_str) ): raise ContextWindowExceededError( - message="AnthropicError - {}".format(error_str), + message=f"AnthropicError - {error_str}", model=model, llm_provider="anthropic", ) elif "overloaded_error" in error_str or "Overloaded" in error_str: raise InternalServerError( - message="AnthropicError - {}".format(error_str), + message=f"AnthropicError - {error_str}", model=model, llm_provider="anthropic", ) if "Invalid API Key" in error_str: raise AuthenticationError( - message="AnthropicError - {}".format(error_str), + message=f"AnthropicError - {error_str}", model=model, llm_provider="anthropic", ) if "content filtering policy" in error_str: raise ContentPolicyViolationError( - message="AnthropicError - {}".format(error_str), + message=f"AnthropicError - {error_str}", model=model, llm_provider="anthropic", ) if "Client error '400 Bad Request'" in error_str: raise BadRequestError( - message="AnthropicError - {}".format(error_str), + message=f"AnthropicError - {error_str}", model=model, llm_provider="anthropic", ) @@ -679,7 +679,7 @@ def _map_replicate_exception( ) raise APIError( status_code=500, - message=f"ReplicateException - {str(original_exception)}", + message=f"ReplicateException - {original_exception!s}", llm_provider="replicate", model=model, request=httpx.Request( @@ -1678,14 +1678,12 @@ def _map_together_ai_exception( model=model, llm_provider="together_ai", ) - elif "error" in error_response and "API key doesn't match expected format." in error_response["error"]: - raise BadRequestError( - message=f"TogetherAIException - {error_response['error']}", - model=model, - llm_provider="together_ai", - response=getattr(original_exception, "response", None), - ) - elif "error_type" in error_response and error_response["error_type"] == "validation": + elif ( + "error" in error_response + and "API key doesn't match expected format." in error_response["error"] + or "error_type" in error_response + and error_response["error_type"] == "validation" + ): raise BadRequestError( message=f"TogetherAIException - {error_response['error']}", model=model, @@ -1869,7 +1867,7 @@ def _map_azure_exception( # Azure OpenAI (especially Images) often nests error details under # body["error"]. Detect content policy violations using the structured # payload in addition to string matching. - azure_error_code: Optional[str] = None + azure_error_code: str | None = None try: body_dict = getattr(original_exception, "body", None) or {} if isinstance(body_dict, dict): @@ -2461,7 +2459,7 @@ def exception_type( # type: ignore ): # deal with edge-case invalid request error bug in openai-python sdk exception_mapping_worked = True raise BadRequestError( - message=f"{exception_provider} BadRequestError : This can happen due to missing AZURE_API_VERSION: {str(original_exception)}", + message=f"{exception_provider} BadRequestError : This can happen due to missing AZURE_API_VERSION: {original_exception!s}", model=model, llm_provider=custom_llm_provider, response=getattr(original_exception, "response", None), @@ -2473,17 +2471,14 @@ def exception_type( # type: ignore exception_mapping_worked = True if hasattr(original_exception, "request"): raise APIConnectionError( - message="{} - {}".format(exception_provider, error_str), + message=f"{exception_provider} - {error_str}", llm_provider=custom_llm_provider, model=model, request=getattr(original_exception, "request", None), ) else: raise APIConnectionError( - message="{}\n{}".format( - str(original_exception), - _redact_string(traceback.format_exc()), - ), + message=f"{original_exception!s}\n{_redact_string(traceback.format_exc())}", llm_provider=custom_llm_provider, model=model, request=httpx.Request(method="POST", url="https://api.openai.com/v1/"), # stub the request @@ -2509,10 +2504,7 @@ def exception_type( # type: ignore setattr(e, "litellm_response_headers", litellm_response_headers) raise e # it's already mapped raised_exc = APIConnectionError( - message="{}\n{}".format( - original_exception, - _redact_string(traceback.format_exc()), - ), + message=f"{original_exception}\n{_redact_string(traceback.format_exc())}", llm_provider="", model="", ) @@ -2548,7 +2540,6 @@ def exception_logging( verbose_logger.debug( f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {traceback.format_exc()}" ) - pass def _add_key_name_and_team_to_alert(request_info: str, metadata: dict) -> str: diff --git a/litellm/litellm_core_utils/fallback_generalizations.py b/litellm/litellm_core_utils/fallback_generalizations.py index 410bb9623fe..ce408dd9d37 100644 --- a/litellm/litellm_core_utils/fallback_generalizations.py +++ b/litellm/litellm_core_utils/fallback_generalizations.py @@ -48,7 +48,7 @@ O(number of rules); callers must only invoke them on a cache miss. import re from dataclasses import dataclass -from typing import Optional, Union +from typing import Union from litellm._logging import verbose_logger @@ -152,14 +152,14 @@ class _FallbackGeneralizations: self.routing_rules: tuple = () self.capability_rules: tuple = () - def set_rules(self, rules: Optional[list]) -> None: + def set_rules(self, rules: list | None) -> None: installed = rules if isinstance(rules, list) else [] compiled = tuple(kind for rule in _resolve_legacy_extends(installed) for kind in _compile_rule(rule)) self.rules = installed self.routing_rules = tuple(rule for rule in compiled if isinstance(rule, _RoutingRule)) self.capability_rules = tuple(rule for rule in compiled if isinstance(rule, _CapabilityRule)) - def match_routing(self, model: str) -> Optional[str]: + def match_routing(self, model: str) -> str | None: if not model: return None return next( @@ -167,7 +167,7 @@ class _FallbackGeneralizations: None, ) - def match_capabilities(self, model: str) -> Optional[dict]: + def match_capabilities(self, model: str) -> dict | None: if not model: return None matched = tuple(rule.model_info for rule in self.capability_rules if rule.pattern.search(model) is not None) @@ -179,7 +179,7 @@ class _FallbackGeneralizations: _registry = _FallbackGeneralizations() -def set_fallback_generalizations(rules: Optional[list]) -> None: +def set_fallback_generalizations(rules: list | None) -> None: """Install the active rule list, compiling and classifying each rule. Legacy ``extends`` inheritance is resolved here, once, before classification; @@ -195,7 +195,7 @@ def get_fallback_generalization_rules() -> list: return _registry.rules -def match_routing_generalization(model: str) -> Optional[str]: +def match_routing_generalization(model: str) -> str | None: """Return the provider of the first routing rule whose regex matches ``model``. O(number of rules). Only call this once exact lookups have missed. @@ -203,7 +203,7 @@ def match_routing_generalization(model: str) -> Optional[str]: return _registry.match_routing(model) -def match_capability_generalizations(model: str) -> Optional[dict]: +def match_capability_generalizations(model: str) -> dict | None: """Return the union of the ``model_info`` of every capability rule matching ``model``. Later rules override earlier ones on key conflicts. Returns ``None`` when no diff --git a/litellm/litellm_core_utils/fallback_utils.py b/litellm/litellm_core_utils/fallback_utils.py index 7aee69ef862..ff4a4c9c74c 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -1,11 +1,9 @@ -from litellm._uuid import uuid -from typing import Optional - import litellm from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.litellm_core_utils.core_helpers import ( - safe_deep_copy, filter_internal_params, + safe_deep_copy, ) from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, @@ -44,7 +42,7 @@ async def async_completion_with_fallbacks(**kwargs): litellm_logging_obj = base_kwargs.pop("litellm_logging_obj", None) # Try each fallback model - most_recent_exception_str: Optional[str] = None + most_recent_exception_str: str | None = None for attempted_fallbacks, fallback in enumerate(fallbacks): try: completion_kwargs = safe_deep_copy(base_kwargs) @@ -72,7 +70,7 @@ async def async_completion_with_fallbacks(**kwargs): ) except Exception as e: - verbose_logger.exception(f"Fallback attempt failed for model {model}: {str(e)}") + verbose_logger.exception(f"Fallback attempt failed for model {model}: {e!s}") most_recent_exception_str = str(e) continue diff --git a/litellm/litellm_core_utils/get_blog_posts.py b/litellm/litellm_core_utils/get_blog_posts.py index 6aea79cb4b3..60026e29c91 100644 --- a/litellm/litellm_core_utils/get_blog_posts.py +++ b/litellm/litellm_core_utils/get_blog_posts.py @@ -14,7 +14,6 @@ import time import xml.etree.ElementTree as ET from email.utils import parsedate_to_datetime from importlib.resources import files -from typing import Dict, List, Optional import httpx from pydantic import BaseModel @@ -32,7 +31,7 @@ class BlogPost(BaseModel): class BlogPostsResponse(BaseModel): - posts: List[BlogPost] + posts: list[BlogPost] class GetBlogPosts: @@ -45,11 +44,11 @@ class GetBlogPosts: - Falls back to the bundled local backup on any failure """ - _cached_posts: Optional[List[Dict[str, str]]] = None + _cached_posts: list[dict[str, str]] | None = None _last_fetch_time: float = 0.0 @staticmethod - def load_local_blog_posts() -> List[Dict[str, str]]: + def load_local_blog_posts() -> list[dict[str, str]]: """Load the bundled local backup blog posts.""" content = json.loads(files("litellm").joinpath("blog_posts.json").read_text(encoding="utf-8")) return content.get("posts", []) @@ -66,7 +65,7 @@ class GetBlogPosts: return response.text @staticmethod - def parse_rss_to_posts(xml_text: str, max_posts: int = 1) -> List[Dict[str, str]]: + def parse_rss_to_posts(xml_text: str, max_posts: int = 1) -> list[dict[str, str]]: """ Parse RSS XML and return a list of blog post dicts. @@ -77,7 +76,7 @@ class GetBlogPosts: if channel is None: raise ValueError("RSS feed missing element") - posts: List[Dict[str, str]] = [] + posts: list[dict[str, str]] = [] for item in channel.findall("item"): if len(posts) >= max_posts: break @@ -111,7 +110,7 @@ class GetBlogPosts: return posts @staticmethod - def validate_blog_posts(posts: List[Dict[str, str]]) -> bool: + def validate_blog_posts(posts: list[dict[str, str]]) -> bool: """Return True if posts is a non-empty list.""" if not isinstance(posts, list) or len(posts) == 0: verbose_logger.warning( @@ -121,7 +120,7 @@ class GetBlogPosts: return True @classmethod - def get_blog_posts(cls, url: str) -> List[Dict[str, str]]: + def get_blog_posts(cls, url: str) -> list[dict[str, str]]: """ Return the blog posts list. @@ -155,6 +154,6 @@ class GetBlogPosts: return posts -def get_blog_posts(url: str) -> List[Dict[str, str]]: +def get_blog_posts(url: str) -> list[dict[str, str]]: """Public entry point — returns the blog posts list.""" return GetBlogPosts.get_blog_posts(url=url) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index b8ef9d8cca7..1fab58b9380 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.openai.data_residency import infer_openai_data_residency AWS_CREDENTIAL_KWARGS_KEYS = frozenset( @@ -55,8 +53,8 @@ _OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS def _get_base_model_from_litellm_call_metadata( - metadata: Optional[dict], -) -> Optional[str]: + metadata: dict | None, +) -> str | None: if metadata is None: return None model_info = metadata.get("model_info") @@ -66,7 +64,7 @@ def _get_base_model_from_litellm_call_metadata( def get_litellm_params( - api_key: Optional[str] = None, + api_key: str | None = None, force_timeout=600, azure=False, logger_fn=None, @@ -74,12 +72,12 @@ def get_litellm_params( hugging_face=False, replicate=False, together_ai=False, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, litellm_call_id=None, model_alias_map=None, completion_call_id=None, - metadata: Optional[dict] = None, + metadata: dict | None = None, model_info=None, proxy_server_request=None, acompletion=None, @@ -96,23 +94,23 @@ def get_litellm_params( text_completion=None, azure_ad_token_provider=None, user_continue_message=None, - base_model: Optional[str] = None, - litellm_trace_id: Optional[str] = None, - litellm_session_id: Optional[str] = None, - hf_model_name: Optional[str] = None, - custom_prompt_dict: Optional[dict] = None, - litellm_metadata: Optional[dict] = None, - disable_add_transform_inline_image_block: Optional[bool] = None, - drop_params: Optional[bool] = None, - prompt_id: Optional[str] = None, - prompt_variables: Optional[dict] = None, - async_call: Optional[bool] = None, - ssl_verify: Optional[bool] = None, - merge_reasoning_content_in_choices: Optional[bool] = None, - use_litellm_proxy: Optional[bool] = None, - api_version: Optional[str] = None, - max_retries: Optional[int] = None, - litellm_request_debug: Optional[bool] = None, + base_model: str | None = None, + litellm_trace_id: str | None = None, + litellm_session_id: str | None = None, + hf_model_name: str | None = None, + custom_prompt_dict: dict | None = None, + litellm_metadata: dict | None = None, + disable_add_transform_inline_image_block: bool | None = None, + drop_params: bool | None = None, + prompt_id: str | None = None, + prompt_variables: dict | None = None, + async_call: bool | None = None, + ssl_verify: bool | None = None, + merge_reasoning_content_in_choices: bool | None = None, + use_litellm_proxy: bool | None = None, + api_version: str | None = None, + max_retries: int | None = None, + litellm_request_debug: bool | None = None, **kwargs, ) -> dict: # Derive litellm_session_id / litellm_trace_id from metadata when not provided (call chaining) @@ -122,7 +120,7 @@ def get_litellm_params( if litellm_trace_id is None: litellm_trace_id = _meta.get("trace_id") or _meta.get("session_id") - data_residency: Optional[str] = infer_openai_data_residency(custom_llm_provider, api_base) + data_residency: str | None = infer_openai_data_residency(custom_llm_provider, api_base) # Build base dict with explicit parameters (always included) litellm_params = { diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 487a7b7e25f..f869909e751 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -1,4 +1,4 @@ -from typing import Optional, Tuple, cast +from typing import cast from urllib.parse import urlparse import litellm @@ -72,8 +72,8 @@ def _is_azure_claude_model(model: str) -> bool: def handle_cohere_chat_model_custom_llm_provider( - model: str, custom_llm_provider: Optional[str] = None -) -> Tuple[str, Optional[str]]: + model: str, custom_llm_provider: str | None = None +) -> tuple[str, str | None]: """ if user sets model = "cohere/command-r" -> use custom_llm_provider = "cohere_chat" @@ -98,8 +98,8 @@ def handle_cohere_chat_model_custom_llm_provider( def handle_anthropic_text_model_custom_llm_provider( - model: str, custom_llm_provider: Optional[str] = None -) -> Tuple[str, Optional[str]]: + model: str, custom_llm_provider: str | None = None +) -> tuple[str, str | None]: """ if user sets model = "anthropic/claude-2" -> use custom_llm_provider = "anthropic_text" @@ -129,11 +129,11 @@ def handle_anthropic_text_model_custom_llm_provider( def get_llm_provider( model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, -) -> Tuple[str, str, Optional[str], Optional[str]]: + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, +) -> tuple[str, str, str | None, str | None]: """ Returns the provider for a given model name - e.g. 'azure/chatgpt-v-2' -> 'azure' @@ -149,7 +149,7 @@ def get_llm_provider( raise ValueError("model parameter is required but was None. Please provide a valid model name.") if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default( - litellm_params=cast(Optional[LiteLLM_Params], litellm_params) + litellm_params=cast(LiteLLM_Params | None, litellm_params) ): return litellm.LiteLLMProxyChatConfig.litellm_proxy_get_custom_llm_provider_info( model=model, api_base=api_base, api_key=api_key @@ -222,11 +222,9 @@ def get_llm_provider( custom_llm_provider = model.split("/", 1)[0] model = model.split("/", 1)[1] if api_base is not None and not isinstance(api_base, str): - raise Exception("api base needs to be a string. api_base={}".format(api_base)) + raise Exception(f"api base needs to be a string. api_base={api_base}") if dynamic_api_key is not None and not isinstance(dynamic_api_key, str): - raise Exception( - "dynamic_api_key needs to be a string. Got type={}".format(type(dynamic_api_key).__name__) - ) + raise Exception(f"dynamic_api_key needs to be a string. Got type={type(dynamic_api_key).__name__}") return model, custom_llm_provider, dynamic_api_key, api_base # check if api base is a known openai compatible endpoint if api_base: @@ -301,10 +299,12 @@ def get_llm_provider( elif endpoint == "api.moonshot.ai/v1": custom_llm_provider = "moonshot" dynamic_api_key = get_secret_str("MOONSHOT_API_KEY") - elif endpoint == "api.minimax.io/anthropic" or endpoint == "api.minimaxi.com/anthropic": - custom_llm_provider = "minimax" - dynamic_api_key = get_secret_str("MINIMAX_API_KEY") - elif endpoint == "api.minimax.io/v1" or endpoint == "api.minimaxi.com/v1": + elif ( + endpoint == "api.minimax.io/anthropic" + or endpoint == "api.minimaxi.com/anthropic" + or endpoint == "api.minimax.io/v1" + or endpoint == "api.minimaxi.com/v1" + ): custom_llm_provider = "minimax" dynamic_api_key = get_secret_str("MINIMAX_API_KEY") elif endpoint == "platform.publicai.co/v1": @@ -351,11 +351,9 @@ def get_llm_provider( dynamic_api_key = get_secret_str("META_API_KEY") if api_base is not None and not isinstance(api_base, str): - raise Exception("api base needs to be a string. api_base={}".format(api_base)) + raise Exception(f"api base needs to be a string. api_base={api_base}") if dynamic_api_key is not None and not isinstance(dynamic_api_key, str): - raise Exception( - "dynamic_api_key needs to be a string. dynamic_api_key={}".format(dynamic_api_key) - ) + raise Exception(f"dynamic_api_key needs to be a string. dynamic_api_key={dynamic_api_key}") return model, custom_llm_provider, dynamic_api_key, api_base # type: ignore # check if model in known model provider list -> for huggingface models, raise exception as they don't have a fixed provider (can be togetherai, anyscale, baseten, runpod, et.) @@ -495,17 +493,17 @@ def get_llm_provider( llm_provider="", ) if api_base is not None and not isinstance(api_base, str): - raise Exception("api base needs to be a string. api_base={}".format(api_base)) + raise Exception(f"api base needs to be a string. api_base={api_base}") if dynamic_api_key is not None and not isinstance(dynamic_api_key, str): - raise Exception("dynamic_api_key needs to be a string. dynamic_api_key={}".format(dynamic_api_key)) + raise Exception(f"dynamic_api_key needs to be a string. dynamic_api_key={dynamic_api_key}") return model, custom_llm_provider, dynamic_api_key, api_base except Exception as e: if isinstance(e, litellm.exceptions.BadRequestError): raise e else: - error_str = f"GetLLMProvider Exception - {str(e)}\n\noriginal model: {model}" + error_str = f"GetLLMProvider Exception - {e!s}\n\noriginal model: {model}" raise litellm.exceptions.BadRequestError( # type: ignore - message=f"GetLLMProvider Exception - {str(e)}\n\noriginal model: {model}", + message=f"GetLLMProvider Exception - {e!s}\n\noriginal model: {model}", model=model, response=None, llm_provider="", @@ -514,11 +512,11 @@ def get_llm_provider( def _get_openai_compatible_provider_info( model: str, - api_base: Optional[str], - api_key: Optional[str], - dynamic_api_key: Optional[str], - litellm_params: Optional[GenericLiteLLMParams] = None, -) -> Tuple[str, str, Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + dynamic_api_key: str | None, + litellm_params: GenericLiteLLMParams | None = None, +) -> tuple[str, str, str | None, str | None]: """ Returns: Tuple[str, str, Optional[str], Optional[str]]: @@ -848,9 +846,9 @@ def _get_openai_compatible_provider_info( dynamic_api_key = api_key or get_secret_str("MANUS_API_KEY") if api_base is not None and not isinstance(api_base, str): - raise Exception("api base needs to be a string. api_base={}".format(api_base)) + raise Exception(f"api base needs to be a string. api_base={api_base}") if dynamic_api_key is not None and not isinstance(dynamic_api_key, str): - raise Exception("dynamic_api_key needs to be a string. dynamic_api_key={}".format(dynamic_api_key)) + raise Exception(f"dynamic_api_key needs to be a string. dynamic_api_key={dynamic_api_key}") if dynamic_api_key is None and api_key is not None: dynamic_api_key = api_key return model, custom_llm_provider, dynamic_api_key, api_base diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index 4c0a01ad645..0addc7586fe 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -11,7 +11,6 @@ export LITELLM_LOCAL_MODEL_COST_MAP=True import json import os from importlib.resources import files -from typing import Dict, List, Optional import httpx @@ -166,9 +165,9 @@ class ModelCostMapSourceInfo: """Tracks the source of the currently loaded model cost map.""" source: str = "local" # "local" or "remote" - url: Optional[str] = None + url: str | None = None is_env_forced: bool = False - fallback_reason: Optional[str] = None + fallback_reason: str | None = None # Module-level singleton tracking the source of the current cost map @@ -204,11 +203,11 @@ def _expand_model_aliases(model_cost: dict) -> dict: If an alias collides with an existing canonical entry the alias is skipped and a warning is logged. """ - aliases_to_add: Dict[str, dict] = {} - keys_with_aliases: List[str] = [] + aliases_to_add: dict[str, dict] = {} + keys_with_aliases: list[str] = [] for model_name, model_info in model_cost.items(): - aliases: Optional[list] = model_info.get("aliases") + aliases: list | None = model_info.get("aliases") if aliases is None: continue keys_with_aliases.append(model_name) @@ -293,7 +292,7 @@ def get_model_cost_map(url: str) -> dict: str(e), ) _cost_map_source_info.source = "local" - _cost_map_source_info.fallback_reason = f"Remote fetch failed: {str(e)}" + _cost_map_source_info.fallback_reason = f"Remote fetch failed: {e!s}" return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map()) # Validate using cached count (cheap int comparison, no file I/O) diff --git a/litellm/litellm_core_utils/get_provider_specific_headers.py b/litellm/litellm_core_utils/get_provider_specific_headers.py index 69a7ec72073..f05c064e0f0 100644 --- a/litellm/litellm_core_utils/get_provider_specific_headers.py +++ b/litellm/litellm_core_utils/get_provider_specific_headers.py @@ -1,14 +1,12 @@ -from typing import Dict, Optional - from litellm.types.utils import ProviderSpecificHeader class ProviderSpecificHeaderUtils: @staticmethod def get_provider_specific_headers( - provider_specific_header: Optional[ProviderSpecificHeader], - custom_llm_provider: Optional[str], - ) -> Dict: + provider_specific_header: ProviderSpecificHeader | None, + custom_llm_provider: str | None, + ) -> dict: """ Get the provider specific headers for the given custom llm provider. diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 19149da0316..6600baf6441 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -1,4 +1,4 @@ -from typing import Literal, Optional +from typing import Literal import litellm from litellm.exceptions import BadRequestError @@ -7,10 +7,10 @@ from litellm.types.utils import LlmProviders, LlmProvidersSet def get_supported_openai_params( model: str, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, request_type: Literal["chat_completion", "embeddings", "transcription"] = "chat_completion", - base_model: Optional[str] = None, -) -> Optional[list]: + base_model: str | None = None, +) -> list | None: """ Returns the supported openai params for a given model + provider @@ -284,9 +284,7 @@ def get_supported_openai_params( ) if provider_config: return provider_config.get_supported_openai_params(model=model) - elif request_type == "embeddings": - return None - elif request_type == "transcription": + elif request_type == "embeddings" or request_type == "transcription": return None return None diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 031c221ad72..d95ea70fc1e 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -3,7 +3,7 @@ Helper functions for health check calls. """ from collections.abc import Callable -from typing import TYPE_CHECKING, Dict, Literal, Optional +from typing import TYPE_CHECKING, Literal from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS @@ -117,9 +117,9 @@ class HealthCheckHelpers: model: str, custom_llm_provider: str, model_params: dict, - prompt: Optional[str] = None, - input: Optional[list] = None, - ) -> Dict[ + prompt: str | None = None, + input: list | None = None, + ) -> dict[ Literal[ "chat", "completion", diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 4e1ca181ecc..11668acb21e 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -1,5 +1,5 @@ from collections.abc import Iterator -from typing import Any, Dict, Optional +from typing import Any from litellm.types.utils import StandardCallbackDynamicParams @@ -86,7 +86,7 @@ _request_blocked_callback_params = { def initialize_standard_callback_dynamic_params( - kwargs: Optional[Dict] = None, + kwargs: dict | None = None, ) -> StandardCallbackDynamicParams: """ Initialize the standard callback dynamic params from the kwargs diff --git a/litellm/litellm_core_utils/json_validation_rule.py b/litellm/litellm_core_utils/json_validation_rule.py index c73b62f8a21..31348889036 100644 --- a/litellm/litellm_core_utils/json_validation_rule.py +++ b/litellm/litellm_core_utils/json_validation_rule.py @@ -1,14 +1,14 @@ import json -from typing import Any, Dict, List, Union +from typing import Any from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH def normalize_json_schema_types( - schema: Union[Dict[str, Any], List[Any], Any], + schema: dict[str, Any] | list[Any] | Any, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH, -) -> Union[Dict[str, Any], List[Any], Any]: +) -> dict[str, Any] | list[Any] | Any: """ Normalize JSON schema types from uppercase to lowercase format. @@ -47,7 +47,7 @@ def normalize_json_schema_types( return [normalize_json_schema_types(item, depth + 1, max_depth) for item in schema] if isinstance(schema, dict): - normalized_schema: Dict[str, Any] = {} + normalized_schema: dict[str, Any] = {} for key, value in schema.items(): if key == "type" and isinstance(value, str) and value in type_mapping: @@ -72,7 +72,7 @@ def normalize_json_schema_types( return schema -def normalize_tool_schema(tool: Dict[str, Any]) -> Dict[str, Any]: +def normalize_tool_schema(tool: dict[str, Any]) -> dict[str, Any]: """ Normalize a tool's parameter schema to use standard JSON Schema lowercase types. diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index db47ab0b367..db10e18e324 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -16,12 +16,8 @@ from functools import lru_cache from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Type, Union, cast, ) @@ -199,11 +195,11 @@ try: from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger - EnterpriseStandardLoggingPayloadSetupVAR: Optional[Type[EnterpriseStandardLoggingPayloadSetup]] = ( + EnterpriseStandardLoggingPayloadSetupVAR: type[EnterpriseStandardLoggingPayloadSetup] | None = ( EnterpriseStandardLoggingPayloadSetup ) except Exception as e: - verbose_logger.debug(f"[Non-Blocking] Unable to import GenericAPILogger - LiteLLM Enterprise Feature - {str(e)}") + verbose_logger.debug(f"[Non-Blocking] Unable to import GenericAPILogger - LiteLLM Enterprise Feature - {e!s}") GenericAPILogger = CustomLogger # type: ignore ResendEmailLogger = CustomLogger # type: ignore SendGridEmailLogger = CustomLogger # type: ignore @@ -211,7 +207,7 @@ except Exception as e: PagerDutyAlerting = CustomLogger # type: ignore EnterpriseCallbackControls = None # type: ignore EnterpriseStandardLoggingPayloadSetupVAR = None -_in_memory_loggers: List[Any] = [] +_in_memory_loggers: list[Any] = [] _STANDARD_LOGGING_METADATA_KEYS: frozenset = frozenset(StandardLoggingMetadata.__annotations__.keys()) @@ -242,10 +238,10 @@ greenscaleLogger = None lunaryLogger = None supabaseClient = None deepevalLogger = None -callback_list: Optional[List[str]] = [] +callback_list: list[str] | None = [] user_logger_fn = None -additional_details: Optional[Dict[str, str]] = {} -local_cache: Optional[Dict[str, str]] = {} +additional_details: dict[str, str] | None = {} +local_cache: dict[str, str] | None = {} last_fetched_at = None last_fetched_at_keys = None @@ -255,15 +251,14 @@ class ServiceTraceIDCache: def __init__(self) -> None: self.cache = InMemoryCache() - def get_cache(self, litellm_call_id: str, service_name: str) -> Optional[str]: - key_name = "{}:{}".format(service_name, litellm_call_id) + def get_cache(self, litellm_call_id: str, service_name: str) -> str | None: + key_name = f"{service_name}:{litellm_call_id}" response = self.cache.get_cache(key=key_name) return response def set_cache(self, litellm_call_id: str, service_name: str, trace_id: str) -> None: - key_name = "{}:{}".format(service_name, litellm_call_id) + key_name = f"{service_name}:{litellm_call_id}" self.cache.set_cache(key=key_name, value=trace_id) - return None in_memory_trace_id_cache = ServiceTraceIDCache() @@ -313,17 +308,17 @@ class Logging(LiteLLMLoggingBaseClass): start_time, litellm_call_id: str, function_id: str, - litellm_trace_id: Optional[str] = None, - dynamic_input_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = None, - dynamic_success_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = None, - dynamic_async_success_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = None, - dynamic_failure_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = None, - dynamic_async_failure_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = None, - applied_guardrails: Optional[List[str]] = None, - kwargs: Optional[Dict] = None, + litellm_trace_id: str | None = None, + dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = None, + dynamic_success_callbacks: list[str | Callable | CustomLogger] | None = None, + dynamic_async_success_callbacks: list[str | Callable | CustomLogger] | None = None, + dynamic_failure_callbacks: list[str | Callable | CustomLogger] | None = None, + dynamic_async_failure_callbacks: list[str | Callable | CustomLogger] | None = None, + applied_guardrails: list[str] | None = None, + kwargs: dict | None = None, log_raw_request_response: bool = False, ): - _input: Optional[str] = messages # save original value of messages + _input: str | None = messages # save original value of messages if messages is not None: if isinstance(messages, str): messages = [ @@ -348,18 +343,18 @@ class Logging(LiteLLMLoggingBaseClass): self.litellm_call_id = litellm_call_id self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4()) self.function_id = function_id - self.streaming_chunks: List[Any] = [] # for generating complete stream response - self.sync_streaming_chunks: List[Any] = [] # for generating complete stream response + self.streaming_chunks: list[Any] = [] # for generating complete stream response + self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response self.log_raw_request_response = log_raw_request_response # Initialize dynamic callbacks - self.dynamic_input_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = dynamic_input_callbacks - self.dynamic_success_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = dynamic_success_callbacks - self.dynamic_async_success_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = ( + self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks + self.dynamic_success_callbacks: list[str | Callable | CustomLogger] | None = dynamic_success_callbacks + self.dynamic_async_success_callbacks: list[str | Callable | CustomLogger] | None = ( dynamic_async_success_callbacks ) - self.dynamic_failure_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = dynamic_failure_callbacks - self.dynamic_async_failure_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = ( + self.dynamic_failure_callbacks: list[str | Callable | CustomLogger] | None = dynamic_failure_callbacks + self.dynamic_async_failure_callbacks: list[str | Callable | CustomLogger] | None = ( dynamic_async_failure_callbacks ) @@ -375,8 +370,8 @@ class Logging(LiteLLMLoggingBaseClass): self.initialize_standard_built_in_tools_params(kwargs) ) ## TIME TO FIRST TOKEN LOGGING ## - self.completion_start_time: Optional[datetime.datetime] = None - self._llm_caching_handler: Optional[LLMCachingHandler] = None + self.completion_start_time: datetime.datetime | None = None + self._llm_caching_handler: LLMCachingHandler | None = None # INITIAL LITELLM_PARAMS litellm_params = {} @@ -387,15 +382,15 @@ class Logging(LiteLLMLoggingBaseClass): self.litellm_params = litellm_params # Initialize cost breakdown field - self.cost_breakdown: Optional[CostBreakdown] = None + self.cost_breakdown: CostBreakdown | None = None # Init Caching related details - self.caching_details: Optional[CachingDetails] = None + self.caching_details: CachingDetails | None = None # Passthrough endpoint guardrails config for field targeting - self.passthrough_guardrails_config: Optional[Dict[str, Any]] = None + self.passthrough_guardrails_config: dict[str, Any] | None = None - self.model_call_details: Dict[str, Any] = { + self.model_call_details: dict[str, Any] = { "litellm_trace_id": self.litellm_trace_id, "litellm_call_id": litellm_call_id, "input": _input, @@ -408,7 +403,7 @@ class Logging(LiteLLMLoggingBaseClass): # post_call guardrails have run; the @client decorator then stores the # enqueue closure here instead of firing it immediately. self._defer_async_logging: bool = False - self._enqueue_deferred_logging: Optional[Callable[[], None]] = None + self._enqueue_deferred_logging: Callable[[], None] | None = None def process_dynamic_callbacks(self): """ @@ -443,9 +438,9 @@ class Logging(LiteLLMLoggingBaseClass): def _process_dynamic_callback_list( self, - callback_list: Optional[List[Union[str, Callable, CustomLogger]]], + callback_list: list[str | Callable | CustomLogger] | None, dynamic_callbacks_type: Literal["input", "success", "failure", "async_success", "async_failure"], - ) -> Optional[List[Union[str, Callable, CustomLogger]]]: + ) -> list[str | Callable | CustomLogger] | None: """ Helper function to initialize CustomLogger compatible callbacks in self.dynamic_* callbacks @@ -457,12 +452,12 @@ class Logging(LiteLLMLoggingBaseClass): if callback_list is None: return None - processed_list: List[Union[str, Callable, CustomLogger]] = [] + processed_list: list[str | Callable | CustomLogger] = [] for callback in callback_list: if isinstance(callback, str) and callback in litellm._known_custom_logger_compatible_callbacks: # For callbacks that support team-scoped credentials (e.g. datadog), # pass only the relevant dynamic params as custom_logger_init_args. - _custom_logger_init_args: Optional[dict] = None + _custom_logger_init_args: dict | None = None if callback == "datadog": _custom_logger_init_args = { k: v for k, v in self.standard_callback_dynamic_params.items() if k.startswith("dd_") @@ -490,9 +485,7 @@ class Logging(LiteLLMLoggingBaseClass): processed_list.append(callback) return processed_list - def initialize_standard_callback_dynamic_params( - self, kwargs: Optional[Dict] = None - ) -> StandardCallbackDynamicParams: + def initialize_standard_callback_dynamic_params(self, kwargs: dict | None = None) -> StandardCallbackDynamicParams: """ Initialize the standard callback dynamic params from the kwargs @@ -501,7 +494,7 @@ class Logging(LiteLLMLoggingBaseClass): return _initialize_standard_callback_dynamic_params(kwargs) - def initialize_standard_built_in_tools_params(self, kwargs: Optional[Dict] = None) -> StandardBuiltInToolsParams: + def initialize_standard_built_in_tools_params(self, kwargs: dict | None = None) -> StandardBuiltInToolsParams: """ Initialize the standard built-in tools params from the kwargs @@ -512,7 +505,7 @@ class Logging(LiteLLMLoggingBaseClass): file_search=StandardBuiltInToolCostTracking._get_file_search_tool_call(kwargs or {}), ) - def get_router_model_id(self) -> Optional[str]: + def get_router_model_id(self) -> str | None: """Extract the router deployment model_id from litellm_params. Checks both litellm_metadata and metadata for model_info.id. @@ -531,10 +524,10 @@ class Logging(LiteLLMLoggingBaseClass): def update_environment_variables( self, - litellm_params: Dict, - optional_params: Dict, - model: Optional[str] = None, - user: Optional[str] = None, + litellm_params: dict, + optional_params: dict, + model: str | None = None, + user: str | None = None, **additional_params, ): self.optional_params = optional_params @@ -580,11 +573,11 @@ class Logging(LiteLLMLoggingBaseClass): def update_from_kwargs( self, - kwargs: Dict, - litellm_params: Optional[Dict] = None, - optional_params: Optional[Dict] = None, - model: Optional[str] = None, - user: Optional[str] = None, + kwargs: dict, + litellm_params: dict | None = None, + optional_params: dict | None = None, + model: str | None = None, + user: str | None = None, **additional_params, ): """ @@ -592,7 +585,7 @@ class Logging(LiteLLMLoggingBaseClass): automatically extracts metadata/litellm_metadata from kwargs, so callers don't need to manually plumb them into litellm_params. """ - base_litellm_params: Dict[str, Any] = {} + base_litellm_params: dict[str, Any] = {} if "metadata" in kwargs: base_litellm_params["metadata"] = kwargs["metadata"] @@ -622,7 +615,7 @@ class Logging(LiteLLMLoggingBaseClass): **additional_params, ) - def update_messages(self, messages: List[AllMessageValues]): + def update_messages(self, messages: list[AllMessageValues]): """ Update the logged value of the messages in the model_call_details @@ -633,9 +626,9 @@ class Logging(LiteLLMLoggingBaseClass): def should_run_prompt_management_hooks( self, - non_default_params: Dict, - prompt_id: Optional[str] = None, - tools: Optional[List[Dict]] = None, + non_default_params: dict, + prompt_id: str | None = None, + tools: list[dict] | None = None, ) -> bool: """ Return True if prompt management hooks should be run @@ -658,8 +651,8 @@ class Logging(LiteLLMLoggingBaseClass): def _should_run_prompt_management_hooks_without_prompt_id( self, - non_default_params: Dict, - tools: Optional[List[Dict]] = None, + non_default_params: dict, + tools: list[dict] | None = None, ) -> bool: """ Certain prompt management hooks don't need a `prompt_id` to be passed in, they are triggered by dynamic params @@ -683,15 +676,15 @@ class Logging(LiteLLMLoggingBaseClass): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], - non_default_params: Dict, - prompt_variables: Optional[dict], - prompt_id: Optional[str] = None, - prompt_spec: Optional[PromptSpec] = None, - prompt_management_logger: Optional[CustomLogger] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ) -> Tuple[str, List[AllMessageValues], dict]: + messages: list[AllMessageValues], + non_default_params: dict, + prompt_variables: dict | None, + prompt_id: str | None = None, + prompt_spec: PromptSpec | None = None, + prompt_management_logger: CustomLogger | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ) -> tuple[str, list[AllMessageValues], dict]: custom_logger = prompt_management_logger or self.get_custom_logger_for_prompt_management( model=model, non_default_params=non_default_params, @@ -722,16 +715,16 @@ class Logging(LiteLLMLoggingBaseClass): async def async_get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], - non_default_params: Dict, - prompt_variables: Optional[dict], - prompt_id: Optional[str] = None, - prompt_spec: Optional[PromptSpec] = None, - prompt_management_logger: Optional[CustomLogger] = None, - tools: Optional[List[Dict]] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ) -> Tuple[str, List[AllMessageValues], dict]: + messages: list[AllMessageValues], + non_default_params: dict, + prompt_variables: dict | None, + prompt_id: str | None = None, + prompt_spec: PromptSpec | None = None, + prompt_management_logger: CustomLogger | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ) -> tuple[str, list[AllMessageValues], dict]: custom_logger = prompt_management_logger or self.get_custom_logger_for_prompt_management( model=model, tools=tools, @@ -765,9 +758,9 @@ class Logging(LiteLLMLoggingBaseClass): def _auto_detect_prompt_management_logger( self, prompt_id: str, - prompt_spec: Optional[PromptSpec], + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, - ) -> Optional[CustomLogger]: + ) -> CustomLogger | None: """ Auto-detect which prompt management system owns the given prompt_id. @@ -803,12 +796,12 @@ class Logging(LiteLLMLoggingBaseClass): def get_custom_logger_for_prompt_management( self, model: str, - non_default_params: Dict, - tools: Optional[List[Dict]] = None, - prompt_id: Optional[str] = None, - prompt_spec: Optional[PromptSpec] = None, - dynamic_callback_params: Optional[StandardCallbackDynamicParams] = None, - ) -> Optional[CustomLogger]: + non_default_params: dict, + tools: list[dict] | None = None, + prompt_id: str | None = None, + prompt_spec: PromptSpec | None = None, + dynamic_callback_params: StandardCallbackDynamicParams | None = None, + ) -> CustomLogger | None: """ Get a custom logger for prompt management based on model name or available callbacks. @@ -879,7 +872,7 @@ class Logging(LiteLLMLoggingBaseClass): return None - def get_custom_logger_for_anthropic_cache_control_hook(self, non_default_params: Dict) -> Optional[CustomLogger]: + def get_custom_logger_for_anthropic_cache_control_hook(self, non_default_params: dict) -> CustomLogger | None: if non_default_params.get("cache_control_injection_points", None): custom_logger = _init_custom_logger_compatible_class( logging_integration="anthropic_cache_control_hook", @@ -889,14 +882,14 @@ class Logging(LiteLLMLoggingBaseClass): return custom_logger return None - def _get_raw_request_body(self, data: Optional[Union[dict, str]]) -> dict: + def _get_raw_request_body(self, data: dict | str | None) -> dict: if data is None: return {"error": "Received empty dictionary for raw request body"} if isinstance(data, str): try: return json.loads(data) except Exception: - return {"error": "Unable to parse raw request body. Got - {}".format(data)} + return {"error": f"Unable to parse raw request body. Got - {data}"} return data def _get_masked_api_base(self, api_base: str) -> str: @@ -974,8 +967,8 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict( error=str(e), ) - _metadata["raw_request"] = "Unable to Log \ - raw request: {}".format(str(e)) + _metadata["raw_request"] = f"Unable to Log \ + raw request: {e!s}" if getattr(self, "logger_fn", None) and callable(self.logger_fn): try: self.logger_fn( @@ -983,7 +976,7 @@ class Logging(LiteLLMLoggingBaseClass): ) # Expectation: any logger function passed in by the user should accept a dict object except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(str(e)) + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {e!s}" ) self.model_call_details["api_call_start_time"] = datetime.datetime.now() @@ -1043,16 +1036,14 @@ class Logging(LiteLLMLoggingBaseClass): callback_func=callback, ) except Exception as e: - verbose_logger.exception("litellm.Logging.pre_call(): Exception occured - {}".format(str(e))) + verbose_logger.exception(f"litellm.Logging.pre_call(): Exception occured - {e!s}") verbose_logger.debug( f"LiteLLM.Logging: is sentry capture exception initialized {capture_exception}" ) if capture_exception: # log this error to sentry for debugging capture_exception(e) except Exception as e: - verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(str(e)) - ) + verbose_logger.exception(f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {e!s}") verbose_logger.error(f"LiteLLM.Logging: is sentry capture exception initialized {capture_exception}") if capture_exception: # log this error to sentry for debugging capture_exception(e) @@ -1103,9 +1094,7 @@ class Logging(LiteLLMLoggingBaseClass): def _get_request_body(self, data: dict) -> str: return str(data) - def _get_request_curl_command( - self, api_base: str, headers: Optional[dict], additional_args: dict, data: dict - ) -> str: + def _get_request_curl_command(self, api_base: str, headers: dict | None, additional_args: dict, data: dict) -> str: masked_api_base = self._get_masked_api_base(api_base) if headers is None: headers = {} @@ -1170,7 +1159,7 @@ class Logging(LiteLLMLoggingBaseClass): ) # Expectation: any logger function passed in by the user should accept a dict object except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(str(e)) + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {e!s}" ) original_response = redact_message_input_output_from_logging( model_call_details=(self.model_call_details if hasattr(self, "model_call_details") else {}), @@ -1207,9 +1196,7 @@ class Logging(LiteLLMLoggingBaseClass): ) except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while post-call logging with integrations {}".format( - str(e) - ) + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while post-call logging with integrations {e!s}" ) verbose_logger.debug( f"LiteLLM.Logging: is sentry capture exception initialized {capture_exception}" @@ -1217,9 +1204,7 @@ class Logging(LiteLLMLoggingBaseClass): if capture_exception: # log this error to sentry for debugging capture_exception(e) except Exception as e: - verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(str(e)) - ) + verbose_logger.exception(f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {e!s}") async def async_post_mcp_tool_call_hook( self, @@ -1246,7 +1231,7 @@ class Logging(LiteLLMLoggingBaseClass): for callback in callbacks: try: if isinstance(callback, CustomLogger): - response: Optional[MCPPostCallResponseObject] = await callback.async_post_mcp_tool_call_hook( + response: MCPPostCallResponseObject | None = await callback.async_post_mcp_tool_call_hook( kwargs=kwargs, response_obj=post_mcp_tool_call_response_obj, start_time=start_time, @@ -1259,12 +1244,10 @@ class Logging(LiteLLMLoggingBaseClass): if response is not None: response_obj = self._parse_post_mcp_call_hook_response(response=response) except Exception as e: - verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(str(e)) - ) + verbose_logger.exception(f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {e!s}") return response_obj - def _parse_post_mcp_call_hook_response(self, response: Optional[MCPPostCallResponseObject]) -> Any: + def _parse_post_mcp_call_hook_response(self, response: MCPPostCallResponseObject | None) -> Any: """ Parse the response from the post_mcp_tool_call_hook @@ -1288,16 +1271,16 @@ class Logging(LiteLLMLoggingBaseClass): output_cost: float, total_cost: float, cost_for_built_in_tools_cost_usd_dollar: float, - additional_costs: Optional[dict] = None, - original_cost: Optional[float] = None, - discount_percent: Optional[float] = None, - discount_amount: Optional[float] = None, - margin_percent: Optional[float] = None, - margin_fixed_amount: Optional[float] = None, - margin_total_amount: Optional[float] = None, - cache_read_cost: Optional[float] = None, - cache_creation_cost: Optional[float] = None, - reasoning_cost: Optional[float] = None, + additional_costs: dict | None = None, + original_cost: float | None = None, + discount_percent: float | None = None, + discount_amount: float | None = None, + margin_percent: float | None = None, + margin_fixed_amount: float | None = None, + margin_total_amount: float | None = None, + cache_read_cost: float | None = None, + cache_creation_cost: float | None = None, + reasoning_cost: float | None = None, ) -> None: """ Helper method to store cost breakdown in the logging object. @@ -1371,10 +1354,10 @@ class Logging(LiteLLMLoggingBaseClass): dict, list, ], - cache_hit: Optional[bool] = None, - litellm_model_name: Optional[str] = None, - router_model_id: Optional[str] = None, - ) -> Optional[float]: + cache_hit: bool | None = None, + litellm_model_name: str | None = None, + router_model_id: str | None = None, + ) -> float | None: """ Calculate response cost using result + logging object variables. @@ -1473,7 +1456,7 @@ class Logging(LiteLLMLoggingBaseClass): return None - def _generate_content_result_as_model_response(self, result: object) -> Optional[ModelResponse]: + def _generate_content_result_as_model_response(self, result: object) -> ModelResponse | None: """ Native Google :generateContent bodies report token usage under ``usageMetadata``, which the cost calculator does not read, so a raw body @@ -1508,20 +1491,18 @@ class Logging(LiteLLMLoggingBaseClass): async def _response_cost_calculator_async( self, - result: Union[ - ModelResponse, - ModelResponseStream, - EmbeddingResponse, - ImageResponse, - TranscriptionResponse, - TextCompletionResponse, - HttpxBinaryResponseContent, - RerankResponse, - Batch, - FineTuningJob, - ], - cache_hit: Optional[bool] = None, - ) -> Optional[float]: + result: ModelResponse + | ModelResponseStream + | EmbeddingResponse + | ImageResponse + | TranscriptionResponse + | TextCompletionResponse + | HttpxBinaryResponseContent + | RerankResponse + | Batch + | FineTuningJob, + cache_hit: bool | None = None, + ) -> float | None: return self._response_cost_calculator(result=result, cache_hit=cache_hit) @staticmethod @@ -1839,7 +1820,7 @@ class Logging(LiteLLMLoggingBaseClass): start_time=None, end_time=None, cache_hit=None, - standard_logging_object: Optional[StandardLoggingPayload] = None, + standard_logging_object: StandardLoggingPayload | None = None, ): try: if start_time is None: @@ -1908,7 +1889,7 @@ class Logging(LiteLLMLoggingBaseClass): return start_time, end_time, result except Exception as e: - raise Exception(f"[Non-Blocking] LiteLLM.Success_Call Error: {str(e)}") + raise Exception(f"[Non-Blocking] LiteLLM.Success_Call Error: {e!s}") def _is_recognized_call_type_for_logging( self, @@ -1948,7 +1929,7 @@ class Logging(LiteLLMLoggingBaseClass): def _flush_passthrough_collected_chunks_helper( self, - raw_bytes: List[bytes], + raw_bytes: list[bytes], provider_config: "BasePassthroughConfig", ) -> Optional["CostResponseTypes"]: all_chunks = provider_config._convert_raw_bytes_to_str_lines(raw_bytes) @@ -1963,7 +1944,7 @@ class Logging(LiteLLMLoggingBaseClass): def flush_passthrough_collected_chunks( self, - raw_bytes: List[bytes], + raw_bytes: list[bytes], provider_config: "BasePassthroughConfig", ): """ @@ -1982,11 +1963,10 @@ class Logging(LiteLLMLoggingBaseClass): if complete_streaming_response is not None: self.success_handler(result=complete_streaming_response) - return async def async_flush_passthrough_collected_chunks( self, - raw_bytes: List[bytes], + raw_bytes: list[bytes], provider_config: "BasePassthroughConfig", ): complete_streaming_response = self._flush_passthrough_collected_chunks_helper( @@ -1996,7 +1976,6 @@ class Logging(LiteLLMLoggingBaseClass): if complete_streaming_response is not None: await self.async_success_handler(result=complete_streaming_response) - return def success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs): verbose_logger.debug(f"Logging Details LiteLLM-Success Call: Cache_hit={cache_hit}") @@ -2013,9 +1992,7 @@ class Logging(LiteLLMLoggingBaseClass): is_sync_request = self._is_sync_litellm_request(litellm_params) try: ## BUILD COMPLETE STREAMED RESPONSE - complete_streaming_response: Optional[ - Union[ModelResponse, TextCompletionResponse, ResponsesAPIResponse] - ] = None + complete_streaming_response: ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None = None if "complete_streaming_response" in self.model_call_details: return # break out of this. complete_streaming_response = self._get_assembled_streaming_response( @@ -2376,7 +2353,7 @@ class Logging(LiteLLMLoggingBaseClass): if ( callable(callback) is True and is_sync_request and customLogger is not None ): # custom logger functions - print_verbose("success callbacks: Running Custom Callback Function - {}".format(callback)) + print_verbose(f"success callbacks: Running Custom Callback Function - {callback}") customLogger.log_event( kwargs=self.model_call_details, @@ -2401,14 +2378,14 @@ class Logging(LiteLLMLoggingBaseClass): pass except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while success logging {}".format(str(e)), + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while success logging {e!s}", ) async def async_success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs): """ Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions. """ - print_verbose("Logging Details LiteLLM-Async Success Call, cache_hit={}".format(cache_hit)) + print_verbose(f"Logging Details LiteLLM-Async Success Call, cache_hit={cache_hit}") if not self._is_assembled_stream_success(result) and not self.should_run_logging( event_type="async_success" ): # prevent double logging (non-streaming) @@ -2469,7 +2446,7 @@ class Logging(LiteLLMLoggingBaseClass): ## BUILD COMPLETE STREAMED RESPONSE if "async_complete_streaming_response" in self.model_call_details: return # break out of this. - complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse, ResponsesAPIResponse]] = ( + complete_streaming_response: ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None = ( self._get_assembled_streaming_response( result=result, start_time=start_time, @@ -2612,7 +2589,7 @@ class Logging(LiteLLMLoggingBaseClass): ) if isinstance(callback, CustomLogger): # custom logger class - model_call_details: Dict = self.model_call_details + model_call_details: dict = self.model_call_details ################################## # call redaction hook for custom logger model_call_details = callback.redact_standard_logging_payload_from_model_call_details( @@ -2696,7 +2673,6 @@ class Logging(LiteLLMLoggingBaseClass): f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while success logging {traceback.format_exc()}" ) self._handle_callback_failure(callback=callback) - pass def _handle_callback_failure(self, callback: Any): """ @@ -2718,7 +2694,7 @@ class Logging(LiteLLMLoggingBaseClass): break # Only increment once except Exception as e: - verbose_logger.debug(f"Error in _handle_callback_failure: {str(e)}") + verbose_logger.debug(f"Error in _handle_callback_failure: {e!s}") def _failure_handler_helper_fn(self, exception, traceback_exception, start_time=None, end_time=None): if start_time is None: @@ -2955,14 +2931,14 @@ class Logging(LiteLLMLoggingBaseClass): except Exception as e: print_verbose( - f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging with integrations {str(e)}" + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging with integrations {e!s}" ) print_verbose(f"LiteLLM.Logging: is sentry capture exception initialized {capture_exception}") if capture_exception: # log this error to sentry for debugging capture_exception(e) except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging {}".format(str(e)) + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging {e!s}" ) async def async_failure_handler(self, exception, traceback_exception, start_time=None, end_time=None): @@ -3018,13 +2994,13 @@ class Logging(LiteLLMLoggingBaseClass): ) except Exception as e: verbose_logger.exception( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure \ - logging {}\nCallback={}".format(str(e), callback) + f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure \ + logging {e!s}\nCallback={callback}" ) # Track callback logging failures in Prometheus self._handle_callback_failure(callback=callback) - def _get_trace_id(self, service_name: Literal["langfuse"]) -> Optional[str]: + def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None: """ For the given service (e.g. langfuse), return the trace_id actually logged. @@ -3034,7 +3010,7 @@ class Logging(LiteLLMLoggingBaseClass): - str: The logged trace id - None: If trace id not yet emitted. """ - trace_id: Optional[str] = None + trace_id: str | None = None if service_name == "langfuse": trace_id = in_memory_trace_id_cache.get_cache( litellm_call_id=self.litellm_call_id, service_name=service_name @@ -3042,7 +3018,7 @@ class Logging(LiteLLMLoggingBaseClass): return trace_id - def _get_callback_object(self, service_name: Literal["langfuse"]) -> Optional[Any]: + def _get_callback_object(self, service_name: Literal["langfuse"]) -> Any | None: """ Return dynamic callback object. @@ -3081,7 +3057,7 @@ class Logging(LiteLLMLoggingBaseClass): result: Any, start_time: datetime.datetime, end_time: datetime.datetime, - cache_hit: Optional[Any] = None, + cache_hit: Any | None = None, ) -> None: """ Handles calling success callbacks for Async calls. @@ -3130,12 +3106,12 @@ class Logging(LiteLLMLoggingBaseClass): _filtered_failure_callbacks = self._remove_internal_litellm_callbacks(_filtered_failure_callbacks) return len(_filtered_failure_callbacks) > 0 - def get_combined_callback_list(self, dynamic_success_callbacks: Optional[List], global_callbacks: List) -> List: + def get_combined_callback_list(self, dynamic_success_callbacks: list | None, global_callbacks: list) -> list: if dynamic_success_callbacks is None: return list(global_callbacks) return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks)) - def _remove_internal_litellm_callbacks(self, callbacks: List) -> List: + def _remove_internal_litellm_callbacks(self, callbacks: list) -> list: """ Creates a filtered list of callbacks, excluding internal LiteLLM callbacks. @@ -3186,38 +3162,32 @@ class Logging(LiteLLMLoggingBaseClass): cb_name = self._get_callback_name(cb) return any(prefix in cb_name for prefix in INTERNAL_PREFIXES) - def _remove_internal_custom_logger_callbacks(self, callbacks: List) -> List: + def _remove_internal_custom_logger_callbacks(self, callbacks: list) -> list: """ Removes internal custom logger callbacks from the list. """ _new_callbacks = [] for _c in callbacks: - if isinstance(_c, CustomLogger): - continue - elif isinstance(_c, str) and _c in litellm._known_custom_logger_compatible_callbacks: + if ( + isinstance(_c, CustomLogger) + or isinstance(_c, str) + and _c in litellm._known_custom_logger_compatible_callbacks + ): continue _new_callbacks.append(_c) return _new_callbacks def _get_assembled_streaming_response( self, - result: Union[ - ModelResponse, - TextCompletionResponse, - ModelResponseStream, - ResponseCompletedEvent, - Any, - ], + result: ModelResponse | TextCompletionResponse | ModelResponseStream | ResponseCompletedEvent | Any, start_time: datetime.datetime, end_time: datetime.datetime, is_async: bool, - streaming_chunks: List[Any], - ) -> Optional[Union[ModelResponse, TextCompletionResponse, ResponsesAPIResponse]]: + streaming_chunks: list[Any], + ) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None: if self.stream is not True: return None - if isinstance(result, ModelResponse): - return result - elif isinstance(result, TextCompletionResponse): + if isinstance(result, ModelResponse) or isinstance(result, TextCompletionResponse): return result elif isinstance( result, @@ -3257,9 +3227,7 @@ class Logging(LiteLLMLoggingBaseClass): """ import httpx - if self.stream and isinstance(result, ModelResponse): - return result - elif isinstance(result, ModelResponse): + if self.stream and isinstance(result, ModelResponse) or isinstance(result, ModelResponse): return result if isinstance( @@ -3396,7 +3364,7 @@ def _get_masked_values( ignore_sensitive_values: bool = False, mask_all_values: bool = False, unmasked_length: int = 4, - number_of_asterisks: Optional[int] = 4, + number_of_asterisks: int | None = 4, _depth: int = 0, _max_depth: int = 20, ) -> dict: @@ -3558,15 +3526,14 @@ def set_callbacks(callback_list, function_id=None): customLogger = CustomLogger() except Exception as e: raise e - return None def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, - internal_usage_cache: Optional[DualCache], - llm_router: Optional[Any], # expect litellm.Router, but typing errors due to circular import - custom_logger_init_args: Optional[dict] = {}, -) -> Optional[CustomLogger]: + internal_usage_cache: DualCache | None, + llm_router: Any | None, # expect litellm.Router, but typing errors due to circular import + custom_logger_init_args: dict | None = {}, +) -> CustomLogger | None: """ Initialize a custom logger compatible class """ @@ -3948,9 +3915,7 @@ def _init_custom_logger_compatible_class( return callback # type: ignore if internal_usage_cache is None: - raise Exception( - "Internal Error: Cache cannot be empty - internal_usage_cache={}".format(internal_usage_cache) - ) + raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") dynamic_rate_limiter_obj = _PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache) @@ -3968,9 +3933,7 @@ def _init_custom_logger_compatible_class( return callback # type: ignore if internal_usage_cache is None: - raise Exception( - "Internal Error: Cache cannot be empty - internal_usage_cache={}".format(internal_usage_cache) - ) + raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") dynamic_rate_limiter_obj_v3 = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=internal_usage_cache) @@ -4180,7 +4143,7 @@ def _init_custom_logger_compatible_class( return None -def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list) -> Optional[Any]: +def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list) -> Any | None: """If ``LITELLM_OTEL_V2`` is on, build (or reuse) a single ``OpenTelemetryV2`` instance configured via the preset for ``callback_name``. @@ -4256,7 +4219,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list) -> None: def get_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, -) -> Optional[CustomLogger]: +) -> CustomLogger | None: try: if logging_integration == "lago": for callback in _in_memory_loggers: @@ -4463,7 +4426,7 @@ def get_custom_logger_compatible_class( return None -def _get_custom_logger_settings_from_proxy_server(callback_name: str) -> Dict: +def _get_custom_logger_settings_from_proxy_server(callback_name: str) -> dict: """ Get the settings for a custom logger from the proxy server config.yaml @@ -4478,7 +4441,7 @@ def _get_custom_logger_settings_from_proxy_server(callback_name: str) -> Dict: return {} -def use_custom_pricing_for_model(litellm_params: Optional[dict]) -> bool: +def use_custom_pricing_for_model(litellm_params: dict | None) -> bool: """ Check if the model uses custom pricing @@ -4516,10 +4479,10 @@ def is_valid_sha256_hash(value: str) -> bool: class StandardLoggingPayloadSetup: @staticmethod def cleanup_timestamps( - start_time: Union[dt_object, float], - end_time: Union[dt_object, float], - completion_start_time: Union[dt_object, float], - ) -> Tuple[float, float, float]: + start_time: dt_object | float, + end_time: dt_object | float, + completion_start_time: dt_object | float, + ) -> tuple[float, float, float]: """ Convert datetime objects to floats @@ -4556,7 +4519,7 @@ class StandardLoggingPayloadSetup: return start_time_float, end_time_float, completion_start_time_float @staticmethod - def append_system_prompt_messages(kwargs: Optional[Dict] = None, messages: Optional[Any] = None): + def append_system_prompt_messages(kwargs: dict | None = None, messages: Any | None = None): """ Append system prompt messages to the messages """ @@ -4614,16 +4577,16 @@ class StandardLoggingPayloadSetup: @staticmethod def get_standard_logging_metadata( - metadata: Optional[Dict[str, Any]], - litellm_params: Optional[dict] = None, - prompt_integration: Optional[str] = None, - applied_guardrails: Optional[List[str]] = None, - mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] = None, - vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]] = None, - usage_object: Optional[dict] = None, - proxy_server_request: Optional[dict] = None, - start_time: Optional[dt_object] = None, - response_id: Optional[str] = None, + metadata: dict[str, Any] | None, + litellm_params: dict | None = None, + prompt_integration: str | None = None, + applied_guardrails: list[str] | None = None, + mcp_tool_call_metadata: StandardLoggingMCPToolCall | None = None, + vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None = None, + usage_object: dict | None = None, + proxy_server_request: dict | None = None, + start_time: dt_object | None = None, + response_id: str | None = None, ) -> StandardLoggingMetadata: """ Clean and filter the metadata dictionary to include only the specified keys in StandardLoggingMetadata. @@ -4639,10 +4602,10 @@ class StandardLoggingPayloadSetup: - If 'user_api_key' is present in metadata and is a valid SHA256 hash, it's stored as 'user_api_key_hash'. """ - prompt_management_metadata: Optional[StandardLoggingPromptManagementMetadata] = None + prompt_management_metadata: StandardLoggingPromptManagementMetadata | None = None if litellm_params is not None: - prompt_id = cast(Optional[str], litellm_params.get("prompt_id", None)) - prompt_variables = cast(Optional[dict], litellm_params.get("prompt_variables", None)) + prompt_id = cast(str | None, litellm_params.get("prompt_id", None)) + prompt_variables = cast(dict | None, litellm_params.get("prompt_variables", None)) if prompt_id is not None and prompt_integration is not None: prompt_management_metadata = StandardLoggingPromptManagementMetadata( @@ -4724,9 +4687,7 @@ class StandardLoggingPayloadSetup: return clean_metadata @staticmethod - def get_usage_from_response_obj( - response_obj: Optional[dict], combined_usage_object: Optional[Usage] = None - ) -> Usage: + def get_usage_from_response_obj(response_obj: dict | None, combined_usage_object: Usage | None = None) -> Usage: ## BASE CASE ## if combined_usage_object is not None: return combined_usage_object @@ -4757,8 +4718,8 @@ class StandardLoggingPayloadSetup: @staticmethod def get_usage_as_dict( - response_obj: Optional[dict], - combined_usage_object: Optional[Usage] = None, + response_obj: dict | None, + combined_usage_object: Usage | None = None, ) -> dict: """ Like get_usage_from_response_obj but returns a plain dict, skipping @@ -4784,11 +4745,11 @@ class StandardLoggingPayloadSetup: @staticmethod def get_model_cost_information( - base_model: Optional[str], - custom_pricing: Optional[bool], - custom_llm_provider: Optional[str], - init_response_obj: Union[Any, BaseModel, dict], - api_base: Optional[str] = None, + base_model: str | None, + custom_pricing: bool | None, + custom_llm_provider: str | None, + init_response_obj: Any | BaseModel | dict, + api_base: str | None = None, ) -> StandardLoggingModelInformation: model_cost_name = _select_model_name_for_cost_calc( model=base_model if custom_pricing else None, @@ -4811,9 +4772,7 @@ class StandardLoggingPayloadSetup: ) except Exception: verbose_logger.debug( # keep in debug otherwise it will trigger on every call - "Model={} is not mapped in model cost map. Defaulting to None model_cost_information for standard_logging_payload".format( - model_cost_name - ) + f"Model={model_cost_name} is not mapped in model cost map. Defaulting to None model_cost_information for standard_logging_payload" ) model_cost_information = StandardLoggingModelInformation( model_map_key=model_cost_name, model_map_value=None @@ -4822,13 +4781,13 @@ class StandardLoggingPayloadSetup: @staticmethod def get_final_response_obj( - response_obj: dict, init_response_obj: Union[Any, BaseModel, dict], kwargs: dict - ) -> Optional[Union[dict, str, list]]: + response_obj: dict, init_response_obj: Any | BaseModel | dict, kwargs: dict + ) -> dict | str | list | None: """ Get final response object after redacting the message input/output from logging """ if response_obj: - final_response_obj: Optional[Union[dict, str, list]] = response_obj + final_response_obj: dict | str | list | None = response_obj elif isinstance(init_response_obj, list) or isinstance(init_response_obj, str): final_response_obj = init_response_obj else: @@ -4848,8 +4807,8 @@ class StandardLoggingPayloadSetup: @staticmethod def get_additional_headers( - additiona_headers: Optional[dict], - ) -> Optional[StandardLoggingAdditionalHeaders]: + additiona_headers: dict | None, + ) -> StandardLoggingAdditionalHeaders | None: if additiona_headers is None: return None @@ -4875,7 +4834,7 @@ class StandardLoggingPayloadSetup: @staticmethod def get_hidden_params( - hidden_params: Optional[dict], + hidden_params: dict | None, ) -> StandardLoggingHiddenParams: clean_hidden_params = StandardLoggingHiddenParams( model_id=None, @@ -4900,7 +4859,7 @@ class StandardLoggingPayloadSetup: return clean_hidden_params @staticmethod - def strip_trailing_slash(api_base: Optional[str]) -> Optional[str]: + def strip_trailing_slash(api_base: str | None) -> str | None: if api_base: if api_base.endswith("//"): return api_base.rstrip("/") @@ -4912,8 +4871,8 @@ class StandardLoggingPayloadSetup: def _generate_cold_storage_object_key( start_time: dt_object, response_id: str, - team_alias: Optional[str] = None, - ) -> Optional[str]: + team_alias: str | None = None, + ) -> str | None: """ Generate cold storage object key in the same format as S3Logger. @@ -4965,8 +4924,8 @@ class StandardLoggingPayloadSetup: @staticmethod def get_error_information( - original_exception: Optional[Exception], - traceback_str: Optional[str] = None, + original_exception: Exception | None, + traceback_str: str | None = None, ) -> StandardLoggingPayloadErrorInformation: from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG @@ -5110,13 +5069,13 @@ class StandardLoggingPayloadSetup: return logging_obj.litellm_trace_id @staticmethod - def _get_user_agent_tags(proxy_server_request: dict) -> Optional[List[str]]: + def _get_user_agent_tags(proxy_server_request: dict) -> list[str] | None: """ Return the user agent tags from the proxy server request for spend tracking """ if litellm.disable_add_user_agent_to_request_tags is True: return None - user_agent_tags: Optional[List[str]] = None + user_agent_tags: list[str] | None = None headers = proxy_server_request.get("headers", {}) if headers is not None and isinstance(headers, dict): if "user-agent" in headers: @@ -5124,7 +5083,7 @@ class StandardLoggingPayloadSetup: if user_agent is not None: if user_agent_tags is None: user_agent_tags = [] - user_agent_part: Optional[str] = None + user_agent_part: str | None = None if "/" in user_agent: user_agent_part = user_agent.split("/")[0] if user_agent_part is not None: @@ -5134,11 +5093,11 @@ class StandardLoggingPayloadSetup: return user_agent_tags @staticmethod - def _get_extra_header_tags(proxy_server_request: dict) -> Optional[List[str]]: + def _get_extra_header_tags(proxy_server_request: dict) -> list[str] | None: """ Extract additional header tags for spend tracking based on config. """ - extra_headers: List[str] = getattr(litellm, "extra_spend_tag_headers", None) or [] + extra_headers: list[str] = getattr(litellm, "extra_spend_tag_headers", None) or [] if not extra_headers: return None @@ -5155,7 +5114,7 @@ class StandardLoggingPayloadSetup: return header_tags if header_tags else None @staticmethod - def _get_request_tags(litellm_params: dict, proxy_server_request: dict) -> List[str]: + def _get_request_tags(litellm_params: dict, proxy_server_request: dict) -> list[str]: # check for 'tags' in both 'metadata' and 'litellm_metadata' metadata = litellm_params.get("metadata") or {} litellm_metadata = litellm_params.get("litellm_metadata") or {} @@ -5176,8 +5135,8 @@ class StandardLoggingPayloadSetup: def _get_status_fields( status: StandardLoggingPayloadStatus, - guardrail_information: Optional[List[dict]], - error_str: Optional[str], + guardrail_information: list[dict] | None, + error_str: str | None, ) -> "StandardLoggingPayloadStatusFields": """ Determine status fields based on request status and guardrail information. @@ -5191,7 +5150,7 @@ def _get_status_fields( StandardLoggingPayloadStatusFields with llm_api_status and guardrail_status """ # Mapping for legacy guardrail status values to new GuardrailStatus values - GUARDRAIL_STATUS_MAP: Dict[str, GuardrailStatus] = { + GUARDRAIL_STATUS_MAP: dict[str, GuardrailStatus] = { "success": "success", "blocked": "guardrail_intervened", # legacy "guardrail_intervened": "guardrail_intervened", # direct @@ -5219,11 +5178,11 @@ def _get_status_fields( def _extract_response_obj_and_hidden_params( - init_response_obj: Union[Any, BaseModel, dict], - original_exception: Optional[Exception], -) -> Tuple[dict, Optional[dict]]: + init_response_obj: Any | BaseModel | dict, + original_exception: Exception | None, +) -> tuple[dict, dict | None]: """Extract response_obj and hidden_params from init_response_obj.""" - hidden_params: Optional[dict] = None + hidden_params: dict | None = None if init_response_obj is None: response_obj = {} elif isinstance(init_response_obj, BaseModel): @@ -5255,16 +5214,16 @@ def _extract_response_obj_and_hidden_params( def get_standard_logging_object_payload( - kwargs: Optional[dict], - init_response_obj: Union[Any, BaseModel, dict], + kwargs: dict | None, + init_response_obj: Any | BaseModel | dict, start_time: dt_object, end_time: dt_object, logging_obj: Logging, status: StandardLoggingPayloadStatus, - error_str: Optional[str] = None, - original_exception: Optional[Exception] = None, - standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None, -) -> Optional[StandardLoggingPayload]: + error_str: str | None = None, + original_exception: Exception | None = None, + standard_built_in_tools_params: StandardBuiltInToolsParams | None = None, +) -> StandardLoggingPayload | None: try: kwargs = kwargs or {} @@ -5283,7 +5242,7 @@ def get_standard_logging_object_payload( # Extract usage as a plain dict, avoiding Pydantic round-trip raw_usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict( response_obj=response_obj, - combined_usage_object=cast(Optional[Usage], kwargs.get("combined_usage_object")), + combined_usage_object=cast(Usage | None, kwargs.get("combined_usage_object")), ) usage_dict = ( {**raw_usage_dict, "output_image_count": len(init_response_obj.data)} @@ -5382,7 +5341,7 @@ def get_standard_logging_object_payload( kwargs=kwargs, ) - stream: Optional[bool] = None + stream: bool | None = None if ( kwargs.get("complete_streaming_response") is not None or kwargs.get("async_complete_streaming_response") is not None @@ -5392,9 +5351,9 @@ def get_standard_logging_object_payload( # Reconstruct full model name with provider prefix for logging # This ensures Bedrock models like "us.anthropic.claude-3-5-sonnet-20240620-v1:0" # are logged as "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" - custom_llm_provider = cast(Optional[str], kwargs.get("custom_llm_provider")) + custom_llm_provider = cast(str | None, kwargs.get("custom_llm_provider")) model_name = reconstruct_model_name(kwargs.get("model", "") or "", custom_llm_provider, metadata) - response_model_name: Optional[str] = None + response_model_name: str | None = None if isinstance(final_response_obj, dict): response_model_name = final_response_obj.get("model") @@ -5467,7 +5426,7 @@ def get_standard_logging_object_payload( return payload except Exception as e: - verbose_logger.exception("Error creating standard logging object - {}".format(str(e))) + verbose_logger.exception(f"Error creating standard logging object - {e!s}") return None @@ -5477,7 +5436,7 @@ def emit_standard_logging_payload(payload: StandardLoggingPayload): def get_standard_logging_metadata( - metadata: Optional[Dict[str, Any]], + metadata: dict[str, Any] | None, ) -> StandardLoggingMetadata: """ Clean and filter the metadata dictionary to include only the specified keys in StandardLoggingMetadata. @@ -5541,7 +5500,7 @@ def get_standard_logging_metadata( return clean_metadata -def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]): +def scrub_sensitive_keys_in_metadata(litellm_params: dict | None): if litellm_params is None: litellm_params = {} @@ -5578,7 +5537,7 @@ def _get_traceback_str_for_error(error_str: str) -> str: from decimal import Decimal # used for unit testing -from typing import Any, Dict, List, Optional, Union +from typing import Any, Optional, Union def create_dummy_standard_logging_payload() -> StandardLoggingPayload: @@ -5586,20 +5545,20 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload: model_info = StandardLoggingModelInformation(model_map_key="gpt-3.5-turbo", model_map_value=None) metadata = StandardLoggingMetadata( # type: ignore - user_api_key_hash=str("test_hash"), - user_api_key_alias=str("test_alias"), - user_api_key_team_id=str("test_team"), - user_api_key_user_id=str("test_user"), - user_api_key_team_alias=str("test_team_alias"), + user_api_key_hash="test_hash", + user_api_key_alias="test_alias", + user_api_key_team_id="test_team", + user_api_key_user_id="test_user", + user_api_key_team_alias="test_team_alias", user_api_key_user_spend=None, user_api_key_user_max_budget=None, user_api_key_team_spend=None, user_api_key_team_max_budget=None, user_api_key_org_id=None, spend_logs_metadata=None, - requester_ip_address=str("127.0.0.1"), + requester_ip_address="127.0.0.1", requester_metadata=None, - user_api_key_end_user_id=str("test_end_user"), + user_api_key_end_user_id="test_end_user", ) hidden_params = StandardLoggingHiddenParams( @@ -5622,17 +5581,17 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload: saved_cache_cost = Decimal("0.0") # Create messages and response with proper typing - messages: List[Dict[str, str]] = [{"role": "user", "content": "Hello, world!"}] - response: Dict[str, List[Dict[str, Dict[str, str]]]] = {"choices": [{"message": {"content": "Hi there!"}}]} + messages: list[dict[str, str]] = [{"role": "user", "content": "Hello, world!"}] + response: dict[str, list[dict[str, dict[str, str]]]] = {"choices": [{"message": {"content": "Hi there!"}}]} # Main payload initialization return StandardLoggingPayload( # type: ignore - id=str("test_id"), - call_type=str("completion"), - stream=bool(False), + id="test_id", + call_type="completion", + stream=False, response_cost=response_cost, response_cost_failure_debug_info=None, - status=str("success"), + status="success", total_tokens=int(DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT), prompt_tokens=int(DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT), completion_tokens=int(DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT), @@ -5640,18 +5599,18 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload: endTime=end_time, completionStartTime=completion_start_time, model_map_information=model_info, - model=str("gpt-3.5-turbo"), - model_id=str("model-123"), - model_group=str("openai-gpt"), - custom_llm_provider=str("openai"), - api_base=str("https://api.openai.com"), + model="gpt-3.5-turbo", + model_id="model-123", + model_group="openai-gpt", + custom_llm_provider="openai", + api_base="https://api.openai.com", metadata=metadata, - cache_hit=bool(False), + cache_hit=False, cache_key=None, saved_cache_cost=saved_cache_cost, request_tags=[], end_user=None, - requester_ip_address=str("127.0.0.1"), + requester_ip_address="127.0.0.1", messages=messages, response=response, error_str=None, diff --git a/litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py b/litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py index 836b02f2049..7a98e0ec67a 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py +++ b/litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py @@ -5,10 +5,8 @@ Shared by provider cost calculators (e.g. Dashscope) and the proxy budget reservation logic so neither has to depend on the other. """ -from typing import List, Optional, Union - -def _coerce_cost_per_token(value: Union[float, int, str, None]) -> float: +def _coerce_cost_per_token(value: float | str | None) -> float: """ Coerce a per-token cost into a float. @@ -27,9 +25,9 @@ def _coerce_cost_per_token(value: Union[float, int, str, None]) -> float: def calculate_tiered_cost( tokens: int, - tiered_pricing: List[dict], + tiered_pricing: list[dict], cost_key: str, - fallback_cost_key: Optional[str] = None, + fallback_cost_key: str | None = None, ) -> float: """ Calculate cost for a given number of tokens based on a true tiered pricing structure. @@ -100,9 +98,9 @@ def calculate_tiered_cost( def select_tier_for_input( - tiered_pricing: List[dict], + tiered_pricing: list[dict], input_tokens: int, -) -> Optional[dict]: +) -> dict | None: """ Select the pricing tier for a request based on its total input token count. @@ -132,7 +130,7 @@ def select_tier_for_input( def tier_rate( tier: dict, cost_key: str, - fallback_cost_key: Optional[str] = None, + fallback_cost_key: str | None = None, ) -> float: """Read a per-token rate from a tier, coercing YAML string costs to float.""" raw = tier.get(cost_key) or tier.get(fallback_cost_key, 0) diff --git a/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py index 221b1ae6eab..2f2fbf2bb89 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py +++ b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py @@ -2,7 +2,7 @@ Helper utilities for tracking the cost of built-in tools. """ -from typing import Any, Dict, List, Literal, Optional, Tuple +from typing import Any, Literal import litellm from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS @@ -34,9 +34,9 @@ class StandardBuiltInToolCostTracking: def get_cost_for_built_in_tools( model: str, response_object: Any, - usage: Optional[Usage] = None, - custom_llm_provider: Optional[str] = None, - standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None, + usage: Usage | None = None, + custom_llm_provider: str | None = None, + standard_built_in_tools_params: StandardBuiltInToolsParams | None = None, ) -> float: """ Get the cost of using built-in tools. @@ -80,8 +80,8 @@ class StandardBuiltInToolCostTracking: @staticmethod def _handle_web_search_cost( model: str, - custom_llm_provider: Optional[str], - usage: Optional[Usage], + custom_llm_provider: str | None, + usage: Usage | None, standard_built_in_tools_params: StandardBuiltInToolsParams, response_object: object = None, ) -> float: @@ -125,7 +125,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def _handle_file_search_cost( model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, standard_built_in_tools_params: StandardBuiltInToolsParams, ) -> float: """Handle file search cost calculation.""" @@ -133,7 +133,7 @@ class StandardBuiltInToolCostTracking: model=model, custom_llm_provider=custom_llm_provider ) file_search_raw: Any = standard_built_in_tools_params.get("file_search", {}) - file_search_usage: Optional[FileSearchTool] = FileSearchTool(**file_search_raw) if file_search_raw else None + file_search_usage: FileSearchTool | None = FileSearchTool(**file_search_raw) if file_search_raw else None # Convert model_info to dict and extract usage parameters model_info_dict = dict(model_info) if model_info is not None else None @@ -150,7 +150,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def _handle_azure_assistant_costs( model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, standard_built_in_tools_params: StandardBuiltInToolsParams, ) -> float: """Handle Azure assistant features cost calculation.""" @@ -177,7 +177,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def _extract_file_search_params( file_search_usage: Any, - ) -> Tuple[Optional[float], Optional[float]]: + ) -> tuple[float | None, float | None]: """Extract and convert file search parameters safely.""" storage_gb = None days = None @@ -202,8 +202,8 @@ class StandardBuiltInToolCostTracking: @staticmethod def _get_vector_store_cost( - model_info: Optional[ModelInfo], - custom_llm_provider: Optional[str], + model_info: ModelInfo | None, + custom_llm_provider: str | None, standard_built_in_tools_params: StandardBuiltInToolsParams, ) -> float: """Calculate vector store cost.""" @@ -222,8 +222,8 @@ class StandardBuiltInToolCostTracking: @staticmethod def _get_computer_use_cost( - model_info: Optional[ModelInfo], - custom_llm_provider: Optional[str], + model_info: ModelInfo | None, + custom_llm_provider: str | None, standard_built_in_tools_params: StandardBuiltInToolsParams, ) -> float: """Calculate computer use cost.""" @@ -246,8 +246,8 @@ class StandardBuiltInToolCostTracking: @staticmethod def _get_code_interpreter_cost( - model_info: Optional[ModelInfo], - custom_llm_provider: Optional[str], + model_info: ModelInfo | None, + custom_llm_provider: str | None, standard_built_in_tools_params: StandardBuiltInToolsParams, ) -> float: """Calculate code interpreter cost.""" @@ -267,7 +267,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def _extract_token_counts( computer_use_usage: Any, - ) -> Tuple[Optional[int], Optional[int]]: + ) -> tuple[int | None, int | None]: """Extract and convert token counts safely.""" input_tokens = None output_tokens = None @@ -282,7 +282,7 @@ class StandardBuiltInToolCostTracking: return input_tokens, output_tokens @staticmethod - def _safe_convert_to_int(value: Any) -> Optional[int]: + def _safe_convert_to_int(value: Any) -> int | None: """Safely convert a value to int.""" if value is not None: try: @@ -312,7 +312,7 @@ class StandardBuiltInToolCostTracking: return usage.model_copy(update={"server_tool_use": server_tool_use}) @staticmethod - def response_object_includes_web_search_call(response_object: Any, usage: Optional[Usage] = None) -> bool: + def response_object_includes_web_search_call(response_object: Any, usage: Usage | None = None) -> bool: """ Check if the response object includes a web search call. @@ -358,14 +358,16 @@ class StandardBuiltInToolCostTracking: response_object=response_object, output_type="web_search_call" ) elif usage is not None: - if hasattr(usage, "server_tool_use") and _get_web_search_requests(usage.server_tool_use) is not None: - return True - elif ( - hasattr(usage, "prompt_tokens_details") - and usage.prompt_tokens_details is not None - and isinstance(usage.prompt_tokens_details, PromptTokensDetailsWrapper) - and hasattr(usage.prompt_tokens_details, "web_search_requests") - and usage.prompt_tokens_details.web_search_requests is not None + if ( + hasattr(usage, "server_tool_use") + and _get_web_search_requests(usage.server_tool_use) is not None + or ( + hasattr(usage, "prompt_tokens_details") + and usage.prompt_tokens_details is not None + and isinstance(usage.prompt_tokens_details, PromptTokensDetailsWrapper) + and hasattr(usage.prompt_tokens_details, "web_search_requests") + and usage.prompt_tokens_details.web_search_requests is not None + ) ): return True @@ -401,7 +403,7 @@ class StandardBuiltInToolCostTracking: ) -> bool: if isinstance(response_object, ModelResponse): for choice in response_object.choices: - message: Optional[Message] = getattr(choice, "message", None) + message: Message | None = getattr(choice, "message", None) if message is None: continue if annotations := getattr(message, "annotations", None): @@ -430,13 +432,13 @@ class StandardBuiltInToolCostTracking: """ output = response_object.output for output_item in output: - _output_type: Optional[str] = getattr(output_item, "type", None) + _output_type: str | None = getattr(output_item, "type", None) if _output_type == output_type: return True return False @staticmethod - def _safe_get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Optional[ModelInfo]: + def _safe_get_model_info(model: str, custom_llm_provider: str | None = None) -> ModelInfo | None: try: return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) except Exception: @@ -444,8 +446,8 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_cost_for_web_search( - web_search_options: Optional[WebSearchOptions] = None, - model_info: Optional[ModelInfo] = None, + web_search_options: WebSearchOptions | None = None, + model_info: ModelInfo | None = None, ) -> float: """ If request includes `web_search_options`, calculate the cost of the web search. @@ -468,7 +470,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_default_cost_for_web_search( - model_info: Optional[ModelInfo] = None, + model_info: ModelInfo | None = None, ) -> float: """ If no web search options are provided, use the `search_context_size_medium` pricing. @@ -485,11 +487,11 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_cost_for_file_search( - file_search: Optional[FileSearchTool] = None, - provider: Optional[str] = None, - model_info: Optional[dict] = None, - storage_gb: Optional[float] = None, - days: Optional[float] = None, + file_search: FileSearchTool | None = None, + provider: str | None = None, + model_info: dict | None = None, + storage_gb: float | None = None, + days: float | None = None, ) -> float: """ " OpenAI: $2.50/1k calls @@ -521,9 +523,9 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_cost_for_vector_store( - vector_store_usage: Optional[dict] = None, - provider: Optional[str] = None, - model_info: Optional[dict] = None, + vector_store_usage: dict | None = None, + provider: str | None = None, + model_info: dict | None = None, ) -> float: """ Calculate cost for vector store usage. @@ -551,10 +553,10 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_cost_for_computer_use( - input_tokens: Optional[int] = None, - output_tokens: Optional[int] = None, - provider: Optional[str] = None, - model_info: Optional[dict] = None, + input_tokens: int | None = None, + output_tokens: int | None = None, + provider: str | None = None, + model_info: dict | None = None, ) -> float: """ Calculate cost for computer use feature. @@ -593,7 +595,7 @@ class StandardBuiltInToolCostTracking: @staticmethod def _get_code_interpreter_cost_from_model_map( provider: str, - ) -> Optional[float]: + ) -> float | None: """ Get code interpreter cost per session from model cost map. """ @@ -614,9 +616,9 @@ class StandardBuiltInToolCostTracking: @staticmethod def get_cost_for_code_interpreter( - sessions: Optional[int] = None, - provider: Optional[str] = None, - model_info: Optional[dict] = None, + sessions: int | None = None, + provider: str | None = None, + model_info: dict | None = None, ) -> float: """ Calculate cost for code interpreter feature. @@ -657,7 +659,7 @@ class StandardBuiltInToolCostTracking: return False @staticmethod - def _get_web_search_options(kwargs: Dict) -> Optional[WebSearchOptions]: + def _get_web_search_options(kwargs: dict) -> WebSearchOptions | None: if "web_search_options" in kwargs: return WebSearchOptions(**kwargs.get("web_search_options", {})) @@ -673,13 +675,13 @@ class StandardBuiltInToolCostTracking: return None @staticmethod - def _get_tools_from_kwargs(kwargs: Dict, tool_type: str) -> Optional[List[Dict]]: + def _get_tools_from_kwargs(kwargs: dict, tool_type: str) -> list[dict] | None: if "tools" in kwargs: return kwargs.get("tools", []) return None @staticmethod - def _get_file_search_tool_call(kwargs: Dict) -> Optional[FileSearchTool]: + def _get_file_search_tool_call(kwargs: dict) -> FileSearchTool | None: tools = StandardBuiltInToolCostTracking._get_tools_from_kwargs(kwargs, "file_search") if tools: for tool in tools: @@ -689,7 +691,7 @@ class StandardBuiltInToolCostTracking: return None @staticmethod - def _is_web_search_tool_call(tool: Dict) -> bool: + def _is_web_search_tool_call(tool: dict) -> bool: if tool.get("type", None) == "web_search_preview": return True if tool.get("type", None) == "web_search": @@ -699,7 +701,7 @@ class StandardBuiltInToolCostTracking: return False @staticmethod - def _is_file_search_tool_call(tool: Dict) -> bool: + def _is_file_search_tool_call(tool: dict) -> bool: if tool.get("type", None) == "file_search": return True return False diff --git a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py index 1c6adbec174..210ac72cd8a 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py +++ b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Optional, Union +from typing import Any from litellm.types.utils import ( PromptTokensDetailsWrapper, @@ -19,8 +19,8 @@ class TranscriptionUsageObjectTransformation: @staticmethod def transform_transcription_usage_object( - usage_object: Union[TranscriptionUsageDurationObject, TranscriptionUsageTokensObject], - ) -> Optional[Usage]: + usage_object: TranscriptionUsageDurationObject | TranscriptionUsageTokensObject, + ) -> Usage | None: if isinstance(usage_object, TranscriptionUsageDurationObject): return None elif isinstance(usage_object, TranscriptionUsageTokensObject): diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 73427769b15..5bc6107dbec 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -4,7 +4,7 @@ from collections.abc import Mapping from dataclasses import dataclass from types import MappingProxyType -from typing import Any, Literal, Optional, Tuple, TypedDict, cast +from typing import Any, Literal, TypedDict, cast import litellm from litellm._logging import verbose_logger @@ -50,7 +50,7 @@ _SERVICE_TIER_TO_COST_KEY_SUFFIX: Mapping[str, str] = MappingProxyType( ) -def _get_token_detail_value(details: object, key: str) -> Optional[int]: +def _get_token_detail_value(details: object, key: str) -> int | None: if isinstance(details, dict): value = details.get(key) else: @@ -58,7 +58,7 @@ def _get_token_detail_value(details: object, key: str) -> Optional[int]: return value if isinstance(value, int) else None -def _get_web_search_requests(server_tool_use: Any) -> Optional[int]: +def _get_web_search_requests(server_tool_use: Any) -> int | None: """ Tolerantly read ``web_search_requests`` from a ``server_tool_use`` value that may be ``None``, a ``dict``, a ``ServerToolUse`` pydantic instance, @@ -115,9 +115,9 @@ def _generic_cost_per_character( custom_llm_provider: str, prompt_characters: float, completion_characters: float, - custom_prompt_cost: Optional[float], - custom_completion_cost: Optional[float], -) -> Tuple[Optional[float], Optional[float]]: + custom_prompt_cost: float | None, + custom_completion_cost: float | None, +) -> tuple[float | None, float | None]: """ Calculates cost per character for aspeech/speech calls. @@ -143,18 +143,14 @@ def _generic_cost_per_character( try: if custom_prompt_cost is None: assert "input_cost_per_character" in model_info and model_info["input_cost_per_character"] is not None, ( - "model info for model={} does not have 'input_cost_per_character'-pricing\nmodel_info={}".format( - model, model_info - ) + f"model info for model={model} does not have 'input_cost_per_character'-pricing\nmodel_info={model_info}" ) custom_prompt_cost = model_info["input_cost_per_character"] prompt_cost = prompt_characters * custom_prompt_cost except Exception as e: verbose_logger.exception( - "litellm.litellm_core_utils.llm_cost_calc.utils.py::cost_per_character(): Exception occured - {}\nDefaulting to None".format( - str(e) - ) + f"litellm.litellm_core_utils.llm_cost_calc.utils.py::cost_per_character(): Exception occured - {e!s}\nDefaulting to None" ) prompt_cost = None @@ -163,17 +159,13 @@ def _generic_cost_per_character( try: if custom_completion_cost is None: assert "output_cost_per_character" in model_info and model_info["output_cost_per_character"] is not None, ( - "model info for model={} does not have 'output_cost_per_character'-pricing\nmodel_info={}".format( - model, model_info - ) + f"model info for model={model} does not have 'output_cost_per_character'-pricing\nmodel_info={model_info}" ) custom_completion_cost = model_info["output_cost_per_character"] completion_cost = completion_characters * custom_completion_cost except Exception as e: verbose_logger.exception( - "litellm.litellm_core_utils.llm_cost_calc.utils.py::cost_per_character(): Exception occured - {}\nDefaulting to None".format( - str(e) - ) + f"litellm.litellm_core_utils.llm_cost_calc.utils.py::cost_per_character(): Exception occured - {e!s}\nDefaulting to None" ) completion_cost = None @@ -181,7 +173,7 @@ def _generic_cost_per_character( return prompt_cost, completion_cost -def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> str: +def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str: """ Get the appropriate cost key based on service tier. @@ -208,8 +200,8 @@ def _parse_above_token_threshold(key: str) -> float: def _get_token_base_cost( - model_info: ModelInfo, usage: Usage, service_tier: Optional[str] = None -) -> Tuple[float, float, float, float, float]: + model_info: ModelInfo, usage: Usage, service_tier: str | None = None +) -> tuple[float, float, float, float, float]: """ Return prompt cost, completion cost, and cache costs for a given model and usage. @@ -260,7 +252,7 @@ def _get_token_base_cost( ) # Only sort the threshold keys (typically 1-2 keys instead of 66+) - threshold: Optional[float] = None + threshold: float | None = None for key in sorted(threshold_keys, key=_parse_above_token_threshold, reverse=True): value = model_info.get(key) if value is not None: @@ -366,7 +358,7 @@ def _get_token_base_cost( ) -def calculate_cost_component(model_info: ModelInfo, cost_key: str, usage_value: Optional[float]) -> float: +def calculate_cost_component(model_info: ModelInfo, cost_key: str, usage_value: float | None) -> float: """ Generic cost calculator for any usage component @@ -384,7 +376,7 @@ def calculate_cost_component(model_info: ModelInfo, cost_key: str, usage_value: return 0.0 -def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: Optional[float] = 0.0) -> Optional[float]: +def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: float | None = 0.0) -> float | None: # Sometimes the cost per unit is a string (e.g.: If a value like "3e-7" was read from the config.yaml) cost_per_unit = model_info.get(cost_key) if isinstance(cost_per_unit, float): @@ -425,7 +417,7 @@ def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: Opti def calculate_cache_writing_cost( cache_creation_tokens: int, - cache_creation_token_details: Optional[CacheCreationTokenDetails], + cache_creation_token_details: CacheCreationTokenDetails | None, cache_creation_cost_above_1hr: float, cache_creation_cost: float, ) -> float: @@ -450,7 +442,7 @@ def calculate_cache_writing_cost( class PromptTokensDetailsResult(TypedDict): cache_hit_tokens: int cache_creation_tokens: int - cache_creation_token_details: Optional[CacheCreationTokenDetails] + cache_creation_token_details: CacheCreationTokenDetails | None text_tokens: int audio_tokens: int image_tokens: int @@ -462,10 +454,10 @@ class PromptTokensDetailsResult(TypedDict): def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: - cache_hit_tokens = cast(Optional[int], getattr(usage.prompt_tokens_details, "cached_tokens", 0)) or 0 + cache_hit_tokens = cast(int | None, getattr(usage.prompt_tokens_details, "cached_tokens", 0)) or 0 cache_creation_tokens = ( cast( - Optional[int], + int | None, getattr(usage.prompt_tokens_details, "cache_write_tokens", 0) or getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0), ) @@ -473,36 +465,36 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: ) cache_creation_token_details = ( cast( - Optional[CacheCreationTokenDetails], + CacheCreationTokenDetails | None, getattr(usage.prompt_tokens_details, "cache_creation_token_details", None), ) or None ) text_tokens = ( - cast(Optional[int], getattr(usage.prompt_tokens_details, "text_tokens", None)) + cast(int | None, getattr(usage.prompt_tokens_details, "text_tokens", None)) or 0 # default to prompt tokens, if this field is not set ) - audio_tokens = cast(Optional[int], getattr(usage.prompt_tokens_details, "audio_tokens", 0)) or 0 - image_tokens = cast(Optional[int], getattr(usage.prompt_tokens_details, "image_tokens", 0)) or 0 + audio_tokens = cast(int | None, getattr(usage.prompt_tokens_details, "audio_tokens", 0)) or 0 + image_tokens = cast(int | None, getattr(usage.prompt_tokens_details, "image_tokens", 0)) or 0 video_tokens = _coerce_token_count(getattr(usage.prompt_tokens_details, "video_tokens", 0)) character_count = ( cast( - Optional[int], + int | None, getattr(usage.prompt_tokens_details, "character_count", 0), ) or 0 ) - image_count = cast(Optional[int], getattr(usage.prompt_tokens_details, "image_count", 0)) or 0 + image_count = cast(int | None, getattr(usage.prompt_tokens_details, "image_count", 0)) or 0 video_length_seconds = ( cast( - Optional[float], + float | None, getattr(usage.prompt_tokens_details, "video_length_seconds", 0), ) or 0.0 ) audio_length_seconds = ( cast( - Optional[float], + float | None, getattr(usage.prompt_tokens_details, "audio_length_seconds", 0), ) or 0.0 @@ -534,28 +526,28 @@ class CompletionTokensDetailsResult(TypedDict): def _parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsResult: audio_tokens = ( cast( - Optional[int], + int | None, getattr(usage.completion_tokens_details, "audio_tokens", 0), ) or 0 ) text_tokens = ( cast( - Optional[int], + int | None, getattr(usage.completion_tokens_details, "text_tokens", None), ) or 0 # default to completion tokens, if this field is not set ) reasoning_tokens = ( cast( - Optional[int], + int | None, getattr(usage.completion_tokens_details, "reasoning_tokens", 0), ) or 0 ) image_tokens = ( cast( - Optional[int], + int | None, getattr(usage.completion_tokens_details, "image_tokens", 0), ) or 0 @@ -578,7 +570,7 @@ def _calculate_input_cost( cache_read_cost: float, cache_creation_cost: float, cache_creation_cost_above_1hr: float, - service_tier: Optional[str] = None, + service_tier: str | None = None, ) -> float: """ Calculates the input cost for a given model, prompt tokens, and completion tokens. @@ -654,7 +646,7 @@ def _calculate_input_cost( return prompt_cost -def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: Optional[str]) -> float: +def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float: """ Resolve the per-model regional-processing uplift multiplier for a given data-residency region. @@ -689,9 +681,9 @@ def generic_cost_per_token( model: str, usage: Usage, custom_llm_provider: str, - service_tier: Optional[str] = None, - data_residency: Optional[str] = None, -) -> Tuple[float, float]: + service_tier: str | None = None, + data_residency: str | None = None, +) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -749,8 +741,7 @@ def generic_cost_per_token( if (text_tokens == 0 and prompt_tokens_details["image_count"] == 0) or has_double_counting: text_tokens = usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens # Clamp to zero: inconsistent streaming usage - if text_tokens < 0: - text_tokens = 0 + text_tokens = max(text_tokens, 0) prompt_tokens_details["text_tokens"] = text_tokens ( @@ -862,10 +853,10 @@ class TokenTypeCostBreakdown: def get_token_type_cost_breakdown( model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, usage: Usage, - service_tier: Optional[str] = None, - data_residency: Optional[str] = None, + service_tier: str | None = None, + data_residency: str | None = None, ) -> TokenTypeCostBreakdown: """ Provider-agnostic cost of reasoning and cache tokens, derived from the usage @@ -910,7 +901,7 @@ def get_token_type_cost_breakdown( cache_read_tokens = 0 cache_creation_tokens = 0 - cache_creation_token_details: Optional[CacheCreationTokenDetails] = None + cache_creation_token_details: CacheCreationTokenDetails | None = None if usage.prompt_tokens_details is not None: prompt_tokens_details = _parse_prompt_tokens_details(usage) cache_read_tokens = prompt_tokens_details["cache_hit_tokens"] @@ -950,7 +941,7 @@ def calculate_image_response_cost_from_usage( model: str, image_response: ImageResponse, custom_llm_provider: str, -) -> Optional[float]: +) -> float | None: """ Calculate image generation cost from usage metadata when available. @@ -975,7 +966,7 @@ def calculate_image_response_cost_from_usage( return None input_tokens_details = getattr(usage, "input_tokens_details", None) - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if input_tokens_details is not None: prompt_tokens_details = PromptTokensDetailsWrapper( text_tokens=getattr(input_tokens_details, "text_tokens", None), @@ -1076,12 +1067,12 @@ class CostCalculatorUtils: def route_image_generation_cost_calculator( model: str, completion_response: ImageResponse, - custom_llm_provider: Optional[str] = None, - quality: Optional[str] = None, - n: Optional[int] = None, - size: Optional[str] = None, - optional_params: Optional[dict] = None, - call_type: Optional[str] = None, + custom_llm_provider: str | None = None, + quality: str | None = None, + n: int | None = None, + size: str | None = None, + optional_params: dict | None = None, + call_type: str | None = None, ) -> float: """ Route the image generation cost calculator based on the custom_llm_provider @@ -1196,29 +1187,10 @@ class CostCalculatorUtils: model=model, image_response=completion_response, ) - elif custom_llm_provider == litellm.LlmProviders.OPENAI.value: - # gpt-image models use token-based pricing. - model_lower = model.lower() - if "gpt-image" in model_lower: - from litellm.llms.openai.image_generation.cost_calculator import ( - cost_calculator as openai_gpt_image_cost_calculator, - ) - - return openai_gpt_image_cost_calculator( - model=model, - image_response=completion_response, - custom_llm_provider=custom_llm_provider, - ) - # Fall through to default for DALL-E models - return default_image_cost_calculator( - model=model, - quality=quality, - custom_llm_provider=custom_llm_provider, - n=n, - size=size, - optional_params=optional_params, - ) - elif custom_llm_provider == litellm.LlmProviders.AZURE.value: + elif ( + custom_llm_provider == litellm.LlmProviders.OPENAI.value + or custom_llm_provider == litellm.LlmProviders.AZURE.value + ): # gpt-image models use token-based pricing. model_lower = model.lower() if "gpt-image" in model_lower: diff --git a/litellm/litellm_core_utils/llm_request_utils.py b/litellm/litellm_core_utils/llm_request_utils.py index 7f76c7aca76..f86017255c0 100644 --- a/litellm/litellm_core_utils/llm_request_utils.py +++ b/litellm/litellm_core_utils/llm_request_utils.py @@ -1,9 +1,7 @@ -from typing import Dict, Optional - import litellm -def _ensure_extra_body_is_safe(extra_body: Optional[Dict]) -> Optional[Dict]: +def _ensure_extra_body_is_safe(extra_body: dict | None) -> dict | None: """ Ensure that the extra_body sent in the request is safe, otherwise users will see this error @@ -64,7 +62,7 @@ def pick_cheapest_chat_models_from_llm_provider(custom_llm_provider: str, n=1): return [model for model, _ in model_costs[:n]] -def get_proxy_server_request_headers(litellm_params: Optional[dict]) -> dict: +def get_proxy_server_request_headers(litellm_params: dict | None) -> dict: """ Get the `proxy_server_request` headers from the litellm_params.\ diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 861c5747963..8177391a74c 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -4,7 +4,7 @@ import re import time import traceback from collections.abc import Iterable -from typing import Dict, List, Literal, Optional, Tuple, Union, cast +from typing import Literal, cast import litellm from litellm._logging import verbose_logger @@ -52,15 +52,15 @@ _MODEL_RESPONSE_FIELDS: frozenset = frozenset(ModelResponse.model_fields.keys()) def _normalize_images_for_message( - images: Optional[List[dict]], -) -> Optional[List[ImageURLListItem]]: + images: list[dict] | None, +) -> list[ImageURLListItem] | None: """ Ensure each image has an 'index' field, as required by ImageURLListItem. Some providers (e.g. OpenRouter) return images without index. """ if not images: - return cast(Optional[List[ImageURLListItem]], images) - normalized: List[ImageURLListItem] = [] + return cast(list[ImageURLListItem] | None, images) + normalized: list[ImageURLListItem] = [] for i, img in enumerate(images): if isinstance(img, dict) and "index" not in img: normalized.append(cast(ImageURLListItem, {**img, "index": i})) @@ -98,15 +98,15 @@ def _safe_convert_created_field(created_value) -> int: def convert_tool_call_to_json_mode( - tool_calls: List[ChatCompletionMessageToolCall], + tool_calls: list[ChatCompletionMessageToolCall], convert_tool_call_to_json_mode: bool, -) -> Tuple[Optional[Message], Optional[str]]: +) -> tuple[Message | None, str | None]: if _should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=convert_tool_call_to_json_mode, ): # to support 'json_schema' logic on older models - json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments") + json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments") if json_mode_content_str is not None: message = litellm.Message(content=json_mode_content_str) finish_reason = "stop" @@ -121,7 +121,7 @@ def convert_tool_call_to_json_mode( _REPLAY_CONTENT_SLICE_RE = re.compile(r"\s*\S+\s*", re.UNICODE) -def _split_assembled_content_for_replay(content: Optional[str]) -> list[str]: +def _split_assembled_content_for_replay(content: str | None) -> list[str]: """ Slice an assembled cached completion's ``content`` into word-shaped pieces for cadence-preserving streaming replay. The split is lossless: @@ -150,7 +150,7 @@ def _clear_later_replay_slice_metadata(choice: StreamingChoices) -> None: async def convert_to_streaming_response_async( - response_object: Optional[dict] = None, + response_object: dict | None = None, ): """ Asynchronously converts a response object to a streaming response. @@ -175,7 +175,7 @@ async def convert_to_streaming_response_async( if model_response_object is None: raise Exception("Error in response creating model response object") - choice_list: List[StreamingChoices] = [] + choice_list: list[StreamingChoices] = [] if not response_object.get("choices"): from litellm.exceptions import APIError @@ -278,14 +278,14 @@ async def convert_to_streaming_response_async( def convert_to_streaming_response( - response_object: Optional[dict] = None, + response_object: dict | None = None, ): # used for yielding Cache hits when stream == True if response_object is None: raise Exception("Error in response object format") model_response_object = ModelResponseStream() - choice_list: List[StreamingChoices] = [] + choice_list: list[StreamingChoices] = [] if not response_object.get("choices"): from litellm.exceptions import APIError @@ -368,7 +368,7 @@ from collections import defaultdict def _handle_invalid_parallel_tool_calls( - tool_calls: List[ChatCompletionMessageToolCall], + tool_calls: list[ChatCompletionMessageToolCall], ): """ Handle hallucinated parallel tool call from openai - https://community.openai.com/t/model-tries-to-call-unknown-function-multi-tool-use-parallel/490653 @@ -379,7 +379,7 @@ def _handle_invalid_parallel_tool_calls( if tool_calls is None: return try: - replacements: Dict[int, List[ChatCompletionMessageToolCall]] = defaultdict(list) + replacements: dict[int, list[ChatCompletionMessageToolCall]] = defaultdict(list) for i, tool_call in enumerate(tool_calls): current_function = tool_call.function.name function_args = json.loads(tool_call.function.arguments) @@ -388,8 +388,7 @@ def _handle_invalid_parallel_tool_calls( for _fake_i, _fake_tool_use in enumerate(function_args["tool_uses"]): _function_args = _fake_tool_use["parameters"] _current_function = _fake_tool_use["recipient_name"] - if _current_function.startswith("functions."): - _current_function = _current_function[len("functions.") :] + _current_function = _current_function.removeprefix("functions.") fixed_tc = ChatCompletionMessageToolCall( id=f"{tool_call.id}_{_fake_i}", @@ -413,8 +412,8 @@ class LiteLLMResponseObjectHandler: @staticmethod def convert_to_image_response( response_object: dict, - model_response_object: Optional[ImageResponse] = None, - hidden_params: Optional[dict] = None, + model_response_object: ImageResponse | None = None, + hidden_params: dict | None = None, ) -> ImageResponse: response_object.update({"hidden_params": hidden_params}) @@ -466,7 +465,7 @@ class LiteLLMResponseObjectHandler: def convert_chat_to_text_completion( response: ModelResponse, text_completion_response: TextCompletionResponse, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> TextCompletionResponse: """ Converts a chat completion response to a text completion response format. @@ -494,7 +493,7 @@ class LiteLLMResponseObjectHandler: text_completion_response["object"] = "text_completion" text_completion_response["created"] = response.get("created", None) text_completion_response["model"] = response.get("model", None) - choices_list: List[TextChoices] = [] + choices_list: list[TextChoices] = [] # Convert each choice to TextChoices for choice in response["choices"]: @@ -513,21 +512,21 @@ class LiteLLMResponseObjectHandler: @staticmethod def _convert_provider_response_logprobs_to_text_completion_logprobs( response: ModelResponse, - custom_llm_provider: Optional[str] = None, - ) -> Optional[TextCompletionLogprobs]: + custom_llm_provider: str | None = None, + ) -> TextCompletionLogprobs | None: """ Convert logprobs from provider to OpenAI.Completion() format Only supported for HF TGI models """ - transformed_logprobs: Optional[TextCompletionLogprobs] = None + transformed_logprobs: TextCompletionLogprobs | None = None return transformed_logprobs def _should_convert_tool_call_to_json_mode( - tool_calls: Optional[Union[List[ChatCompletionMessageToolCall], List[DatabricksTool]]] = None, - convert_tool_call_to_json_mode: Optional[bool] = None, + tool_calls: list[ChatCompletionMessageToolCall] | list[DatabricksTool] | None = None, + convert_tool_call_to_json_mode: bool | None = None, ) -> bool: """ Determine if tool calls should be converted to JSON mode @@ -543,25 +542,22 @@ def _should_convert_tool_call_to_json_mode( def convert_to_model_response_object( - response_object: Optional[dict] = None, - model_response_object: Optional[ - Union[ - ModelResponse, - EmbeddingResponse, - ImageResponse, - TranscriptionResponse, - RerankResponse, - ] - ] = None, + response_object: dict | None = None, + model_response_object: ModelResponse + | EmbeddingResponse + | ImageResponse + | TranscriptionResponse + | RerankResponse + | None = None, response_type: Literal[ "completion", "embedding", "image_generation", "audio_transcription", "rerank" ] = "completion", stream=False, start_time=None, end_time=None, - hidden_params: Optional[dict] = None, - _response_headers: Optional[dict] = None, - convert_tool_call_to_json_mode: Optional[bool] = None, # used for supporting 'json_schema' on older models + hidden_params: dict | None = None, + _response_headers: dict | None = None, + convert_tool_call_to_json_mode: bool | None = None, # used for supporting 'json_schema' on older models ): additional_headers = get_response_headers(_response_headers) @@ -625,7 +621,7 @@ def convert_to_model_response_object( if stream is True: # for returning cached responses, we need to yield a generator return convert_to_streaming_response(response_object=response_object) - choice_list: List[Choices] = [] + choice_list: list[Choices] = [] if not response_object.get("choices") or not isinstance(response_object["choices"], Iterable): from litellm.exceptions import APIError @@ -653,14 +649,14 @@ def convert_to_model_response_object( if fixed_tool_calls is not None: tool_calls = fixed_tool_calls - message: Optional[Message] = None - finish_reason: Optional[str] = None + message: Message | None = None + finish_reason: str | None = None if tool_calls is not None and _should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=convert_tool_call_to_json_mode, ): # to support 'json_schema' logic on older models - json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments") + json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments") if json_mode_content_str is not None: message = litellm.Message(content=json_mode_content_str) finish_reason = "stop" @@ -675,14 +671,9 @@ def convert_to_model_response_object( reasoning_content, content = _extract_reasoning_content(choice["message"]) # Handle thinking models that display `thinking_blocks` within `content` - thinking_blocks: Optional[ - List[ - Union[ - ChatCompletionThinkingBlock, - ChatCompletionRedactedThinkingBlock, - ] - ] - ] = None + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = ( + None + ) if "thinking_blocks" in choice["message"]: thinking_blocks = choice["message"]["thinking_blocks"] provider_specific_fields["thinking_blocks"] = thinking_blocks @@ -826,9 +817,7 @@ def convert_to_model_response_object( setattr(model_response_object, key, response_object[key]) if "usage" in response_object and response_object["usage"] is not None: - tr_usage_object: Optional[Union[TranscriptionUsageDurationObject, TranscriptionUsageTokensObject]] = ( - None - ) + tr_usage_object: TranscriptionUsageDurationObject | TranscriptionUsageTokensObject | None = None if response_object["usage"].get("type", None) == "duration": tr_usage_object = TranscriptionUsageDurationObject(**response_object["usage"]) diff --git a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py index cc61ef0c899..5e332f4c8d6 100644 --- a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py +++ b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import litellm from litellm import verbose_logger @@ -7,7 +5,7 @@ from ...litellm_core_utils.get_llm_provider_logic import get_llm_provider from ...types.router import LiteLLM_Params -def get_api_base(model: str, optional_params: Union[dict, LiteLLM_Params]) -> Optional[str]: +def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | None: """ Returns the api base used for calling the model. @@ -55,7 +53,7 @@ def get_api_base(model: str, optional_params: Union[dict, LiteLLM_Params]) -> Op api_key=_optional_params.api_key, ) except Exception as e: - verbose_logger.debug("Error occurred in getting api base - {}".format(str(e))) + verbose_logger.debug(f"Error occurred in getting api base - {e!s}") custom_llm_provider = None dynamic_api_base = None @@ -78,19 +76,9 @@ def get_api_base(model: str, optional_params: Union[dict, LiteLLM_Params]) -> Op ) else: if stream: - _api_base = "{}-aiplatform.googleapis.com/v1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent".format( - _optional_params.vertex_location, - _optional_params.vertex_project, - _optional_params.vertex_location, - model, - ) + _api_base = f"{_optional_params.vertex_location}-aiplatform.googleapis.com/v1/projects/{_optional_params.vertex_project}/locations/{_optional_params.vertex_location}/publishers/google/models/{model}:streamGenerateContent" else: - _api_base = "{}-aiplatform.googleapis.com/v1/projects/{}/locations/{}/publishers/google/models/{}:generateContent".format( - _optional_params.vertex_location, - _optional_params.vertex_project, - _optional_params.vertex_location, - model, - ) + _api_base = f"{_optional_params.vertex_location}-aiplatform.googleapis.com/v1/projects/{_optional_params.vertex_project}/locations/{_optional_params.vertex_location}/publishers/google/models/{model}:generateContent" return _api_base if custom_llm_provider is None: @@ -98,9 +86,9 @@ def get_api_base(model: str, optional_params: Union[dict, LiteLLM_Params]) -> Op if custom_llm_provider == "gemini": if stream: - _api_base = "https://generativelanguage.googleapis.com/v1beta/models/{}:streamGenerateContent".format(model) + _api_base = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:streamGenerateContent" else: - _api_base = "https://generativelanguage.googleapis.com/v1beta/models/{}:generateContent".format(model) + _api_base = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent" return _api_base elif custom_llm_provider == "openai": _api_base = "https://api.openai.com" diff --git a/litellm/litellm_core_utils/llm_response_utils/get_formatted_prompt.py b/litellm/litellm_core_utils/llm_response_utils/get_formatted_prompt.py index f7406398a46..549a2d153a2 100644 --- a/litellm/litellm_core_utils/llm_response_utils/get_formatted_prompt.py +++ b/litellm/litellm_core_utils/llm_response_utils/get_formatted_prompt.py @@ -1,4 +1,4 @@ -from typing import List, Literal +from typing import Literal def get_formatted_prompt( @@ -25,7 +25,7 @@ def get_formatted_prompt( content = message.get("content") if isinstance(content, str): prompt += message["content"] - elif isinstance(content, List): + elif isinstance(content, list): for c in content: if c["type"] == "text": prompt += c["text"] diff --git a/litellm/litellm_core_utils/llm_response_utils/get_headers.py b/litellm/litellm_core_utils/llm_response_utils/get_headers.py index f4bbfae3039..f43da2ee401 100644 --- a/litellm/litellm_core_utils/llm_response_utils/get_headers.py +++ b/litellm/litellm_core_utils/llm_response_utils/get_headers.py @@ -1,7 +1,4 @@ -from typing import Optional - - -def get_response_headers(_response_headers: Optional[dict] = None) -> dict: +def get_response_headers(_response_headers: dict | None = None) -> dict: """ Sets the Appropriate OpenAI headers for the response and forward all headers as llm_provider-{header} diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index 5ac2dca9ccf..27f4b257808 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -1,5 +1,5 @@ import datetime -from typing import Any, Optional, Union +from typing import Any from litellm.constants import LITELLM_DETAILED_TIMING from litellm.litellm_core_utils.core_helpers import process_response_headers @@ -20,7 +20,7 @@ class ResponseMetadata: def __init__(self, result: Any): self.result = result - self._hidden_params: Union[HiddenParams, dict] = getattr(result, "_hidden_params", {}) or {} + self._hidden_params: HiddenParams | dict = getattr(result, "_hidden_params", {}) or {} @property def supports_response_time(self) -> bool: @@ -31,7 +31,7 @@ class ResponseMetadata: or isinstance(self.result, TranscriptionResponse) ) - def set_hidden_params(self, logging_obj: LiteLLMLoggingObject, model: Optional[str], kwargs: dict) -> None: + def set_hidden_params(self, logging_obj: LiteLLMLoggingObject, model: str | None, kwargs: dict) -> None: """Set hidden parameters on the response""" ## ADD OTHER HIDDEN PARAMS @@ -64,7 +64,7 @@ class ResponseMetadata: for key, value in new_params.items(): setattr(self._hidden_params, key, value) - def _get_value_from_hidden_params(self, key: str) -> Optional[Any]: + def _get_value_from_hidden_params(self, key: str) -> Any | None: """Get value from hidden params - handles when self._hidden_params is a dict or HiddenParams object""" if isinstance(self._hidden_params, dict): return self._hidden_params.get(key, None) @@ -166,7 +166,7 @@ class ResponseMetadata: def update_response_metadata( result: Any, logging_obj: LiteLLMLoggingObject, - model: Optional[str], + model: str | None, kwargs: dict, start_time: datetime.datetime, end_time: datetime.datetime, diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 71ad4c3b4e9..be732adfbe1 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import TYPE_CHECKING, Dict, List, Optional, Set, Type, Union +from typing import TYPE_CHECKING import litellm from litellm._logging import verbose_logger @@ -14,7 +14,7 @@ if TYPE_CHECKING: else: _custom_logger_compatible_callbacks_literal = str -_generic_api_logger_cache: Dict[str, GenericAPILogger] = {} +_generic_api_logger_cache: dict[str, GenericAPILogger] = {} class LoggingCallbackManager: @@ -38,7 +38,7 @@ class LoggingCallbackManager: except Exception: return False - def add_litellm_input_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_input_callback(self, callback: CustomLogger | str | Callable): """ Add a input callback to litellm.input_callback. Auto-routes async callbacks to litellm._async_input_callback. @@ -48,13 +48,13 @@ class LoggingCallbackManager: else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.input_callback) - def add_litellm_service_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_service_callback(self, callback: CustomLogger | str | Callable): """ Add a service callback to litellm.service_callback """ self._safe_add_callback_to_list(callback=callback, parent_list=litellm.service_callback) - def add_litellm_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_callback(self, callback: CustomLogger | str | Callable): """ Add a callback to litellm.callbacks @@ -65,20 +65,23 @@ class LoggingCallbackManager: parent_list=litellm.callbacks, # type: ignore ) - def add_litellm_success_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_success_callback(self, callback: CustomLogger | str | Callable): """ Add a success callback to `litellm.success_callback`. Auto-routes async callbacks to litellm._async_success_callback. Special-cases 'dynamodb' and 'openmeter' as async callbacks. """ - if isinstance(callback, str) and callback in ("dynamodb", "openmeter"): - self._safe_add_callback_to_list(callback=callback, parent_list=litellm._async_success_callback) - elif not isinstance(callback, str) and self._is_async_callable(callback): + if ( + isinstance(callback, str) + and callback in ("dynamodb", "openmeter") + or not isinstance(callback, str) + and self._is_async_callable(callback) + ): self._safe_add_callback_to_list(callback=callback, parent_list=litellm._async_success_callback) else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.success_callback) - def add_litellm_failure_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_failure_callback(self, callback: CustomLogger | str | Callable): """ Add a failure callback to `litellm.failure_callback`. Auto-routes async callbacks to litellm._async_failure_callback. @@ -88,13 +91,13 @@ class LoggingCallbackManager: else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.failure_callback) - def add_litellm_async_success_callback(self, callback: Union[CustomLogger, Callable, str]): + def add_litellm_async_success_callback(self, callback: CustomLogger | Callable | str): """ Add a success callback to litellm._async_success_callback """ self._safe_add_callback_to_list(callback=callback, parent_list=litellm._async_success_callback) - def add_litellm_async_failure_callback(self, callback: Union[CustomLogger, Callable, str]): + def add_litellm_async_failure_callback(self, callback: CustomLogger | Callable | str): """ Add a failure callback to litellm._async_failure_callback """ @@ -136,7 +139,7 @@ class LoggingCallbackManager: for c in remove_list: callback_list.remove(c) - def _add_string_callback_to_list(self, callback: str, parent_list: List[Union[CustomLogger, Callable, str]]): + def _add_string_callback_to_list(self, callback: str, parent_list: list[CustomLogger | Callable | str]): """ Add a string callback to a list, if the callback is already in the list, do not add it again. """ @@ -145,7 +148,7 @@ class LoggingCallbackManager: else: verbose_logger.debug(f"Callback {callback} already exists in {parent_list}, not adding again..") - def _check_callback_list_size(self, parent_list: List[Union[CustomLogger, Callable, str]]) -> bool: + def _check_callback_list_size(self, parent_list: list[CustomLogger | Callable | str]) -> bool: """ Check if adding another callback would exceed MAX_CALLBACKS Returns True if safe to add, False if would exceed limit @@ -160,7 +163,7 @@ class LoggingCallbackManager: @staticmethod def _add_custom_callback_generic_api_str( callback: str, - ) -> Union[GenericAPILogger, str]: + ) -> GenericAPILogger | str: """ litellm_settings: success_callback: ["custom_callback_name"] @@ -241,8 +244,8 @@ class LoggingCallbackManager: def _safe_add_callback_to_list( self, - callback: Union[CustomLogger, Callable, str], - parent_list: List[Union[CustomLogger, Callable, str]], + callback: CustomLogger | Callable | str, + parent_list: list[CustomLogger | Callable | str], ): """ Safe add a callback to a list, if the callback is already in the list, do not add it again. @@ -269,7 +272,7 @@ class LoggingCallbackManager: elif callable(callback): self._add_callback_function_to_list(callback=callback, parent_list=parent_list) - def _add_callback_function_to_list(self, callback: Callable, parent_list: List[Union[CustomLogger, Callable, str]]): + def _add_callback_function_to_list(self, callback: Callable, parent_list: list[CustomLogger | Callable | str]): """ Add a callback function to a list, if the callback is already in the list, do not add it again. """ @@ -284,7 +287,7 @@ class LoggingCallbackManager: def _add_custom_logger_to_list( self, custom_logger: CustomLogger, - parent_list: List[Union[CustomLogger, Callable, str]], + parent_list: list[CustomLogger | Callable | str], ): """ Add a custom logger to a list, if another instance of the same custom logger exists in the list, do not add it again. @@ -333,7 +336,7 @@ class LoggingCallbackManager: litellm._async_failure_callback = [] litellm.callbacks = [] - def _get_all_callbacks(self) -> List[Union[CustomLogger, Callable, str]]: + def _get_all_callbacks(self) -> list[CustomLogger | Callable | str]: """ Get all callbacks from litellm.callbacks, litellm.success_callback, litellm.failure_callback, litellm._async_success_callback, litellm._async_failure_callback """ @@ -361,7 +364,7 @@ class LoggingCallbackManager: def get_active_additional_logging_utils_from_custom_logger( self, - ) -> Set[AdditionalLoggingUtils]: + ) -> set[AdditionalLoggingUtils]: """ Get all custom loggers that are instances of the given class type @@ -372,13 +375,13 @@ class LoggingCallbackManager: Set[CustomLogger]: Set of custom loggers that are instances of the given class type """ all_callbacks = self._get_all_callbacks() - matched_callbacks: Set[AdditionalLoggingUtils] = set() + matched_callbacks: set[AdditionalLoggingUtils] = set() for callback in all_callbacks: if isinstance(callback, CustomLogger) and isinstance(callback, AdditionalLoggingUtils): matched_callbacks.add(callback) return matched_callbacks - def get_custom_loggers_for_type(self, callback_type: Type[CustomLogger]) -> List[CustomLogger]: + def get_custom_loggers_for_type(self, callback_type: type[CustomLogger]) -> list[CustomLogger]: """ Get all custom loggers that are instances of the given class type """ @@ -389,7 +392,7 @@ class LoggingCallbackManager: all_callbacks.append(callback) return all_callbacks - def callback_is_active(self, callback_type: Type[CustomLogger]) -> bool: + def callback_is_active(self, callback_type: type[CustomLogger]) -> bool: """ Returns True if any of the active callbacks are of the given type """ @@ -433,7 +436,7 @@ class LoggingCallbackManager: return result - def _get_callback_string(self, callback: Union[CustomLogger, Callable, str]) -> str: + def _get_callback_string(self, callback: CustomLogger | Callable | str) -> str: from litellm.litellm_core_utils.custom_logger_registry import ( CustomLoggerRegistry, ) @@ -452,7 +455,7 @@ class LoggingCallbackManager: def get_active_custom_logger_for_callback_name( self, callback_name: _custom_logger_compatible_callbacks_literal, - ) -> Optional[CustomLogger]: + ) -> CustomLogger | None: """ Get the active custom logger for a given callback name """ diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index 720a850b47f..32e2abc53b0 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -4,7 +4,7 @@ import inspect import re import time from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union from litellm._logging import verbose_logger from litellm.constants import MAX_BASE64_LENGTH_FOR_LOGGING @@ -101,7 +101,7 @@ def _truncate_base64_in_value(value: Any) -> Any: if isinstance(v, str): container[k] = _truncate_base64_in_string(v) elif isinstance(v, dict): - copy: Union[dict, list] = {ck: cv for ck, cv in v.items()} + copy: dict | list = {ck: cv for ck, cv in v.items()} container[k] = copy stack.append((copy, depth + 1)) elif isinstance(v, list): @@ -125,8 +125,8 @@ def _truncate_base64_in_value(value: Any) -> Any: def truncate_base64_in_messages( - messages: Optional[Union[str, list, dict]], -) -> Optional[Union[str, list, dict]]: + messages: str | list | dict | None, +) -> str | list | dict | None: """ Return a copy of *messages* with long base64 data-URI payloads replaced by human-readable size placeholders. @@ -155,8 +155,8 @@ def _get_service_logger(): def _get_parent_otel_span_from_logging_obj( - logging_obj: Optional[LiteLLMLoggingObject] = None, -) -> Optional[Span]: + logging_obj: LiteLLMLoggingObject | None = None, +) -> Span | None: """ Extract the parent OTEL span from the logging object using existing helper. @@ -178,13 +178,13 @@ def _get_parent_otel_span_from_logging_obj( return _get_parent_otel_span_from_kwargs(logging_obj.model_call_details) except Exception as e: - verbose_logger.exception(f"Error in _get_parent_otel_span_from_logging_obj: {str(e)}") + verbose_logger.exception(f"Error in _get_parent_otel_span_from_logging_obj: {e!s}") return None def convert_litellm_response_object_to_str( - response_obj: Union[Any, LiteLLMModelResponse], -) -> Optional[str]: + response_obj: Any | LiteLLMModelResponse, +) -> str | None: """ Get the string of the response object from LiteLLM @@ -201,11 +201,11 @@ def convert_litellm_response_object_to_str( def _assemble_complete_response_from_streaming_chunks( - result: Union[ModelResponse, TextCompletionResponse, ModelResponseStream], + result: ModelResponse | TextCompletionResponse | ModelResponseStream, start_time: datetime, end_time: datetime, request_kwargs: dict, - streaming_chunks: List[Any], + streaming_chunks: list[Any], is_async: bool, ): """ @@ -227,7 +227,7 @@ def _assemble_complete_response_from_streaming_chunks( Optional[Union[ModelResponse, TextCompletionResponse]]: Complete streaming response """ - complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse]] = None + complete_streaming_response: ModelResponse | TextCompletionResponse | None = None if isinstance(result, ModelResponse): return result @@ -265,7 +265,7 @@ def _set_duration_in_model_call_details( else: verbose_logger.debug("`logging_obj` not found - unable to track `llm_api_duration_ms") except Exception as e: - verbose_logger.warning(f"Error setting `llm_api_duration_ms`: {str(e)}") + verbose_logger.warning(f"Error setting `llm_api_duration_ms`: {e!s}") def track_llm_api_timing(): @@ -321,7 +321,7 @@ def track_llm_api_timing(): ) ) except Exception as e: - verbose_logger.debug(f"Error in service logging: {str(e)}") + verbose_logger.debug(f"Error in service logging: {e!s}") @functools.wraps(func) def sync_wrapper(*args, **kwargs): @@ -366,7 +366,7 @@ def track_llm_api_timing(): parent_otel_span=parent_otel_span, ) except Exception as e: - verbose_logger.debug(f"Error in service logging: {str(e)}") + verbose_logger.debug(f"Error in service logging: {e!s}") # Check if the function is async or sync if inspect.iscoroutinefunction(func): diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 3e1f1a9818e..b0d1de32c3c 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -6,7 +6,6 @@ import atexit import contextvars import logging from collections.abc import Coroutine -from typing import Optional from typing_extensions import TypedDict @@ -50,11 +49,11 @@ class LoggingWorker: self.timeout = timeout self.max_queue_size = max_queue_size self.concurrency = concurrency - self._queue: Optional[asyncio.Queue[LoggingTask]] = None - self._worker_task: Optional[asyncio.Task] = None + self._queue: asyncio.Queue[LoggingTask] | None = None + self._worker_task: asyncio.Task | None = None self._running_tasks: set[asyncio.Task] = set() - self._sem: Optional[asyncio.Semaphore] = None - self._bound_loop: Optional[asyncio.AbstractEventLoop] = None + self._sem: asyncio.Semaphore | None = None + self._bound_loop: asyncio.AbstractEventLoop | None = None self._last_aggressive_clear_time: float = 0.0 self._aggressive_clear_in_progress: bool = False @@ -278,7 +277,7 @@ class LoggingWorker: return extracted_tasks - async def _aggressively_clear_queue_async(self, new_task: Optional[LoggingTask] = None) -> None: + async def _aggressively_clear_queue_async(self, new_task: LoggingTask | None = None) -> None: """ Aggressively clear the queue by extracting and processing items. This is called when the queue is full to prevent dropping logs. diff --git a/litellm/litellm_core_utils/mock_functions.py b/litellm/litellm_core_utils/mock_functions.py index 0083a2b1454..ffbe5f72357 100644 --- a/litellm/litellm_core_utils/mock_functions.py +++ b/litellm/litellm_core_utils/mock_functions.py @@ -1,5 +1,3 @@ -from typing import List, Optional - from ..types.utils import ( Embedding, EmbeddingResponse, @@ -9,7 +7,7 @@ from ..types.utils import ( ) -def mock_embedding(model: str, mock_response: Optional[List[float]]): +def mock_embedding(model: str, mock_response: list[float] | None): if mock_response is None: mock_response = [0.0] * 1536 elif mock_response == "error": diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py index 39b3f0d5376..480d86a7e12 100644 --- a/litellm/litellm_core_utils/model_param_helper.py +++ b/litellm/litellm_core_utils/model_param_helper.py @@ -1,5 +1,4 @@ from functools import lru_cache -from typing import Set from openai.types.chat.completion_create_params import ( CompletionCreateParamsNonStreaming, @@ -39,11 +38,11 @@ class ModelParamHelper: return standard_logging_model_parameters @staticmethod - def get_exclude_params_for_model_parameters() -> Set[str]: + def get_exclude_params_for_model_parameters() -> set[str]: return set(["messages", "prompt", "input"]) @staticmethod - def _get_relevant_args_to_use_for_logging() -> Set[str]: + def _get_relevant_args_to_use_for_logging() -> set[str]: """ Gets all relevant llm api params besides the ones with prompt content """ @@ -56,7 +55,7 @@ class ModelParamHelper: @staticmethod @lru_cache(maxsize=1) - def _get_all_llm_api_params() -> Set[str]: + def _get_all_llm_api_params() -> set[str]: """ Gets the supported kwargs for each call type and combines them. @@ -86,28 +85,28 @@ class ModelParamHelper: return combined_kwargs @staticmethod - def get_litellm_provider_specific_params_for_chat_params() -> Set[str]: + def get_litellm_provider_specific_params_for_chat_params() -> set[str]: return set(["thinking"]) @staticmethod - def _get_litellm_supported_chat_completion_kwargs() -> Set[str]: + def _get_litellm_supported_chat_completion_kwargs() -> set[str]: """ Get the litellm supported chat completion kwargs This follows the OpenAI API Spec """ - non_streaming_params: Set[str] = set(getattr(CompletionCreateParamsNonStreaming, "__annotations__", {}).keys()) - streaming_params: Set[str] = set(getattr(CompletionCreateParamsStreaming, "__annotations__", {}).keys()) - litellm_provider_specific_params: Set[str] = ( + non_streaming_params: set[str] = set(getattr(CompletionCreateParamsNonStreaming, "__annotations__", {}).keys()) + streaming_params: set[str] = set(getattr(CompletionCreateParamsStreaming, "__annotations__", {}).keys()) + litellm_provider_specific_params: set[str] = ( ModelParamHelper.get_litellm_provider_specific_params_for_chat_params() ) - all_chat_completion_kwargs: Set[str] = non_streaming_params.union(streaming_params).union( + all_chat_completion_kwargs: set[str] = non_streaming_params.union(streaming_params).union( litellm_provider_specific_params ) return all_chat_completion_kwargs @staticmethod - def _get_litellm_supported_text_completion_kwargs() -> Set[str]: + def _get_litellm_supported_text_completion_kwargs() -> set[str]: """ Get the litellm supported text completion kwargs @@ -119,14 +118,14 @@ class ModelParamHelper: return all_text_completion_kwargs @staticmethod - def _get_litellm_supported_rerank_kwargs() -> Set[str]: + def _get_litellm_supported_rerank_kwargs() -> set[str]: """ Get the litellm supported rerank kwargs """ return set(RerankRequest.model_fields.keys()) @staticmethod - def _get_litellm_supported_embedding_kwargs() -> Set[str]: + def _get_litellm_supported_embedding_kwargs() -> set[str]: """ Get the litellm supported embedding kwargs @@ -135,7 +134,7 @@ class ModelParamHelper: return set(getattr(EmbeddingCreateParams, "__annotations__", {}).keys()) @staticmethod - def _get_litellm_supported_transcription_kwargs() -> Set[str]: + def _get_litellm_supported_transcription_kwargs() -> set[str]: """ Get the litellm supported transcription kwargs @@ -157,18 +156,18 @@ class ModelParamHelper: return set() @staticmethod - def _get_litellm_supported_responses_api_kwargs() -> Set[str]: + def _get_litellm_supported_responses_api_kwargs() -> set[str]: """ Get the litellm supported responses API kwargs This follows the OpenAI API Spec """ - non_streaming_params: Set[str] = set(getattr(ResponseCreateParamsNonStreaming, "__annotations__", {}).keys()) - streaming_params: Set[str] = set(getattr(ResponseCreateParamsStreaming, "__annotations__", {}).keys()) + non_streaming_params: set[str] = set(getattr(ResponseCreateParamsNonStreaming, "__annotations__", {}).keys()) + streaming_params: set[str] = set(getattr(ResponseCreateParamsStreaming, "__annotations__", {}).keys()) return non_streaming_params.union(streaming_params) @staticmethod - def _get_exclude_kwargs() -> Set[str]: + def _get_exclude_kwargs() -> set[str]: """ Get the kwargs to exclude from the cache key """ diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index e70520998ca..639c93dfb80 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -12,12 +12,7 @@ from pathlib import Path from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, - Union, cast, ) @@ -55,7 +50,7 @@ if TYPE_CHECKING: def handle_any_messages_to_chat_completion_str_messages_conversion( messages: Any, -) -> List[Dict[str, str]]: +) -> list[dict[str, str]]: """ Handles any messages to chat completion str messages conversion @@ -66,7 +61,7 @@ def handle_any_messages_to_chat_completion_str_messages_conversion( if isinstance(messages, list): try: return cast( - List[Dict[str, str]], + list[dict[str, str]], handle_messages_with_content_list_to_str_conversion(messages), ) except Exception: @@ -83,8 +78,8 @@ def handle_any_messages_to_chat_completion_str_messages_conversion( def handle_messages_with_content_list_to_str_conversion( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: """ Handles messages with content list conversion """ @@ -95,7 +90,7 @@ def handle_messages_with_content_list_to_str_conversion( return messages -def strip_name_from_message(message: AllMessageValues, allowed_name_roles: List[str] = ["user"]) -> AllMessageValues: +def strip_name_from_message(message: AllMessageValues, allowed_name_roles: list[str] = ["user"]) -> AllMessageValues: """ Removes 'name' from message """ @@ -106,8 +101,8 @@ def strip_name_from_message(message: AllMessageValues, allowed_name_roles: List[ def strip_name_from_messages( - messages: List[AllMessageValues], allowed_name_roles: List[str] = ["user"] -) -> List[AllMessageValues]: + messages: list[AllMessageValues], allowed_name_roles: list[str] = ["user"] +) -> list[AllMessageValues]: """ Removes 'name' from messages """ @@ -162,7 +157,7 @@ def extract_search_results_text(search_results: object) -> str: def convert_content_list_to_str( - message: Union[AllMessageValues, ChatCompletionResponseMessage], + message: AllMessageValues | ChatCompletionResponseMessage, ) -> str: """ - handles scenario where content is list and not string @@ -185,7 +180,7 @@ def convert_content_list_to_str( return texts -def get_str_from_messages(messages: List[AllMessageValues]) -> str: +def get_str_from_messages(messages: list[AllMessageValues]) -> str: """ Converts a list of messages to a string """ @@ -214,8 +209,8 @@ def _audio_or_image_in_message_content(message: AllMessageValues) -> bool: def convert_openai_message_to_only_content_messages( - messages: List[AllMessageValues], -) -> List[Dict[str, str]]: + messages: list[AllMessageValues], +) -> list[dict[str, str]]: """ Converts OpenAI messages to only content messages @@ -231,7 +226,7 @@ def convert_openai_message_to_only_content_messages( return converted_messages -def get_content_from_model_response(response: Union[ModelResponse, dict]) -> str: +def get_content_from_model_response(response: ModelResponse | dict) -> str: """ Gets content from model response """ @@ -256,8 +251,8 @@ def get_content_from_model_response(response: Union[ModelResponse, dict]) -> str def detect_first_expected_role( - messages: List[AllMessageValues], -) -> Optional[Literal["user", "assistant"]]: + messages: list[AllMessageValues], +) -> Literal["user", "assistant"] | None: """ Detect the first expected role based on the message sequence. @@ -292,10 +287,10 @@ def _counts_for_alternation(message: AllMessageValues) -> bool: def _insert_user_continue_message( - messages: List[AllMessageValues], - user_continue_message: Optional[ChatCompletionUserMessage], + messages: list[AllMessageValues], + user_continue_message: ChatCompletionUserMessage | None, ensure_alternating_roles: bool, -) -> List[AllMessageValues]: +) -> list[AllMessageValues]: """ Inserts a user continue message into the messages list. Handles three cases: @@ -352,10 +347,10 @@ def _insert_user_continue_message( def _insert_assistant_continue_message( - messages: List[AllMessageValues], - assistant_continue_message: Optional[ChatCompletionAssistantMessage] = None, + messages: list[AllMessageValues], + assistant_continue_message: ChatCompletionAssistantMessage | None = None, ensure_alternating_roles: bool = True, -) -> List[AllMessageValues]: +) -> list[AllMessageValues]: """ Add assistant continuation messages between consecutive user messages. @@ -383,7 +378,7 @@ def _insert_assistant_continue_message( j -= 1 # Build the result with assistant_continue inserted at the right positions - modified_messages: List[AllMessageValues] = [] + modified_messages: list[AllMessageValues] = [] for i, message in enumerate(messages): if i in insert_before_indexes: modified_messages.append(continue_message) @@ -393,11 +388,11 @@ def _insert_assistant_continue_message( def get_completion_messages( - messages: List[AllMessageValues], - assistant_continue_message: Optional[ChatCompletionAssistantMessage], - user_continue_message: Optional[ChatCompletionUserMessage], + messages: list[AllMessageValues], + assistant_continue_message: ChatCompletionAssistantMessage | None, + user_continue_message: ChatCompletionUserMessage | None, ensure_alternating_roles: bool, -) -> List[AllMessageValues]: +) -> list[AllMessageValues]: """ Ensures messages alternate between user and assistant roles by adding placeholders only when there are consecutive messages of the same role. @@ -416,7 +411,7 @@ def get_completion_messages( return messages -def get_format_from_file_id(file_id: Optional[str]) -> Optional[str]: +def get_format_from_file_id(file_id: str | None) -> str | None: """ Gets format from file id @@ -445,10 +440,10 @@ def get_format_from_file_id(file_id: Optional[str]) -> Optional[str]: def update_messages_with_model_file_ids( - messages: List[AllMessageValues], + messages: list[AllMessageValues], model_id: str | None, - model_file_id_mapping: Dict[str, Dict[str, str]], -) -> List[AllMessageValues]: + model_file_id_mapping: dict[str, dict[str, str]], +) -> list[AllMessageValues]: """ Updates messages with model file ids. @@ -512,9 +507,9 @@ def update_messages_with_model_file_ids( def update_responses_input_with_model_file_ids( input: Any, - model_id: Optional[str] = None, - model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, -) -> Union[str, List[Dict[str, Any]]]: + model_id: str | None = None, + model_file_id_mapping: dict[str, dict[str, str]] | None = None, +) -> str | list[dict[str, Any]]: """ Updates responses API input with provider-specific file IDs. File IDs are always inside the content array, not as direct input_file items. @@ -589,8 +584,8 @@ def update_responses_input_with_model_file_ids( def _decode_vector_store_ids_in_tools( - tools: Optional[List[Dict[str, Any]]], -) -> Optional[List[Dict[str, Any]]]: + tools: list[dict[str, Any]] | None, +) -> list[dict[str, Any]] | None: """ Decodes unified (LiteLLM-managed) vector_store_ids in file_search tools to provider-native IDs. Non-unified IDs are passed through unchanged. @@ -642,10 +637,10 @@ def _decode_vector_store_ids_in_tools( def update_responses_tools_with_model_file_ids( - tools: Optional[List[Dict[str, Any]]], - model_id: Optional[str] = None, - model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, -) -> Optional[List[Dict[str, Any]]]: + tools: list[dict[str, Any]] | None, + model_id: str | None = None, + model_file_id_mapping: dict[str, dict[str, str]] | None = None, +) -> list[dict[str, Any]] | None: """ Updates responses API tools with provider-specific file IDs. @@ -705,7 +700,7 @@ def update_responses_tools_with_model_file_ids( return updated_tools -def extract_file_metadata(file_data: FileTypes) -> Tuple[Optional[str], Optional[str]]: +def extract_file_metadata(file_data: FileTypes) -> tuple[str | None, str | None]: """ Resolve (filename, content_type) without reading the file body. @@ -713,8 +708,8 @@ def extract_file_metadata(file_data: FileTypes) -> Tuple[Optional[str], Optional it stays O(1) on large uploads. Use this when only metadata is needed (batch detection, GCS object naming) and the body must remain a streamable Path/handle. """ - filename: Optional[str] = None - content_type: Optional[str] = None + filename: str | None = None + content_type: str | None = None file_content: Any = None if isinstance(file_data, tuple): @@ -877,7 +872,7 @@ def _estimate_json_bytes(obj: Any) -> int: def unpack_defs( schema: dict, defs: dict, - max_inlined_bytes: Optional[int] = None, + max_inlined_bytes: int | None = None, ) -> None: """Expand *all* ``$ref`` entries pointing into ``$defs`` / ``definitions``. @@ -913,7 +908,7 @@ def unpack_defs( # Use iterative approach with queue to avoid recursion # Each item in queue is (node, parent_container, key/index, active_defs, ref_chain) - queue: deque[tuple[Any, Union[dict, list, None], Union[str, int, None], dict, set]] = deque( + queue: deque[tuple[Any, dict | list | None, str | int | None, dict, set]] = deque( [(schema, None, None, root_defs, set())] ) inlined_bytes = 0 @@ -957,9 +952,12 @@ def unpack_defs( # Replace the reference with resolved copy resolved = copy.deepcopy(target_schema) if parent is not None and key is not None: - if isinstance(parent, dict) and isinstance(key, str): - parent[key] = resolved - elif isinstance(parent, list) and isinstance(key, int): + if ( + isinstance(parent, dict) + and isinstance(key, str) + or isinstance(parent, list) + and isinstance(key, int) + ): parent[key] = resolved else: # This is the root schema itself @@ -1072,7 +1070,7 @@ def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSc return AnthropicInputSchema(**filtered) -def _get_image_mime_type_from_url(url: str) -> Optional[str]: +def _get_image_mime_type_from_url(url: str) -> str | None: """ Get mime type for common image URLs See gemini mime types: https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-understanding#image-requirements @@ -1140,7 +1138,7 @@ def _get_image_mime_type_from_url(url: str) -> Optional[str]: def infer_content_type_from_url_and_content( url: str, content: bytes, - current_content_type: Optional[str] = None, + current_content_type: str | None = None, ) -> str: """ Infer content type from URL extension and binary content when content-type header is missing or generic. @@ -1230,11 +1228,11 @@ def infer_content_type_from_url_and_content( raise ValueError(f"Unable to determine content type from URL: {url}. Response content-type: {current_content_type}") -def get_tool_call_names(tools: List[ChatCompletionToolParam]) -> List[str]: +def get_tool_call_names(tools: list[ChatCompletionToolParam]) -> list[str]: """ Get tool call names from tools """ - tool_call_names: List[str] = [] + tool_call_names: list[str] = [] for tool in tools: if tool.get("type") == "function": tool_call_name = tool.get("function", {}).get("name") @@ -1252,7 +1250,7 @@ def is_function_call(optional_params: dict) -> bool: return False -def get_file_ids_from_messages(messages: List[AllMessageValues]) -> List[str]: +def get_file_ids_from_messages(messages: list[AllMessageValues]) -> list[str]: """ Gets file ids from messages """ @@ -1350,7 +1348,7 @@ def migrate_file_to_image_url( return image_url_object -def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]: +def get_last_user_message(messages: list[AllMessageValues]) -> str | None: """ Get the last consecutive block of messages from the user. @@ -1389,7 +1387,7 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]: return result if result else None -def set_last_user_message(messages: List[AllMessageValues], content: str) -> List[AllMessageValues]: +def set_last_user_message(messages: list[AllMessageValues], content: str) -> list[AllMessageValues]: """ Set the last user message @@ -1411,10 +1409,10 @@ def set_last_user_message(messages: List[AllMessageValues], content: str) -> Lis def add_system_prompt_to_messages( - messages: List[AllMessageValues], + messages: list[AllMessageValues], system_prompt: str, merge_with_first_system: bool = False, -) -> List[AllMessageValues]: +) -> list[AllMessageValues]: """ Add a system prompt to the messages list. @@ -1434,7 +1432,7 @@ def add_system_prompt_to_messages( if merge_with_first_system and messages and messages[0].get("role") == "system": first = dict(messages[0]) existing_content = first.get("content", "") - merged_content: Union[str, List[Dict[str, str]]] + merged_content: str | list[dict[str, str]] if isinstance(existing_content, str): merged_content = f"{system_prompt.strip()}\n\n{existing_content}" elif isinstance(existing_content, list): @@ -1449,8 +1447,8 @@ def add_system_prompt_to_messages( def convert_prefix_message_to_non_prefix_messages( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: """ For models that don't support {prefix: true} in messages, we need to convert the prefix message to a non-prefix message. @@ -1469,7 +1467,7 @@ def convert_prefix_message_to_non_prefix_messages( do this in place """ - new_messages: List[AllMessageValues] = [] + new_messages: list[AllMessageValues] = [] for message in messages: if message.get("prefix"): new_messages.append( @@ -1486,7 +1484,7 @@ def convert_prefix_message_to_non_prefix_messages( return new_messages -def _extract_reasoning_content(message: dict) -> Tuple[Optional[str], Optional[str]]: +def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]: """ Extract reasoning content and main content from a message. @@ -1507,8 +1505,8 @@ def _extract_reasoning_content(message: dict) -> Tuple[Optional[str], Optional[s def _parse_content_for_reasoning( - message_text: Optional[str], -) -> Tuple[Optional[str], Optional[str]]: + message_text: str | None, +) -> tuple[str | None, str | None]: """ Parse the content for reasoning @@ -1553,7 +1551,7 @@ def _extract_base64_data(image_url: str) -> str: return image_url -def extract_images_from_message(message: AllMessageValues) -> List[str]: +def extract_images_from_message(message: AllMessageValues) -> list[str]: """ Extract images from a message. @@ -1574,7 +1572,7 @@ def extract_images_from_message(message: AllMessageValues) -> List[str]: return images -def _attempt_json_repair(s: str) -> Optional[Any]: +def _attempt_json_repair(s: str) -> Any | None: """ Attempt to repair truncated JSON produced by LLM tool calls. @@ -1633,9 +1631,9 @@ def _attempt_json_repair(s: str) -> Optional[Any]: def parse_tool_call_arguments( - arguments: Optional[str], - tool_name: Optional[str] = None, - context: Optional[str] = None, + arguments: str | None, + tool_name: str | None = None, + context: str | None = None, ) -> Any: """ Parse tool call arguments from a JSON string. @@ -1685,12 +1683,12 @@ def parse_tool_call_arguments( if context: error_parts.append(f"({context})") - error_message = " ".join(error_parts) + f". Error: {str(original_error)}. Arguments: {arguments}" + error_message = " ".join(error_parts) + f". Error: {original_error!s}. Arguments: {arguments}" raise ValueError(error_message) from original_error -def split_concatenated_json_objects(raw: str) -> List[Dict[str, Any]]: +def split_concatenated_json_objects(raw: str) -> list[dict[str, Any]]: """ Split a string that contains one or more concatenated JSON objects into a list of parsed dicts. @@ -1723,7 +1721,7 @@ def split_concatenated_json_objects(raw: str) -> List[Dict[str, Any]]: return [] decoder = json.JSONDecoder() - results: List[Dict[str, Any]] = [] + results: list[dict[str, Any]] = [] idx = 0 length = len(raw) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 0752bf2d771..147280af1b1 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -5,9 +5,9 @@ import json import mimetypes import re import xml.etree.ElementTree as ET -from enum import Enum from collections.abc import Iterator, Mapping, Sequence -from typing import Any, Dict, List, Optional, Set, Tuple, TypedDict, Union, cast, overload +from enum import Enum +from typing import Any, TypedDict, cast, overload from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -185,7 +185,7 @@ def convert_to_ollama_image(openai_image_url: str): ) -def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> Tuple[str, int]: +def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> tuple[str, int]: system_content_str = "" ## MERGE CONSECUTIVE SYSTEM CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "system": @@ -199,9 +199,9 @@ def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> Tu def ollama_pt( model: str, messages: list -) -> Union[ - str, OllamaVisionModelObject -]: # https://github.com/ollama/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template +) -> ( + str | OllamaVisionModelObject +): # https://github.com/ollama/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template user_message_types = {"user", "tool", "function"} msg_i = 0 images = [] @@ -439,13 +439,13 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st return rendered_text except Exception as e: raise Exception( - f"Error rendering template - {str(e)}" + f"Error rendering template - {e!s}" ) # don't use verbose_logger.exception, if exception is raised async def _afetch_and_extract_template( - model: str, chat_template: Optional[Any], get_config_fn, get_template_fn -) -> Tuple[str, str, str]: + model: str, chat_template: Any | None, get_config_fn, get_template_fn +) -> tuple[str, str, str]: """ Async version: Fetch template and tokens from HuggingFace. @@ -498,8 +498,8 @@ async def _afetch_and_extract_template( def _fetch_and_extract_template( - model: str, chat_template: Optional[Any], get_config_fn, get_template_fn -) -> Tuple[str, str, str]: + model: str, chat_template: Any | None, get_config_fn, get_template_fn +) -> tuple[str, str, str]: """ Sync version: Fetch template and tokens from HuggingFace. @@ -551,7 +551,7 @@ def _fetch_and_extract_template( return chat_template, bos_token, eos_token # type: ignore -async def ahf_chat_template(model: str, messages: list, chat_template: Optional[Any] = None): +async def ahf_chat_template(model: str, messages: list, chat_template: Any | None = None): """HuggingFace chat template (async version)""" from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( _aget_chat_template_file, @@ -578,7 +578,7 @@ async def ahf_chat_template(model: str, messages: list, chat_template: Optional[ ) -def hf_chat_template(model: str, messages: list, chat_template: Optional[Any] = None): +def hf_chat_template(model: str, messages: list, chat_template: Any | None = None): """HuggingFace chat template (sync version)""" from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( _get_chat_template_file, @@ -826,7 +826,7 @@ def convert_generic_image_chunk_to_openai_image_obj( return "data:{};{},{}".format(media_type, image_chunk["type"], image_chunk["data"]) -def convert_to_anthropic_image_obj(openai_image_url: str, format: Optional[str]) -> GenericImageParsingChunk: +def convert_to_anthropic_image_obj(openai_image_url: str, format: str | None) -> GenericImageParsingChunk: """ Input: "image_url": "data:image/jpeg;base64,{base64_image}", @@ -858,13 +858,13 @@ def convert_to_anthropic_image_obj(openai_image_url: str, format: Optional[str]) raise except Exception as e: raise Exception( - f"""Image url not in expected format. Example Expected input - "image_url": "data:image/jpeg;base64,{{base64_image}}". Supported formats - ['image/jpeg', 'image/png', 'image/gif', 'image/webp']. Error: {str(e)}""" + f"""Image url not in expected format. Example Expected input - "image_url": "data:image/jpeg;base64,{{base64_image}}". Supported formats - ['image/jpeg', 'image/png', 'image/gif', 'image/webp']. Error: {e!s}""" ) def create_anthropic_image_param( - image_url_input: Union[str, dict], - format: Optional[str] = None, + image_url_input: str | dict, + format: str | None = None, is_bedrock_invoke: bool = False, ) -> AnthropicMessagesImageParam: """ @@ -1100,7 +1100,7 @@ def anthropic_messages_pt_xml(messages: list): def _azure_tool_call_invoke_helper( function_call_params: ChatCompletionToolCallFunctionChunk, -) -> Optional[ChatCompletionToolCallFunctionChunk]: +) -> ChatCompletionToolCallFunctionChunk | None: """ Azure requires 'arguments' to be a string. """ @@ -1112,12 +1112,11 @@ def _azure_tool_call_invoke_helper( def _azure_image_url_helper(content: ChatCompletionImageObject): if isinstance(content["image_url"], str): content["image_url"] = {"url": content["image_url"]} - return def convert_to_azure_openai_messages( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: for m in messages: if m["role"] == "assistant": function_call = m.get("function_call", None) @@ -1163,8 +1162,8 @@ def infer_protocol_value( def _gemini_tool_call_invoke_helper( function_call_params: ChatCompletionToolCallFunctionChunk, - tool_call_id: Optional[str] = None, -) -> Optional[VertexFunctionCall]: + tool_call_id: str | None = None, +) -> VertexFunctionCall | None: name = function_call_params.get("name", "") or "" arguments = function_call_params.get("arguments", "") if ( @@ -1186,7 +1185,7 @@ def _gemini_tool_call_invoke_helper( return function_call -def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: Optional[str]) -> str: +def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: str | None) -> str: """ Embed thought signature into tool call ID for OpenAI client compatibility. @@ -1205,7 +1204,7 @@ def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: Op return tool_call_id -def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> Optional[str]: +def _get_thought_signature_from_tool(tool: dict, model: str | None = None) -> str | None: """Extract thought signature from tool call's provider_specific_fields. If not provided try to extract thought signature from tool call id @@ -1266,9 +1265,9 @@ def _get_dummy_thought_signature() -> str: def convert_to_gemini_tool_call_invoke( message: ChatCompletionAssistantMessage, - model: Optional[str] = None, + model: str | None = None, forward_function_call_id: bool = False, -) -> List[VertexPartType]: +) -> list[VertexPartType]: """ OpenAI tool invokes: { @@ -1309,7 +1308,7 @@ def convert_to_gemini_tool_call_invoke( - json.load the arguments """ try: - _parts_list: List[VertexPartType] = [] + _parts_list: list[VertexPartType] = [] tool_calls = message.get("tool_calls", None) function_call = message.get("function_call", None) @@ -1320,7 +1319,7 @@ def convert_to_gemini_tool_call_invoke( if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: - gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper( + gemini_function_call: VertexFunctionCall | None = _gemini_tool_call_invoke_helper( function_call_params=tool["function"], tool_call_id=(tool.get("id") if forward_function_call_id else None), ) @@ -1333,9 +1332,7 @@ def convert_to_gemini_tool_call_invoke( _parts_list.append(part_dict) else: # don't silently drop params. Make it clear to user what's happening. raise Exception( - "function_call missing. Received tool call with 'type': 'function'. No function call in argument - {}".format( - tool - ) + f"function_call missing. Received tool call with 'type': 'function'. No function call in argument - {tool}" ) elif function_call is not None: gemini_function_call = _gemini_tool_call_invoke_helper(function_call_params=function_call) @@ -1360,22 +1357,18 @@ def convert_to_gemini_tool_call_invoke( _parts_list.append(part_dict_function) else: # don't silently drop params. Make it clear to user what's happening. raise Exception( - "function_call missing. Received tool call with 'type': 'function'. No function call in argument - {}".format( - message - ) + f"function_call missing. Received tool call with 'type': 'function'. No function call in argument - {message}" ) return _parts_list except Exception as e: - raise Exception( - "Unable to convert openai tool calls={} to gemini tool calls. Received error={}".format(message, str(e)) - ) + raise Exception(f"Unable to convert openai tool calls={message} to gemini tool calls. Received error={e!s}") def convert_to_gemini_tool_call_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], - last_message_with_tool_calls: Optional[dict], + message: ChatCompletionToolMessage | ChatCompletionFunctionMessage, + last_message_with_tool_calls: dict | None, forward_function_call_id: bool = False, -) -> Union[VertexPartType, List[VertexPartType]]: +) -> VertexPartType | list[VertexPartType]: """ OpenAI message with a tool result looks like: { @@ -1405,7 +1398,7 @@ def convert_to_gemini_tool_call_result( from litellm.types.llms.vertex_ai import BlobType content_str: str = "" - inline_data_list: List[BlobType] = [] + inline_data_list: list[BlobType] = [] if "content" in message: if isinstance(message["content"], str): @@ -1423,7 +1416,7 @@ def convert_to_gemini_tool_call_result( content_str = "" except Exception as e: verbose_logger.warning(f"Failed to parse data URL in tool response: {e}") - elif isinstance(message["content"], List): + elif isinstance(message["content"], list): content_list = message["content"] for content in content_list: content_type = content.get("type", "") @@ -1484,7 +1477,7 @@ def convert_to_gemini_tool_call_result( ) except Exception as e: verbose_logger.warning(f"Failed to process file in tool response: {e}") - name: Optional[str] = message.get("name", "") # type: ignore + name: str | None = message.get("name", "") # type: ignore # Recover name from last message with tool calls if last_message_with_tool_calls: @@ -1496,7 +1489,7 @@ def convert_to_gemini_tool_call_result( name = tool.get("function", {}).get("name", "") # Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix). - gemini_call_id: Optional[str] = None + gemini_call_id: str | None = None if forward_function_call_id: raw_tool_call_id = message.get("tool_call_id") if raw_tool_call_id and isinstance(raw_tool_call_id, str): @@ -1506,9 +1499,7 @@ def convert_to_gemini_tool_call_result( if not name: raise Exception( - "Missing corresponding tool call for tool response message. Received - message={}, last_message_with_tool_calls={}".format( - message, last_message_with_tool_calls - ) + f"Missing corresponding tool call for tool response message. Received - message={message}, last_message_with_tool_calls={last_message_with_tool_calls}" ) # Parse response data - support both JSON string and plain string @@ -1578,7 +1569,7 @@ def _is_anthropic_document_data_uri(url: str) -> bool: def convert_to_anthropic_tool_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], + message: ChatCompletionToolMessage | ChatCompletionFunctionMessage, force_base64: bool = False, ) -> AnthropicMessagesToolResultParam: """ @@ -1612,26 +1603,15 @@ def convert_to_anthropic_tool_result( ] } """ - anthropic_content: Union[ - str, - List[ - Union[ - AnthropicMessagesToolResultContent, - AnthropicMessagesImageParam, - AnthropicMessagesDocumentParam, - ] - ], - ] = "" + anthropic_content: ( + str | list[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam] + ) = "" if isinstance(message["content"], str): anthropic_content = message["content"] - elif isinstance(message["content"], List): + elif isinstance(message["content"], list): content_list = message["content"] - anthropic_content_list: List[ - Union[ - AnthropicMessagesToolResultContent, - AnthropicMessagesImageParam, - AnthropicMessagesDocumentParam, - ] + anthropic_content_list: list[ + AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam ] = [] for content in content_list: if content["type"] == "text": @@ -1684,7 +1664,7 @@ def convert_to_anthropic_tool_result( anthropic_content_list.append(_file_block) anthropic_content = anthropic_content_list - anthropic_tool_result: Optional[AnthropicMessagesToolResultParam] = None + anthropic_tool_result: AnthropicMessagesToolResultParam | None = None ## PROMPT CACHING CHECK ## cache_control = message.get("cache_control", None) if message["role"] == "tool": @@ -1720,8 +1700,8 @@ def convert_to_anthropic_tool_result( def convert_function_to_anthropic_tool_invoke( - function_call: Union[dict, ChatCompletionToolCallFunctionChunk], -) -> List[AnthropicMessagesToolUseParam]: + function_call: dict | ChatCompletionToolCallFunctionChunk, +) -> list[AnthropicMessagesToolUseParam]: try: _name = get_attribute_or_key(function_call, "name") or "" _arguments = get_attribute_or_key(function_call, "arguments") @@ -1742,10 +1722,10 @@ def convert_function_to_anthropic_tool_invoke( def convert_to_anthropic_tool_invoke( - tool_calls: List[ChatCompletionAssistantToolCall], - web_search_results: Optional[List[Any]] = None, - tool_results: Optional[List[Any]] = None, -) -> List[Union[AnthropicMessagesToolUseParam, Dict[str, Any]]]: + tool_calls: list[ChatCompletionAssistantToolCall], + web_search_results: list[Any] | None = None, + tool_results: list[Any] | None = None, +) -> list[AnthropicMessagesToolUseParam | dict[str, Any]]: """ OpenAI tool invokes: { @@ -1788,7 +1768,7 @@ def convert_to_anthropic_tool_invoke( Fixes: https://github.com/BerriAI/litellm/issues/17737 """ - anthropic_tool_invoke: List[Union[AnthropicMessagesToolUseParam, Dict[str, Any]]] = [] + anthropic_tool_invoke: list[AnthropicMessagesToolUseParam | dict[str, Any]] = [] for tool in tool_calls: if not get_attribute_or_key(tool, "type") == "function": @@ -1809,7 +1789,7 @@ def convert_to_anthropic_tool_invoke( # Server tool IDs start with "srvtoolu_" if tool_id.startswith("srvtoolu_"): # Create server_tool_use block instead of tool_use - _anthropic_server_tool_use: Dict[str, Any] = { + _anthropic_server_tool_use: dict[str, Any] = { "type": "server_tool_use", "id": tool_id, "name": tool_name, @@ -1820,7 +1800,7 @@ def convert_to_anthropic_tool_invoke( # Add corresponding tool result if available. # Check both web_search_results (web_search_tool_result / web_fetch_tool_result) # and tool_results (bash_code_execution_tool_result, etc.) - _all_tool_results: List[Any] = [] + _all_tool_results: list[Any] = [] if web_search_results: _all_tool_results.extend(web_search_results) if tool_results: @@ -1853,15 +1833,13 @@ def convert_to_anthropic_tool_invoke( def add_cache_control_to_content( - anthropic_content_element: Union[ - dict, - AnthropicMessagesImageParam, - AnthropicMessagesTextParam, - AnthropicMessagesDocumentParam, - AnthropicMessagesToolUseParam, - ChatCompletionThinkingBlock, - ], - original_content_element: Union[dict, AllMessageValues], + anthropic_content_element: dict + | AnthropicMessagesImageParam + | AnthropicMessagesTextParam + | AnthropicMessagesDocumentParam + | AnthropicMessagesToolUseParam + | ChatCompletionThinkingBlock, + original_content_element: dict | AllMessageValues, ): cache_control_param = original_content_element.get("cache_control") if cache_control_param is not None and isinstance(cache_control_param, dict): @@ -1874,9 +1852,9 @@ def add_cache_control_to_content( def _anthropic_content_element_factory( image_chunk: GenericImageParsingChunk, -) -> Union[AnthropicMessagesImageParam, AnthropicMessagesDocumentParam]: +) -> AnthropicMessagesImageParam | AnthropicMessagesDocumentParam: if image_chunk["media_type"] == "application/pdf": - _anthropic_content_element: Union[AnthropicMessagesDocumentParam, AnthropicMessagesImageParam] = ( + _anthropic_content_element: AnthropicMessagesDocumentParam | AnthropicMessagesImageParam = ( AnthropicMessagesDocumentParam( type="document", source=AnthropicContentParamSource( @@ -1927,11 +1905,7 @@ def anthropic_infer_file_id_content_type( def anthropic_process_openai_file_message( message: ChatCompletionFileObject, -) -> Union[ - AnthropicMessagesDocumentParam, - AnthropicMessagesImageParam, - AnthropicMessagesContainerUploadParam, -]: +) -> AnthropicMessagesDocumentParam | AnthropicMessagesImageParam | AnthropicMessagesContainerUploadParam: file_message = cast(ChatCompletionFileObject, message) file_sub = file_message.get("file") if file_sub is None: @@ -1963,13 +1937,9 @@ def anthropic_process_openai_file_message( if format else anthropic_infer_file_id_content_type(file_id) ) - return_block_param: Optional[ - Union[ - AnthropicMessagesDocumentParam, - AnthropicMessagesImageParam, - AnthropicMessagesContainerUploadParam, - ] - ] = None + return_block_param: ( + AnthropicMessagesDocumentParam | AnthropicMessagesImageParam | AnthropicMessagesContainerUploadParam | None + ) = None if content_block_type == "document": return_block_param = AnthropicMessagesDocumentParam( type="document", @@ -2037,7 +2007,7 @@ def _sanitize_empty_text_content( # Walk the blocks and rewrite any empty text blocks. We rewrite (rather # than drop) so callers don't end up with an entirely empty content # list, which Anthropic also rejects. - new_blocks: List[Any] = [] + new_blocks: list[Any] = [] rewrote_any = False for block in content: if isinstance(block, dict) and block.get("type") == "text": @@ -2062,9 +2032,9 @@ def _sanitize_empty_text_content( def _add_missing_tool_results( current_message: AllMessageValues, - messages: List[AllMessageValues], + messages: list[AllMessageValues], current_index: int, -) -> Tuple[List[AllMessageValues], int]: +) -> tuple[list[AllMessageValues], int]: """ Case A: Missing tool_result for tool_use (orphaned tool calls) - If an assistant message has tool_calls but no corresponding tool result follows, @@ -2076,7 +2046,7 @@ def _add_missing_tool_results( followed by any dummy tool results needed - Number of original messages consumed (to adjust iteration index) """ - result_messages: List[AllMessageValues] = [] + result_messages: list[AllMessageValues] = [] tool_calls = current_message.get("tool_calls") if not tool_calls or len(cast(list, tool_calls)) == 0: @@ -2095,7 +2065,7 @@ def _add_missing_tool_results( # Collect actual tool result messages that follow this assistant message found_tool_call_ids = set() - actual_tool_results: List[AllMessageValues] = [] + actual_tool_results: list[AllMessageValues] = [] j = current_index + 1 while j < len(messages): @@ -2164,7 +2134,7 @@ def _add_missing_tool_results( def _is_orphaned_tool_result( current_message: AllMessageValues, - sanitized_messages: List[AllMessageValues], + sanitized_messages: list[AllMessageValues], ) -> bool: """ Case B: Orphaned tool_result (unexpected result) @@ -2254,8 +2224,8 @@ def _iter_tool_exchange_groups(messages: Sequence[Mapping[str, Any]]) -> Iterato def sanitize_messages_for_tool_calling( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: """ Sanitize messages for tool calling to handle common issues when modify_params=True: @@ -2281,7 +2251,7 @@ def sanitize_messages_for_tool_calling( if not litellm.modify_params: return messages - sanitized_messages: List[AllMessageValues] = [] + sanitized_messages: list[AllMessageValues] = [] i = 0 while i < len(messages): @@ -2322,8 +2292,8 @@ def sanitize_messages_for_tool_calling( # which keeps the *first*. The Bedrock case handles provider-side content # block duplication where the first is authoritative; here the duplicate # arises from history replay where the last entry is the final state. - duplicates_to_remove: Set[int] = set() - seen_in_block: Dict[str, int] = {} # tool_call_id -> index (reset per block) + duplicates_to_remove: set[int] = set() + seen_in_block: dict[str, int] = {} # tool_call_id -> index (reset per block) for idx, msg in enumerate(sanitized_messages): role = msg.get("role") tcid = msg.get("tool_call_id") if role in ["tool", "function"] else None @@ -2367,21 +2337,16 @@ def _is_unsignable_thinking_block(block: object) -> bool: def _drop_unsignable_thinking_blocks( - thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]], -) -> list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]: + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock], +) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]: return [block for block in thinking_blocks if not _is_unsignable_thinking_block(block)] def anthropic_messages_pt( - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, llm_provider: str, -) -> List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] -]: +) -> list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]: """ format messages for anthropic 1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately) @@ -2409,12 +2374,7 @@ def anthropic_messages_pt( # add role=tool support to allow function call result/error submission user_message_types = {"user", "tool", "function"} # reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them. - new_messages: List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ] = [] + new_messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam] = [] if len(messages) == 0: if not litellm.modify_params: @@ -2434,17 +2394,15 @@ def anthropic_messages_pt( msg_i = 0 while msg_i < len(messages): - user_content: List[AnthropicMessagesUserMessageValues] = [] + user_content: list[AnthropicMessagesUserMessageValues] = [] init_msg_i = msg_i if isinstance(messages[msg_i], BaseModel): messages[msg_i] = dict(messages[msg_i]) # type: ignore ## MERGE CONSECUTIVE USER CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types: - user_message_types_block: Union[ - ChatCompletionToolMessage, - ChatCompletionUserMessage, - ChatCompletionFunctionMessage, - ] = messages[msg_i] # type: ignore + user_message_types_block: ( + ChatCompletionToolMessage | ChatCompletionUserMessage | ChatCompletionFunctionMessage + ) = messages[msg_i] # type: ignore if user_message_types_block["role"] == "user": if isinstance(user_message_types_block["content"], list): for m in user_message_types_block["content"]: @@ -2454,7 +2412,7 @@ def anthropic_messages_pt( # Convert ChatCompletionImageUrlObject to dict if needed image_url_value = m["image_url"] if isinstance(image_url_value, str): - image_url_input: Union[str, dict[str, Any]] = image_url_value + image_url_input: str | dict[str, Any] = image_url_value else: # ChatCompletionImageUrlObject or dict case - convert to dict image_url_input = { @@ -2545,9 +2503,9 @@ def anthropic_messages_pt( new_messages.append({"role": "user", "content": user_content}) # Track unique tool IDs in this merge block to avoid duplication - unique_tool_ids: Set[str] = set() + unique_tool_ids: set[str] = set() - assistant_content: List[AnthropicMessagesAssistantMessageValues] = [] + assistant_content: list[AnthropicMessagesAssistantMessageValues] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i] # type: ignore @@ -2592,9 +2550,9 @@ def anthropic_messages_pt( # Build the tool call groups (server_tool_use + its result) _provider_specific_fields_raw_tc = assistant_content_block.get("provider_specific_fields") - _provider_specific_fields_tc: Dict[str, Any] = {} + _provider_specific_fields_tc: dict[str, Any] = {} if isinstance(_provider_specific_fields_raw_tc, dict): - _provider_specific_fields_tc = cast(Dict[str, Any], _provider_specific_fields_raw_tc) + _provider_specific_fields_tc = cast(dict[str, Any], _provider_specific_fields_raw_tc) _web_search_results_tc = _provider_specific_fields_tc.get("web_search_results") _tool_results_tc = _provider_specific_fields_tc.get("tool_results") tool_invoke_results = convert_to_anthropic_tool_invoke( @@ -2605,9 +2563,9 @@ def anthropic_messages_pt( # Group tool invoke results into (server_tool_use, result) pairs # and separate regular tool_use blocks - server_tool_groups: List[List[Any]] = [] - regular_tool_uses: List[Any] = [] - _current_group: List[Any] = [] + server_tool_groups: list[list[Any]] = [] + regular_tool_uses: list[Any] = [] + _current_group: list[Any] = [] for item in tool_invoke_results: item_type = item.get("type", "") if isinstance(item, dict) else getattr(item, "type", "") if item_type == "server_tool_use": @@ -2731,10 +2689,9 @@ def anthropic_messages_pt( and len(thinking_block) > 0 and not _is_unsignable_thinking_block(m) ): # don't pass empty text blocks. anthropic api raises errors. - anthropic_message: Union[ - ChatCompletionThinkingBlock, - AnthropicMessagesTextParam, - ] = cast(ChatCompletionThinkingBlock, m) + anthropic_message: ChatCompletionThinkingBlock | AnthropicMessagesTextParam = cast( + ChatCompletionThinkingBlock, m + ) assistant_content.append(anthropic_message) # handle text elif ( @@ -2749,12 +2706,7 @@ def anthropic_messages_pt( assistant_content.append(cast(AnthropicMessagesTextParam, _cached_message)) # handle server_tool_use blocks (tool search, web search, etc.) # Pass through as-is since these are Anthropic-native content types - elif m.get("type", "") == "server_tool_use": - assistant_content.append(m) # type: ignore - # handle all *_tool_result blocks (tool_search_tool_result, - # web_search_tool_result, bash_code_execution_tool_result, etc.) - # Pass through as-is since these are Anthropic-native content types - elif m.get("type", "").endswith("_tool_result"): + elif m.get("type", "") == "server_tool_use" or m.get("type", "").endswith("_tool_result"): assistant_content.append(m) # type: ignore elif ( "content" in assistant_content_block @@ -2781,9 +2733,9 @@ def anthropic_messages_pt( # for server_tool_use reconstruction. # Fixes: https://github.com/BerriAI/litellm/issues/17737 _provider_specific_fields_raw = assistant_content_block.get("provider_specific_fields") - _provider_specific_fields: Dict[str, Any] = {} + _provider_specific_fields: dict[str, Any] = {} if isinstance(_provider_specific_fields_raw, dict): - _provider_specific_fields = cast(Dict[str, Any], _provider_specific_fields_raw) + _provider_specific_fields = cast(dict[str, Any], _provider_specific_fields_raw) _web_search_results = _provider_specific_fields.get("web_search_results") _tool_results = _provider_specific_fields.get("tool_results") tool_invoke_results = convert_to_anthropic_tool_invoke( @@ -2833,7 +2785,7 @@ def anthropic_messages_pt( return new_messages -def extract_between_tags(tag: str, string: str, strip: bool = False) -> List[str]: +def extract_between_tags(tag: str, string: str, strip: bool = False) -> list[str]: ext_list = re.findall(f"<{tag}>(.+?)", string, re.DOTALL) if strip: ext_list = [e.strip() for e in ext_list] @@ -2844,7 +2796,7 @@ def contains_tag(tag: str, string: str) -> bool: return bool(re.search(f"<{tag}>(.+?)", string, re.DOTALL)) -def parse_xml_params(xml_content, json_schema: Optional[dict] = None): +def parse_xml_params(xml_content, json_schema: dict | None = None): """ Compare the xml output to the json schema @@ -2920,8 +2872,8 @@ from litellm.types.llms.cohere import ( def convert_openai_message_to_cohere_tool_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], - tool_calls: List, + message: ChatCompletionToolMessage | ChatCompletionFunctionMessage, + tool_calls: list, ) -> ToolResultObject: """ OpenAI message with a tool result looks like: @@ -2961,7 +2913,7 @@ def convert_openai_message_to_cohere_tool_result( content_str: str = "" if isinstance(message["content"], str): content_str = message["content"] - elif isinstance(message["content"], List): + elif isinstance(message["content"], list): content_list = message["content"] for content in content_list: if content["type"] == "text": @@ -3006,13 +2958,13 @@ def convert_openai_message_to_cohere_tool_result( return cohere_tool_result -def get_all_tool_calls(messages: List) -> List: +def get_all_tool_calls(messages: list) -> list: """ Returns extracted list of `tool_calls`. Done to handle openai no longer returning tool call 'name' in tool results. """ - tool_calls: List = [] + tool_calls: list = [] for m in messages: if m.get("tool_calls", None) is not None: if isinstance(m["tool_calls"], list): @@ -3021,7 +2973,7 @@ def get_all_tool_calls(messages: List) -> List: return tool_calls -def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: +def convert_to_cohere_tool_invoke(tool_calls: list) -> list[ToolCallObject]: """ OpenAI tool invokes: { @@ -3048,7 +3000,7 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: } """ - cohere_tool_invoke: List[ToolCallObject] = [ + cohere_tool_invoke: list[ToolCallObject] = [ { "name": get_attribute_or_key(get_attribute_or_key(tool, "function"), "name"), "parameters": json.loads(get_attribute_or_key(get_attribute_or_key(tool, "function"), "arguments")), @@ -3061,10 +3013,10 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: def cohere_messages_pt_v2( - messages: List, + messages: list, model: str, llm_provider: str, -) -> Tuple[Union[str, ToolResultObject], ChatHistory]: +) -> tuple[str | ToolResultObject, ChatHistory]: """ Returns a tuple(Union[tool_result, message], chat_history) @@ -3078,16 +3030,16 @@ def cohere_messages_pt_v2( - message must be at least 1 token long or tool results must be specified. - cannot specify tool_results if the last entry in chat history contains a user message """ - tool_calls: List = get_all_tool_calls(messages=messages) + tool_calls: list = get_all_tool_calls(messages=messages) ## GET MOST RECENT MESSAGE most_recent_message = messages.pop(-1) - returned_message: Union[ToolResultObject, str] = "" + returned_message: ToolResultObject | str = "" if most_recent_message.get("role", "") is not None and most_recent_message["role"] == "tool": # tool result returned_message = convert_openai_message_to_cohere_tool_result(most_recent_message, tool_calls) else: - content: Union[str, List] = most_recent_message.get("content") + content: str | list = most_recent_message.get("content") if isinstance(content, str): returned_message = content else: @@ -3133,7 +3085,7 @@ def cohere_messages_pt_v2( new_messages.append(ChatHistorySystem(role="SYSTEM", message=system_content)) assistant_content: str = "" - assistant_tool_calls: List[ToolCallObject] = [] + assistant_tool_calls: list[ToolCallObject] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": if messages[msg_i].get("content", None) is not None and isinstance(messages[msg_i]["content"], list): @@ -3160,7 +3112,7 @@ def cohere_messages_pt_v2( ) ## MERGE CONSECUTIVE TOOL RESULTS - tool_results: List[ToolResultObject] = [] + tool_results: list[ToolResultObject] = [] while msg_i < len(messages) and messages[msg_i]["role"] in tool_message_types: tool_results.append(convert_openai_message_to_cohere_tool_result(messages[msg_i], tool_calls)) @@ -3180,7 +3132,7 @@ def cohere_messages_pt_v2( def cohere_message_pt(messages: list): - tool_calls: List = get_all_tool_calls(messages=messages) + tool_calls: list = get_all_tool_calls(messages=messages) prompt = "" tool_results = [] for message in messages: @@ -3369,7 +3321,7 @@ def azure_text_pt(messages: list): ###### AZURE AI ####### -def stringify_json_tool_call_content(messages: List) -> List: +def stringify_json_tool_call_content(messages: list) -> list: """ - Check 'content' in tool role -> convert to dict (if not) -> stringify @@ -3397,14 +3349,14 @@ import httpx from litellm.types.llms.bedrock import ( BedrockConverseReasoningContentBlock, BedrockConverseReasoningTextBlock, + BedrockToolSpec, + SearchResultBlock, ) from litellm.types.llms.bedrock import ContentBlock as BedrockContentBlock from litellm.types.llms.bedrock import DocumentBlock as BedrockDocumentBlock from litellm.types.llms.bedrock import ImageBlock as BedrockImageBlock from litellm.types.llms.bedrock import SourceBlock as BedrockSourceBlock -from litellm.types.llms.bedrock import BedrockToolSpec from litellm.types.llms.bedrock import ToolBlock as BedrockToolBlock -from litellm.types.llms.bedrock import SearchResultBlock from litellm.types.llms.bedrock import ToolResultBlock as BedrockToolResultBlock from litellm.types.llms.bedrock import ( ToolResultContentBlock as BedrockToolResultContentBlock, @@ -3419,7 +3371,7 @@ def _parse_content_type(content_type: str) -> str: return m.get_content_type() -def _parse_mime_type(base64_data: str) -> Optional[str]: +def _parse_mime_type(base64_data: str) -> str | None: mime_type_match = re.match(r"data:(.*?);base64", base64_data) if mime_type_match: return mime_type_match.group(1) @@ -3431,7 +3383,7 @@ class BedrockImageProcessor: """Handles both sync and async image processing for Bedrock conversations.""" @staticmethod - def _post_call_image_processing(response: httpx.Response, image_url: str = "") -> Tuple[str, str]: + def _post_call_image_processing(response: httpx.Response, image_url: str = "") -> tuple[str, str]: # Check the response's content type to ensure it is an image content_type = response.headers.get("content-type") @@ -3450,7 +3402,7 @@ class BedrockImageProcessor: return base64_bytes, content_type @staticmethod - async def get_image_details_async(image_url) -> Tuple[str, str]: + async def get_image_details_async(image_url) -> tuple[str, str]: try: client = get_async_httpx_client( llm_provider=httpxSpecialProvider.PromptFactory, @@ -3466,7 +3418,7 @@ class BedrockImageProcessor: raise e @staticmethod - def get_image_details(image_url) -> Tuple[str, str]: + def get_image_details(image_url) -> tuple[str, str]: try: client = HTTPHandler(concurrent_limit=1) # Send a GET request to the image URL @@ -3479,7 +3431,7 @@ class BedrockImageProcessor: raise e @staticmethod - def _parse_base64_image(image_url: str) -> Tuple[str, str, str]: + def _parse_base64_image(image_url: str) -> tuple[str, str, str]: """Parse base64 encoded image data.""" image_metadata, img_without_base_64 = image_url.split(",") @@ -3507,7 +3459,7 @@ class BedrockImageProcessor: document_types = ["application", "text"] is_document = any(mime_type.startswith(doc_type) for doc_type in document_types) - supported_image_and_video_formats: List[str] = supported_video_formats + supported_image_formats + supported_image_and_video_formats: list[str] = supported_video_formats + supported_image_formats if is_document: return BedrockImageProcessor._get_document_format( @@ -3525,7 +3477,7 @@ class BedrockImageProcessor: return image_format @staticmethod - def _get_document_format(mime_type: str, supported_doc_formats: List[str]) -> str: + def _get_document_format(mime_type: str, supported_doc_formats: list[str]) -> str: """ Get the document format from the mime type @@ -3543,7 +3495,7 @@ class BedrockImageProcessor: Returns: The document format """ - valid_extensions: Optional[List[str]] = None + valid_extensions: list[str] | None = None potential_extensions = mimetypes.guess_all_extensions(mime_type, strict=False) valid_extensions = [ext[1:] for ext in potential_extensions if ext[1:] in supported_doc_formats] @@ -3619,7 +3571,7 @@ class BedrockImageProcessor: return BedrockContentBlock(image=BedrockImageBlock(source=_blob, format=image_format)) @classmethod - def process_image_sync(cls, image_url: str, format: Optional[str] = None) -> BedrockContentBlock: + def process_image_sync(cls, image_url: str, format: str | None = None) -> BedrockContentBlock: """Synchronous image processing.""" if "base64" in image_url: @@ -3638,7 +3590,7 @@ class BedrockImageProcessor: return cls._create_bedrock_block(img_bytes, mime_type, image_format) @classmethod - async def process_image_async(cls, image_url: str, format: Optional[str]) -> BedrockContentBlock: + async def process_image_async(cls, image_url: str, format: str | None) -> BedrockContentBlock: """Asynchronous image processing.""" if "base64" in image_url: @@ -3659,8 +3611,8 @@ class BedrockImageProcessor: def _convert_to_bedrock_tool_call_invoke( tool_calls: list, - model: Optional[str] = None, -) -> List[BedrockContentBlock]: + model: str | None = None, +) -> list[BedrockContentBlock]: """ OpenAI tool invokes: { @@ -3701,7 +3653,7 @@ def _convert_to_bedrock_tool_call_invoke( ) try: - _parts_list: List[BedrockContentBlock] = [] + _parts_list: list[BedrockContentBlock] = [] for tool in tool_calls: if "function" in tool: tool_id = tool["id"] @@ -3761,13 +3713,11 @@ def _convert_to_bedrock_tool_call_invoke( _parts_list.append(cache_point_block) return _parts_list except Exception as e: - raise Exception( - "Unable to convert openai tool calls={} to bedrock tool calls. Received error={}".format(tool_calls, str(e)) - ) + raise Exception(f"Unable to convert openai tool calls={tool_calls} to bedrock tool calls. Received error={e!s}") def _append_bedrock_tool_result_media_block( - tool_result_content_blocks: List[BedrockToolResultContentBlock], + tool_result_content_blocks: list[BedrockToolResultContentBlock], processed_block: BedrockContentBlock, content: dict, content_type: str, @@ -3786,10 +3736,10 @@ def _append_bedrock_tool_result_media_block( def _append_bedrock_tool_result_image_url_block( - tool_result_content_blocks: List[BedrockToolResultContentBlock], + tool_result_content_blocks: list[BedrockToolResultContentBlock], content: dict, ) -> None: - format: Optional[str] = None + format: str | None = None if isinstance(content["image_url"], dict): image_url = content["image_url"]["url"] format = content["image_url"].get("format") @@ -3803,7 +3753,7 @@ def _append_bedrock_tool_result_image_url_block( def _append_bedrock_tool_result_file_block( - tool_result_content_blocks: List[BedrockToolResultContentBlock], + tool_result_content_blocks: list[BedrockToolResultContentBlock], content: dict, ) -> None: # Match the user-message path (_process_file_message): accept either @@ -3813,7 +3763,7 @@ def _append_bedrock_tool_result_file_block( file_id = file_obj.get("file_id") if file_data is None and file_id is None: raise litellm.BadRequestError( - message="file_data and file_id cannot both be None. Got={}".format(content), + message=f"file_data and file_id cannot both be None. Got={content}", model="", llm_provider="bedrock", ) @@ -3825,9 +3775,9 @@ def _append_bedrock_tool_result_file_block( def _parse_bedrock_tool_result_content_list( - content_list: List, -) -> List[BedrockToolResultContentBlock]: - tool_result_content_blocks: List[BedrockToolResultContentBlock] = [] + content_list: list, +) -> list[BedrockToolResultContentBlock]: + tool_result_content_blocks: list[BedrockToolResultContentBlock] = [] for content in content_list: if content["type"] == "text": tool_result_content_blocks.append(BedrockToolResultContentBlock(text=content["text"])) @@ -3839,8 +3789,8 @@ def _parse_bedrock_tool_result_content_list( def _build_bedrock_tool_result_content_blocks( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], -) -> tuple[List[BedrockToolResultContentBlock], bool]: + message: ChatCompletionToolMessage | ChatCompletionFunctionMessage, +) -> tuple[list[BedrockToolResultContentBlock], bool]: # Optional OpenAI tool-message extension: # allow structured Bedrock search results on tool messages and map them # directly to toolResult.content[].searchResult for Converse API. @@ -3849,7 +3799,7 @@ def _build_bedrock_tool_result_content_blocks( # to avoid generating mixed text + searchResult blocks. search_results = message.get("search_results") if isinstance(search_results, list): - tool_result_content_blocks: List[BedrockToolResultContentBlock] = [] + tool_result_content_blocks: list[BedrockToolResultContentBlock] = [] for result in search_results: if not isinstance(result, dict): continue @@ -3862,13 +3812,13 @@ def _build_bedrock_tool_result_content_blocks( message_content = message["content"] if isinstance(message_content, str): return [BedrockToolResultContentBlock(text=message_content)], False - if isinstance(message_content, List): + if isinstance(message_content, list): return _parse_bedrock_tool_result_content_list(message_content), False return [], False def _convert_to_bedrock_tool_call_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], + message: ChatCompletionToolMessage | ChatCompletionFunctionMessage, ) -> BedrockContentBlock: """ OpenAI message with a tool result looks like: @@ -3925,10 +3875,10 @@ def _convert_to_bedrock_tool_call_result( def _deduplicate_bedrock_content_blocks( - blocks: List[BedrockContentBlock], + blocks: list[BedrockContentBlock], block_key: str, id_key: str = "toolUseId", -) -> List[BedrockContentBlock]: +) -> list[BedrockContentBlock]: """ Remove duplicate content blocks that share the same ID under ``block_key``. @@ -3948,8 +3898,8 @@ def _deduplicate_bedrock_content_blocks( block_key: The dict key to inspect (e.g. ``"toolResult"`` or ``"toolUse"``). id_key: The nested key that holds the unique ID (default ``"toolUseId"``). """ - seen_ids: Set[str] = set() - deduplicated: List[BedrockContentBlock] = [] + seen_ids: set[str] = set() + deduplicated: list[BedrockContentBlock] = [] for block in blocks: keyed = block.get(block_key) if keyed is not None and isinstance(keyed, dict): @@ -3971,15 +3921,15 @@ def _deduplicate_bedrock_content_blocks( def _deduplicate_bedrock_tool_content( - tool_content: List[BedrockContentBlock], -) -> List[BedrockContentBlock]: + tool_content: list[BedrockContentBlock], +) -> list[BedrockContentBlock]: """Convenience wrapper: deduplicate ``toolResult`` blocks by ``toolUseId``.""" return _deduplicate_bedrock_content_blocks(tool_content, "toolResult") def _rename_duplicate_bedrock_document_names( - contents: List[BedrockMessageBlock], -) -> List[BedrockMessageBlock]: + contents: list[BedrockMessageBlock], +) -> list[BedrockMessageBlock]: """ Rename duplicate document names across all messages in a Bedrock request. @@ -3991,14 +3941,14 @@ def _rename_duplicate_bedrock_document_names( (``_2``, ``_3``, ...), bumped further if the suffixed name already belongs to another document (e.g. an organic name ending in ``_2``). """ - used_names: Set[str] = set() + used_names: set[str] = set() for message in contents: for block in message.get("content") or []: document = block.get("document") if isinstance(document, dict) and document.get("name"): used_names.add(document["name"]) - name_counts: Dict[str, int] = {} + name_counts: dict[str, int] = {} for message in contents: for block in message.get("content") or []: document = block.get("document") @@ -4021,8 +3971,8 @@ def _rename_duplicate_bedrock_document_names( def _sort_bedrock_assistant_content_blocks( - blocks: List[BedrockContentBlock], -) -> List[BedrockContentBlock]: + blocks: list[BedrockContentBlock], +) -> list[BedrockContentBlock]: """ Sort assistant content blocks so that ``text`` blocks appear before ``toolUse`` blocks. @@ -4055,9 +4005,9 @@ def _sort_bedrock_assistant_content_blocks( def _insert_assistant_continue_message( - messages: List[BedrockMessageBlock], - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, -) -> List[BedrockMessageBlock]: + messages: list[BedrockMessageBlock], + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, +) -> list[BedrockMessageBlock]: """ Add dummy message between user/tool result blocks. @@ -4094,7 +4044,7 @@ def _insert_assistant_continue_message( def get_user_message_block_or_continue_message( message: ChatCompletionUserMessage, - user_continue_message: Optional[ChatCompletionUserMessage] = None, + user_continue_message: ChatCompletionUserMessage | None = None, ) -> ChatCompletionUserMessage: """ Returns the user content block @@ -4156,7 +4106,7 @@ def get_user_message_block_or_continue_message( def return_assistant_continue_message( - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, ) -> ChatCompletionAssistantMessage: if assistant_continue_message and isinstance(assistant_continue_message, str): return ChatCompletionAssistantMessage( @@ -4169,7 +4119,7 @@ def return_assistant_continue_message( return DEFAULT_ASSISTANT_CONTINUE_MESSAGE -def _skip_empty_dict_blocks(blocks: List[dict]) -> List[dict]: +def _skip_empty_dict_blocks(blocks: list[dict]) -> list[dict]: """ Filter out empty text blocks from a list of dictionaries. @@ -4197,8 +4147,8 @@ def skip_empty_text_blocks( def skip_empty_text_blocks( - message: Union[ChatCompletionAssistantMessage, ChatCompletionUserMessage], -) -> Union[ChatCompletionAssistantMessage, ChatCompletionUserMessage]: + message: ChatCompletionAssistantMessage | ChatCompletionUserMessage, +) -> ChatCompletionAssistantMessage | ChatCompletionUserMessage: """ Skips empty text blocks in message content text blocks. @@ -4217,7 +4167,7 @@ def skip_empty_text_blocks( modified_message["content"] = None # user message content cannot be None return modified_message elif isinstance(content_block, list): - modified_content_block = _skip_empty_dict_blocks(cast(List[dict], content_block)) + modified_content_block = _skip_empty_dict_blocks(cast(list[dict], content_block)) # If no content remains and it's an assistant message, set content to None if not modified_content_block and message["role"] == "assistant": @@ -4230,12 +4180,12 @@ def skip_empty_text_blocks( # Type-specific casting based on message role if message["role"] == "assistant": modified_message_alt["content"] = cast( # type: ignore - Optional[List[OpenAIMessageContentListBlock]], + list[OpenAIMessageContentListBlock] | None, modified_content_block or None, ) elif message["role"] == "user" and modified_content_block is not None: modified_message_alt["content"] = cast( # type: ignore - Optional[List[ChatCompletionTextObject]], modified_content_block + list[ChatCompletionTextObject] | None, modified_content_block ) return modified_message_alt @@ -4245,7 +4195,7 @@ def skip_empty_text_blocks( def process_empty_text_blocks( message: ChatCompletionAssistantMessage, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, ) -> ChatCompletionAssistantMessage: modified_content_block = message.get("content", None) ## BASE CASE ## @@ -4270,7 +4220,7 @@ def process_empty_text_blocks( modified_message = message.copy() modified_message["content"] = cast( - Union[List[ChatCompletionTextObject], List[ChatCompletionThinkingBlock]], + list[ChatCompletionTextObject] | list[ChatCompletionThinkingBlock], modified_content_block, ) return modified_message @@ -4278,7 +4228,7 @@ def process_empty_text_blocks( def get_assistant_message_block_or_continue_message( message: ChatCompletionAssistantMessage, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, ) -> ChatCompletionAssistantMessage: """ Returns the user content block @@ -4324,11 +4274,11 @@ def get_assistant_message_block_or_continue_message( class BedrockConverseMessagesProcessor: @staticmethod def _initial_message_setup( - messages: List, + messages: list, model: str, llm_provider: str, - user_continue_message: Optional[ChatCompletionUserMessage] = None, - ) -> List: + user_continue_message: ChatCompletionUserMessage | None = None, + ) -> list: # gracefully handle base case of no messages at all if len(messages) == 0: if user_continue_message is not None: @@ -4361,13 +4311,13 @@ class BedrockConverseMessagesProcessor: @staticmethod async def _bedrock_converse_messages_pt_async( - messages: List, + messages: list, model: str, llm_provider: str, - user_continue_message: Optional[ChatCompletionUserMessage] = None, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, - ) -> List[BedrockMessageBlock]: - contents: List[BedrockMessageBlock] = [] + user_continue_message: ChatCompletionUserMessage | None = None, + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, + ) -> list[BedrockMessageBlock]: + contents: list[BedrockMessageBlock] = [] msg_i = 0 messages = BedrockConverseMessagesProcessor._initial_message_setup( @@ -4375,7 +4325,7 @@ class BedrockConverseMessagesProcessor: ) while msg_i < len(messages): - user_content: List[BedrockContentBlock] = [] + user_content: list[BedrockContentBlock] = [] init_msg_i = msg_i ## MERGE CONSECUTIVE USER CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "user": @@ -4384,7 +4334,7 @@ class BedrockConverseMessagesProcessor: user_continue_message=user_continue_message, ) if isinstance(message_block["content"], list): - _parts: List[BedrockContentBlock] = [] + _parts: list[BedrockContentBlock] = [] for element in message_block["content"]: if isinstance(element, dict): if element["type"] == "text": @@ -4401,7 +4351,7 @@ class BedrockConverseMessagesProcessor: _part = BedrockContentBlock(text=element["text"]) _parts.append(_part) elif element["type"] == "image_url": - format: Optional[str] = None + format: str | None = None if isinstance(element["image_url"], dict): image_url = element["image_url"]["url"] format = element["image_url"].get("format") @@ -4455,7 +4405,7 @@ class BedrockConverseMessagesProcessor: contents.append(BedrockMessageBlock(role="user", content=user_content)) ## MERGE CONSECUTIVE TOOL CALL MESSAGES ## - tool_content: List[BedrockContentBlock] = [] + tool_content: list[BedrockContentBlock] = [] while msg_i < len(messages) and messages[msg_i]["role"] == "tool": current_message = messages[msg_i] tool_call_result = _convert_to_bedrock_tool_call_result(current_message) @@ -4504,7 +4454,7 @@ class BedrockConverseMessagesProcessor: contents[-1]["content"].extend(tool_content) else: contents.append(BedrockMessageBlock(role="user", content=tool_content)) - assistant_content: List[BedrockContentBlock] = [] + assistant_content: list[BedrockContentBlock] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": assistant_message_block = get_assistant_message_block_or_continue_message( @@ -4513,7 +4463,7 @@ class BedrockConverseMessagesProcessor: ) _assistant_content = assistant_message_block.get("content", None) thinking_blocks = cast( - Optional[List[ChatCompletionThinkingBlock]], + list[ChatCompletionThinkingBlock] | None, assistant_message_block.get("thinking_blocks"), ) @@ -4529,7 +4479,7 @@ class BedrockConverseMessagesProcessor: ) if _assistant_content is not None and isinstance(_assistant_content, list): - assistants_parts: List[BedrockContentBlock] = [] + assistants_parts: list[BedrockContentBlock] = [] for element in _assistant_content: if isinstance(element, dict): if element["type"] == "thinking": @@ -4600,9 +4550,9 @@ class BedrockConverseMessagesProcessor: @staticmethod def translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks: List[ChatCompletionThinkingBlock], - ) -> List[BedrockContentBlock]: - reasoning_content_blocks: List[BedrockContentBlock] = [] + thinking_blocks: list[ChatCompletionThinkingBlock], + ) -> list[BedrockContentBlock]: + reasoning_content_blocks: list[BedrockContentBlock] = [] for thinking_block in thinking_blocks: reasoning_text = thinking_block.get("thinking") reasoning_signature = thinking_block.get("signature") @@ -4632,7 +4582,7 @@ class BedrockConverseMessagesProcessor: if file_data is None and file_id is None: raise litellm.BadRequestError( - message="file_data and file_id cannot both be None. Got={}".format(message), + message=f"file_data and file_id cannot both be None. Got={message}", model="", llm_provider="bedrock", ) @@ -4655,7 +4605,7 @@ class BedrockConverseMessagesProcessor: format = file_message.get("format") if file_data is None and file_id is None: raise litellm.BadRequestError( - message="file_data and file_id cannot both be None. Got={}".format(message), + message=f"file_data and file_id cannot both be None. Got={message}", model="", llm_provider="bedrock", ) @@ -4699,9 +4649,9 @@ class BedrockConverseMessagesProcessor: @staticmethod def add_thinking_blocks_to_assistant_content( - thinking_blocks: List[BedrockContentBlock], - assistant_parts: List[BedrockContentBlock], - ) -> List[BedrockContentBlock]: + thinking_blocks: list[BedrockContentBlock], + assistant_parts: list[BedrockContentBlock], + ) -> list[BedrockContentBlock]: """ If contains 'signature', it is a thinking block. If missing 'signature', it is a text block - e.g. when using a non-anthropic model. @@ -4727,12 +4677,12 @@ class BedrockConverseMessagesProcessor: def _bedrock_converse_messages_pt( - messages: List, + messages: list, model: str, llm_provider: str, - user_continue_message: Optional[ChatCompletionUserMessage] = None, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, -) -> List[BedrockMessageBlock]: + user_continue_message: ChatCompletionUserMessage | None = None, + assistant_continue_message: str | ChatCompletionAssistantMessage | None = None, +) -> list[BedrockMessageBlock]: """ Converts given messages from OpenAI format to Bedrock format @@ -4741,7 +4691,7 @@ def _bedrock_converse_messages_pt( - Conversation blocks and tool result blocks cannot be provided in the same turn. Issue: https://github.com/BerriAI/litellm/issues/6053 """ - contents: List[BedrockMessageBlock] = [] + contents: list[BedrockMessageBlock] = [] msg_i = 0 messages = BedrockConverseMessagesProcessor._initial_message_setup( @@ -4749,7 +4699,7 @@ def _bedrock_converse_messages_pt( ) while msg_i < len(messages): - user_content: List[BedrockContentBlock] = [] + user_content: list[BedrockContentBlock] = [] init_msg_i = msg_i ## MERGE CONSECUTIVE USER CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "user": @@ -4758,7 +4708,7 @@ def _bedrock_converse_messages_pt( user_continue_message=user_continue_message, ) if isinstance(message_block["content"], list): - _parts: List[BedrockContentBlock] = [] + _parts: list[BedrockContentBlock] = [] for element in message_block["content"]: if isinstance(element, dict): if element["type"] == "text": @@ -4775,7 +4725,7 @@ def _bedrock_converse_messages_pt( _part = BedrockContentBlock(text=element["text"]) _parts.append(_part) elif element["type"] == "image_url": - format: Optional[str] = None + format: str | None = None if isinstance(element["image_url"], dict): image_url = element["image_url"]["url"] format = element["image_url"].get("format") @@ -4830,7 +4780,7 @@ def _bedrock_converse_messages_pt( contents.append(BedrockMessageBlock(role="user", content=user_content)) ## MERGE CONSECUTIVE TOOL CALL MESSAGES ## - tool_content: List[BedrockContentBlock] = [] + tool_content: list[BedrockContentBlock] = [] while msg_i < len(messages) and messages[msg_i]["role"] == "tool": tool_call_result = _convert_to_bedrock_tool_call_result(messages[msg_i]) current_message = messages[msg_i] @@ -4881,7 +4831,7 @@ def _bedrock_converse_messages_pt( contents[-1]["content"].extend(tool_content) else: contents.append(BedrockMessageBlock(role="user", content=tool_content)) - assistant_content: List[BedrockContentBlock] = [] + assistant_content: list[BedrockContentBlock] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": assistant_message_block = get_assistant_message_block_or_continue_message( @@ -4890,7 +4840,7 @@ def _bedrock_converse_messages_pt( ) _assistant_content = assistant_message_block.get("content", None) thinking_blocks = cast( - Optional[List[ChatCompletionThinkingBlock]], + list[ChatCompletionThinkingBlock] | None, assistant_message_block.get("thinking_blocks"), ) @@ -4906,7 +4856,7 @@ def _bedrock_converse_messages_pt( ) if _assistant_content is not None and isinstance(_assistant_content, list): - assistants_parts: List[BedrockContentBlock] = [] + assistants_parts: list[BedrockContentBlock] = [] for element in _assistant_content: if isinstance(element, dict): if element["type"] == "thinking": @@ -5004,7 +4954,7 @@ def make_valid_bedrock_tool_name(input_tool_name: str) -> str: return valid_string -def add_cache_point_tool_block(tool: dict, model: Optional[str] = None) -> Optional[BedrockToolBlock]: +def add_cache_point_tool_block(tool: dict, model: str | None = None) -> BedrockToolBlock | None: from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock cache_control = tool.get("cache_control", None) @@ -5044,7 +4994,7 @@ def _is_bedrock_tool_block(tool: dict) -> bool: return isinstance(tool, dict) and ("systemTool" in tool or "toolSpec" in tool or "cachePoint" in tool) -def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockToolBlock]: +def _bedrock_tools_pt(tools: list, model: str | None = None) -> list[BedrockToolBlock]: """ OpenAI tools looks like: tools = [ @@ -5093,11 +5043,11 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT } ] """ + from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs from litellm.llms.bedrock.common_utils import ( bedrock_converse_supports_strict_tools, normalize_json_schema_custom_types_to_object, ) - from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs _valid_json_schema_root_types = frozenset(("array", "boolean", "integer", "null", "number", "object", "string")) # Only Claude on Bedrock honours strict tool schemas; other families @@ -5106,7 +5056,7 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT # maps toolSpec to the native Anthropic tool shape, which has no strict # field, even though Anthropic's native API accepts it as a top-level key. supports_strict_tools = bool(model and bedrock_converse_supports_strict_tools(model)) - tool_block_list: List[BedrockToolBlock] = [] + tool_block_list: list[BedrockToolBlock] = [] for tool_idx, tool in enumerate(tools): # Check if tool is already a BedrockToolBlock (e.g., systemTool for Nova grounding) if _is_bedrock_tool_block(tool): @@ -5198,8 +5148,8 @@ def response_schema_prompt(model: str, response_schema: dict) -> str: Returns the prompt str that's passed to the model as a user message """ - custom_prompt_details: Optional[dict] = None - response_schema_as_message = [{"role": "user", "content": "{}".format(response_schema)}] + custom_prompt_details: dict | None = None + response_schema_as_message = [{"role": "user", "content": f"{response_schema}"}] if f"{model}/response_schema_prompt" in litellm.custom_prompt_dict: custom_prompt_details = litellm.custom_prompt_dict[ f"{model}/response_schema_prompt" @@ -5224,10 +5174,10 @@ def default_response_schema_prompt(response_schema: dict) -> str: This is the default prompt. Allow user to override this with a custom_prompt. """ - prompt_str = """Use this JSON schema: + prompt_str = f"""Use this JSON schema: ```json - {} - ```""".format(response_schema) + {response_schema} + ```""" return prompt_str @@ -5277,8 +5227,8 @@ def custom_prompt( def prompt_factory( model: str, messages: list, - custom_llm_provider: Optional[str] = None, - api_key: Optional[str] = None, + custom_llm_provider: str | None = None, + api_key: str | None = None, ): original_model_name = model model = model.lower() @@ -5389,12 +5339,12 @@ def get_attribute_or_key(tool_or_function, attribute, default=None): class NormalizedToolCall(TypedDict): - id: Optional[str] - name: Optional[str] + id: str | None + name: str | None arguments: dict[str, Any] -def _parse_tool_call_arguments(raw: Any, tool_name: Optional[str], context: str) -> dict[str, Any]: +def _parse_tool_call_arguments(raw: Any, tool_name: str | None, context: str) -> dict[str, Any]: # Anthropic's tool_use blocks already carry a parsed dict in "input"; # chat completions and the Responses API carry a JSON string that may be # truncated by the model, so route those through the repair-aware parser. diff --git a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py index fc8a0d28583..1b960b84058 100644 --- a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py +++ b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py @@ -1,6 +1,6 @@ import json from datetime import datetime -from typing import Any, Dict, Union +from typing import Any from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, @@ -22,7 +22,7 @@ def strftime_now(fmt: str) -> str: return datetime.now().strftime(fmt) -def _get_tokenizer_config(hf_model_name: str) -> Dict[str, Any]: +def _get_tokenizer_config(hf_model_name: str) -> dict[str, Any]: """ Fetch tokenizer_config.json from HuggingFace (sync) @@ -45,7 +45,7 @@ def _get_tokenizer_config(hf_model_name: str) -> Dict[str, Any]: return {"status": "failure"} -async def _aget_tokenizer_config(hf_model_name: str) -> Dict[str, Any]: +async def _aget_tokenizer_config(hf_model_name: str) -> dict[str, Any]: """ Fetch tokenizer_config.json from HuggingFace (async) @@ -70,7 +70,7 @@ async def _aget_tokenizer_config(hf_model_name: str) -> Dict[str, Any]: return {"status": "failure"} -def _get_chat_template_file(hf_model_name: str) -> Dict[str, Any]: +def _get_chat_template_file(hf_model_name: str) -> dict[str, Any]: """ Fetch chat template from separate .jinja file (sync) @@ -98,7 +98,7 @@ def _get_chat_template_file(hf_model_name: str) -> Dict[str, Any]: return {"status": "failure"} -async def _aget_chat_template_file(hf_model_name: str) -> Dict[str, Any]: +async def _aget_chat_template_file(hf_model_name: str) -> dict[str, Any]: """ Fetch chat template from separate .jinja file (async) @@ -128,7 +128,7 @@ async def _aget_chat_template_file(hf_model_name: str) -> Dict[str, Any]: return {"status": "failure"} -def _extract_token_value(token_value: Union[None, str, Dict[str, Any]]) -> str: +def _extract_token_value(token_value: None | str | dict[str, Any]) -> str: """ Extract token string from various formats (string, dict, etc.) diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index 92a4296c432..a71fce974c1 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -120,7 +120,6 @@ def convert_url_to_base64(url: str) -> str: raise except Exception as e: verbose_logger.exception(e) - pass raise litellm.ImageFetchError( f"Error: Unable to fetch image from URL after 3 attempts. url={url}", ) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 220d1caa3d2..c401741f9dc 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,6 +1,6 @@ import asyncio import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, Union, cast +from typing import TYPE_CHECKING, Any, Protocol, cast import litellm from litellm._logging import verbose_logger @@ -46,22 +46,22 @@ class RealTimeStreaming: websocket: Any, backend_ws: CLIENT_CONNECTION_CLASS, logging_obj: LiteLLMLogging, - provider_config: Optional[BaseRealtimeConfig] = None, + provider_config: BaseRealtimeConfig | None = None, model: str = "", - user_api_key_dict: Optional[Any] = None, - request_data: Optional[Dict] = None, - backend_uses_beta_protocol: Optional[bool] = None, - force_transcription_model: Optional[str] = None, - event_normalizer: Optional[RealtimeEventNormalizer] = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, + backend_uses_beta_protocol: bool | None = None, + force_transcription_model: str | None = None, + event_normalizer: RealtimeEventNormalizer | None = None, ): self.websocket = websocket self.backend_ws = backend_ws self.logging_obj = logging_obj - self.messages: List[OpenAIRealtimeEvents] = [] - self.input_message: Dict = {} - self.input_messages: List[Dict[str, str]] = [] - self.session_tools: List[Dict] = [] - self.tool_calls: List[Dict] = [] + self.messages: list[OpenAIRealtimeEvents] = [] + self.input_message: dict = {} + self.input_messages: list[dict[str, str]] = [] + self.session_tools: list[dict] = [] + self.tool_calls: list[dict] = [] # Detect whether the client is explicitly opting into the beta protocol. self._client_wants_beta = self._detect_beta_header(websocket) @@ -76,20 +76,20 @@ class RealTimeStreaming: self.logged_real_time_event_types = _logged_real_time_event_types self.provider_config = provider_config self.model = model - self.current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] = None - self.current_output_item_id: Optional[str] = None - self.current_response_id: Optional[str] = None - self.current_conversation_id: Optional[str] = None - self.current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] = None - self.current_delta_type: Optional[ALL_DELTA_TYPES] = None - self.session_configuration_request: Optional[str] = None + self.current_delta_chunks: list[OpenAIRealtimeResponseDelta] | None = None + self.current_output_item_id: str | None = None + self.current_response_id: str | None = None + self.current_conversation_id: str | None = None + self.current_item_chunks: list[OpenAIRealtimeOutputItemDone] | None = None + self.current_delta_type: ALL_DELTA_TYPES | None = None + self.session_configuration_request: str | None = None self.user_api_key_dict = user_api_key_dict - self.request_data: Dict = request_data or {} + self.request_data: dict = request_data or {} # Violation counter for end_session_after_n_fails support self._violation_count: int = 0 # When a text message is blocked, hold the guardrail reason so the next # response.create can be rewritten to include the failure context. - self._pending_guardrail_message: Optional[str] = None + self._pending_guardrail_message: str | None = None # Track whether session.created has already been sent to the client # (e.g. synthetic event in deferred setup mode). self._session_created_sent_to_client: bool = False @@ -100,7 +100,7 @@ class RealTimeStreaming: # Buffer client audio until the backend acknowledges setup (setupComplete). self._backend_setup_complete: bool = provider_config is None or provider_config.requires_session_configuration() self._flushing_pending_messages_until_setup: bool = False - self._pending_messages_until_setup: List[str] = [] + self._pending_messages_until_setup: list[str] = [] self._pending_messages_byte_total: int = 0 # Gemini Live rejects a follow-up BidiGenerateContentSetup once any # content (realtimeInput / clientContent / toolResponse) has been sent. @@ -127,13 +127,13 @@ class RealTimeStreaming: ] ) _CLIENT_AUDIO_BUFFER_COMMIT_TYPES = frozenset(["input_audio_buffer.commit", "input_audio_buffer.end"]) - _AUDIO_FORMAT_MAP: Dict[str, Dict[str, Any]] = { + _AUDIO_FORMAT_MAP: dict[str, dict[str, Any]] = { "pcm16": {"type": "audio/pcm", "rate": 24000}, "g711_ulaw": {"type": "audio/G711-ulaw", "rate": 8000}, "g711_alaw": {"type": "audio/G711-alaw", "rate": 8000}, } # GA name → beta name (when client WebSocket includes OpenAI-Beta: realtime=v1) - _GA_TO_BETA_EVENT_TYPES: Dict[str, str] = { + _GA_TO_BETA_EVENT_TYPES: dict[str, str] = { "conversation.item.added": "conversation.item.created", "response.output_text.delta": "response.text.delta", "response.output_audio.delta": "response.audio.delta", @@ -142,14 +142,14 @@ class RealTimeStreaming: "response.output_audio.done": "response.audio.done", "response.output_audio_transcript.done": "response.audio_transcript.done", } - _GA_TO_BETA_CONTENT_TYPES: Dict[str, str] = { + _GA_TO_BETA_CONTENT_TYPES: dict[str, str] = { "output_text": "text", "output_audio": "audio", } def _should_store_message( self, - message_obj: Union[dict, OpenAIRealtimeEvents], + message_obj: dict | OpenAIRealtimeEvents, ) -> bool: _msg_type = message_obj["type"] if "type" in message_obj else None if self.logged_real_time_event_types == "*": @@ -158,15 +158,15 @@ class RealTimeStreaming: return True return False - def store_message(self, message: Union[str, bytes, dict, OpenAIRealtimeEvents]): + def store_message(self, message: str | bytes | dict | OpenAIRealtimeEvents): """Store message in list""" if isinstance(message, bytes): message = message.decode("utf-8") if isinstance(message, dict): # TypedDict union members do not narrow to plain dict for mypy. - message_obj: Dict[str, Any] = cast(Dict[str, Any], message) + message_obj: dict[str, Any] = cast(dict[str, Any], message) else: - message_obj = cast(Dict[str, Any], json.loads(cast(str, message))) + message_obj = cast(dict[str, Any], json.loads(cast(str, message))) self._collect_tool_calls_from_response_done(cast(dict, message_obj)) if not self._should_store_message(message_obj): return @@ -183,7 +183,7 @@ class RealTimeStreaming: return self.messages.append(typed_obj) - def _collect_user_input_from_client_event(self, message: Union[str, dict]) -> None: + def _collect_user_input_from_client_event(self, message: str | dict) -> None: """Extract user text content from client WebSocket events for spend logging.""" try: if isinstance(message, str): @@ -219,7 +219,7 @@ class RealTimeStreaming: except (json.JSONDecodeError, AttributeError, TypeError): pass - def _collect_user_input_from_backend_event(self, event_obj: Union[dict, OpenAIRealtimeEvents]) -> None: + def _collect_user_input_from_backend_event(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """Extract user voice transcription from backend events for spend logging.""" try: event_type = event_obj.get("type", "") @@ -230,7 +230,7 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass - def _detect_transcription_session_from_backend(self, event_obj: Union[dict, OpenAIRealtimeEvents]) -> None: + def _detect_transcription_session_from_backend(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """Flag transcription-only sessions from backend session events.""" try: event_type = event_obj.get("type", "") @@ -246,7 +246,7 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass - def _capture_transcription_usage(self, event_obj: Union[dict, OpenAIRealtimeEvents]) -> None: + def _capture_transcription_usage(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """ Append a usage-only transcription completed event to the logged results so the cost calculator can bill it by audio duration. The default logged event @@ -275,12 +275,12 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass - def _collect_tool_calls_from_response_done(self, event_obj: Union[dict, OpenAIRealtimeEvents]) -> None: + def _collect_tool_calls_from_response_done(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """Extract function_call items from response.done events for spend logging.""" try: if event_obj.get("type") != "response.done": return - response = cast(Dict[str, Any], event_obj.get("response", {})) + response = cast(dict[str, Any], event_obj.get("response", {})) for item in response.get("output", []): if item.get("type") == "function_call": self.tool_calls.append( @@ -296,7 +296,7 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass - def store_input(self, message: Union[str, dict]): + def store_input(self, message: str | dict): """Store input message""" self.input_message = message if isinstance(message, dict) else {} self._collect_user_input_from_client_event(message) @@ -441,15 +441,15 @@ class RealTimeStreaming: return not self.provider_config.requires_session_configuration() @staticmethod - def _collapse_buffered_audio_messages(messages: List[str]) -> List[str]: + def _collapse_buffered_audio_messages(messages: list[str]) -> list[str]: """Apply ``input_audio_buffer.clear`` semantics before replaying buffered frames. During deferred Gemini Live setup, ``clear`` is buffered alongside appends. On flush each append becomes a provider ``realtimeInput``; ``clear`` must drop preceding uncommitted appends instead of being forwarded as a no-op. """ - collapsed: List[str] = [] - pending_appends: List[str] = [] + collapsed: list[str] = [] + pending_appends: list[str] = [] for message in messages: try: @@ -595,12 +595,12 @@ class RealTimeStreaming: def _make_disable_auto_response_message(self) -> str: """Return a session.update that disables VAD auto-response.""" - turn_detection: Dict[str, Any] = { + turn_detection: dict[str, Any] = { "type": "server_vad", "create_response": False, } if self._backend_uses_beta_protocol: - session: Dict[str, Any] = {"turn_detection": turn_detection} + session: dict[str, Any] = {"turn_detection": turn_detection} else: session = { "type": "realtime", @@ -654,7 +654,7 @@ class RealTimeStreaming: def _has_realtime_guardrails_for_event_hooks( self, - event_hooks: List[Any], + event_hooks: list[Any], ) -> bool: """Return True if any callback would run for one of ``event_hooks``.""" from litellm.integrations.custom_guardrail import CustomGuardrail @@ -697,9 +697,9 @@ class RealTimeStreaming: async def run_realtime_guardrails( self, transcript: str, - item_id: Optional[str] = None, - pre_block_backend_message: Optional[str] = None, - event_hooks: Optional[List[Any]] = None, + item_id: str | None = None, + pre_block_backend_message: str | None = None, + event_hooks: list[Any] | None = None, ) -> bool: """ Run registered guardrails on realtime text (transcript, user message, tool output). @@ -807,7 +807,7 @@ class RealTimeStreaming: await self._send_to_backend(json.dumps({"type": "response.create"})) self._violation_count += 1 - end_session_after: Optional[int] = getattr(callback, "end_session_after_n_fails", None) + end_session_after: int | None = getattr(callback, "end_session_after_n_fails", None) should_end = getattr(callback, "on_violation", None) == "end_session" or ( end_session_after is not None and self._violation_count >= end_session_after ) @@ -900,7 +900,7 @@ class RealTimeStreaming: await self._send_event_to_client(event, event_str) blocked = await self.run_realtime_guardrails( cast(str, transcript), - item_id=cast(Optional[str], event.get("item_id")), + item_id=cast(str | None, event.get("item_id")), ) if not blocked: await self._send_to_backend(json.dumps({"type": "response.create"})) @@ -910,7 +910,7 @@ class RealTimeStreaming: await self._send_event_to_client(event, event_str) @staticmethod - def _parse_backend_event(raw_response: str) -> Optional[dict]: + def _parse_backend_event(raw_response: str) -> dict | None: """Parse a backend frame once. Returns None for non-JSON or non-object frames.""" try: event = json.loads(raw_response) @@ -1073,9 +1073,9 @@ class RealTimeStreaming: session["output_modalities"] = ["text"] # 3-7. Lift flat audio fields into the nested audio object - audio: Dict[str, Any] = {} - inp: Dict[str, Any] = {} - out: Dict[str, Any] = {} + audio: dict[str, Any] = {} + inp: dict[str, Any] = {} + out: dict[str, Any] = {} # voice → audio.output.voice if "voice" in session: @@ -1118,7 +1118,7 @@ class RealTimeStreaming: return session @staticmethod - def _translate_event_to_beta(event: dict) -> Optional[dict]: + def _translate_event_to_beta(event: dict) -> dict | None: """Translate a single GA event dict to its beta equivalent. Returns None when the event must be dropped (the GA-only @@ -1174,7 +1174,7 @@ class RealTimeStreaming: ## GUARDRAIL: intercept conversation.item.create for text-based injection. guardrail_turn_detection_injected = False - msg_type: Optional[str] = None + msg_type: str | None = None try: from litellm.types.guardrails import GuardrailEventHooks diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 43181e7f5ff..117d891f156 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -10,7 +10,7 @@ import asyncio import copy import inspect -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import litellm from litellm.integrations.custom_logger import CustomLogger @@ -338,7 +338,7 @@ def should_redact_message_logging(model_call_details: dict) -> bool: return litellm.turn_off_message_logging is True -def redact_message_input_output_from_logging(model_call_details: dict, result, input: Optional[Any] = None) -> Any: +def redact_message_input_output_from_logging(model_call_details: dict, result, input: Any | None = None) -> Any: """ Removes messages, prompts, input, response from logging. This modifies the data in-place only redacts when litellm.turn_off_message_logging == True @@ -350,13 +350,13 @@ def redact_message_input_output_from_logging(model_call_details: dict, result, i def _get_turn_off_message_logging_from_dynamic_params( model_call_details: dict, -) -> Optional[bool]: +) -> bool | None: """ gets the value of `turn_off_message_logging` from the dynamic params, if it exists. handles boolean and string values of `turn_off_message_logging` """ - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = model_call_details.get( + standard_callback_dynamic_params: StandardCallbackDynamicParams | None = model_call_details.get( "standard_callback_dynamic_params", None ) if standard_callback_dynamic_params: diff --git a/litellm/litellm_core_utils/request_timeout_resolver.py b/litellm/litellm_core_utils/request_timeout_resolver.py index 146c39ce9f3..c6d9fa1afc2 100644 --- a/litellm/litellm_core_utils/request_timeout_resolver.py +++ b/litellm/litellm_core_utils/request_timeout_resolver.py @@ -12,12 +12,10 @@ tell "user asked for this" from "nobody set it". This resolver answers that: from __future__ import annotations -from typing import Optional - from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS -def get_configured_request_timeout() -> Optional[float]: +def get_configured_request_timeout() -> float | None: """Return the explicitly-configured ``litellm.request_timeout``, else ``None``.""" import litellm diff --git a/litellm/litellm_core_utils/rules.py b/litellm/litellm_core_utils/rules.py index 425c3a80e26..82edc39a799 100644 --- a/litellm/litellm_core_utils/rules.py +++ b/litellm/litellm_core_utils/rules.py @@ -1,5 +1,3 @@ -from typing import Optional - import litellm @@ -40,7 +38,7 @@ class Rules: ) # type: ignore return True - def post_call_rules(self, input: Optional[str], model: str) -> bool: + def post_call_rules(self, input: str | None, model: str) -> bool: if input is None: return True for rule in litellm.post_call_rules: diff --git a/litellm/litellm_core_utils/safe_json_dumps.py b/litellm/litellm_core_utils/safe_json_dumps.py index 81cd8e57798..d26921f679b 100644 --- a/litellm/litellm_core_utils/safe_json_dumps.py +++ b/litellm/litellm_core_utils/safe_json_dumps.py @@ -1,5 +1,5 @@ import json -from typing import Any, Union +from typing import Any from pydantic import BaseModel @@ -31,7 +31,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: if id(obj) in seen: return "CircularReference Detected" seen.add(id(obj)) - result: Union[dict, list, tuple, set, str] + result: dict | list | tuple | set | str if isinstance(obj, dict): result = {} for k, v in obj.items(): diff --git a/litellm/litellm_core_utils/safe_json_loads.py b/litellm/litellm_core_utils/safe_json_loads.py index b0a8e57d552..bb973064134 100644 --- a/litellm/litellm_core_utils/safe_json_loads.py +++ b/litellm/litellm_core_utils/safe_json_loads.py @@ -2,8 +2,8 @@ Helper for safe JSON loading in LiteLLM. """ -from typing import Any import json +from typing import Any def safe_json_loads(data: str, default: Any = None) -> Any: diff --git a/litellm/litellm_core_utils/secret_redaction.py b/litellm/litellm_core_utils/secret_redaction.py index 455d0f00c35..be86fec6c16 100644 --- a/litellm/litellm_core_utils/secret_redaction.py +++ b/litellm/litellm_core_utils/secret_redaction.py @@ -7,7 +7,6 @@ secrets from strings without depending on the logging-configuration module. """ import re -from typing import List from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH @@ -15,7 +14,7 @@ _REDACTED = "REDACTED" def _build_secret_patterns() -> "re.Pattern[str]": - patterns: List[str] = [ + patterns: list[str] = [ # PEM private key / certificate blocks r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----", # GCP OAuth2 access tokens (ya29.*) diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 7861e13bae5..4be60bef3e0 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -1,5 +1,5 @@ from collections.abc import Mapping -from typing import Any, Dict, List, Optional, Set +from typing import Any from pydantic import BaseModel @@ -9,8 +9,8 @@ from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER class SensitiveDataMasker: def __init__( self, - sensitive_patterns: Optional[Set[str]] = None, - non_sensitive_overrides: Optional[Set[str]] = None, + sensitive_patterns: set[str] | None = None, + non_sensitive_overrides: set[str] | None = None, visible_prefix: int = 4, visible_suffix: int = 4, mask_char: str = "*", @@ -60,7 +60,7 @@ class SensitiveDataMasker: f"{value_str[: self.visible_prefix]}{self.mask_char * masked_length}{value_str[-self.visible_suffix :]}" ) - def is_sensitive_key(self, key: str, excluded_keys: Optional[Set[str]] = None) -> bool: + def is_sensitive_key(self, key: str, excluded_keys: set[str] | None = None) -> bool: # Check if key is in excluded_keys first (exact match) if excluded_keys and key in excluded_keys: return False @@ -82,13 +82,13 @@ class SensitiveDataMasker: def _mask_sequence( self, - values: List[Any], + values: list[Any], depth: int, max_depth: int, - excluded_keys: Optional[Set[str]], + excluded_keys: set[str] | None, key_is_sensitive: bool, - ) -> List[Any]: - masked_items: List[Any] = [] + ) -> list[Any]: + masked_items: list[Any] = [] if depth >= max_depth: return values @@ -105,15 +105,15 @@ class SensitiveDataMasker: def mask_dict( self, - data: Dict[str, Any], + data: dict[str, Any], depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER, - excluded_keys: Optional[Set[str]] = None, - ) -> Dict[str, Any]: + excluded_keys: set[str] | None = None, + ) -> dict[str, Any]: if depth >= max_depth: return data - masked_data: Dict[str, Any] = {} + masked_data: dict[str, Any] = {} for k, v in data.items(): try: key_is_sensitive = self.is_sensitive_key(k, excluded_keys) @@ -188,7 +188,7 @@ def _walk_payload(node: object, key_is_sensitive: bool, depth: int) -> object: return node -def mask_sensitive_keys(data: Dict[str, Any], sensitive_fields: Set[str]) -> Dict[str, Any]: +def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dict[str, Any]: """Return a new dict with values masked for keys listed in ``sensitive_fields``. Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name @@ -200,7 +200,7 @@ def mask_sensitive_keys(data: Dict[str, Any], sensitive_fields: Set[str]) -> Dic range and are replaced with a fixed-length all-mask string, so a short credential is never returned verbatim. """ - masked: Dict[str, Any] = {} + masked: dict[str, Any] = {} mask_char = _default_masker.mask_char min_visible = _default_masker.visible_prefix + _default_masker.visible_suffix for key, value in data.items(): diff --git a/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py b/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py index e71f64bc900..2286da0cedf 100644 --- a/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py +++ b/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py @@ -10,7 +10,7 @@ This ensures we do import hashlib import json -from typing import Any, Optional +from typing import Any import litellm from litellm.constants import _DEFAULT_TTL_FOR_HTTPX_CLIENTS @@ -67,7 +67,7 @@ class DynamicLoggingCache: cache_key = hashlib.sha256(args_str.encode("utf-8")).hexdigest() return cache_key - def get_cache(self, credentials: dict, service_name: str) -> Optional[Any]: + def get_cache(self, credentials: dict, service_name: str) -> Any | None: key_name = self.get_cache_key(args={**credentials, "service_name": service_name}) response = self.cache.get_cache(key=key_name) return response @@ -75,4 +75,3 @@ class DynamicLoggingCache: def set_cache(self, credentials: dict, service_name: str, logging_obj: Any) -> None: key_name = self.get_cache_key(args={**credentials, "service_name": service_name}) self.cache.set_cache(key=key_name, value=logging_obj) - return None diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index d52d9849310..2e62a151f98 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -1,7 +1,8 @@ import base64 import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Union, cast +from litellm._logging import verbose_logger from litellm.types.llms.openai import ( ChatCompletionAssistantContentValue, ChatCompletionAudioDelta, @@ -21,7 +22,6 @@ from litellm.types.utils import ( ServerToolUse, Usage, ) -from litellm._logging import verbose_logger from litellm.utils import print_verbose, token_counter if TYPE_CHECKING: @@ -35,7 +35,7 @@ if TYPE_CHECKING: class ChunkProcessor: - def __init__(self, chunks: List, messages: Optional[list] = None): + def __init__(self, chunks: list, messages: list | None = None): self.chunks = self._sort_chunks(chunks) self.messages = messages self.first_chunk = chunks[0] @@ -45,7 +45,7 @@ class ChunkProcessor: return [] first_chunk = chunks[0] - first_hidden_params: Dict[str, Any] = {} + first_hidden_params: dict[str, Any] = {} if isinstance(first_chunk, dict): candidate = first_chunk.get("_hidden_params", {}) if isinstance(candidate, dict): @@ -57,20 +57,20 @@ class ChunkProcessor: if first_hidden_params.get("created_at"): - def _created_at(chunk: Any) -> Union[int, float]: + def _created_at(chunk: Any) -> int | float: if isinstance(chunk, dict): params = chunk.get("_hidden_params", {}) else: params = getattr(chunk, "_hidden_params", {}) if isinstance(params, dict): - return cast(Union[int, float], params.get("created_at", float("inf"))) + return cast(int | float, params.get("created_at", float("inf"))) return float("inf") return sorted(chunks, key=_created_at) return chunks def update_model_response_with_hidden_params( - self, model_response: ModelResponse, chunk: Optional[Dict[str, Any]] = None + self, model_response: ModelResponse, chunk: dict[str, Any] | None = None ) -> ModelResponse: if chunk is None: return model_response @@ -82,8 +82,8 @@ class ChunkProcessor: @staticmethod def apply_provider_assembled_streaming_metadata( response: ModelResponse, - chunks: List[Any], - logging_obj: Optional[Any] = None, + chunks: list[Any], + logging_obj: Any | None = None, ) -> None: if not chunks: return @@ -126,7 +126,7 @@ class ChunkProcessor: ) @staticmethod - def _get_chunk_id(chunks: List[Dict[str, Any]]) -> str: + def _get_chunk_id(chunks: list[dict[str, Any]]) -> str: """ Chunks: [{"id": ""}, {"id": "1"}, {"id": "1"}] @@ -137,7 +137,7 @@ class ChunkProcessor: return "" @staticmethod - def _get_model_from_chunks(chunks: List[Dict[str, Any]], first_chunk_model: str) -> str: + def _get_model_from_chunks(chunks: list[dict[str, Any]], first_chunk_model: str) -> str: """ Get the actual model from chunks, preferring a model that differs from the first chunk. @@ -153,7 +153,7 @@ class ChunkProcessor: # Fall back to first chunk's model if no different model found return first_chunk_model - def build_base_response(self, chunks: List[Dict[str, Any]]) -> ModelResponse: + def build_base_response(self, chunks: list[dict[str, Any]]) -> ModelResponse: chunk = self.first_chunk id = ChunkProcessor._get_chunk_id(chunks) object = chunk["object"] @@ -202,9 +202,9 @@ class ChunkProcessor: response = self.update_model_response_with_hidden_params(model_response=response, chunk=chunk) return response - def get_combined_tool_content(self, tool_call_chunks: List[Dict[str, Any]]) -> List[ChatCompletionMessageToolCall]: - tool_calls_list: List[ChatCompletionMessageToolCall] = [] - tool_call_map: Dict[int, Dict[str, Any]] = {} # Map to store tool calls by index + def get_combined_tool_content(self, tool_call_chunks: list[dict[str, Any]]) -> list[ChatCompletionMessageToolCall]: + tool_calls_list: list[ChatCompletionMessageToolCall] = [] + tool_call_map: dict[int, dict[str, Any]] = {} # Map to store tool calls by index for chunk in tool_call_chunks: choices = chunk["choices"] @@ -324,7 +324,7 @@ class ChunkProcessor: return tool_calls_list - def get_combined_function_call_content(self, function_call_chunks: List[Dict[str, Any]]) -> FunctionCall: + def get_combined_function_call_content(self, function_call_chunks: list[dict[str, Any]]) -> FunctionCall: argument_list = [] delta = function_call_chunks[0]["choices"][0]["delta"] function_call = delta.get("function_call", "") @@ -350,9 +350,9 @@ class ChunkProcessor: ) def get_combined_content( - self, chunks: List[Dict[str, Any]], delta_key: str = "content" + self, chunks: list[dict[str, Any]], delta_key: str = "content" ) -> ChatCompletionAssistantContentValue: - content_list: List[str] = [] + content_list: list[str] = [] for chunk in chunks: choices = chunk["choices"] for choice in choices: @@ -369,16 +369,16 @@ class ChunkProcessor: return combined_content def get_combined_thinking_content( - self, chunks: List[Dict[str, Any]] - ) -> Optional[List[Union["ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]]]: + self, chunks: list[dict[str, Any]] + ) -> list[Union["ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]] | None: from litellm.types.llms.openai import ( ChatCompletionRedactedThinkingBlock, ChatCompletionThinkingBlock, ) - thinking_blocks: List[Union["ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]] = [] - current_thinking_text_parts: List[str] = [] - current_signature: Optional[str] = None + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = [] + current_thinking_text_parts: list[str] = [] + current_signature: str | None = None def _flush_thinking_block() -> None: nonlocal current_thinking_text_parts, current_signature @@ -426,20 +426,20 @@ class ChunkProcessor: return thinking_blocks return None - def get_combined_reasoning_content(self, chunks: List[Dict[str, Any]]) -> ChatCompletionAssistantContentValue: + def get_combined_reasoning_content(self, chunks: list[dict[str, Any]]) -> ChatCompletionAssistantContentValue: return self.get_combined_content(chunks, delta_key="reasoning_content") - def get_combined_audio_content(self, chunks: List[Dict[str, Any]]) -> ChatCompletionAudioResponse: - base64_data_list: List[str] = [] - transcript_list: List[str] = [] - expires_at: Optional[int] = None - id: Optional[str] = None + def get_combined_audio_content(self, chunks: list[dict[str, Any]]) -> ChatCompletionAudioResponse: + base64_data_list: list[str] = [] + transcript_list: list[str] = [] + expires_at: int | None = None + id: str | None = None for chunk in chunks: choices = chunk["choices"] for choice in choices: delta = choice.get("delta") or {} - audio: Optional[ChatCompletionAudioDelta] = delta.get("audio") + audio: ChatCompletionAudioDelta | None = delta.get("audio") if audio is not None: for k, v in audio.items(): if k == "data" and v is not None and isinstance(v, str): @@ -463,11 +463,11 @@ class ChunkProcessor: prompt_tokens = 0 completion_tokens = 0 ## anthropic prompt caching information ## - cache_creation_input_tokens: Optional[int] = None - cache_read_input_tokens: Optional[int] = None - completion_tokens_details: Optional[CompletionTokensDetails] = None - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None - cost: Optional[float] = None + cache_creation_input_tokens: int | None = None + cache_read_input_tokens: int | None = None + completion_tokens_details: CompletionTokensDetails | None = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None + cost: float | None = None if "prompt_tokens" in usage_chunk: prompt_tokens = usage_chunk.get("prompt_tokens", 0) or 0 @@ -500,8 +500,8 @@ class ChunkProcessor: "cost": cost, } - def count_reasoning_tokens(self, response: ModelResponse) -> Optional[int]: - reasoning_tokens: Optional[int] = None + def count_reasoning_tokens(self, response: ModelResponse) -> int | None: + reasoning_tokens: int | None = None for choice in response.choices: if ( hasattr(cast(Choices, choice).message, "reasoning_content") @@ -534,7 +534,7 @@ class ChunkProcessor: def _calculate_usage_per_chunk( self, - chunks: List[Union[Dict[str, Any], ModelResponse]], + chunks: list[dict[str, Any] | ModelResponse], ) -> "UsagePerChunk": from litellm.types.litellm_core_utils.streaming_chunk_builder_utils import ( UsagePerChunk, @@ -555,20 +555,20 @@ class ChunkProcessor: # arrived) from a stale lone cursor. completion_usage_updates = 0 ## anthropic prompt caching information ## - cache_creation_input_tokens: Optional[int] = None - cache_read_input_tokens: Optional[int] = None + cache_creation_input_tokens: int | None = None + cache_read_input_tokens: int | None = None - server_tool_use: Optional[ServerToolUse] = None - web_search_requests: Optional[int] = None - completion_tokens_details: Optional[CompletionTokensDetails] = None - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + server_tool_use: ServerToolUse | None = None + web_search_requests: int | None = None + completion_tokens_details: CompletionTokensDetails | None = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None # Anthropic emits the cache-creation TTL breakdown (5m/1h split) only on # the `message_start` event; the later `message_delta` carries the flat # cache-creation count but drops the nested breakdown. prompt_tokens_details # is last-wins, so without preserving this separately the 1h breakdown is # lost and 1h cache writes get billed at the 5m rate. - cache_creation_token_details: Optional[CacheCreationTokenDetails] = None - cost: Optional[float] = None + cache_creation_token_details: CacheCreationTokenDetails | None = None + cost: float | None = None for chunk in chunks: usage_chunk = self._extract_usage_chunk(chunk) @@ -616,7 +616,7 @@ class ChunkProcessor: ) prompt_tokens_details = cast( - Optional[PromptTokensDetailsWrapper], + PromptTokensDetailsWrapper | None, usage_chunk_dict["prompt_tokens_details"], ) @@ -651,11 +651,11 @@ class ChunkProcessor: @staticmethod def _capture_cache_creation_token_details( - prompt_tokens_details: Optional[PromptTokensDetailsWrapper], - current: Optional[CacheCreationTokenDetails], - ) -> Optional[CacheCreationTokenDetails]: + prompt_tokens_details: PromptTokensDetailsWrapper | None, + current: CacheCreationTokenDetails | None, + ) -> CacheCreationTokenDetails | None: incoming = cast( - Optional[CacheCreationTokenDetails], + CacheCreationTokenDetails | None, getattr(prompt_tokens_details, "cache_creation_token_details", None), ) if incoming is not None: @@ -664,13 +664,13 @@ class ChunkProcessor: @staticmethod def _attach_cache_creation_token_details( - prompt_tokens_details: Optional[PromptTokensDetailsWrapper], - cache_creation_token_details: Optional[CacheCreationTokenDetails], - ) -> Optional[PromptTokensDetailsWrapper]: + prompt_tokens_details: PromptTokensDetailsWrapper | None, + cache_creation_token_details: CacheCreationTokenDetails | None, + ) -> PromptTokensDetailsWrapper | None: if prompt_tokens_details is None or cache_creation_token_details is None: return prompt_tokens_details existing = cast( - Optional[CacheCreationTokenDetails], + CacheCreationTokenDetails | None, getattr(prompt_tokens_details, "cache_creation_token_details", None), ) if existing is not None: @@ -702,7 +702,7 @@ class ChunkProcessor: if saw_non_cursor_completion: return completion_tokens - custom_llm_provider: Optional[str] = None + custom_llm_provider: str | None = None if chunks: first_chunk = chunks[0] if isinstance(first_chunk, dict): @@ -718,11 +718,11 @@ class ChunkProcessor: def calculate_usage( self, - chunks: List[Union[Dict[str, Any], ModelResponse]], + chunks: list[dict[str, Any] | ModelResponse], model: str, completion_output: str, - messages: Optional[List] = None, - reasoning_tokens: Optional[int] = None, + messages: list | None = None, + reasoning_tokens: int | None = None, ) -> Usage: """ Calculate usage for the given chunks. @@ -734,18 +734,16 @@ class ChunkProcessor: prompt_tokens = calculated_usage_per_chunk["prompt_tokens"] completion_tokens = calculated_usage_per_chunk["completion_tokens"] ## anthropic prompt caching information ## - cache_creation_input_tokens: Optional[int] = calculated_usage_per_chunk["cache_creation_input_tokens"] - cache_read_input_tokens: Optional[int] = calculated_usage_per_chunk["cache_read_input_tokens"] + cache_creation_input_tokens: int | None = calculated_usage_per_chunk["cache_creation_input_tokens"] + cache_read_input_tokens: int | None = calculated_usage_per_chunk["cache_read_input_tokens"] - server_tool_use: Optional[ServerToolUse] = calculated_usage_per_chunk["server_tool_use"] - web_search_requests: Optional[int] = calculated_usage_per_chunk["web_search_requests"] - completion_tokens_details: Optional[CompletionTokensDetails] = calculated_usage_per_chunk[ + server_tool_use: ServerToolUse | None = calculated_usage_per_chunk["server_tool_use"] + web_search_requests: int | None = calculated_usage_per_chunk["web_search_requests"] + completion_tokens_details: CompletionTokensDetails | None = calculated_usage_per_chunk[ "completion_tokens_details" ] - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = calculated_usage_per_chunk[ - "prompt_tokens_details" - ] - cost: Optional[float] = calculated_usage_per_chunk["cost"] + prompt_tokens_details: PromptTokensDetailsWrapper | None = calculated_usage_per_chunk["prompt_tokens_details"] + cost: float | None = calculated_usage_per_chunk["cost"] try: returned_usage.prompt_tokens = prompt_tokens or token_counter(model=model, messages=messages) @@ -813,7 +811,7 @@ class ChunkProcessor: return returned_usage -def concatenate_base64_list(base64_strings: List[str]) -> str: +def concatenate_base64_list(base64_strings: list[str]) -> str: """ Concatenates a list of base64-encoded strings. diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 0310d85b810..fb7d06bee93 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -10,10 +10,7 @@ from collections.abc import AsyncIterator, Callable, Iterator from dataclasses import dataclass from typing import ( Any, - Dict, - List, NoReturn, - Optional, Union, cast, ) @@ -113,10 +110,10 @@ class CustomStreamWrapper: completion_stream, model, logging_obj: Any, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, stream_options=None, - make_call: Optional[Callable] = None, - _response_headers: Optional[dict] = None, + make_call: Callable | None = None, + _response_headers: dict | None = None, ): self.model = model self.make_call = make_call @@ -135,9 +132,9 @@ class CustomStreamWrapper: self.sent_last_thinking_block = False self.thinking_content = "" - self.system_fingerprint: Optional[str] = None - self.received_finish_reason: Optional[str] = None - self.intermittent_finish_reason: Optional[str] = None # finish reasons that show up mid-stream + self.system_fingerprint: str | None = None + self.received_finish_reason: str | None = None + self.intermittent_finish_reason: str | None = None # finish reasons that show up mid-stream self.special_tokens = [ "<|assistant|>", "<|system|>", @@ -150,7 +147,7 @@ class CustomStreamWrapper: self.holding_chunk = "" self.complete_response = "" self.response_uptil_now = "" - _model_info: Dict = litellm_params.model_info or {} + _model_info: dict = litellm_params.model_info or {} _api_base = get_api_base( model=model or "", @@ -167,7 +164,7 @@ class CustomStreamWrapper: ) # GUARANTEE OPENAI HEADERS IN RESPONSE self._response_headers = _response_headers - self.response_id: Optional[str] = None + self.response_id: str | None = None self.logging_loop = None self.rules = Rules() self.stream_options = stream_options or getattr(logging_obj, "stream_options", None) @@ -175,28 +172,28 @@ class CustomStreamWrapper: self.sent_stream_usage = False self.send_stream_usage = True if self.check_send_stream_usage(self.stream_options) else False self.tool_call = False - self.chunks: List = [] # keep track of the returned chunks - used for calculating the input/output tokens for stream options + self.chunks: list = [] # keep track of the returned chunks - used for calculating the input/output tokens for stream options self._repeated_messages_count = 1 self.is_function_call = self.check_is_function_call(logging_obj=logging_obj) - self.created: Optional[int] = None - self._last_returned_hidden_params: Optional[dict] = None + self.created: int | None = None + self._last_returned_hidden_params: dict | None = None _cached_logging_provider = self.logging_obj.model_call_details.get("custom_llm_provider", None) - self._cached_logging_llm_provider: Optional[str] = _cached_logging_provider + self._cached_logging_llm_provider: str | None = _cached_logging_provider _effective_model = model or "" if custom_llm_provider == "openai" and custom_llm_provider != _cached_logging_provider: - _effective_model = "{}/{}".format(_cached_logging_provider, _effective_model) + _effective_model = f"{_cached_logging_provider}/{_effective_model}" self._cached_model_name: str = _effective_model # Snapshot assumes self._hidden_params is populated from litellm_params # at init and never mutated during the stream. If that ever changes, # this cache must be removed. - self._base_hidden_params: Dict[str, Any] = { + self._base_hidden_params: dict[str, Any] = { **self._hidden_params, "response_cost": None, } - self._post_streaming_hooks: Optional[List] = None + self._post_streaming_hooks: list | None = None def _check_max_streaming_duration(self) -> None: """Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS.""" @@ -239,7 +236,7 @@ class CustomStreamWrapper: e, ) - def check_send_stream_usage(self, stream_options: Optional[dict]): + def check_send_stream_usage(self, stream_options: dict | None): return stream_options is not None and stream_options.get("include_usage", False) is True def check_is_function_call(self, logging_obj) -> bool: @@ -305,12 +302,12 @@ class CustomStreamWrapper: if self._repeated_messages_count >= litellm.REPEATED_STREAMING_CHUNK_LIMIT: # All last n chunks are identical raise litellm.InternalServerError( - message="The model is repeating the same chunk = {}.".format(last_content), + message=f"The model is repeating the same chunk = {last_content}.", model="", llm_provider="", ) - def check_special_tokens(self, chunk: str, finish_reason: Optional[str]): + def check_special_tokens(self, chunk: str, finish_reason: str | None): """ Output parse / special tokens for sagemaker + hf streaming. """ @@ -621,9 +618,7 @@ class CustomStreamWrapper: else: return "" except Exception as e: - verbose_logger.exception( - "litellm.CustomStreamWrapper.handle_baseten_chunk(): Exception occured - {}".format(str(e)) - ) + verbose_logger.exception(f"litellm.CustomStreamWrapper.handle_baseten_chunk(): Exception occured - {e!s}") return "" def handle_triton_stream(self, chunk): @@ -661,12 +656,12 @@ class CustomStreamWrapper: except Exception as e: raise e - def model_response_creator(self, chunk: Optional[dict] = None, hidden_params: Optional[dict] = None): + def model_response_creator(self, chunk: dict | None = None, hidden_params: dict | None = None): _model = self._cached_model_name _logging_obj_llm_provider = self._cached_logging_llm_provider if chunk is None: - args: Dict[str, Any] = {"model": _model} + args: dict[str, Any] = {"model": _model} else: chunk.pop("model", None) args = {"model": _model} @@ -712,11 +707,7 @@ class CustomStreamWrapper: def is_delta_empty(self, delta: Delta) -> bool: is_empty = True - if delta.content: - is_empty = False - elif delta.tool_calls is not None: - is_empty = False - elif delta.function_call is not None: + if delta.content or delta.tool_calls is not None or delta.function_call is not None: is_empty = False return is_empty @@ -740,7 +731,7 @@ class CustomStreamWrapper: def copy_model_response_level_provider_specific_fields( self, - original_chunk: Union[ModelResponseStream, OpenAIChatCompletionChunk], + original_chunk: ModelResponseStream | OpenAIChatCompletionChunk, model_response: ModelResponseStream, ) -> ModelResponseStream: """ @@ -755,9 +746,9 @@ class CustomStreamWrapper: def is_chunk_non_empty( self, - completion_obj: Dict[str, Any], + completion_obj: dict[str, Any], model_response: ModelResponseStream, - response_obj: Dict[str, Any], + response_obj: dict[str, Any], ) -> bool: if ( "content" in completion_obj @@ -881,9 +872,9 @@ class CustomStreamWrapper: def return_processed_chunk_logic( # noqa: C901 self, - completion_obj: Dict[str, Any], + completion_obj: dict[str, Any], model_response: ModelResponseStream, - response_obj: Dict[str, Any], + response_obj: dict[str, Any], ): from litellm.litellm_core_utils.core_helpers import ( preserve_upstream_non_openai_attributes, @@ -943,7 +934,7 @@ class CustomStreamWrapper: if response_obj.get("provider_specific_fields") is not None: completion_obj["provider_specific_fields"] = response_obj["provider_specific_fields"] model_response.choices[0].delta = Delta(**completion_obj) - _index: Optional[int] = completion_obj.get("index") + _index: int | None = completion_obj.get("index") if _index is not None: model_response.choices[0].index = _index @@ -1036,7 +1027,6 @@ class CustomStreamWrapper: if hasattr(model_response.choices[0].delta, "reasoning_content"): del model_response.choices[0].delta.reasoning_content - return def _dispatch_provider_chunk( self, @@ -1189,7 +1179,7 @@ class CustomStreamWrapper: content=None, tool_calls=[ { - "id": f"call_{str(uuid.uuid4())}", + "id": f"call_{uuid.uuid4()!s}", "function": { "arguments": args_str, "name": function_call.name, @@ -1214,7 +1204,7 @@ class CustomStreamWrapper: ) except Exception: if chunk.candidates[0].finish_reason.name == "SAFETY": # type: ignore - raise Exception(f"The response was blocked by VertexAI. {str(chunk)}") + raise Exception(f"The response was blocked by VertexAI. {chunk!s}") else: completion_obj["content"] = str(chunk) elif self.custom_llm_provider == "petals": @@ -1330,9 +1320,7 @@ class CustomStreamWrapper: if response_obj["is_finished"]: if response_obj["finish_reason"] == "error": raise Exception( - "{} raised a streaming error - finish_reason: error, no content string given. Received Chunk={}".format( - self.custom_llm_provider, response_obj - ) + f"{self.custom_llm_provider} raised a streaming error - finish_reason: error, no content string given. Received Chunk={response_obj}" ) self.received_finish_reason = response_obj["finish_reason"] if response_obj.get("original_chunk", None) is not None: @@ -1442,7 +1430,7 @@ class CustomStreamWrapper: model_response.choices[0].delta = Delta(**_json_delta) except Exception as e: verbose_logger.exception( - "litellm.CustomStreamWrapper.chunk_creator(): Exception occured - {}".format(str(e)) + f"litellm.CustomStreamWrapper.chunk_creator(): Exception occured - {e!s}" ) model_response.choices[0].delta = Delta() elif self._has_any_special_delta_attributes(delta): @@ -1550,7 +1538,7 @@ class CustomStreamWrapper: except Exception as e: from litellm._logging import verbose_logger - verbose_logger.exception(f"Error in post-call streaming deployment hook: {str(e)}") + verbose_logger.exception(f"Error in post-call streaming deployment hook: {e!s}") return chunk def _add_mcp_list_tools_to_first_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: @@ -1590,7 +1578,7 @@ class CustomStreamWrapper: except Exception as e: from litellm._logging import verbose_logger - verbose_logger.exception(f"Error adding MCP list tools to first chunk: {str(e)}") + verbose_logger.exception(f"Error adding MCP list tools to first chunk: {e!s}") return chunk @@ -1627,7 +1615,7 @@ class CustomStreamWrapper: except Exception as e: from litellm._logging import verbose_logger - verbose_logger.exception(f"Error adding MCP metadata to final chunk: {str(e)}") + verbose_logger.exception(f"Error adding MCP metadata to final chunk: {e!s}") return chunk @@ -1724,7 +1712,7 @@ class CustomStreamWrapper: print_verbose( f"PROCESSED CHUNK PRE CHUNK CREATOR: {chunk.decode('utf-8', errors='replace') if isinstance(chunk, bytes) else chunk}; custom_llm_provider: {self.custom_llm_provider}" ) - response: Optional[ModelResponseStream] = self.chunk_creator(chunk=chunk) + response: ModelResponseStream | None = self.chunk_creator(chunk=chunk) print_verbose(f"PROCESSED CHUNK POST CHUNK CREATOR: {response}") if response is None: @@ -1912,7 +1900,7 @@ class CustomStreamWrapper: elif self.custom_llm_provider == "gemini" and hasattr(chunk, "parts") and len(chunk.parts) == 0: continue - processed_chunk: Optional[ModelResponseStream] = self.chunk_creator(chunk=chunk) + processed_chunk: ModelResponseStream | None = self.chunk_creator(chunk=chunk) if processed_chunk is None: continue @@ -2000,7 +1988,7 @@ class CustomStreamWrapper: except httpx.TimeoutException as e: # if httpx read timeout error occues traceback_exception = traceback.format_exc() ## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT - traceback_exception += "\nLiteLLM Default Request Timeout - {}".format(litellm.request_timeout) + traceback_exception += f"\nLiteLLM Default Request Timeout - {litellm.request_timeout}" if self.logging_obj is not None: self._record_partial_usage_for_failure() ## LOGGING @@ -2133,7 +2121,7 @@ class CustomStreamWrapper: return try: partial_response = litellm.stream_chunk_builder(chunks=self.chunks) - usage = cast(Optional[Usage], getattr(partial_response, "usage", None)) + usage = cast(Usage | None, getattr(partial_response, "usage", None)) if usage is None: return self.logging_obj.model_call_details["combined_usage_object"] = usage @@ -2174,7 +2162,7 @@ class CustomStreamWrapper: except Exception as mapping_error: mapped_exception = mapping_error - def _normalize_status_code(exc: Exception) -> Optional[int]: + def _normalize_status_code(exc: Exception) -> int | None: """Best-effort status_code extraction.""" try: code = getattr(exc, "status_code", None) @@ -2214,7 +2202,7 @@ class CustomStreamWrapper: ) @staticmethod - def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]: + def _strip_sse_data_from_chunk(chunk: str | None) -> str | None: """ Strips the 'data: ' prefix from Server-Sent Events (SSE) chunks. @@ -2250,7 +2238,7 @@ class CustomStreamWrapper: return chunk -def calculate_total_usage(chunks: List[ModelResponse]) -> Usage: +def calculate_total_usage(chunks: list[ModelResponse]) -> Usage: """Assume most recent usage chunk has total usage uptil then.""" prompt_tokens: int = 0 completion_tokens: int = 0 diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 939b8d2b60f..fbd19b43f3e 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -6,11 +6,7 @@ import struct from collections.abc import Callable, Mapping from typing import ( Any, - List, Literal, - Optional, - Tuple, - Union, cast, ) @@ -47,11 +43,11 @@ from litellm.types.utils import Message, SelectTokenizerResponse def get_modified_max_tokens( model: str, base_model: str, - messages: Optional[List[AllMessageValues]], - user_max_tokens: Optional[int], - buffer_perc: Optional[float], - buffer_num: Optional[float], -) -> Optional[int]: + messages: list[AllMessageValues] | None, + user_max_tokens: int | None, + buffer_perc: float | None, + buffer_num: float | None, +) -> int | None: """ Params: @@ -108,9 +104,7 @@ def get_modified_max_tokens( return user_max_tokens except Exception as e: verbose_logger.debug( - "litellm.litellm_core_utils.token_counter.py::get_modified_max_tokens() - Error while checking max token limit: {}\nmodel={}, base_model={}".format( - str(e), model, base_model - ) + f"litellm.litellm_core_utils.token_counter.py::get_modified_max_tokens() - Error while checking max token limit: {e!s}\nmodel={model}, base_model={base_model}" ) return user_max_tokens @@ -118,7 +112,7 @@ def get_modified_max_tokens( def resize_image_high_res( width: int, height: int, -) -> Tuple[int, int]: +) -> tuple[int, int]: # Maximum dimensions for high res mode max_short_side = MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES max_long_side = MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES @@ -166,7 +160,7 @@ def calculate_tiles_needed( return total_tiles -def get_image_type(image_data: bytes) -> Union[str, None]: +def get_image_type(image_data: bytes) -> str | None: """take an image (really only the first ~100 bytes max are needed) and return 'png' 'gif' 'jpeg' 'webp' 'heic' or None. method added to allow deprecation of imghdr in 3.13""" @@ -191,7 +185,7 @@ def get_image_type(image_data: bytes) -> Union[str, None]: def get_image_dimensions( data: str, -) -> Tuple[int, int]: +) -> tuple[int, int]: """ Async Function to get the dimensions of an image from a URL or base64 encoded string. @@ -286,7 +280,7 @@ def calculate_img_tokens( int: The number of tokens for the image. """ if use_default_image_token_count: - verbose_logger.debug("Using default image token count: {}".format(DEFAULT_IMAGE_TOKEN_COUNT)) + verbose_logger.debug(f"Using default image token count: {DEFAULT_IMAGE_TOKEN_COUNT}") return DEFAULT_IMAGE_TOKEN_COUNT if mode == "low" or mode == "auto": return base_tokens @@ -316,7 +310,7 @@ class _MessageCountParams: def __init__( self, model: str, - custom_tokenizer: Optional[Union[dict, SelectTokenizerResponse]], + custom_tokenizer: dict | SelectTokenizerResponse | None, ): from litellm.utils import print_verbose @@ -324,10 +318,7 @@ class _MessageCountParams: if actual_model == "gpt-3.5-turbo-0301": self.tokens_per_message = 4 # every message follows <|start|>{role/name}\n{content}<|end|>\n self.tokens_per_name = -1 # if there's a name, the role is omitted - elif actual_model in litellm.open_ai_chat_completion_models: - self.tokens_per_message = 3 - self.tokens_per_name = 1 - elif actual_model in litellm.azure_llms: + elif actual_model in litellm.open_ai_chat_completion_models or actual_model in litellm.azure_llms: self.tokens_per_message = 3 self.tokens_per_name = 1 else: @@ -339,14 +330,14 @@ class _MessageCountParams: def token_counter( model="", - custom_tokenizer: Optional[Union[dict, SelectTokenizerResponse]] = None, - text: Optional[Union[str, List[str]]] = None, - messages: Optional[List[Union[AllMessageValues, Message]]] = None, - count_response_tokens: Optional[bool] = False, - tools: Optional[List[ChatCompletionToolParam]] = None, - tool_choice: Optional[ChatCompletionNamedToolChoiceParam] = None, - use_default_image_token_count: Optional[bool] = False, - default_token_count: Optional[int] = None, + custom_tokenizer: dict | SelectTokenizerResponse | None = None, + text: str | list[str] | None = None, + messages: list[AllMessageValues | Message] | None = None, + count_response_tokens: bool | None = False, + tools: list[ChatCompletionToolParam] | None = None, + tool_choice: ChatCompletionNamedToolChoiceParam | None = None, + use_default_image_token_count: bool | None = False, + default_token_count: int | None = None, ) -> int: """ Count the number of tokens in a given text using a specified model. @@ -385,7 +376,7 @@ def token_counter( if text is not None: if tools or tool_choice: raise ValueError("tools or tool_choice cannot be set if using text") - if isinstance(text, List): + if isinstance(text, list): text_to_count = "".join(t for t in text if isinstance(t, str)) elif isinstance(text, str): text_to_count = text @@ -393,7 +384,7 @@ def token_counter( num_tokens = count_function(text_to_count) elif messages is not None: - new_messages = cast(List[AllMessageValues], convert_list_message_to_dict(messages)) + new_messages = cast(list[AllMessageValues], convert_list_message_to_dict(messages)) params = _MessageCountParams(model, custom_tokenizer) num_tokens = _count_messages(params, new_messages, use_default_image_token_count, default_token_count) if count_response_tokens is False: @@ -421,7 +412,7 @@ def _count_function_call_tokens( tool/function definitions and `tool_choice`. """ if key == "tool_calls": - if not isinstance(value, List): + if not isinstance(value, list): raise ValueError(f"Unsupported type {type(value)} for key tool_calls in message {message}") total = 0 for tool_call in value: @@ -439,9 +430,9 @@ def _count_function_call_tokens( def _count_messages( params: _MessageCountParams, - messages: List[AllMessageValues], + messages: list[AllMessageValues], use_default_image_token_count: bool, - default_token_count: Optional[int], + default_token_count: int | None, ) -> int: """ Count the number of tokens in a list of messages. @@ -466,7 +457,7 @@ def _count_messages( num_tokens += params.count_function(value) if key == "name": num_tokens += params.tokens_per_name - elif key == "content" and isinstance(value, List): + elif key == "content" and isinstance(value, list): num_tokens += _count_content_list( params.count_function, value, @@ -489,8 +480,8 @@ def _count_messages( def _count_extra( count_function: TokenCounterFunction, - tools: Optional[List[ChatCompletionToolParam]], - tool_choice: Optional[ChatCompletionNamedToolChoiceParam], + tools: list[ChatCompletionToolParam] | None, + tool_choice: ChatCompletionNamedToolChoiceParam | None, includes_system_message: bool, ) -> int: """Count extra tokens for function definitions and tool choices. @@ -522,8 +513,8 @@ def _count_extra( def _get_count_function( - model: Optional[str], - custom_tokenizer: Optional[Union[dict, SelectTokenizerResponse]] = None, + model: str | None, + custom_tokenizer: dict | SelectTokenizerResponse | None = None, ) -> TokenCounterFunction: """ Get the function to count tokens based on the model and custom tokenizer.""" @@ -643,7 +634,7 @@ def _count_anthropic_content( content: Mapping[str, Any], count_function: TokenCounterFunction, use_default_image_token_count: bool, - default_token_count: Optional[int], + default_token_count: int | None, ) -> int: """ Count tokens in Anthropic-specific content blocks (tool_use, tool_result, etc.). @@ -692,7 +683,7 @@ def _count_content_list( count_function: TokenCounterFunction, content_list: OpenAIMessageContent, use_default_image_token_count: bool, - default_token_count: Optional[int], + default_token_count: int | None, ) -> int: """ Recursively count tokens from a list of content blocks. diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index a83cb3bc69e..9ef6c43d9e5 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -21,7 +21,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config): import socket from ipaddress import ip_address, ip_network -from typing import Any, List, Optional, Set, Tuple +from typing import Any from urllib.parse import quote, urlparse, urlunparse import httpx @@ -43,8 +43,6 @@ _ALLOWED_SCHEMES = ("http", "https") class SSRFError(ValueError): """Raised when a URL targets a blocked network.""" - pass - def encode_url_path_segment(value: Any, *, field_name: str = "path parameter") -> str: """Percent-encode one user-controlled URL path segment. @@ -116,7 +114,7 @@ def _default_port_for_scheme(scheme: str) -> int: def _parse_url_destination_allowlist_entry( entry: str, -) -> Optional[Tuple[str, Optional[str], Optional[int]]]: +) -> tuple[str, str | None, int | None] | None: """Parse an admin allowlist entry into host, optional scheme, optional port. Entries may be bare hosts (``api.example.com``), host+port @@ -141,14 +139,14 @@ def _parse_url_destination_allowlist_entry( except ValueError: return None - scheme: Optional[str] = parsed.scheme if has_scheme else None + scheme: str | None = parsed.scheme if has_scheme else None if scheme is not None and port is None: port = _default_port_for_scheme(scheme) return _normalize_host(parsed.hostname), scheme, port -def provider_url_destination_candidates(value: str) -> Tuple[str, ...]: +def provider_url_destination_candidates(value: str) -> tuple[str, ...]: return tuple( candidate for part in value.split(",") @@ -157,7 +155,7 @@ def provider_url_destination_candidates(value: str) -> Tuple[str, ...]: ) -def is_url_destination_allowed_by_host(url: str, allowed_hosts: List[str]) -> bool: +def is_url_destination_allowed_by_host(url: str, allowed_hosts: list[str]) -> bool: """Return True when a credential-bearing provider URL is admin-allowlisted. This does not fetch, resolve, or rewrite URLs. It only answers whether the @@ -227,17 +225,17 @@ def _is_host_allowlisted(hostname: str, effective_port: int) -> bool: literals are written bracketed (``[::1]`` / ``[::1]:8080``). Matching is case-insensitive on the hostname. """ - configured: List[str] = getattr(litellm, "user_url_allowed_hosts", []) or [] + configured: list[str] = getattr(litellm, "user_url_allowed_hosts", []) or [] if not configured: return False normalized_host = _normalize_host(hostname) host_repr = f"[{normalized_host}]" if ":" in normalized_host else normalized_host - candidates: Set[str] = {host_repr, f"{host_repr}:{effective_port}"} - allowlist: Set[str] = {_normalize_host(entry) for entry in configured if entry} + candidates: set[str] = {host_repr, f"{host_repr}:{effective_port}"} + allowlist: set[str] = {_normalize_host(entry) for entry in configured if entry} return bool(candidates & allowlist) -def validate_url(url: str) -> Tuple[str, str]: +def validate_url(url: str) -> tuple[str, str]: """ Validate a user-supplied URL and rewrite it to connect to a validated IP. diff --git a/litellm/llms/__init__.py b/litellm/llms/__init__.py index 6aec359b7b7..60715ac9bbf 100644 --- a/litellm/llms/__init__.py +++ b/litellm/llms/__init__.py @@ -1,6 +1,6 @@ import importlib import os -from typing import TYPE_CHECKING, Dict, Optional, Type +from typing import TYPE_CHECKING from litellm._logging import verbose_logger from litellm.types.utils import CallTypes @@ -14,9 +14,7 @@ if TYPE_CHECKING: from litellm.types.utils import ModelInfo, Usage -def get_cost_for_web_search_request( - custom_llm_provider: str, usage: "Usage", model_info: "ModelInfo" -) -> Optional[float]: +def get_cost_for_web_search_request(custom_llm_provider: str, usage: "Usage", model_info: "ModelInfo") -> float | None: """ Get the cost for a web search request for a given model. @@ -61,7 +59,7 @@ def get_cost_for_web_search_request( return None -def discover_guardrail_translation_mappings() -> Dict[CallTypes, Type["BaseTranslation"]]: +def discover_guardrail_translation_mappings() -> dict[CallTypes, type["BaseTranslation"]]: """ Discover guardrail translation mappings by scanning the llms directory structure. @@ -70,7 +68,7 @@ def discover_guardrail_translation_mappings() -> Dict[CallTypes, Type["BaseTrans Returns: Dict[CallTypes, Type[BaseTranslation]]: A dictionary mapping call types to their translation handler classes """ - discovered_mappings: Dict[CallTypes, Type["BaseTranslation"]] = {} + discovered_mappings: dict[CallTypes, type[BaseTranslation]] = {} try: # Get the path to the llms directory @@ -138,7 +136,7 @@ def discover_guardrail_translation_mappings() -> Dict[CallTypes, Type["BaseTrans # Cache the discovered mappings -endpoint_guardrail_translation_mappings: Optional[Dict[CallTypes, Type["BaseTranslation"]]] = None +endpoint_guardrail_translation_mappings: dict[CallTypes, type["BaseTranslation"]] | None = None def load_guardrail_translation_mappings(): @@ -148,7 +146,7 @@ def load_guardrail_translation_mappings(): return endpoint_guardrail_translation_mappings -def get_guardrail_translation_mapping(call_type: CallTypes) -> Type["BaseTranslation"]: +def get_guardrail_translation_mapping(call_type: CallTypes) -> type["BaseTranslation"]: """ Get the guardrail translation handler for a given call type. diff --git a/litellm/llms/a2a/chat/guardrail_translation/handler.py b/litellm/llms/a2a/chat/guardrail_translation/handler.py index 740b0fff50c..9660a8fa367 100644 --- a/litellm/llms/a2a/chat/guardrail_translation/handler.py +++ b/litellm/llms/a2a/chat/guardrail_translation/handler.py @@ -11,7 +11,7 @@ A2A Protocol Format: """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -64,8 +64,8 @@ class A2AGuardrailHandler(BaseTranslation): verbose_proxy_logger.debug("A2A: No parts in message, skipping guardrail") return data - texts_to_check: List[str] = [] - text_part_indices: List[int] = [] # Track which parts contain text + texts_to_check: list[str] = [] + text_part_indices: list[int] = [] # Track which parts contain text # Step 1: Extract text from all text parts for part_idx, part in enumerate(parts): @@ -111,7 +111,7 @@ class A2AGuardrailHandler(BaseTranslation): guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, - request_data: Optional[dict] = None, + request_data: dict | None = None, ) -> Any: """ Process A2A output response by applying guardrails to text content. @@ -148,10 +148,10 @@ class A2AGuardrailHandler(BaseTranslation): return response # Find all text-containing parts in the response - texts_to_check: List[str] = [] + texts_to_check: list[str] = [] # Each mapping is (path_to_parts_list, part_index) # path_to_parts_list is a tuple of keys to navigate to the parts list - task_mappings: List[Tuple[Tuple[str, ...], int]] = [] + task_mappings: list[tuple[tuple[str, ...], int]] = [] # Extract texts from all possible locations self._extract_texts_from_result( @@ -214,12 +214,12 @@ class A2AGuardrailHandler(BaseTranslation): async def process_output_streaming_response( self, - responses_so_far: List[Any], + responses_so_far: list[Any], guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, - request_data: Optional[dict] = None, - ) -> List[Any]: + request_data: dict | None = None, + ) -> list[Any]: """ Process A2A streaming output by applying guardrails to accumulated text. @@ -262,14 +262,14 @@ class A2AGuardrailHandler(BaseTranslation): guardrailed_text = guardrailed_texts[0] # Find first chunk (by original index) that has text; put full guardrailed text there and clear rest - first_chunk_with_text: Optional[int] = chunk_indices_with_text[0] if chunk_indices_with_text else None + first_chunk_with_text: int | None = chunk_indices_with_text[0] if chunk_indices_with_text else None for orig_i, obj in valid_parsed: result = obj.get("result", {}) if not isinstance(result, dict): continue - texts_in_chunk: List[str] = [] - mappings: List[Tuple[Tuple[str, ...], int]] = [] + texts_in_chunk: list[str] = [] + mappings: list[tuple[tuple[str, ...], int]] = [] self._extract_texts_from_result( result=result, texts_to_check=texts_in_chunk, @@ -305,10 +305,10 @@ class A2AGuardrailHandler(BaseTranslation): def _parse_streaming_responses( self, - responses_so_far: List[Any], - ) -> Tuple[List[Optional[Dict[str, Any]]], List[Tuple[int, Dict[str, Any]]]]: + responses_so_far: list[Any], + ) -> tuple[list[dict[str, Any] | None], list[tuple[int, dict[str, Any]]]]: """Parse JSON-RPC items, returning aligned parsed list and valid entries.""" - parsed: List[Optional[Dict[str, Any]]] = [None] * len(responses_so_far) + parsed: list[dict[str, Any] | None] = [None] * len(responses_so_far) for i, item in enumerate(responses_so_far): if isinstance(item, dict): obj = item @@ -326,13 +326,13 @@ class A2AGuardrailHandler(BaseTranslation): def _collect_text_from_parsed_chunks( self, - valid_parsed: List[Tuple[int, Dict[str, Any]]], - ) -> Tuple[str, List[int]]: + valid_parsed: list[tuple[int, dict[str, Any]]], + ) -> tuple[str, list[int]]: """Collect text from parsed chunks, returning combined text and indices.""" from litellm.llms.a2a.common_utils import extract_text_from_a2a_response - text_parts: List[str] = [] - chunk_indices_with_text: List[int] = [] + text_parts: list[str] = [] + chunk_indices_with_text: list[int] = [] for _idx, (orig_i, obj) in enumerate(valid_parsed): t = extract_text_from_a2a_response(obj) if t: @@ -342,9 +342,9 @@ class A2AGuardrailHandler(BaseTranslation): def _extract_texts_from_result( self, - result: Dict[str, Any], - texts_to_check: List[str], - task_mappings: List[Tuple[Tuple[str, ...], int]], + result: dict[str, Any], + texts_to_check: list[str], + task_mappings: list[tuple[tuple[str, ...], int]], ) -> None: """ Extract text from all possible locations in an A2A result. @@ -411,10 +411,10 @@ class A2AGuardrailHandler(BaseTranslation): def _extract_texts_from_parts( self, - parts: List[Dict[str, Any]], - path: Tuple[str, ...], - texts_to_check: List[str], - task_mappings: List[Tuple[Tuple[str, ...], int]], + parts: list[dict[str, Any]], + path: tuple[str, ...], + texts_to_check: list[str], + task_mappings: list[tuple[tuple[str, ...], int]], ) -> None: """Extract text from message parts.""" for part_idx, part in enumerate(parts): @@ -426,8 +426,8 @@ class A2AGuardrailHandler(BaseTranslation): def _apply_text_to_path( self, - result: Dict[Union[str, int], Any], - path: Tuple[str, ...], + result: dict[str | int, Any], + path: tuple[str, ...], part_idx: int, text: str, ) -> None: diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index a7302ac2f0b..da5a2f41a9f 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -2,8 +2,6 @@ A2A Streaming Response Iterator """ -from typing import Optional, Union - from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.types.utils import GenericStreamingChunk, ModelResponseStream @@ -21,7 +19,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): self, streaming_response, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, model: str = "a2a/agent", ): super().__init__( @@ -31,7 +29,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): ) self.model = model - def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]: + def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: """ Parse A2A streaming chunk to OpenAI format. @@ -83,7 +81,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): tool_use=None, ) - def _get_finish_reason(self, chunk: dict) -> Optional[str]: + def _get_finish_reason(self, chunk: dict) -> str | None: """Extract finish reason from A2A chunk""" result = chunk.get("result", {}) diff --git a/litellm/llms/a2a/chat/transformation.py b/litellm/llms/a2a/chat/transformation.py index d8b67bcf22f..3de584d1d5f 100644 --- a/litellm/llms/a2a/chat/transformation.py +++ b/litellm/llms/a2a/chat/transformation.py @@ -4,7 +4,7 @@ A2A Protocol Transformation for LiteLLM import uuid from collections.abc import Iterator -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -31,11 +31,11 @@ class A2AConfig(BaseConfig): @staticmethod def resolve_agent_config_from_registry( model: str, - api_base: Optional[str], - api_key: Optional[str], - headers: Optional[Dict[str, Any]], - optional_params: Dict[str, Any], - ) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]: + 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]: """ Resolve agent configuration from registry if model format is "a2a/". @@ -90,7 +90,7 @@ class A2AConfig(BaseConfig): return api_base, api_key, headers - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """Return list of supported OpenAI parameters""" return [ "stream", @@ -123,11 +123,11 @@ class A2AConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set headers for A2A requests. @@ -156,12 +156,12 @@ class A2AConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete A2A agent endpoint URL. @@ -191,7 +191,7 @@ class A2AConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -248,12 +248,12 @@ class A2AConfig(BaseConfig): model_response: ModelResponse, logging_obj: Any, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform A2A JSON-RPC 2.0 response to OpenAI format. @@ -279,7 +279,7 @@ class A2AConfig(BaseConfig): except Exception as e: raise A2AError( status_code=raw_response.status_code, - message=f"Failed to parse A2A response: {str(e)}", + message=f"Failed to parse A2A response: {e!s}", headers=dict(raw_response.headers), ) @@ -336,9 +336,9 @@ class A2AConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator, Any], + streaming_response: Iterator | Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> BaseModelResponseIterator: """ Get streaming iterator for A2A responses. @@ -357,7 +357,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: dict[str, Any]) -> dict[str, Any]: """ Convert OpenAI message to A2A message format. @@ -376,9 +376,7 @@ class A2AConfig(BaseConfig): "messageId": str(uuid.uuid4()), } - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """Return appropriate error class for A2A errors""" # Convert headers to dict if needed headers_dict = dict(headers) if isinstance(headers, httpx.Headers) else headers diff --git a/litellm/llms/a2a/common_utils.py b/litellm/llms/a2a/common_utils.py index 4fc0ff2623e..3366ee873ee 100644 --- a/litellm/llms/a2a/common_utils.py +++ b/litellm/llms/a2a/common_utils.py @@ -2,7 +2,7 @@ Common utilities for A2A (Agent-to-Agent) Protocol """ -from typing import Any, Dict, List +from typing import Any from pydantic import BaseModel @@ -20,7 +20,7 @@ class A2AError(BaseLLMException): self, status_code: int, message: str, - headers: Dict[str, Any] = {}, + headers: dict[str, Any] = {}, ): super().__init__( status_code=status_code, @@ -29,7 +29,7 @@ class A2AError(BaseLLMException): ) -def convert_messages_to_prompt(messages: List[AllMessageValues]) -> str: +def convert_messages_to_prompt(messages: list[AllMessageValues]) -> str: """ Convert OpenAI messages to a single prompt string for A2A agent. @@ -61,7 +61,7 @@ def convert_messages_to_prompt(messages: List[AllMessageValues]) -> str: return "\n".join(conversation_parts) -def extract_text_from_a2a_message(message: Dict[str, Any], depth: int = 0, max_depth: int = 10) -> str: +def extract_text_from_a2a_message(message: dict[str, Any], depth: int = 0, max_depth: int = 10) -> str: """ Extract text content from A2A message parts. @@ -77,7 +77,7 @@ def extract_text_from_a2a_message(message: Dict[str, Any], depth: int = 0, max_d return "" parts = message.get("parts", []) - text_parts: List[str] = [] + text_parts: list[str] = [] for part in parts: if part.get("kind") == "text": @@ -91,7 +91,7 @@ def extract_text_from_a2a_message(message: Dict[str, Any], depth: int = 0, max_d return " ".join(text_parts) -def extract_text_from_a2a_response(response_dict: Dict[str, Any], max_depth: int = 10) -> str: +def extract_text_from_a2a_response(response_dict: dict[str, Any], max_depth: int = 10) -> str: """ Extract text content from A2A response result. diff --git a/litellm/llms/ai21/chat/transformation.py b/litellm/llms/ai21/chat/transformation.py index 1a07b50de5b..bd0ab247748 100644 --- a/litellm/llms/ai21/chat/transformation.py +++ b/litellm/llms/ai21/chat/transformation.py @@ -4,8 +4,6 @@ AI21 Chat Completions API this is OpenAI compatible - no translation needed / occurs """ -from typing import Optional, Union - from ...openai_like.chat.transformation import OpenAILikeChatConfig @@ -16,30 +14,30 @@ class AI21ChatConfig(OpenAILikeChatConfig): Below are the parameters: """ - tools: Optional[list] = None - response_format: Optional[dict] = None - documents: Optional[list] = None - max_tokens: Optional[int] = None - stop: Optional[Union[str, list]] = None - n: Optional[int] = None - stream: Optional[bool] = None - seed: Optional[int] = None - tool_choice: Optional[str] = None - user: Optional[str] = None + tools: list | None = None + response_format: dict | None = None + documents: list | None = None + max_tokens: int | None = None + stop: str | list | None = None + n: int | None = None + stream: bool | None = None + seed: int | None = None + tool_choice: str | None = None + user: str | None = None def __init__( self, - tools: Optional[list] = None, - response_format: Optional[dict] = None, - max_tokens: Optional[int] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - stop: Optional[Union[str, list]] = None, - n: Optional[int] = None, - stream: Optional[bool] = None, - seed: Optional[int] = None, - tool_choice: Optional[str] = None, - user: Optional[str] = None, + tools: list | None = None, + response_format: dict | None = None, + max_tokens: int | None = None, + temperature: float | None = None, + top_p: float | None = None, + stop: str | list | None = None, + n: int | None = None, + stream: bool | None = None, + seed: int | None = None, + tool_choice: str | None = None, + user: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/aiml/chat/transformation.py b/litellm/llms/aiml/chat/transformation.py index e62aa6238d7..e258367061a 100644 --- a/litellm/llms/aiml/chat/transformation.py +++ b/litellm/llms/aiml/chat/transformation.py @@ -1,22 +1,18 @@ -from typing import Optional, Tuple - from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.secret_managers.main import get_secret_str class AIMLChatConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "aiml" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # AIML is openai compatible, we just need to set the api_base api_base = ( api_base or get_secret_str("AIML_API_BASE") or "https://api.aimlapi.com/v1" # Default AIML API base URL ) # type: ignore dynamic_api_key = api_key or get_secret_str("AIML_API_KEY") return api_base, dynamic_api_key - - pass diff --git a/litellm/llms/aiml/image_generation/transformation.py b/litellm/llms/aiml/image_generation/transformation.py index b1ab443eb84..b6b7100306e 100644 --- a/litellm/llms/aiml/image_generation/transformation.py +++ b/litellm/llms/aiml/image_generation/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -37,7 +37,7 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): """ return model.startswith(OPENAI_STYLE_IMAGE_MODEL_PREFIXES) - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ https://api.aimlapi.com/v1/images/generations """ @@ -64,8 +64,8 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): supported_params = self.get_supported_openai_params(model) is_openai_style = self._is_openai_style_model(model) - for k in non_default_params.keys(): - if k in optional_params.keys(): + for k in non_default_params: + if k in optional_params: continue if k not in supported_params: if drop_params: @@ -99,12 +99,12 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -113,8 +113,7 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): complete_url = complete_url.rstrip("/") # Strip /v1 suffix if present since IMAGE_GENERATION_ENDPOINT already includes v1 - if complete_url.endswith("/v1"): - complete_url = complete_url[:-3] + complete_url = complete_url.removesuffix("/v1") complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}" return complete_url @@ -122,13 +121,13 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = ( + final_api_key: str | None = ( api_key or get_secret_str("AIML_API_KEY") or get_secret_str("AIMLAPI_KEY") # Alternative name ) if not final_api_key: @@ -171,8 +170,8 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the image generation response to the litellm image response diff --git a/litellm/llms/aiohttp_openai/chat/transformation.py b/litellm/llms/aiohttp_openai/chat/transformation.py index 346b565b6f5..d75cb92c1ac 100644 --- a/litellm/llms/aiohttp_openai/chat/transformation.py +++ b/litellm/llms/aiohttp_openai/chat/transformation.py @@ -7,7 +7,7 @@ https://github.com/BerriAI/litellm/issues/6592 New config to ensure we introduce this without causing breaking changes for users """ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any from aiohttp import ClientResponse @@ -26,12 +26,12 @@ else: class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Ensure - /v1/chat/completions is at the end of the url @@ -48,11 +48,11 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {"Authorization": f"Bearer {api_key}"} @@ -63,12 +63,12 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: _json_response = await raw_response.json() model_response.id = _json_response.get("id") diff --git a/litellm/llms/amazon_nova/chat/transformation.py b/litellm/llms/amazon_nova/chat/transformation.py index 8afcbd40ffc..e40a6af3d0b 100644 --- a/litellm/llms/amazon_nova/chat/transformation.py +++ b/litellm/llms/amazon_nova/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to Amazon Nova's `/v1/chat/completions` """ -from typing import Any, List, Optional, Tuple +from typing import Any import httpx @@ -18,22 +18,22 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig class AmazonNovaChatConfig(OpenAILikeChatConfig): - max_completion_tokens: Optional[int] = None - max_tokens: Optional[int] = None - metadata: Optional[int] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - tools: Optional[list] = None - reasoning_effort: Optional[list] = None + max_completion_tokens: int | None = None + max_tokens: int | None = None + metadata: int | None = None + temperature: int | None = None + top_p: int | None = None + tools: list | None = None + reasoning_effort: list | None = None def __init__( self, - max_completion_tokens: Optional[int] = None, - max_tokens: Optional[int] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - tools: Optional[list] = None, - reasoning_effort: Optional[list] = None, + max_completion_tokens: int | None = None, + max_tokens: int | None = None, + temperature: int | None = None, + top_p: int | None = None, + tools: list | None = None, + reasoning_effort: list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -41,7 +41,7 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "amazon_nova" @classmethod @@ -49,8 +49,8 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig): return super().get_config() def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # Amazon Nova is openai compatible, we just need to set this to custom_openai and have the api_base be Nova's endpoint api_base = api_base or get_secret_str("AMAZON_NOVA_API_BASE") or "https://api.nova.amazon.com/v1" # type: ignore @@ -58,7 +58,7 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig): key = api_key or litellm.amazon_nova_api_key or get_secret_str("AMAZON_NOVA_API_KEY") or litellm.api_key return api_base, key - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "top_p", "temperature", @@ -80,12 +80,12 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: model_response = super().transform_response( model=model, diff --git a/litellm/llms/amazon_nova/cost_calculation.py b/litellm/llms/amazon_nova/cost_calculation.py index 3b1121f1f8c..6e691cf279e 100644 --- a/litellm/llms/amazon_nova/cost_calculation.py +++ b/litellm/llms/amazon_nova/cost_calculation.py @@ -3,7 +3,7 @@ Helper util for handling amazon nova cost calculation - e.g.: prompt caching """ -from typing import TYPE_CHECKING, Tuple +from typing import TYPE_CHECKING from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token @@ -11,7 +11,7 @@ if TYPE_CHECKING: from litellm.types.utils import Usage -def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]: +def cost_per_token(model: str, usage: "Usage") -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. Follows the same logic as Anthropic's cost per token calculation. diff --git a/litellm/llms/anthropic/__init__.py b/litellm/llms/anthropic/__init__.py index 341fc8d1628..709dd08d212 100644 --- a/litellm/llms/anthropic/__init__.py +++ b/litellm/llms/anthropic/__init__.py @@ -1,5 +1,3 @@ -from typing import Type, Union - from .batches.transformation import AnthropicBatchesConfig from .chat.transformation import AnthropicConfig @@ -8,7 +6,7 @@ __all__ = ["AnthropicBatchesConfig", "AnthropicConfig"] def get_anthropic_config( url_route: str, -) -> Union[Type[AnthropicBatchesConfig], Type[AnthropicConfig]]: +) -> type[AnthropicBatchesConfig] | type[AnthropicConfig]: if "messages/batches" in url_route and "results" in url_route: return AnthropicBatchesConfig else: diff --git a/litellm/llms/anthropic/batches/__init__.py b/litellm/llms/anthropic/batches/__init__.py index dd9ae5273b8..4ae6beddd51 100644 --- a/litellm/llms/anthropic/batches/__init__.py +++ b/litellm/llms/anthropic/batches/__init__.py @@ -1,4 +1,4 @@ from .handler import AnthropicBatchesHandler from .transformation import AnthropicBatchesConfig -__all__ = ["AnthropicBatchesHandler", "AnthropicBatchesConfig"] +__all__ = ["AnthropicBatchesConfig", "AnthropicBatchesHandler"] diff --git a/litellm/llms/anthropic/batches/handler.py b/litellm/llms/anthropic/batches/handler.py index c83c96e50c4..3735742903b 100644 --- a/litellm/llms/anthropic/batches/handler.py +++ b/litellm/llms/anthropic/batches/handler.py @@ -4,7 +4,7 @@ Anthropic Batches API Handler import asyncio from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -37,11 +37,11 @@ class AnthropicBatchesHandler: async def aretrieve_batch( self, batch_id: str, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - logging_obj: Optional[LiteLLMLoggingObj] = None, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> LiteLLMBatch: """ Async: Retrieve a batch from Anthropic. @@ -125,12 +125,12 @@ class AnthropicBatchesHandler: self, _is_async: bool, batch_id: str, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - logging_obj: Optional[LiteLLMLoggingObj] = None, - ) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + logging_obj: LiteLLMLoggingObj | None = None, + ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: """ Retrieve a batch from Anthropic. diff --git a/litellm/llms/anthropic/batches/transformation.py b/litellm/llms/anthropic/batches/transformation.py index bfae42f96cf..851ccb0b943 100644 --- a/litellm/llms/anthropic/batches/transformation.py +++ b/litellm/llms/anthropic/batches/transformation.py @@ -1,6 +1,6 @@ import json import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Literal, cast import httpx from httpx import Headers, Response @@ -36,11 +36,11 @@ class AnthropicBatchesConfig(BaseBatchesConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """Validate and prepare environment-specific headers and parameters.""" if api_base is None and isinstance(litellm_params, dict): @@ -64,11 +64,11 @@ class AnthropicBatchesConfig(BaseBatchesConfig): def get_complete_batch_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, data: CreateBatchRequest, ) -> str: """Get the complete URL for batch creation request.""" @@ -83,7 +83,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): create_batch_data: CreateBatchRequest, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, Dict[str, Any]]: + ) -> bytes | str | dict[str, Any]: """ Transform the batch creation request to Anthropic format. @@ -93,7 +93,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): def transform_create_batch_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LoggingClass, litellm_params: dict, @@ -107,10 +107,10 @@ class AnthropicBatchesConfig(BaseBatchesConfig): def get_retrieve_batch_url( self, - api_base: Optional[str], + api_base: str | None, batch_id: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, ) -> str: """ Get the complete URL for batch retrieval request. @@ -133,7 +133,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): batch_id: str, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, Dict[str, Any]]: + ) -> bytes | str | dict[str, Any]: """ Transform batch retrieval request for Anthropic. @@ -145,7 +145,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): def transform_retrieve_batch_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LoggingClass, litellm_params: dict, @@ -161,7 +161,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): processing_status = response_data.get("processing_status", "in_progress") # Map Anthropic processing_status to OpenAI status - status_mapping: Dict[ + status_mapping: dict[ str, Literal[ "validating", @@ -181,7 +181,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig): openai_status = status_mapping.get(processing_status, "in_progress") # Parse timestamps - def parse_timestamp(ts_str: Optional[str]) -> Optional[int]: + def parse_timestamp(ts_str: str | None) -> int | None: if not ts_str: return None try: @@ -239,15 +239,13 @@ class AnthropicBatchesConfig(BaseBatchesConfig): metadata={}, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, Headers] - ) -> "BaseLLMException": + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> "BaseLLMException": """Get the appropriate error class for Anthropic.""" from ..common_utils import AnthropicError # Convert Dict to Headers if needed if isinstance(headers, dict): - headers_obj: Optional[Headers] = Headers(headers) + headers_obj: Headers | None = Headers(headers) else: headers_obj = headers if isinstance(headers, Headers) else None @@ -259,19 +257,19 @@ class AnthropicBatchesConfig(BaseBatchesConfig): raw_response: Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: from litellm.cost_calculator import BaseTokenUsageProcessor from litellm.types.utils import Usage response_text = raw_response.text.strip() - all_usage: List[Usage] = [] + all_usage: list[Usage] = [] try: # Split by newlines and try to parse each line as JSON diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index a549db94224..0fb7d7802a0 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -13,7 +13,7 @@ Pattern Overview: """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast from litellm._logging import verbose_proxy_logger from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -76,8 +76,8 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _build_streaming_usage_response( responses_so_far: list[Any], - request_data: Optional[dict], - ) -> Optional[ModelResponse]: + request_data: dict | None, + ) -> ModelResponse | None: chunks = tuple(response for response in responses_so_far if isinstance(response, (str, bytes))) if not chunks: return None @@ -93,7 +93,7 @@ class AnthropicMessagesHandler(BaseTranslation): self, exc: "ModifyResponseException", stream_started: bool = False, - responses_so_far: Optional[list[Any]] = None, + responses_so_far: list[Any] | None = None, ) -> list[bytes]: """ Build an Anthropic SSE sequence delivering the guardrail block message @@ -187,7 +187,7 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _content_block_state( responses_so_far: list[Any], - ) -> tuple[Optional[int], Optional[int]]: + ) -> tuple[int | None, int | None]: """From the SSE chunks already sent to the client, return (open content-block index or None, highest content-block index seen or None). @@ -196,7 +196,7 @@ class AnthropicMessagesHandler(BaseTranslation): considered -- matching how ``get_streaming_string_so_far`` reads the same stream.""" open_indices: set[int] = set() - max_index: Optional[int] = None + max_index: int | None = None for item in responses_so_far: for data in AnthropicMessagesHandler._iter_sse_events(item): event_type = data.get("type") @@ -247,7 +247,7 @@ class AnthropicMessagesHandler(BaseTranslation): ) return chat_completion_compatible_request - def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]: + def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None: """ Convert Anthropic messages request data to OpenAI-spec structured messages. @@ -258,7 +258,7 @@ class AnthropicMessagesHandler(BaseTranslation): return None chat_completion_compatible_request = self._translate_to_openai(data) result = cast( - List[AllMessageValues], + list[AllMessageValues], chat_completion_compatible_request.get("messages", []), ) return result if result else None @@ -267,7 +267,7 @@ class AnthropicMessagesHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input messages by applying guardrails to text content. @@ -282,7 +282,7 @@ class AnthropicMessagesHandler(BaseTranslation): chat_completion_compatible_request = self._translate_to_openai(data) structured_messages = cast( - List[AllMessageValues], + list[AllMessageValues], chat_completion_compatible_request.get("messages", []), ) if skip_system: @@ -290,10 +290,10 @@ class AnthropicMessagesHandler(BaseTranslation): if skip_tool: structured_messages = openai_messages_without_tool(structured_messages) - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - tools_to_check: List[ChatCompletionToolParam] = chat_completion_compatible_request.get("tools", []) - task_mappings: List[Tuple[int, Optional[int]]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + tools_to_check: list[ChatCompletionToolParam] = chat_completion_compatible_request.get("tools", []) + task_mappings: list[tuple[int, int | None]] = [] # Step 1: Extract all text content and images for msg_idx, message in enumerate(messages): @@ -333,7 +333,7 @@ class AnthropicMessagesHandler(BaseTranslation): if guardrailed_tools is not None: # Convert tools back from OpenAI format to Anthropic format anthropic_config = AnthropicConfig() - anthropic_tools: List[AllAnthropicToolsValues] = [] + anthropic_tools: list[AllAnthropicToolsValues] = [] for tool in guardrailed_tools: converted_tool, mcp_server = anthropic_config._map_tool_helper(tool) if converted_tool is not None: @@ -397,9 +397,9 @@ class AnthropicMessagesHandler(BaseTranslation): block.pop("cache_control", None) data["messages"] = converted - def extract_request_tool_names(self, data: dict) -> List[str]: + def extract_request_tool_names(self, data: dict) -> list[str]: """Extract tool names from Anthropic messages request (tools[].name).""" - names: List[str] = [] + names: list[str] = [] for tool in data.get("tools") or []: if isinstance(tool, dict) and tool.get("name"): names.append(str(tool["name"])) @@ -407,11 +407,11 @@ class AnthropicMessagesHandler(BaseTranslation): def _extract_input_text_and_images( self, - message: Dict[str, Any], + message: dict[str, Any], msg_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - task_mappings: List[Tuple[int, Optional[int]]], + texts_to_check: list[str], + images_to_check: list[str], + task_mappings: list[tuple[int, int | None]], skip_system_message: bool = False, skip_tool_message: bool = False, ) -> None: @@ -457,8 +457,8 @@ class AnthropicMessagesHandler(BaseTranslation): def _extract_input_tools( self, - tools: List[Dict[str, Any]], - tools_to_check: List[ChatCompletionToolParam], + tools: list[dict[str, Any]], + tools_to_check: list[ChatCompletionToolParam], ) -> None: """ Extract tools from a message. @@ -467,15 +467,15 @@ class AnthropicMessagesHandler(BaseTranslation): if tools is not None and isinstance(tools, list): # TRANSFORM ANTHROPIC TOOLS TO OPENAI TOOLS openai_tools = self.adapter.translate_anthropic_tools_to_openai( - tools=cast(List[AllAnthropicToolsValues], tools) + tools=cast(list[AllAnthropicToolsValues], tools) ) tools_to_check.extend(openai_tools) # type: ignore async def _apply_guardrail_responses_to_input( self, - messages: List[Dict[str, Any]], - responses: List[str], - task_mappings: List[Tuple[int, Optional[int]]], + messages: list[dict[str, Any]], + responses: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to input messages. @@ -485,7 +485,7 @@ class AnthropicMessagesHandler(BaseTranslation): for task_idx, guardrail_response in enumerate(responses): mapping = task_mappings[task_idx] msg_idx = cast(int, mapping[0]) - content_idx_optional = cast(Optional[int], mapping[1]) + content_idx_optional = cast(int | None, mapping[1]) content = messages[msg_idx].get("content", None) if content is None: @@ -503,9 +503,9 @@ class AnthropicMessagesHandler(BaseTranslation): self, response: "AnthropicMessagesResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response by applying guardrails to text content and tool calls. @@ -526,10 +526,10 @@ class AnthropicMessagesHandler(BaseTranslation): ... ] """ - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - tool_calls_to_check: List[ChatCompletionToolCallChunk] = [] - task_mappings: List[Tuple[int, Optional[int]]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + tool_calls_to_check: list[ChatCompletionToolCallChunk] = [] + task_mappings: list[tuple[int, int | None]] = [] response_content = self._get_response_content(response) if not response_content: @@ -582,12 +582,12 @@ class AnthropicMessagesHandler(BaseTranslation): async def process_output_streaming_response( self, - responses_so_far: List[Any], + responses_so_far: list[Any], guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, - ) -> List[Any]: + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, + ) -> list[Any]: """ Process output streaming response by applying guardrails to text content. @@ -609,7 +609,7 @@ class AnthropicMessagesHandler(BaseTranslation): model_response = cast(ModelResponse, built_response) first_choice = cast(Choices, model_response.choices[0]) tool_calls_list = cast( - Optional[List[ChatCompletionMessageToolCall]], + list[ChatCompletionMessageToolCall] | None, first_choice.message.tool_calls, ) string_so_far = first_choice.message.content @@ -664,9 +664,9 @@ class AnthropicMessagesHandler(BaseTranslation): def _prepare_request_data( self, - request_data: Optional[dict], + request_data: dict | None, response: Any, - user_api_key_dict: Optional[Any], + user_api_key_dict: Any | None, key: str, ) -> dict: """Ensure request_data has the response/responses_so_far key and metadata.""" @@ -683,7 +683,7 @@ class AnthropicMessagesHandler(BaseTranslation): return request_data @staticmethod - def _get_response_content(response: Any) -> List[Any]: + def _get_response_content(response: Any) -> list[Any]: """Extract content list from a dict or object response.""" if isinstance(response, dict): return response.get("content", []) or [] @@ -693,18 +693,18 @@ class AnthropicMessagesHandler(BaseTranslation): def _extract_from_content_blocks( self, - response_content: List[Any], - texts_to_check: List[str], - images_to_check: List[str], - task_mappings: List[Tuple[int, Optional[int]]], - tool_calls_to_check: List["ChatCompletionToolCallChunk"], + response_content: list[Any], + texts_to_check: list[str], + images_to_check: list[str], + task_mappings: list[tuple[int, int | None]], + tool_calls_to_check: list["ChatCompletionToolCallChunk"], ) -> None: """Extract text, images, and tool calls from content blocks.""" for content_idx, content_block in enumerate(response_content): - block_dict: Dict[str, Any] = {} + block_dict: dict[str, Any] = {} if isinstance(content_block, dict): block_type = content_block.get("type") - block_dict = cast(Dict[str, Any], content_block) + block_dict = cast(dict[str, Any], content_block) elif hasattr(content_block, "type"): block_type = getattr(content_block, "type", None) if hasattr(content_block, "model_dump"): @@ -729,9 +729,9 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _build_guardrail_inputs( - texts_to_check: List[str], - images_to_check: List[str], - tool_calls_to_check: List["ChatCompletionToolCallChunk"], + texts_to_check: list[str], + images_to_check: list[str], + tool_calls_to_check: list["ChatCompletionToolCallChunk"], response: Any, ) -> "GenericGuardrailAPIInputs": """Build GenericGuardrailAPIInputs with optional images, tool calls, model.""" @@ -749,7 +749,7 @@ class AnthropicMessagesHandler(BaseTranslation): inputs["model"] = response_model return inputs - def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str: + def get_streaming_string_so_far(self, responses_so_far: list[Any]) -> str: """ Parse streaming responses and extract accumulated text content. @@ -832,7 +832,7 @@ class AnthropicMessagesHandler(BaseTranslation): return text - def _check_streaming_has_ended(self, responses_so_far: List[Any]) -> bool: + def _check_streaming_has_ended(self, responses_so_far: list[Any]) -> bool: """ Check if streaming response has ended by looking for non-null stop_reason. @@ -927,12 +927,12 @@ class AnthropicMessagesHandler(BaseTranslation): def _extract_output_text_and_images( self, - content_block: Dict[str, Any], + content_block: dict[str, Any], content_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - task_mappings: List[Tuple[int, Optional[int]]], - tool_calls_to_check: Optional[List[ChatCompletionToolCallChunk]] = None, + texts_to_check: list[str], + images_to_check: list[str], + task_mappings: list[tuple[int, int | None]], + tool_calls_to_check: list[ChatCompletionToolCallChunk] | None = None, ) -> None: """ Extract text content, images, and tool calls from a response content block. @@ -962,8 +962,8 @@ class AnthropicMessagesHandler(BaseTranslation): async def _apply_guardrail_responses_to_output( self, response: "AnthropicMessagesResponse", - responses: List[str], - task_mappings: List[Tuple[int, Optional[int]]], + responses: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to output response. @@ -975,7 +975,7 @@ class AnthropicMessagesHandler(BaseTranslation): content_idx = cast(int, mapping[0]) # Handle both dict and object responses - response_content: List[Any] = [] + response_content: list[Any] = [] if isinstance(response, dict): response_content = response.get("content", []) or [] elif hasattr(response, "content"): @@ -997,7 +997,7 @@ class AnthropicMessagesHandler(BaseTranslation): # Handle both dict and Pydantic object content blocks if isinstance(content_block, dict): if content_block.get("type") == "text": - cast(Dict[str, Any], content_block)["text"] = guardrail_response + cast(dict[str, Any], content_block)["text"] = guardrail_response elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text": # Update Pydantic object's text attribute if hasattr(content_block, "text"): diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 39780ac5df4..111dae52d90 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -8,10 +8,7 @@ from collections.abc import Callable from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Tuple, Union, cast, ) @@ -79,11 +76,11 @@ async def make_call( model: str, messages: list, logging_obj, - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, json_mode: bool, speed: str | None = None, - tool_name_reverse_map: Dict[str, str] | None = None, -) -> Tuple[Any, httpx.Headers]: + tool_name_reverse_map: dict[str, str] | None = None, +) -> tuple[Any, httpx.Headers]: if client is None: client = litellm.module_level_aclient @@ -139,11 +136,11 @@ def make_sync_call( model: str, messages: list, logging_obj, - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, json_mode: bool, speed: str | None = None, - tool_name_reverse_map: Dict[str, str] | None = None, -) -> Tuple[Any, httpx.Headers]: + tool_name_reverse_map: dict[str, str] | None = None, +) -> tuple[Any, httpx.Headers]: if client is None: client = litellm.module_level_client # re-use a module level client @@ -211,7 +208,7 @@ class AnthropicChatCompletion(BaseLLM): custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, client: AsyncHTTPHandler | None, encoding, api_key, @@ -261,7 +258,7 @@ class AnthropicChatCompletion(BaseLLM): custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, encoding, api_key, logging_obj, @@ -335,7 +332,7 @@ class AnthropicChatCompletion(BaseLLM): api_key, logging_obj, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, acompletion=None, logger_fn=None, @@ -530,11 +527,11 @@ class ModelResponseIterator: sync_stream: bool, json_mode: bool | None = False, speed: str | None = None, - tool_name_reverse_map: Dict[str, str] | None = None, + tool_name_reverse_map: dict[str, str] | None = None, ): self.streaming_response = streaming_response self.response_iterator = self.streaming_response - self.content_blocks: List[ContentBlockDelta] = [] + self.content_blocks: list[ContentBlockDelta] = [] self.tool_index = -1 self.json_mode = json_mode self.speed = speed @@ -544,7 +541,7 @@ class ModelResponseIterator: # `foo_bar` is *not* reverse-mapped just because some other tool was # rewritten to `foo_bar` in a different request. Empty/None is the # common case (no '/' or other invalid chars in any tool name). - self.tool_name_reverse_map: Dict[str, str] = tool_name_reverse_map or {} + self.tool_name_reverse_map: dict[str, str] = tool_name_reverse_map or {} # Generate response ID once per stream to match OpenAI-compatible behavior self.response_id = _generate_id() @@ -564,18 +561,18 @@ class ModelResponseIterator: # Accumulate web_search_tool_result blocks for multi-turn reconstruction # See: https://github.com/BerriAI/litellm/issues/17737 - self.web_search_results: List[Dict[str, Any]] = [] + self.web_search_results: list[dict[str, Any]] = [] # Accumulate compaction blocks for multi-turn reconstruction - self.compaction_blocks: List[Dict[str, Any]] = [] + self.compaction_blocks: list[dict[str, Any]] = [] # Accumulate streamed thinking text so final usage can split reasoning # tokens from regular output tokens. - self.reasoning_content_chunks: List[str] = [] + self.reasoning_content_chunks: list[str] = [] # Track server tool use inputs and results for code_interpreter_results - self._server_tool_inputs: Dict[str, Any] = {} - self.tool_results: List[Dict[str, Any]] = [] + self._server_tool_inputs: dict[str, Any] = {} + self.tool_results: list[dict[str, Any]] = [] self._current_server_tool_id: str | None = None self._container_id: str | None = None @@ -602,7 +599,7 @@ class ModelResponseIterator: return True return False - def _handle_usage(self, anthropic_usage_chunk: Union[dict, UsageDelta]) -> Usage: + def _handle_usage(self, anthropic_usage_chunk: dict | UsageDelta) -> Usage: reasoning_content = "".join(self.reasoning_content_chunks) if self.reasoning_content_chunks else None return AnthropicConfig().calculate_usage( usage_object=cast(dict, anthropic_usage_chunk), @@ -612,11 +609,11 @@ class ModelResponseIterator: def _content_block_delta_helper( self, chunk: dict - ) -> Tuple[ + ) -> tuple[ str, ChatCompletionToolCallChunk | None, - List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]], - Dict[str, Any], + list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock], + dict[str, Any], str | None, ]: """ @@ -627,7 +624,7 @@ class ModelResponseIterator: provider_specific_fields = {} reasoning_content: str | None = None content_block = ContentBlockDelta(**chunk) # type: ignore - thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = [] + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = [] self.content_blocks.append(content_block) if "text" in content_block["delta"]: @@ -698,8 +695,8 @@ class ModelResponseIterator: def _handle_redacted_thinking_content( self, content_block_start: ContentBlockStart, - provider_specific_fields: Dict[str, Any], - ) -> Tuple[List[ChatCompletionRedactedThinkingBlock], Dict[str, Any]]: + provider_specific_fields: dict[str, Any], + ) -> tuple[list[ChatCompletionRedactedThinkingBlock], dict[str, Any]]: """ Handle the redacted thinking content """ @@ -767,9 +764,9 @@ class ModelResponseIterator: tool_use: ChatCompletionToolCallChunk | None = None finish_reason = "" usage: Usage | None = None - provider_specific_fields: Dict[str, Any] = {} + provider_specific_fields: dict[str, Any] = {} reasoning_content: str | None = None - thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None # Always use index=0 for OpenAI choice format (fixes multi-choice errors) index = 0 @@ -831,7 +828,7 @@ class ModelResponseIterator: if "caller" in content_block_start["content_block"]: caller_data = content_block_start["content_block"]["caller"] if caller_data: - tool_use["caller"] = cast(Dict[str, Any], caller_data) # type: ignore[typeddict-item] + tool_use["caller"] = cast(dict[str, Any], caller_data) # type: ignore[typeddict-item] elif content_block_start["content_block"]["type"] == "redacted_thinking": ( thinking_blocks, @@ -986,7 +983,7 @@ class ModelResponseIterator: def _handle_json_mode_chunk( self, text: str, tool_use: ChatCompletionToolCallChunk | None - ) -> Tuple[str, ChatCompletionToolCallChunk | None]: + ) -> tuple[str, ChatCompletionToolCallChunk | None]: """ If JSON mode is enabled, convert the tool call to a message. @@ -1030,7 +1027,7 @@ class ModelResponseIterator: return text, tool_use - def _handle_message_delta(self, chunk: dict) -> Tuple[str, Usage | None, Dict[str, Any] | None]: + def _handle_message_delta(self, chunk: dict) -> tuple[str, Usage | None, dict[str, Any] | None]: """ Handle message_delta event for finish_reason, usage, and container. diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index e99f356f8f2..40d1dbac187 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -4,12 +4,7 @@ import time from typing import ( TYPE_CHECKING, Any, - Dict, - List, NoReturn, - Optional, - Tuple, - Union, cast, ) @@ -74,12 +69,10 @@ from litellm.types.responses.main import ( from litellm.types.utils import ( CacheCreationTokenDetails, CompletionTokensDetailsWrapper, -) -from litellm.types.utils import Message as LitellmMessage -from litellm.types.utils import ( PromptTokensDetailsWrapper, ServerToolUse, ) +from litellm.types.utils import Message as LitellmMessage from litellm.utils import ( ModelResponse, Usage, @@ -152,8 +145,8 @@ def _basic_sanitize_anthropic_tool_name(name: str) -> str: def _build_anthropic_tool_name_maps( - original_names: List[str], -) -> Tuple[Dict[str, str], Dict[str, str]]: + original_names: list[str], +) -> tuple[dict[str, str], dict[str, str]]: """Build (forward, reverse) tool-name maps for a single request. forward[original] = sanitized -- only present when name was rewritten @@ -173,7 +166,7 @@ def _build_anthropic_tool_name_maps( seen gets the disambiguating suffix. Callers should preserve the caller's tool order (we do). """ - forward: Dict[str, str] = {} + forward: dict[str, str] = {} used: set = set() # First pass: reserve slots for names that are already valid so they @@ -214,7 +207,7 @@ def _build_anthropic_tool_name_maps( return forward, reverse -REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = { +REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: dict[str, str] = { "low": "low", "minimal": "low", "medium": "medium", @@ -245,23 +238,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): to pass metadata to anthropic, it's {"user_id": "any-relevant-information"} """ - max_tokens: Optional[int] = None - stop_sequences: Optional[list] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - metadata: Optional[dict] = None - system: Optional[str] = None + max_tokens: int | None = None + stop_sequences: list | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + metadata: dict | None = None + system: str | None = None def __init__( self, - max_tokens: Optional[int] = None, - stop_sequences: Optional[list] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - metadata: Optional[dict] = None, - system: Optional[str] = None, + max_tokens: int | None = None, + stop_sequences: list | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + metadata: dict | None = None, + system: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -269,7 +262,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "anthropic" @property @@ -277,7 +270,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return self.custom_llm_provider or "anthropic" @classmethod - def get_config(cls, *, model: Optional[str] = None): + def get_config(cls, *, model: str | None = None): config = super().get_config() # anthropic requires a default value for max_tokens @@ -287,7 +280,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return config @staticmethod - def get_max_tokens_for_model(model: Optional[str] = None) -> int: + def get_max_tokens_for_model(model: str | None = None) -> int: """ Get the max output tokens for a given model. Falls back to DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS (configurable via env var) if model is not found. @@ -304,7 +297,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def convert_tool_use_to_openai_format( - anthropic_tool_content: Dict[str, Any], + anthropic_tool_content: dict[str, Any], index: int, ) -> ChatCompletionToolCallChunk: """ @@ -329,7 +322,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) # Include caller information if present (for programmatic tool calling) if "caller" in anthropic_tool_content: - tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item] + tool_call["caller"] = cast(dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item] return tool_call @staticmethod @@ -352,7 +345,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) @staticmethod - def _validate_effort_for_model(model: str, effort: Optional[str], custom_llm_provider: str) -> Optional[str]: + def _validate_effort_for_model(model: str, effort: str | None, custom_llm_provider: str) -> str | None: """Return ``None`` if ``effort`` is allowed on ``model``, else an error message.""" if effort == "max" and not ( AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider) @@ -380,7 +373,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) @staticmethod - def _model_supports_speed_param(model: str, custom_llm_provider: Optional[str] = None) -> bool: + def _model_supports_speed_param(model: str, custom_llm_provider: str | None = None) -> bool: """Whether the model accepts Anthropic's ``speed`` parameter (fast mode). Fast mode is direct Anthropic API-only (not Bedrock, Vertex, or Azure). @@ -397,7 +390,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): model: str, optional_params: dict, drop_params: bool, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> None: if "speed" not in optional_params: return @@ -476,7 +469,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return params @staticmethod - def filter_anthropic_output_schema(schema: Dict[str, Any]) -> Dict[str, Any]: + def filter_anthropic_output_schema(schema: dict[str, Any]) -> dict[str, Any]: """ Filter out unsupported fields from JSON schema for Anthropic's output_format API. @@ -563,7 +556,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): note_value = json.dumps(value) if isinstance(value, (dict, list)) else value constraint_descriptions.append(label.format(note_value)) - result: Dict[str, Any] = {} + result: dict[str, Any] = {} # Update description with removed constraint info if constraint_descriptions: @@ -609,7 +602,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return result - def get_json_schema_from_pydantic_object(self, response_format: Union[Any, Dict, None]) -> Optional[dict]: + def get_json_schema_from_pydantic_object(self, response_format: Any | dict | None) -> dict | None: return type_to_response_format_param( response_format, ref_template="/$defs/{model}" ) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755 @@ -624,10 +617,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_choice( self, - tool_choice: Optional[str], - parallel_tool_use: Optional[bool], - ) -> Optional[AnthropicMessagesToolChoice]: - _tool_choice: Optional[AnthropicMessagesToolChoice] = None + tool_choice: str | None, + parallel_tool_use: bool | None, + ) -> AnthropicMessagesToolChoice | None: + _tool_choice: AnthropicMessagesToolChoice | None = None if tool_choice == "auto": _tool_choice = AnthropicMessagesToolChoice( type="auto", @@ -668,9 +661,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_helper( self, tool: ChatCompletionToolParam, - ) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]: - returned_tool: Optional[AllAnthropicToolsValues] = None - mcp_server: Optional[AnthropicMcpServerTool] = None + ) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]: + returned_tool: AllAnthropicToolsValues | None = None + mcp_server: AnthropicMcpServerTool | None = None if tool["type"] == "function" or tool["type"] == "custom": _input_schema = tool["function"].get( @@ -700,8 +693,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if "parameters" not in tool["function"]: raise ValueError("Missing required parameter: parameters") - _display_width_px: Optional[int] = tool["function"]["parameters"].get("display_width_px") - _display_height_px: Optional[int] = tool["function"]["parameters"].get("display_height_px") + _display_width_px: int | None = tool["function"]["parameters"].get("display_width_px") + _display_height_px: int | None = tool["function"]["parameters"].get("display_height_px") if _display_width_px is None or _display_height_px is None: raise ValueError("Missing required parameter: display_width_px or display_height_px") @@ -862,14 +855,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration allowed_tools = tool.get("allowed_tools", None) - tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None + tool_configuration: AnthropicMcpServerToolConfiguration | None = None if allowed_tools is not None: tool_configuration = AnthropicMcpServerToolConfiguration( allowed_tools=tool.get("allowed_tools", None), ) headers = tool.get("headers", {}) - authorization_token: Optional[str] = None + authorization_token: str | None = None if headers is not None: bearer_token = headers.get("Authorization", None) if bearer_token is not None: @@ -889,8 +882,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tools( self, - tools: List, - ) -> Tuple[List[AllAnthropicToolsValues], List[AnthropicMcpServerTool]]: + tools: list, + ) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]: anthropic_tools = [] mcp_servers = [] for tool in tools: @@ -935,9 +928,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _rewrite_tool_names_in_messages( - messages: List[AllMessageValues], - name_forward_map: Dict[str, str], - ) -> List[AllMessageValues]: + messages: list[AllMessageValues], + name_forward_map: dict[str, str], + ) -> list[AllMessageValues]: """Return a copy of `messages` with tool_call/function_call names rewritten using the per-request forward map. @@ -948,7 +941,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): """ if not name_forward_map: return messages - new_messages: List[AllMessageValues] = [] + new_messages: list[AllMessageValues] = [] for msg in messages: if not isinstance(msg, dict): new_messages.append(msg) @@ -986,8 +979,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _build_request_tool_name_maps( - tools: List, - ) -> Tuple[Dict[str, str], Dict[str, str]]: + tools: list, + ) -> tuple[dict[str, str], dict[str, str]]: """Build the (forward, reverse) tool-name maps for an OpenAI tools list. Operates on **OpenAI-format** tool dicts (pre-``_map_tools``). The @@ -1001,7 +994,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): the original name out of either ``{"function": {"name": ...}}`` (legacy OpenAI shape) or ``{"name": ...}`` (rare top-level shape). """ - original_names: List[str] = [] + original_names: list[str] = [] for tool in tools or []: if not isinstance(tool, dict): continue @@ -1014,8 +1007,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _sanitize_tool_names_in_request( - optional_params: Dict[str, Any], - ) -> Tuple[Dict[str, str], Dict[str, str]]: + optional_params: dict[str, Any], + ) -> tuple[dict[str, str], dict[str, str]]: """Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']`` in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``. @@ -1039,7 +1032,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # Order matters: the first occurrence wins the canonical slot; # later collisions get numeric suffixes (see # ``_build_anthropic_tool_name_maps``). - original_names: List[str] = [] + original_names: list[str] = [] for t in tools: if not isinstance(t, dict): continue @@ -1061,7 +1054,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # so a caller reusing the same tool list/dicts across requests # doesn't see its inputs permanently rewritten (which would also # drop the original key from `forward` on the next request). - new_tools: List[Any] = [] + new_tools: list[Any] = [] for t in tools: if ( isinstance(t, dict) @@ -1087,7 +1080,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return forward, reverse - def _detect_tool_search_tools(self, tools: Optional[List]) -> bool: + def _detect_tool_search_tools(self, tools: list | None) -> bool: """Check if tool search tools are present in the tools list.""" if not tools: return False @@ -1101,7 +1094,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return True return False - def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]: + def _separate_deferred_tools(self, tools: list) -> tuple[list, list]: """ Separate tools into deferred and non-deferred lists. @@ -1121,9 +1114,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _expand_tool_references( self, - content: List, - deferred_tools: List, - ) -> List: + content: list, + deferred_tools: list, + ) -> list: """ Expand tool_reference blocks to full tool definitions. @@ -1164,8 +1157,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return expanded_content - def _map_stop_sequences(self, stop: Optional[Union[str, List[str]]]) -> Optional[List[str]]: - new_stop: Optional[List[str]] = None + def _map_stop_sequences(self, stop: str | list[str] | None) -> list[str] | None: + new_stop: list[str] | None = None if isinstance(stop, str): if ( stop.isspace() and litellm.drop_params is True @@ -1186,11 +1179,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _map_reasoning_effort( - reasoning_effort: Optional[Union[REASONING_EFFORT, str]], + reasoning_effort: REASONING_EFFORT | str | None, model: str, custom_llm_provider: str, llm_provider: str = "anthropic", - ) -> Optional[AnthropicThinkingParam]: + ) -> AnthropicThinkingParam | None: """Capability probes read the cost map under ``custom_llm_provider``; ``llm_provider`` only tags raised exceptions.""" if reasoning_effort is None or reasoning_effort == "none": return None @@ -1244,8 +1237,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _cap_thinking_budget_to_max_tokens( - thinking: AnthropicThinkingParam, max_tokens: Optional[int] - ) -> Optional[AnthropicThinkingParam]: + thinking: AnthropicThinkingParam, max_tokens: int | None + ) -> AnthropicThinkingParam | None: """Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic requires ``max_tokens > budget_tokens``). Returns the (possibly capped) thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the @@ -1259,10 +1252,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return thinking return AnthropicThinkingParam(type=thinking.get("type", "enabled"), budget_tokens=max_tokens - 1) - def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]: + def _extract_json_schema_from_response_format(self, value: dict | None) -> dict | None: if value is None: return None - json_schema: Optional[dict] = None + json_schema: dict | None = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: @@ -1270,8 +1263,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return json_schema - def map_response_format_to_anthropic_output_format(self, value: Optional[dict]) -> Optional[AnthropicOutputSchema]: - json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value) + def map_response_format_to_anthropic_output_format(self, value: dict | None) -> AnthropicOutputSchema | None: + json_schema: dict | None = self._extract_json_schema_from_response_format(value) if json_schema is None: return None @@ -1297,13 +1290,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) def map_response_format_to_anthropic_tool( - self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool - ) -> Optional[AnthropicMessagesTool]: + self, value: dict | None, optional_params: dict, is_thinking_enabled: bool + ) -> AnthropicMessagesTool | None: ignore_response_format_types = ["text"] if value is None or value["type"] in ignore_response_format_types: # value is a no-op return None - json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value) + json_schema: dict | None = self._extract_json_schema_from_response_format(value) if json_schema is None: return None """ @@ -1348,8 +1341,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def map_openai_context_management_to_anthropic( - context_management: Union[List[Dict[str, Any]], Dict[str, Any]], - ) -> Optional[Dict[str, Any]]: + context_management: list[dict[str, Any]] | dict[str, Any], + ) -> dict[str, Any] | None: """ OpenAI format: [{"type": "compaction", "compact_threshold": 200000}] Anthropic format: { @@ -1380,7 +1373,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): entry_type = entry.get("type") if entry_type == "compaction": - anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"} + anthropic_edit: dict[str, Any] = {"type": "compact_20260112"} compact_threshold = entry.get("compact_threshold") # Rewrite to 'trigger' with correct nesting if threshold exists if compact_threshold is not None and isinstance(compact_threshold, (int, float)): @@ -1424,9 +1417,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # ``Extra inputs are not permitted``). for param, value in non_default_params.items(): - if param == "max_tokens": - optional_params["max_tokens"] = value if isinstance(value, int) else max(1, int(round(value))) - elif param == "max_completion_tokens": + if param == "max_tokens" or param == "max_completion_tokens": optional_params["max_tokens"] = value if isinstance(value, int) else max(1, int(round(value))) elif param == "tools": anthropic_tools, mcp_servers = self._map_tools(value) @@ -1436,7 +1427,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if mcp_servers: optional_params["mcp_servers"] = mcp_servers elif param == "tool_choice" or param == "parallel_tool_calls": - _tool_choice: Optional[AnthropicMessagesToolChoice] = self._map_tool_choice( + _tool_choice: AnthropicMessagesToolChoice | None = self._map_tool_choice( tool_choice=non_default_params.get("tool_choice"), parallel_tool_use=non_default_params.get("parallel_tool_calls"), ) @@ -1585,7 +1576,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _create_json_tool_call_for_response_format( self, - json_schema: Optional[dict] = None, + json_schema: dict | None = None, ) -> AnthropicMessagesTool: """ Handles creating a tool call for getting responses in JSON format. @@ -1620,7 +1611,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): """ return False - def translate_system_message(self, messages: List[AllMessageValues]) -> List[AnthropicSystemMessageContent]: + def translate_system_message(self, messages: list[AllMessageValues]) -> list[AnthropicSystemMessageContent]: """ Translate system message to anthropic format. @@ -1628,7 +1619,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped. """ system_prompt_indices = [] - anthropic_system_message_list: List[AnthropicSystemMessageContent] = [] + anthropic_system_message_list: list[AnthropicSystemMessageContent] = [] for idx, message in enumerate(messages): if message["role"] == "system": system_prompt_indices.append(idx) @@ -1678,9 +1669,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def add_code_execution_tool( self, - messages: List[AllAnthropicMessageValues], - tools: List[Union[AllAnthropicToolsValues, Dict]], - ) -> List[Union[AllAnthropicToolsValues, Dict]]: + messages: list[AllAnthropicMessageValues], + tools: list[AllAnthropicToolsValues | dict], + ) -> list[AllAnthropicToolsValues | dict]: """if 'container_upload' in messages, add code_execution tool""" add_code_execution_tool = False for message in messages: @@ -1795,7 +1786,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -1884,7 +1875,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): except Exception as e: raise AnthropicError( status_code=400, - message="{}\nReceived Messages={}".format(str(e), messages), + message=f"{e!s}\nReceived Messages={messages}", ) # don't use verbose_logger.exception, if exception is raised ## Auto-strip advisor blocks from history if advisor tool is absent. @@ -1897,7 +1888,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ## Add code_execution tool if container_upload is in messages _tools = ( cast( - Optional[List[Union[AllAnthropicToolsValues, Dict]]], + list[AllAnthropicToolsValues | dict] | None, optional_params.get("tools"), ) or [] @@ -1995,12 +1986,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _resolve_json_mode_non_streaming( self, - json_mode: Optional[bool], - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Tuple[ - Optional[LitellmMessage], - List[ChatCompletionToolCallChunk], - Optional[str], + json_mode: bool | None, + tool_calls: list[ChatCompletionToolCallChunk], + ) -> tuple[ + LitellmMessage | None, + list[ChatCompletionToolCallChunk], + str | None, ]: """Strip internal response_format tool calls; merge payload into content when mixed with user tools.""" if json_mode is not True or not tool_calls: @@ -2021,30 +2012,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): first_json = tool_calls[json_indices[0]] json_msg = AnthropicConfig._convert_tool_response_to_message([first_json]) - extra_content: Optional[str] = json_msg.content if json_msg is not None else None + extra_content: str | None = json_msg.content if json_msg is not None else None filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices] return None, filtered_tools, extra_content def extract_response_content( self, completion_response: dict - ) -> Tuple[ + ) -> tuple[ str, - Optional[List[Any]], - Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], - Optional[str], - List[ChatCompletionToolCallChunk], - Optional[List[Any]], - Optional[List[Any]], - Optional[List[Any]], + list[Any] | None, + list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None, + str | None, + list[ChatCompletionToolCallChunk], + list[Any] | None, + list[Any] | None, + list[Any] | None, ]: text_content = "" - citations: Optional[List[Any]] = None - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None - reasoning_content: Optional[str] = None - tool_calls: List[ChatCompletionToolCallChunk] = [] - web_search_results: Optional[List[Any]] = None - tool_results: Optional[List[Any]] = None - compaction_blocks: Optional[List[Any]] = None + citations: list[Any] | None = None + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None + reasoning_content: str | None = None + tool_calls: list[ChatCompletionToolCallChunk] = [] + web_search_results: list[Any] | None = None + tool_results: list[Any] | None = None + compaction_blocks: list[Any] | None = None for idx, content in enumerate(completion_response["content"]): if content["type"] == "text": text_content += content["text"] @@ -2062,11 +2053,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if content["type"] == "tool_search_tool_result": continue # Handle web_search_tool_result separately for backwards compatibility - if content["type"] == "web_search_tool_result": - if web_search_results is None: - web_search_results = [] - web_search_results.append(content) - elif content["type"] == "web_fetch_tool_result": + if content["type"] == "web_search_tool_result" or content["type"] == "web_fetch_tool_result": if web_search_results is None: web_search_results = [] web_search_results.append(content) @@ -2107,7 +2094,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if thinking_blocks is not None: reasoning_content = "" for block in thinking_blocks: - thinking_content = cast(Optional[str], block.get("thinking")) + thinking_content = cast(str | None, block.get("thinking")) if thinking_content is not None: reasoning_content += thinking_content @@ -2125,9 +2112,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def calculate_usage( self, usage_object: dict, - reasoning_content: Optional[str], - completion_response: Optional[dict] = None, - speed: Optional[str] = None, + reasoning_content: str | None, + completion_response: dict | None = None, + speed: str | None = None, ) -> Usage: # NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0 @@ -2137,10 +2124,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage = usage_object cache_creation_input_tokens: int = 0 cache_read_input_tokens: int = 0 - cache_creation_token_details: Optional[CacheCreationTokenDetails] = None - web_search_requests: Optional[int] = None - tool_search_requests: Optional[int] = None - inference_geo: Optional[str] = None + cache_creation_token_details: CacheCreationTokenDetails | None = None + web_search_requests: int | None = None + tool_search_requests: int | None = None + inference_geo: str | None = None if "inference_geo" in _usage and _usage["inference_geo"] is not None: inference_geo = _usage["inference_geo"] service_tier = cast( @@ -2148,7 +2135,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage.get("service_tier"), ) - iterations: Optional[List[Any]] = _usage.get("iterations") + iterations: list[Any] | None = _usage.get("iterations") if iterations: prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations) completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations) @@ -2206,7 +2193,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) reasoning_tokens = min(estimated_reasoning_tokens, completion_tokens) completion_token_details = CompletionTokensDetailsWrapper( - reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else 0, + reasoning_tokens=max(0, reasoning_tokens), text_tokens=(completion_tokens - reasoning_tokens if reasoning_tokens > 0 else completion_tokens), ) total_tokens = prompt_tokens + completion_tokens @@ -2234,8 +2221,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return usage - def _build_code_by_id_map(self, tool_calls: List[ChatCompletionToolCallChunk]) -> Dict[str, str]: - code_by_id: Dict[str, str] = {} + def _build_code_by_id_map(self, tool_calls: list[ChatCompletionToolCallChunk]) -> dict[str, str]: + code_by_id: dict[str, str] = {} for tc in tool_calls: try: args = json.loads(tc.get("function", {}).get("arguments", "{}")) @@ -2249,10 +2236,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_code_interpreter_results( self, - tool_results: List[Any], - code_by_id: Dict[str, str], - container_id: Optional[str], - ) -> List[OutputCodeInterpreterCall]: + tool_results: list[Any], + code_by_id: dict[str, str], + container_id: str | None, + ) -> list[OutputCodeInterpreterCall]: code_interpreter_results = [] for tr in tool_results: if tr.get("type") != "bash_code_execution_tool_result": @@ -2275,14 +2262,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_provider_specific_fields( self, completion_response: dict, - citations: Optional[List[Any]], - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], - web_search_results: Optional[List[Any]], - tool_results: Optional[List[Any]], - compaction_blocks: Optional[List[Any]], - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Dict[str, Any]: - provider_specific_fields: Dict[str, Any] = { + citations: list[Any] | None, + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None, + web_search_results: list[Any] | None, + tool_results: list[Any] | None, + compaction_blocks: list[Any] | None, + tool_calls: list[ChatCompletionToolCallChunk], + ) -> dict[str, Any]: + provider_specific_fields: dict[str, Any] = { "citations": citations, "thinking_blocks": thinking_blocks, } @@ -2319,12 +2306,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): completion_response: dict, raw_response: httpx.Response, model_response: ModelResponse, - json_mode: Optional[bool] = None, - prefix_prompt: Optional[str] = None, - speed: Optional[str] = None, - tool_name_reverse_map: Optional[Dict[str, str]] = None, + json_mode: bool | None = None, + prefix_prompt: str | None = None, + speed: str | None = None, + tool_name_reverse_map: dict[str, str] | None = None, ): - _hidden_params: Dict = {} + _hidden_params: dict = {} _hidden_params["additional_headers"] = process_anthropic_headers(dict(raw_response.headers)) if "error" in completion_response: response_headers = getattr(raw_response, "headers", None) @@ -2420,7 +2407,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): model_response._hidden_params = _hidden_params return model_response - def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]: + def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None: """ Get the prefix prompt from the messages. @@ -2445,13 +2432,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2467,14 +2454,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): except Exception as e: response_headers = getattr(raw_response, "headers", None) raise AnthropicError( - message="Unable to get json response - {}, Original Response: {}".format(str(e), raw_response.text), + message=f"Unable to get json response - {e!s}, Original Response: {raw_response.text}", status_code=raw_response.status_code, headers=response_headers, ) prefix_prompt = self.get_prefix_prompt(messages=messages) speed = optional_params.get("speed") - tool_name_reverse_map: Optional[Dict[str, str]] = None + tool_name_reverse_map: dict[str, str] | None = None if isinstance(litellm_params, dict): _candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY) if isinstance(_candidate, dict): @@ -2493,14 +2480,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _convert_tool_response_to_message( - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Optional[LitellmMessage]: + tool_calls: list[ChatCompletionToolCallChunk], + ) -> LitellmMessage | None: """ In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format """ ## HANDLE JSON MODE - anthropic returns single function call - json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments") + json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments") try: if json_mode_content_str is not None: args = json.loads(json_mode_content_str) @@ -2517,9 +2504,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return litellm.Message(content=json_mode_content_str) return None - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return AnthropicError( status_code=status_code, message=error_message, diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 256fee6b166..73d897c8d28 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -4,7 +4,7 @@ This file contains common utils for anthropic calls. import copy import re -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -44,17 +44,16 @@ def _strip_bedrock_id_suffixes(model: str) -> str: ) -def is_anthropic_oauth_key(value: Optional[str]) -> bool: +def is_anthropic_oauth_key(value: str | None) -> bool: """Check if a value contains an Anthropic OAuth token (sk-ant-oat*).""" if value is None: return False # Handle both raw token and "Bearer " format - if value.startswith("Bearer "): - value = value[7:] + value = value.removeprefix("Bearer ") return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX) -def _merge_beta_headers(existing: Optional[str], new_beta: str) -> str: +def _merge_beta_headers(existing: str | None, new_beta: str) -> str: """Merge a new beta value into an existing comma-separated anthropic-beta header.""" if not existing: return new_beta @@ -63,7 +62,7 @@ def _merge_beta_headers(existing: Optional[str], new_beta: str) -> str: return ",".join(sorted(betas)) -def optionally_handle_anthropic_oauth(headers: dict, api_key: Optional[str]) -> tuple[dict, Optional[str]]: +def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tuple[dict, str | None]: """ Handle Anthropic OAuth token detection and header setup. @@ -99,13 +98,13 @@ class AnthropicError(BaseLLMException): self, status_code: int, message, - headers: Optional[httpx.Headers] = None, + headers: httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) class AnthropicModelInfo(BaseLLMModelInfo): - def is_cache_control_set(self, messages: List[AllMessageValues]) -> bool: + def is_cache_control_set(self, messages: list[AllMessageValues]) -> bool: """ Return if {"cache_control": ..} in message content block @@ -122,21 +121,21 @@ class AnthropicModelInfo(BaseLLMModelInfo): return False - def is_file_id_used(self, messages: List[AllMessageValues]) -> bool: + def is_file_id_used(self, messages: list[AllMessageValues]) -> bool: """ Return if {"source": {"type": "file", "file_id": ..}} in message content block """ file_ids = get_file_ids_from_messages(messages) return len(file_ids) > 0 - def is_mcp_server_used(self, mcp_servers: Optional[List[AnthropicMcpServerTool]]) -> bool: + def is_mcp_server_used(self, mcp_servers: list[AnthropicMcpServerTool] | None) -> bool: if mcp_servers is None: return False if mcp_servers: return True return False - def is_computer_tool_used(self, tools: Optional[List[AllAnthropicToolsValues]]) -> Optional[str]: + def is_computer_tool_used(self, tools: list[AllAnthropicToolsValues] | None) -> str | None: """Returns the computer tool version if used, e.g. 'computer_20250124' or None""" if tools is None: return None @@ -145,7 +144,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return tool["type"] return None - def is_web_search_tool_used(self, tools: Optional[List[AllAnthropicToolsValues]]) -> bool: + def is_web_search_tool_used(self, tools: list[AllAnthropicToolsValues] | None) -> bool: """Returns True if web_search tool is used""" if tools is None: return False @@ -154,7 +153,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return True return False - def is_pdf_used(self, messages: List[AllMessageValues]) -> bool: + def is_pdf_used(self, messages: list[AllMessageValues]) -> bool: """ Set to true if media passed into messages. @@ -166,7 +165,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return True return False - def is_tool_search_used(self, tools: Optional[List]) -> bool: + def is_tool_search_used(self, tools: list | None) -> bool: """ Check if tool search tools are present in the tools list. """ @@ -182,7 +181,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return True return False - def is_programmatic_tool_calling_used(self, tools: Optional[List]) -> bool: + def is_programmatic_tool_calling_used(self, tools: list | None) -> bool: """ Check if programmatic tool calling is being used (tools with allowed_callers field). @@ -208,7 +207,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return False - def is_input_examples_used(self, tools: Optional[List]) -> bool: + def is_input_examples_used(self, tools: list | None) -> bool: """ Check if input_examples is being used in any tools. @@ -297,7 +296,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return model @staticmethod - def _model_map_lookup_candidates(model: str) -> List[str]: + def _model_map_lookup_candidates(model: str) -> list[str]: """Model-map keys to try for ``model``: the id itself, the same id with a bedrock/vertex routing prefix removed, the Bedrock base model, and each of those normalized by stripping a Bedrock version suffix (``-v1:0`` fully or @@ -337,7 +336,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return list(dict.fromkeys((*primary, *normalized))) @staticmethod - def _get_model_capability(model: str, key: str) -> Optional[bool]: + def _get_model_capability(model: str, key: str) -> bool | None: """Read boolean capability ``key`` from the model map, or None when no entry declares it.""" from litellm.utils import _get_bundled_model_cost_map @@ -354,7 +353,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return None @staticmethod - def _get_exact_model_capability(model: str, key: str) -> Optional[bool]: + def _get_exact_model_capability(model: str, key: str) -> bool | None: """Read boolean capability ``key`` from the exact model-map entry only. Unlike ``_get_model_capability``, does not walk stripped provider aliases. @@ -364,7 +363,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return value if isinstance(value, bool) else None @staticmethod - def _get_provider_resolved_capability(model: str, key: str, custom_llm_provider: str) -> Optional[bool]: + def _get_provider_resolved_capability(model: str, key: str, custom_llm_provider: str) -> bool | None: """Resolve boolean capability ``key`` for ``model`` under the caller's provider. Returns the flag when the provider-aware lookup resolves ``model`` to an @@ -422,8 +421,8 @@ class AnthropicModelInfo(BaseLLMModelInfo): def is_effort_used( self, - optional_params: Optional[dict], - model: Optional[str] = None, + optional_params: dict | None, + model: str | None = None, *, custom_llm_provider: str, ) -> bool: @@ -456,7 +455,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return False - def is_code_execution_tool_used(self, tools: Optional[List]) -> bool: + def is_code_execution_tool_used(self, tools: list | None) -> bool: """ Check if code execution tool is being used. @@ -471,7 +470,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return True return False - def is_container_with_skills_used(self, optional_params: Optional[dict]) -> bool: + def is_container_with_skills_used(self, optional_params: dict | None) -> bool: """ Check if container with skills is being used. @@ -487,7 +486,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return True return False - def _get_user_anthropic_beta_headers(self, anthropic_beta_header: Optional[str]) -> Optional[List[str]]: + def _get_user_anthropic_beta_headers(self, anthropic_beta_header: str | None) -> list[str] | None: if anthropic_beta_header is None: return None return anthropic_beta_header.split(",") @@ -514,14 +513,14 @@ class AnthropicModelInfo(BaseLLMModelInfo): def get_anthropic_beta_list( self, model: str, - optional_params: Optional[dict] = None, - computer_tool_used: Optional[str] = None, + optional_params: dict | None = None, + computer_tool_used: str | None = None, prompt_caching_set: bool = False, file_id_used: bool = False, mcp_server_used: bool = False, *, custom_llm_provider: str, - ) -> List[str]: + ) -> list[str]: """ Get list of common beta headers based on the features that are active. @@ -566,10 +565,10 @@ class AnthropicModelInfo(BaseLLMModelInfo): def get_anthropic_headers( self, - api_key: Optional[str] = None, - auth_token: Optional[str] = None, - anthropic_version: Optional[str] = None, - computer_tool_used: Optional[str] = None, + api_key: str | None = None, + auth_token: str | None = None, + anthropic_version: str | None = None, + computer_tool_used: str | None = None, prompt_caching_set: bool = False, pdf_used: bool = False, file_id_used: bool = False, @@ -580,7 +579,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): input_examples_used: bool = False, effort_used: bool = False, is_vertex_request: bool = False, - user_anthropic_beta_headers: Optional[List[str]] = None, + user_anthropic_beta_headers: list[str] | None = None, code_execution_tool_used: bool = False, container_with_skills_used: bool = False, api_base: str | None = None, @@ -654,12 +653,12 @@ class AnthropicModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: if api_base is None and isinstance(litellm_params, dict): api_base = litellm_params.get("api_base") use_bearer_for_custom_base: bool = bool( @@ -669,7 +668,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) api_key = AnthropicModelInfo.get_api_key(api_key) # Resolve auth_token from ANTHROPIC_AUTH_TOKEN if api_key is not set - auth_token: Optional[str] = None + auth_token: str | None = None if api_key is None: auth_token = AnthropicModelInfo.get_auth_token() if api_key is None and auth_token is None: @@ -721,7 +720,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return headers @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: from litellm.secret_managers.main import get_secret_str return ( @@ -732,13 +731,13 @@ class AnthropicModelInfo(BaseLLMModelInfo): ) @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: from litellm.secret_managers.main import get_secret_str return api_key or get_secret_str("ANTHROPIC_API_KEY") @staticmethod - def get_auth_token(auth_token: Optional[str] = None) -> Optional[str]: + def get_auth_token(auth_token: str | None = None) -> str | None: """Get auth token from ANTHROPIC_AUTH_TOKEN env var. Unlike api_key (which uses X-Api-Key header), auth_token uses @@ -771,10 +770,10 @@ class AnthropicModelInfo(BaseLLMModelInfo): return None @staticmethod - def get_base_model(model: Optional[str] = None) -> Optional[str]: + def get_base_model(model: str | None = None) -> str | None: return model.replace("anthropic/", "") if model else None - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base = AnthropicModelInfo.get_api_base(api_base) auth_header = AnthropicModelInfo.get_auth_header(api_key, api_base) if api_base is None or auth_header is None: @@ -804,7 +803,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): litellm_model_names.append(litellm_model_name) return litellm_model_names - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create an Anthropic token counter. @@ -818,7 +817,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return AnthropicTokenCounter() -def strip_advisor_blocks_from_messages(messages: List[Any], replace_with_text: bool = False) -> List[Any]: +def strip_advisor_blocks_from_messages(messages: list[Any], replace_with_text: bool = False) -> list[Any]: """ Remove (or replace) server_tool_use (name='advisor') and advisor_tool_result blocks from assistant message content. @@ -919,7 +918,7 @@ def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool: return "thinking" in lower and "signature" in lower and ("invalid" in lower or "valid string" in lower) -def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]: +def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[Any]: """ Return a new message list with thinking / redacted_thinking content blocks removed from each message. Used to recover from invalid thinking signatures on retry. @@ -927,7 +926,7 @@ def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[A Messages whose content is a list and becomes empty after stripping are omitted, since Anthropic rejects empty content arrays. """ - out: List[Any] = [] + out: list[Any] = [] for m in messages: if not isinstance(m, dict): out.append(m) @@ -946,7 +945,7 @@ def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[A def strip_thinking_blocks_from_anthropic_messages_request_dict( - data: Dict[str, Any], + data: dict[str, Any], ) -> None: """ Mutate an Anthropic Messages-style request dict: strip thinking blocks from @@ -959,8 +958,8 @@ def strip_thinking_blocks_from_anthropic_messages_request_dict( def strip_empty_text_blocks_from_anthropic_messages( - messages: List[Any], -) -> List[Any]: + messages: list[Any], +) -> list[Any]: """ Return a new message list with empty or whitespace-only ``{"type": "text"}`` content blocks removed. @@ -980,7 +979,7 @@ def strip_empty_text_blocks_from_anthropic_messages( The caller's list and its content blocks are never mutated; modified messages are returned as shallow copies with a fresh content list. """ - out: List[Any] = [] + out: list[Any] = [] for m in messages: if not isinstance(m, dict) or not isinstance(m.get("content"), list): out.append(m) @@ -1058,7 +1057,7 @@ def sanitize_tool_use_ids_in_anthropic_messages(messages: list[Any]) -> list[Any return out -def process_anthropic_headers(headers: Union[httpx.Headers, dict]) -> dict: +def process_anthropic_headers(headers: httpx.Headers | dict) -> dict: openai_headers = {} if "anthropic-ratelimit-requests-limit" in headers: openai_headers["x-ratelimit-limit-requests"] = headers["anthropic-ratelimit-requests-limit"] diff --git a/litellm/llms/anthropic/completion/transformation.py b/litellm/llms/anthropic/completion/transformation.py index 4b652512322..7fa9f189560 100644 --- a/litellm/llms/anthropic/completion/transformation.py +++ b/litellm/llms/anthropic/completion/transformation.py @@ -7,7 +7,6 @@ Litellm provider slug: `anthropic_text/` import json import time from collections.abc import AsyncIterator, Iterator -from typing import Dict, List, Optional, Union import httpx @@ -54,21 +53,21 @@ class AnthropicTextConfig(BaseConfig): to pass metadata to anthropic, it's {"user_id": "any-relevant-information"} """ - max_tokens_to_sample: Optional[int] = litellm.max_tokens # anthropic requires a default - stop_sequences: Optional[list] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - metadata: Optional[dict] = None + max_tokens_to_sample: int | None = litellm.max_tokens # anthropic requires a default + stop_sequences: list | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + metadata: dict | None = None def __init__( self, - max_tokens_to_sample: Optional[int] = DEFAULT_MAX_TOKENS, # anthropic requires a default - stop_sequences: Optional[list] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - metadata: Optional[dict] = None, + max_tokens_to_sample: int | None = DEFAULT_MAX_TOKENS, # anthropic requires a default + stop_sequences: list | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + metadata: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -80,11 +79,11 @@ class AnthropicTextConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: raise ValueError( @@ -102,7 +101,7 @@ class AnthropicTextConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -179,12 +178,12 @@ class AnthropicTextConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: completion_response = raw_response.json() @@ -220,9 +219,7 @@ class AnthropicTextConfig(BaseConfig): setattr(model_response, "usage", usage) return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return AnthropicTextError( status_code=status_code, message=error_message, @@ -232,7 +229,7 @@ class AnthropicTextConfig(BaseConfig): def _is_anthropic_text_model(model: str) -> bool: return model == "claude-2" or model == "claude-instant-1" - def _get_anthropic_text_prompt_from_messages(self, messages: List[AllMessageValues], model: str) -> str: + def _get_anthropic_text_prompt_from_messages(self, messages: list[AllMessageValues], model: str) -> str: custom_prompt_dict = litellm.custom_prompt_dict if model in custom_prompt_dict: # check if the model has a registered custom prompt @@ -250,9 +247,9 @@ class AnthropicTextConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return AnthropicTextCompletionResponseIterator( streaming_response=streaming_response, @@ -265,10 +262,10 @@ class AnthropicTextCompletionResponseIterator(BaseModelResponseIterator): def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: try: text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None provider_specific_fields = None index = int(chunk.get("index", 0)) _chunk_text = chunk.get("completion", None) diff --git a/litellm/llms/anthropic/cost_calculation.py b/litellm/llms/anthropic/cost_calculation.py index 82a97b53d28..3f028c05778 100644 --- a/litellm/llms/anthropic/cost_calculation.py +++ b/litellm/llms/anthropic/cost_calculation.py @@ -3,7 +3,7 @@ Helper util for handling anthropic-specific cost calculation - e.g.: prompt caching """ -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING, Optional from pydantic import BaseModel, ValidationError @@ -56,7 +56,7 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage", service_ti return cache_cost -def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) -> Tuple[float, float]: +def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/anthropic/count_tokens/__init__.py b/litellm/llms/anthropic/count_tokens/__init__.py index ef46862bda6..10ef4a9366c 100644 --- a/litellm/llms/anthropic/count_tokens/__init__.py +++ b/litellm/llms/anthropic/count_tokens/__init__.py @@ -9,7 +9,7 @@ from litellm.llms.anthropic.count_tokens.transformation import ( ) __all__ = [ - "AnthropicCountTokensHandler", "AnthropicCountTokensConfig", + "AnthropicCountTokensHandler", "AnthropicTokenCounter", ] diff --git a/litellm/llms/anthropic/count_tokens/handler.py b/litellm/llms/anthropic/count_tokens/handler.py index e70e0f19b33..b1584b98456 100644 --- a/litellm/llms/anthropic/count_tokens/handler.py +++ b/litellm/llms/anthropic/count_tokens/handler.py @@ -4,7 +4,7 @@ Anthropic CountTokens API handler. Uses httpx for HTTP requests instead of the Anthropic SDK. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -27,13 +27,13 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): async def handle_count_tokens_request( self, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], api_key: str, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Dict[str, Any]: + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> dict[str, Any]: """ Handle a CountTokens request using httpx. @@ -109,14 +109,14 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): raise except httpx.HTTPStatusError as e: # HTTP errors - preserve the actual status code - verbose_logger.error(f"HTTP error in CountTokens handler: {str(e)}") + verbose_logger.error(f"HTTP error in CountTokens handler: {e!s}") raise AnthropicError( status_code=e.response.status_code, message=e.response.text, ) except Exception as e: - verbose_logger.error(f"Error in CountTokens handler: {str(e)}") + verbose_logger.error(f"Error in CountTokens handler: {e!s}") raise AnthropicError( status_code=500, - message=f"CountTokens processing error: {str(e)}", + message=f"CountTokens processing error: {e!s}", ) diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index 89249ec42f0..8cc9d2ec0a9 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -3,7 +3,7 @@ Anthropic Token Counter implementation using the CountTokens API. """ import os -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler @@ -19,20 +19,20 @@ class AnthropicTokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: return custom_llm_provider == LlmProviders.ANTHROPIC.value async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: """ Count tokens using Anthropic's CountTokens API. diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index ad5bbbda25f..7c70d9ed0eb 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -4,7 +4,7 @@ Anthropic CountTokens API transformation logic. This module handles the transformation of requests to Anthropic's CountTokens API format. """ -from typing import Any, Dict, List, Optional +from typing import Any from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION @@ -31,16 +31,16 @@ class AnthropicCountTokensConfig: def transform_request_to_count_tokens( self, model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Dict[str, Any]: + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> dict[str, Any]: """ Transform request to Anthropic CountTokens format. Includes optional system and tools fields for accurate token counting. """ - request: Dict[str, Any] = { + request: dict[str, Any] = { "model": model, "messages": messages, } @@ -53,7 +53,7 @@ class AnthropicCountTokensConfig: return request - def get_required_headers(self, api_key: str) -> Dict[str, str]: + def get_required_headers(self, api_key: str) -> dict[str, str]: """ Get the required headers for the CountTokens API. @@ -67,7 +67,7 @@ class AnthropicCountTokensConfig: optionally_handle_anthropic_oauth, ) - headers: Dict[str, str] = { + headers: dict[str, str] = { "Content-Type": "application/json", "x-api-key": api_key, "anthropic-version": "2023-06-01", @@ -76,7 +76,7 @@ class AnthropicCountTokensConfig: headers, _ = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) return headers - def validate_request(self, model: str, messages: List[Dict[str, Any]]) -> None: + def validate_request(self, model: str, messages: list[dict[str, Any]]) -> None: """ Validate the incoming count tokens request. diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py index 5870a3bf7de..a5aa1509969 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py @@ -1,12 +1,6 @@ from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( - TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Union, cast, ) @@ -30,15 +24,11 @@ from litellm.types.llms.anthropic_messages.anthropic_response import ( from litellm.types.utils import ModelResponse from litellm.utils import get_model_info -if TYPE_CHECKING: - pass - - # Anthropic-only keys already mapped by the translator; strip on extra_kwargs re-merge. ANTHROPIC_ONLY_REQUEST_KEYS: frozenset[str] = frozenset({"output_config"}) -def _messages_have_compaction_block(messages: List[Dict]) -> bool: +def _messages_have_compaction_block(messages: list[dict]) -> bool: """Return True when any message carries a ``compaction`` content block.""" for msg in messages: content = msg.get("content") @@ -50,7 +40,7 @@ def _messages_have_compaction_block(messages: List[Dict]) -> bool: return False -def _extract_proxy_litellm_metadata(kwargs: Dict[str, Any]) -> Optional[Dict[str, Any]]: +def _extract_proxy_litellm_metadata(kwargs: dict[str, Any]) -> dict[str, Any] | None: """Return ``kwargs["litellm_metadata"]`` when it's a dict; ``None`` otherwise. The proxy attaches its auth/spend-attribution fields (``user_api_key``, @@ -71,15 +61,15 @@ def _extract_proxy_litellm_metadata(kwargs: Dict[str, Any]) -> Optional[Dict[str async def _prepare_context_managed_request( *, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - system: Optional[Any], + messages: list[dict], + tools: list[dict] | None, + system: Any | None, context_management_spec: Any, - litellm_metadata: Optional[Dict], - additional_drop_params: Optional[list[str]], + litellm_metadata: dict | None, + additional_drop_params: list[str] | None, llm_router: Any, user_api_key_auth: Any = None, -) -> Optional[PolyfillResult]: +) -> PolyfillResult | None: """Apply client compaction history, then optional context_management polyfill.""" from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import ( apply_client_compaction_block_history, @@ -97,12 +87,12 @@ async def _prepare_context_managed_request( ) if polyfill_will_run: - history_result: Optional[PolyfillResult] = None - working_messages: List[Dict] = messages - working_system: Optional[Any] = system + history_result: PolyfillResult | None = None + working_messages: list[dict] = messages + working_system: Any | None = system else: history_result = apply_client_compaction_block_history( - messages=cast(List[Dict[str, Any]], messages), + messages=cast(list[dict[str, Any]], messages), system=system, ) working_messages = history_result.messages if history_result is not None else messages @@ -132,7 +122,7 @@ async def _prepare_context_managed_request( # to non-Anthropic backends that would reject them. if polyfill_will_run and history_result is None: history_result = apply_client_compaction_block_history( - messages=cast(List[Dict[str, Any]], messages), + messages=cast(list[dict[str, Any]], messages), system=system, ) return history_result @@ -141,7 +131,7 @@ async def _prepare_context_managed_request( def _polyfill_will_run( *, context_management_spec: Any, - additional_drop_params: Optional[list[str]], + additional_drop_params: list[str] | None, ) -> bool: """Return True when ``compact_20260112`` will run via the polyfill dispatcher. @@ -168,7 +158,7 @@ def _polyfill_will_run( def _spec_has_non_compact_edits( *, context_management_spec: Any, - additional_drop_params: Optional[list[str]], + additional_drop_params: list[str] | None, ) -> bool: """Return True when the spec includes edits other than ``compact_20260112``. @@ -194,7 +184,7 @@ def _spec_has_non_compact_edits( ) -def _context_management_explicitly_dropped(additional_drop_params: Optional[list[str]]) -> bool: +def _context_management_explicitly_dropped(additional_drop_params: list[str] | None) -> bool: """True when the caller opted out of context_management via ``additional_drop_params``. ``drop_params`` deliberately does NOT gate the polyfill: ``context_management`` @@ -209,8 +199,8 @@ def _context_management_explicitly_dropped(additional_drop_params: Optional[list def _normalize_spec_edits( *, context_management_spec: Any, - additional_drop_params: Optional[list[str]], -) -> Optional[List[Dict[str, Any]]]: + additional_drop_params: list[str] | None, +) -> list[dict[str, Any]] | None: """Return the normalized ``edits`` list, or ``None`` if the polyfill won't run. Delegates spec-shape normalization to the dispatcher's ``_normalize_spec`` @@ -235,15 +225,15 @@ def _normalize_spec_edits( async def _run_polyfill_if_enabled( *, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - system: Optional[Any], + messages: list[dict], + tools: list[dict] | None, + system: Any | None, context_management_spec: Any, - litellm_metadata: Optional[Dict], - additional_drop_params: Optional[list[str]], + litellm_metadata: dict | None, + additional_drop_params: list[str] | None, llm_router: Any, user_api_key_auth: Any = None, -) -> Optional[PolyfillResult]: +) -> PolyfillResult | None: """Run the async context_management polyfill if a spec is present. Returns ``None`` when the spec is empty or ``context_management`` is @@ -303,9 +293,9 @@ ANTHROPIC_ADAPTER = AnthropicAdapter() class LiteLLMMessagesToCompletionTransformationHandler: @staticmethod def _route_openai_thinking_to_responses_api_if_needed( - completion_kwargs: Dict[str, Any], + completion_kwargs: dict[str, Any], *, - thinking: Optional[Dict[str, Any]], + thinking: dict[str, Any] | None, ) -> None: """ When users call `litellm.anthropic.messages.*` with a non-Anthropic model and @@ -352,7 +342,7 @@ class LiteLLMMessagesToCompletionTransformationHandler: reasoning_effort = completion_kwargs.get("reasoning_effort") summary = thinking.get("summary") if isinstance(reasoning_effort, str) and reasoning_effort: - reasoning_dict: Dict[str, Any] = {"effort": reasoning_effort} + reasoning_dict: dict[str, Any] = {"effort": reasoning_effort} if summary: reasoning_dict["summary"] = summary elif auto_summary: @@ -368,7 +358,7 @@ class LiteLLMMessagesToCompletionTransformationHandler: @staticmethod def _normalize_reasoning_effort( - completion_kwargs: Dict[str, Any], + completion_kwargs: dict[str, Any], ) -> None: """ Normalize reasoning_effort values based on target model capabilities. @@ -406,21 +396,21 @@ class LiteLLMMessagesToCompletionTransformationHandler: def _prepare_completion_kwargs( *, max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[Union[str, List[Dict[str, Any]]]] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, - extra_kwargs: Optional[Dict[str, Any]] = None, - ) -> Tuple[Dict[str, Any], Dict[str, str]]: + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | list[dict[str, Any]] | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, + extra_kwargs: dict[str, Any] | None = None, + ) -> tuple[dict[str, Any], dict[str, str]]: """Prepare kwargs for litellm.completion/acompletion. Returns: @@ -477,7 +467,7 @@ class LiteLLMMessagesToCompletionTransformationHandler: if openai_request is None: raise ValueError("Failed to translate request to OpenAI format") - completion_kwargs: Dict[str, Any] = dict(openai_request) + completion_kwargs: dict[str, Any] = dict(openai_request) if stream: completion_kwargs["stream"] = stream @@ -527,24 +517,24 @@ class LiteLLMMessagesToCompletionTransformationHandler: @staticmethod async def async_anthropic_messages_handler( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, **kwargs, - ) -> Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]]: + ) -> AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]: """Handle non-Anthropic models asynchronously using the adapter""" context_management = kwargs.pop("context_management", None) - additional_drop_params: Optional[list[str]] = kwargs.get("additional_drop_params", None) + additional_drop_params: list[str] | None = kwargs.get("additional_drop_params", None) litellm_router = kwargs.pop("litellm_router", None) if litellm_router is None: try: @@ -621,31 +611,27 @@ class LiteLLMMessagesToCompletionTransformationHandler: @staticmethod def anthropic_messages_handler( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, _is_async: bool = False, **kwargs, - ) -> Union[ - AnthropicMessagesResponse, - Iterator[bytes], - AsyncIterator[Any], - Coroutine[ - Any, - Any, - Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]], - ], - ]: + ) -> ( + AnthropicMessagesResponse + | Iterator[bytes] + | AsyncIterator[Any] + | Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]] + ): """Handle non-Anthropic models using the adapter.""" if _is_async is True: return LiteLLMMessagesToCompletionTransformationHandler.async_anthropic_messages_handler( @@ -672,7 +658,7 @@ class LiteLLMMessagesToCompletionTransformationHandler: # ``compact_20260112`` editor can ``await`` the summarization model); # bridge to it via ``run_async_function``. context_management = kwargs.pop("context_management", None) - additional_drop_params: Optional[list[str]] = kwargs.get("additional_drop_params", None) + additional_drop_params: list[str] | None = kwargs.get("additional_drop_params", None) # Deliberately do NOT auto-attach the proxy ``llm_router`` here: # ``run_async_function`` spawns a new event loop in a worker thread # to bridge to the async dispatcher, but the proxy router's httpx @@ -693,7 +679,7 @@ class LiteLLMMessagesToCompletionTransformationHandler: # and bridging through a worker-thread event loop just to discover # there is no work is pure overhead. if context_management is None and not _messages_have_compaction_block(messages): - polyfill_result: Optional[PolyfillResult] = None + polyfill_result: PolyfillResult | None = None else: proxy_litellm_metadata = _extract_proxy_litellm_metadata(kwargs) user_api_key_auth = ( diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index f1261ef1413..bb61043742d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -8,10 +8,7 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, get_args, ) @@ -72,8 +69,8 @@ class _CombinedChunkSplitter: def __init__(self, completion_stream: Any): self._stream = completion_stream - self._sync_iter: Optional[Iterator[Any]] = None - self._async_iter: Optional[AsyncIterator[Any]] = None + self._sync_iter: Iterator[Any] | None = None + self._async_iter: AsyncIterator[Any] | None = None self._buffer: deque = deque() @staticmethod @@ -180,7 +177,7 @@ class _CombinedChunkSplitter: return {"reasoning_content": thinking_text} @staticmethod - def _split(chunk: Any) -> List[Any]: + def _split(chunk: Any) -> list[Any]: """Return ``[chunk]``, or ``[content_chunk, finish_chunk]`` if combined.""" if not _CombinedChunkSplitter._is_combined(chunk): return [chunk] @@ -254,8 +251,8 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): sent_content_block_finish: bool = False current_content_block_type: Literal["text", "tool_use", "thinking"] = "text" sent_last_message: bool = False - holding_chunk: Optional[Any] = None - holding_stop_reason_chunk: Optional[Any] = None + holding_chunk: Any | None = None + holding_stop_reason_chunk: Any | None = None queued_usage_chunk: bool = False current_content_block_index: int = 0 @@ -263,10 +260,10 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): self, completion_stream: Any, model: str, - tool_name_mapping: Optional[Dict[str, str]] = None, - applied_edits: Optional[List[AppliedEdit]] = None, - compaction_block: Optional[CompactionBlock] = None, - iterations_usage: Optional[List[UsageIteration]] = None, + tool_name_mapping: dict[str, str] | None = None, + applied_edits: list[AppliedEdit] | None = None, + compaction_block: CompactionBlock | None = None, + iterations_usage: list[UsageIteration] | None = None, ): # Wrap the upstream stream so chunks that carry both content and a # finish_reason (fake-streamed providers) are split into two — see @@ -276,7 +273,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # Mapping of truncated tool names to original names (for OpenAI's 64-char limit) self.tool_name_mapping = tool_name_mapping or {} # Polyfill applied_edits on final message_delta. - self.applied_edits: List[AppliedEdit] = list(applied_edits or []) + self.applied_edits: list[AppliedEdit] = list(applied_edits or []) # Synthesized compaction block from compact_20260112 polyfill (streaming). self.compaction_block = compaction_block self.iterations_usage = iterations_usage @@ -297,12 +294,12 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # class level) so concurrent streams don't share the same mutable dict # — `_should_start_new_content_block` mutates `tool_block["name"]` in # place, which would otherwise leak across streams. - self.current_content_block_start: "AnthropicStreamWrapper.ContentBlockContentBlockDict" = self.TextBlock( + self.current_content_block_start: AnthropicStreamWrapper.ContentBlockContentBlockDict = self.TextBlock( type="text", text="", ) - def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> Dict[str, Any]: + def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> dict[str, Any]: """Merge usage data from ``chunk`` into the held ``message_delta`` chunk. Shared by both the sync ``__next__`` and async ``__anext__`` paths so @@ -328,7 +325,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): merged_chunk["context_management"] = ContextManagementResponse(applied_edits=list(self.applied_edits)) return self._augment_message_delta_usage(merged_chunk) - def _ensure_context_management_attached(self, message_delta_chunk: Dict[str, Any]) -> Dict[str, Any]: + def _ensure_context_management_attached(self, message_delta_chunk: dict[str, Any]) -> dict[str, Any]: """Attach ``context_management`` to a ``message_delta`` chunk if ``self.applied_edits`` is non-empty and the chunk does not already carry it. Returns the (possibly new) chunk dict. @@ -343,7 +340,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): augmented["context_management"] = ContextManagementResponse(applied_edits=list(self.applied_edits)) return augmented - def _augment_message_delta_usage(self, message_delta_chunk: Dict[str, Any]) -> Dict[str, Any]: + def _augment_message_delta_usage(self, message_delta_chunk: dict[str, Any]) -> dict[str, Any]: """Attach polyfill compaction iteration usage to the final message_delta. Also defensively re-attaches ``context_management`` so the direct @@ -361,7 +358,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): output_tokens = usage.get("output_tokens", 0) or 0 augmented = message_delta_chunk.copy() augmented_usage = dict(usage) - iterations: List[UsageIteration] = list(self.iterations_usage) + iterations: list[UsageIteration] = list(self.iterations_usage) # Only emit a ``message`` iteration when we have real token data. # Without a separate usage chunk (e.g. provider sent finish_reason # alone), the held ``message_delta`` carries placeholder zeros from @@ -378,7 +375,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): augmented["usage"] = augmented_usage return augmented - def _next_compaction_event(self) -> Optional[Dict[str, Any]]: + def _next_compaction_event(self) -> dict[str, Any] | None: """Return the next compaction content-block SSE event, or ``None``. Anthropic delivers compaction as a single delta (no token-by-token @@ -467,7 +464,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): { "type": "message_start", "message": { - "id": "msg_{}".format(uuid.uuid4()), + "id": f"msg_{uuid.uuid4()}", "type": "message", "role": "assistant", "content": [], @@ -672,7 +669,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): return {"type": "message_stop"} raise StopIteration except Exception as e: - verbose_logger.error("Anthropic Adapter - {}\n{}".format(e, traceback.format_exc())) + verbose_logger.error(f"Anthropic Adapter - {e}\n{traceback.format_exc()}") raise StopIteration async def __anext__(self): @@ -690,7 +687,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): { "type": "message_start", "message": { - "id": "msg_{}".format(uuid.uuid4()), + "id": f"msg_{uuid.uuid4()}", "type": "message", "role": "assistant", "content": [], @@ -930,7 +927,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): self.current_content_block_index += 1 @staticmethod - def _delta_has_content(processed_chunk: Dict[str, Any]) -> bool: + def _delta_has_content(processed_chunk: dict[str, Any]) -> bool: """Return True if a translated chunk carries a non-empty ``content_block_delta`` payload. diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 9826bcb2490..d9f24ff6c67 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -5,12 +5,7 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, - Union, cast, ) @@ -47,8 +42,8 @@ def truncate_tool_name(name: str) -> str: def create_tool_name_mapping( - tools: List[Dict[str, Any]], -) -> Dict[str, str]: + tools: list[dict[str, Any]], +) -> dict[str, str]: """ Create a mapping of truncated tool names to original names. @@ -58,7 +53,7 @@ def create_tool_name_mapping( Returns: Dict mapping truncated names to original names (only for truncated tools) """ - mapping: Dict[str, str] = {} + mapping: dict[str, str] = {} for tool in tools: original_name = tool.get("name", "") truncated_name = truncate_tool_name(original_name) @@ -143,7 +138,7 @@ class AnthropicAdapter: def __init__(self) -> None: pass - def translate_completion_input_params(self, kwargs) -> Optional[ChatCompletionRequest]: + def translate_completion_input_params(self, kwargs) -> ChatCompletionRequest | None: """ Translate Anthropic request params to OpenAI format. @@ -158,7 +153,7 @@ class AnthropicAdapter: def translate_completion_input_params_with_tool_mapping( self, kwargs - ) -> Tuple[Optional[ChatCompletionRequest], Dict[str, str]]: + ) -> tuple[ChatCompletionRequest | None, dict[str, str]]: """ Translate Anthropic request params to OpenAI format, returning tool name mapping. @@ -195,9 +190,9 @@ class AnthropicAdapter: def translate_completion_output_params( self, response: ModelResponse, - tool_name_mapping: Optional[Dict[str, str]] = None, - polyfill_result: Optional[PolyfillResult] = None, - ) -> Optional[AnthropicMessagesResponse]: + tool_name_mapping: dict[str, str] | None = None, + polyfill_result: PolyfillResult | None = None, + ) -> AnthropicMessagesResponse | None: """ Translate OpenAI response to Anthropic format. @@ -218,10 +213,10 @@ class AnthropicAdapter: self, completion_stream: Any, model: str, - tool_name_mapping: Optional[Dict[str, str]] = None, - polyfill_result: Optional[PolyfillResult] = None, + tool_name_mapping: dict[str, str] | None = None, + polyfill_result: PolyfillResult | None = None, is_async: bool = True, - ) -> Union[AsyncIterator[bytes], Iterator[bytes], None]: + ) -> AsyncIterator[bytes] | Iterator[bytes] | None: """ Translate OpenAI streaming response to Anthropic format. @@ -260,7 +255,7 @@ class LiteLLMAnthropicMessagesAdapter: ### FOR [BETA] `/v1/messages` endpoint support - def _extract_signature_from_tool_call(self, tool_call: Any) -> Optional[str]: + def _extract_signature_from_tool_call(self, tool_call: Any) -> str | None: """ Extract signature from a tool call's provider_specific_fields. Only checks provider_specific_fields, not thinking blocks. @@ -276,7 +271,7 @@ class LiteLLMAnthropicMessagesAdapter: return signature - def _extract_signature_from_tool_use_content(self, content: Dict[str, Any]) -> Optional[str]: + def _extract_signature_from_tool_use_content(self, content: dict[str, Any]) -> str | None: """ Extract signature from a tool_use content block's provider_specific_fields. """ @@ -289,7 +284,7 @@ class LiteLLMAnthropicMessagesAdapter: self, source: Any, target: Any, - model: Optional[str], + model: str | None, ) -> None: """ Extract cache_control from source and add to target if it should be preserved. @@ -315,9 +310,9 @@ class LiteLLMAnthropicMessagesAdapter: target["cache_control"] = cache_control # type: ignore[typeddict-item] else: # Fallback for non-dict objects (shouldn't happen in practice) - cast(Dict[str, Any], target)["cache_control"] = cache_control + cast(dict[str, Any], target)["cache_control"] = cache_control - def translatable_anthropic_params(self) -> List: + def translatable_anthropic_params(self) -> list: """ Which anthropic params, we need to translate to the openai format. """ @@ -333,7 +328,7 @@ class LiteLLMAnthropicMessagesAdapter: "stop_sequences", ] - def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool: + def _is_web_search_tool(self, tool: dict[str, Any]) -> bool: """ Check if a tool is an Anthropic web search tool. @@ -353,19 +348,14 @@ class LiteLLMAnthropicMessagesAdapter: def translate_anthropic_messages_to_openai( self, - messages: List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], - model: Optional[str] = None, - ) -> List: - new_messages: List[AllMessageValues] = [] + messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam], + model: str | None = None, + ) -> list: + new_messages: list[AllMessageValues] = [] for m in messages: - user_message: Optional[ChatCompletionUserMessage] = None - tool_message_list: List[ChatCompletionToolMessage] = [] - new_user_content_list: List[Union[ChatCompletionTextObject, ChatCompletionImageObject]] = [] + user_message: ChatCompletionUserMessage | None = None + tool_message_list: list[ChatCompletionToolMessage] = [] + new_user_content_list: list[ChatCompletionTextObject | ChatCompletionImageObject] = [] ## USER MESSAGE ## if m["role"] == "user": ## translate user message @@ -456,11 +446,8 @@ class LiteLLMAnthropicMessagesAdapter: else: # For multiple content items, combine into a single tool message # with list content to preserve all items while having one tool_use_id - combined_content_parts: List[ - Union[ - ChatCompletionTextObject, - ChatCompletionImageObject, - ] + combined_content_parts: list[ + ChatCompletionTextObject | ChatCompletionImageObject ] = [] for c in content_items: if isinstance(c, str): @@ -507,11 +494,11 @@ class LiteLLMAnthropicMessagesAdapter: new_messages.append({"role": "user", "content": new_user_content_list}) # type: ignore ## ASSISTANT MESSAGE ## - assistant_message_str: Optional[str] = None - assistant_content_list: List[Dict[str, Any]] = [] # For content blocks with cache_control + assistant_message_str: str | None = None + assistant_content_list: list[dict[str, Any]] = [] # For content blocks with cache_control has_cache_control_in_text = False - tool_calls: List[ChatCompletionAssistantToolCall] = [] - thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = [] + tool_calls: list[ChatCompletionAssistantToolCall] = [] + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = [] if m["role"] == "assistant": if isinstance(m.get("content"), str): assistant_message_str = str(m.get("content", "")) @@ -521,7 +508,7 @@ class LiteLLMAnthropicMessagesAdapter: assistant_message_str = str(content) elif isinstance(content, dict): if content.get("type") == "text": - text_block: Dict[str, Any] = { + text_block: dict[str, Any] = { "type": "text", "text": content.get("text", ""), } @@ -536,10 +523,10 @@ class LiteLLMAnthropicMessagesAdapter: "name": tool_name, "arguments": json.dumps(content.get("input", {})), } - signature = self._extract_signature_from_tool_use_content(cast(Dict[str, Any], content)) + signature = self._extract_signature_from_tool_use_content(cast(dict[str, Any], content)) if signature: - provider_specific_fields: Dict[str, Any] = ( + provider_specific_fields: dict[str, Any] = ( function_chunk.get("provider_specific_fields") or {} ) provider_specific_fields["thought_signature"] = signature @@ -598,8 +585,8 @@ class LiteLLMAnthropicMessagesAdapter: @staticmethod def translate_anthropic_thinking_to_reasoning_effort( - thinking: Dict[str, Any], - ) -> Optional[str]: + thinking: dict[str, Any], + ) -> str | None: """ Translate Anthropic's thinking parameter to OpenAI's reasoning_effort. @@ -655,9 +642,9 @@ class LiteLLMAnthropicMessagesAdapter: @staticmethod def translate_thinking_for_model( - thinking: Dict[str, Any], + thinking: dict[str, Any], model: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Translate Anthropic thinking parameter based on the target model. @@ -693,7 +680,7 @@ class LiteLLMAnthropicMessagesAdapter: @staticmethod def _apply_reasoning_summary_wrapping( reasoning_effort: str, - thinking: Dict[str, Any], + thinking: dict[str, Any], ) -> Any: """ Apply the reasoning_effort/summary wrapping rules shared by every @@ -730,11 +717,11 @@ class LiteLLMAnthropicMessagesAdapter: elif tool_choice["type"] == "none": return "none" else: - raise ValueError("Incompatible tool choice param submitted - {}".format(tool_choice)) + raise ValueError(f"Incompatible tool choice param submitted - {tool_choice}") def translate_anthropic_tools_to_openai( - self, tools: List[AllAnthropicToolsValues], model: Optional[str] = None - ) -> Tuple[List[ChatCompletionToolParam], Dict[str, str]]: + self, tools: list[AllAnthropicToolsValues], model: str | None = None + ) -> tuple[list[ChatCompletionToolParam], dict[str, str]]: """ Translate Anthropic tools to OpenAI format. @@ -743,8 +730,8 @@ class LiteLLMAnthropicMessagesAdapter: - tool_name_mapping maps truncated names back to original names for tools that exceeded OpenAI's 64-char limit """ - new_tools: List[ChatCompletionToolParam] = [] - tool_name_mapping: Dict[str, str] = {} + new_tools: list[ChatCompletionToolParam] = [] + tool_name_mapping: dict[str, str] = {} # "type" is the Anthropic tool type (e.g. "custom"); it must not be # merged into the OpenAI function `parameters` schema below, or it # overwrites the real parameters.type ("object") and the provider @@ -793,7 +780,7 @@ class LiteLLMAnthropicMessagesAdapter: return new_tools, tool_name_mapping # type: ignore[return-value] - def translate_anthropic_output_format_to_openai(self, output_format: Any) -> Optional[Dict[str, Any]]: + def translate_anthropic_output_format_to_openai(self, output_format: Any) -> dict[str, Any] | None: """ Translate Anthropic's output_format to OpenAI's response_format. @@ -868,7 +855,7 @@ class LiteLLMAnthropicMessagesAdapter: def _add_system_message_to_messages( self, - new_messages: List[AllMessageValues], + new_messages: list[AllMessageValues], anthropic_message_request: AnthropicMessagesRequest, ) -> None: """Add system message to messages list if present in request.""" @@ -885,11 +872,11 @@ class LiteLLMAnthropicMessagesAdapter: ) elif isinstance(system_content, list): # Convert Anthropic system content blocks to OpenAI format - openai_system_content: List[Dict[str, Any]] = [] + openai_system_content: list[dict[str, Any]] = [] model_name = anthropic_message_request.get("model", "") for block in system_content: if isinstance(block, dict) and block.get("type") == "text": - text_block: Dict[str, Any] = { + text_block: dict[str, Any] = { "type": "text", "text": block.get("text", ""), } @@ -947,7 +934,7 @@ class LiteLLMAnthropicMessagesAdapter: self, anthropic_message_request: AnthropicMessagesRequest, new_kwargs: ChatCompletionRequest, - ) -> Dict[str, str]: + ) -> dict[str, str]: """Translate tools and extract web_search_options when needed.""" if "tools" not in anthropic_message_request: return {} @@ -956,10 +943,10 @@ class LiteLLMAnthropicMessagesAdapter: if not tools: return {} - web_search_tools: List[AllAnthropicToolsValues] = [] - regular_tools: List[AllAnthropicToolsValues] = [] + web_search_tools: list[AllAnthropicToolsValues] = [] + regular_tools: list[AllAnthropicToolsValues] = [] for tool in tools: - cast_tool = cast(Dict[str, Any], tool) + cast_tool = cast(dict[str, Any], tool) if self._is_web_search_tool(cast_tool): web_search_tools.append(cast(AllAnthropicToolsValues, tool)) else: @@ -996,7 +983,7 @@ class LiteLLMAnthropicMessagesAdapter: new_kwargs["thinking"] = thinking # type: ignore return - reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(Dict[str, Any], thinking)) + reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(dict[str, Any], thinking)) if not reasoning_effort: return @@ -1009,7 +996,7 @@ class LiteLLMAnthropicMessagesAdapter: reasoning_effort = output_config["effort"] new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping( - reasoning_effort, cast(Dict[str, Any], thinking) + reasoning_effort, cast(dict[str, Any], thinking) ) def _translate_output_format_to_openai( @@ -1053,7 +1040,7 @@ class LiteLLMAnthropicMessagesAdapter: def translate_anthropic_to_openai( self, anthropic_message_request: AnthropicMessagesRequest - ) -> Tuple[ChatCompletionRequest, Dict[str, str]]: + ) -> tuple[ChatCompletionRequest, dict[str, str]]: """ This is used by the beta Anthropic Adapter, for translating anthropic `/v1/messages` requests to the openai format. @@ -1063,17 +1050,12 @@ class LiteLLMAnthropicMessagesAdapter: for tools that exceeded OpenAI's 64-char limit """ # Debug: Processing Anthropic message request - new_messages: List[AllMessageValues] = [] - tool_name_mapping: Dict[str, str] = {} + new_messages: list[AllMessageValues] = [] + tool_name_mapping: dict[str, str] = {} ## CONVERT ANTHROPIC MESSAGES TO OPENAI - messages_list: List[Union[AnthropicMessagesUserMessageParam, AnthopicMessagesAssistantMessageParam]] = cast( - List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + messages_list: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam] = cast( + list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam], anthropic_message_request["messages"], ) new_messages = self.translate_anthropic_messages_to_openai( @@ -1124,7 +1106,7 @@ class LiteLLMAnthropicMessagesAdapter: return new_kwargs, tool_name_mapping - def _translate_anthropic_image_to_openai(self, image_source: dict) -> Optional[str]: + def _translate_anthropic_image_to_openai(self, image_source: dict) -> str | None: """ Translate Anthropic image source format to OpenAI-compatible image URL. @@ -1153,10 +1135,10 @@ class LiteLLMAnthropicMessagesAdapter: def _translate_openai_content_to_anthropic( self, - choices: List[Choices], - tool_name_mapping: Optional[Dict[str, str]] = None, - ) -> List[Dict[str, Any]]: - new_content: List[Dict[str, Any]] = [] + choices: list[Choices], + tool_name_mapping: dict[str, str] | None = None, + ) -> list[dict[str, Any]]: + new_content: list[dict[str, Any]] = [] for choice in choices: # Handle thinking blocks first if hasattr(choice.message, "thinking_blocks") and choice.message.thinking_blocks: @@ -1317,8 +1299,8 @@ class LiteLLMAnthropicMessagesAdapter: def translate_openai_response_to_anthropic( self, response: ModelResponse, - tool_name_mapping: Optional[Dict[str, str]] = None, - polyfill_result: Optional[PolyfillResult] = None, + tool_name_mapping: dict[str, str] | None = None, + polyfill_result: PolyfillResult | None = None, ) -> AnthropicMessagesResponse: """ Translate OpenAI response to Anthropic format. @@ -1373,8 +1355,8 @@ class LiteLLMAnthropicMessagesAdapter: return translated_obj def _translate_streaming_openai_chunk_to_anthropic_content_block( - self, choices: List[Union[OpenAIStreamingChoice, StreamingChoices]] - ) -> Tuple[ + self, choices: list[OpenAIStreamingChoice | StreamingChoices] + ) -> tuple[ Literal["text", "tool_use", "thinking"], "ContentBlockContentBlockDict", ]: @@ -1389,11 +1371,11 @@ class LiteLLMAnthropicMessagesAdapter: ): raw_id = choice.delta.tool_calls[0].id or str(uuid.uuid4()) tool_name = choice.delta.tool_calls[0].function.name or "" - thought_sig: Optional[str] = None + thought_sig: str | None = None if THOUGHT_SIGNATURE_SEPARATOR in raw_id: parts = raw_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1) thought_sig = parts[1] if len(parts) > 1 else None - tool_block: Dict[str, Any] = { + tool_block: dict[str, Any] = { "type": "tool_use", "id": normalize_anthropic_tool_use_id(raw_id), "name": tool_name, @@ -1431,20 +1413,15 @@ class LiteLLMAnthropicMessagesAdapter: return "text", TextBlock(type="text", text="") def _translate_streaming_openai_chunk_to_anthropic( - self, choices: List[Union[OpenAIStreamingChoice, StreamingChoices]] - ) -> Tuple[ + self, choices: list[OpenAIStreamingChoice | StreamingChoices] + ) -> tuple[ StreamingContentBlockDeltaType, - Union[ - ContentTextBlockDelta, - ContentJsonBlockDelta, - ContentThinkingBlockDelta, - ContentThinkingSignatureBlockDelta, - ], + ContentTextBlockDelta | ContentJsonBlockDelta | ContentThinkingBlockDelta | ContentThinkingSignatureBlockDelta, ]: text: str = "" reasoning_content: str = "" reasoning_signature: str = "" - partial_json: Optional[str] = None + partial_json: str | None = None for choice in choices: if choice.delta.content is not None and len(choice.delta.content) > 0: text += choice.delta.content @@ -1487,15 +1464,15 @@ class LiteLLMAnthropicMessagesAdapter: self, response: ModelResponse, current_content_block_index: int, - applied_edits: Optional[List[AppliedEdit]] = None, - ) -> Union[ContentBlockDelta, MessageBlockDelta]: + applied_edits: list[AppliedEdit] | None = None, + ) -> ContentBlockDelta | MessageBlockDelta: ## base case - final chunk w/ finish reason if response.choices[0].finish_reason is not None: delta = MessageDelta( stop_reason=self._translate_openai_finish_reason_to_anthropic(response.choices[0].finish_reason), ) if getattr(response, "usage", None) is not None: - litellm_usage_chunk: Optional[Usage] = response.usage # type: ignore + litellm_usage_chunk: Usage | None = response.usage # type: ignore elif hasattr(response, "_hidden_params") and "usage" in response._hidden_params: litellm_usage_chunk = response._hidden_params["usage"] else: diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py b/litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py index 729b2864524..032d4baadac 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py @@ -4,8 +4,8 @@ from .errors import AnthropicContextManagementError from .result import PolyfillResult __all__ = [ - "apply_context_management", - "AnthropicContextManagementError", "CLEARED_TOOL_RESULT_PLACEHOLDER", + "AnthropicContextManagementError", "PolyfillResult", + "apply_context_management", ] diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/dispatcher.py b/litellm/llms/anthropic/experimental_pass_through/context_management/dispatcher.py index d86210c0ebb..ac925ed6f22 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/dispatcher.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/dispatcher.py @@ -2,7 +2,7 @@ import inspect from collections.abc import Awaitable, Callable -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, cast from litellm._logging import verbose_logger from litellm.types.llms.anthropic import AppliedEdit @@ -13,15 +13,15 @@ from .result import PolyfillResult EditorFn = Callable[..., Any] -_EDITOR_REGISTRY: Dict[str, EditorFn] = { +_EDITOR_REGISTRY: dict[str, EditorFn] = { CLEAR_TOOL_USES_EDIT_TYPE: apply_clear_tool_uses_20250919, COMPACT_EDIT_TYPE: apply_compact_20260112, } def _normalize_spec( - spec: Union[Dict[str, Any], List[Dict[str, Any]], None], -) -> Optional[List[Dict[str, Any]]]: + spec: dict[str, Any] | list[dict[str, Any]] | None, +) -> list[dict[str, Any]] | None: """Accept Anthropic-native dict form or OpenAI list form; return edits list.""" if isinstance(spec, list): # Local import to avoid an import cycle at module load. @@ -46,7 +46,7 @@ def _wrap_editor_return(raw: Any, *, fallback_system: Any) -> PolyfillResult: return raw # Legacy 2-tuple return — sync editors don't mutate ``system``, so # carry the caller's value forward. - messages, applied = cast(Tuple[List[Dict[str, Any]], Any], raw) + messages, applied = cast(tuple[list[dict[str, Any]], Any], raw) return PolyfillResult( messages=messages, system=fallback_system, @@ -57,11 +57,11 @@ def _wrap_editor_return(raw: Any, *, fallback_system: Any) -> PolyfillResult: async def apply_context_management( *, model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]], + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, system: Any, - context_management_spec: Union[Dict[str, Any], List[Dict[str, Any]], None], - litellm_metadata: Optional[Dict[str, Any]] = None, + context_management_spec: dict[str, Any] | list[dict[str, Any]] | None, + litellm_metadata: dict[str, Any] | None = None, llm_router: Any = None, user_api_key_auth: Any = None, ) -> PolyfillResult: @@ -78,7 +78,7 @@ async def apply_context_management( current_messages = messages current_system = system - aggregated_applied: List[AppliedEdit] = [] + aggregated_applied: list[AppliedEdit] = [] aggregated_compaction_block = None aggregated_iterations_usage = None @@ -92,7 +92,7 @@ async def apply_context_management( ) continue - kwargs: Dict[str, Any] = { + kwargs: dict[str, Any] = { "model": model, "messages": current_messages, "tools": tools, diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/clear_tool_uses.py b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/clear_tool_uses.py index 8bcf8acfff6..abfc45859ef 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/clear_tool_uses.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/clear_tool_uses.py @@ -1,6 +1,6 @@ """``clear_tool_uses_20250919`` polyfill (v0: ``trigger`` and ``keep`` only).""" -from typing import Any, Dict, List, Optional, Tuple, cast +from typing import Any, cast import litellm from litellm._logging import verbose_logger @@ -14,7 +14,7 @@ from ..constants import ( from ..placeholders import build_cleared_tool_result_content -def _count_tool_uses(messages: List[Dict[str, Any]]) -> int: +def _count_tool_uses(messages: list[dict[str, Any]]) -> int: """Return the number of tool_use content blocks across all messages. Only counts blocks with a string ``id`` to stay consistent with @@ -32,9 +32,9 @@ def _count_tool_uses(messages: List[Dict[str, Any]]) -> int: return count -def _collect_tool_use_ids_in_order(messages: List[Dict[str, Any]]) -> List[str]: +def _collect_tool_use_ids_in_order(messages: list[dict[str, Any]]) -> list[str]: """Return tool_use ids in the chronological order they appear in messages.""" - ids: List[str] = [] + ids: list[str] = [] for msg in messages: content = msg.get("content") if isinstance(content, list): @@ -47,11 +47,11 @@ def _collect_tool_use_ids_in_order(messages: List[Dict[str, Any]]) -> List[str]: def _trigger_met( - trigger: Dict[str, Any], + trigger: dict[str, Any], model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]], -) -> Tuple[bool, Optional[int]]: + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, +) -> tuple[bool, int | None]: """Return (trigger_met, input_tokens if counted for reuse).""" trigger_type = trigger.get("type", "input_tokens") threshold = trigger.get("value") @@ -73,7 +73,7 @@ def _trigger_met( return current_tokens > threshold, current_tokens -def _resolve_keep_count(keep: Dict[str, Any]) -> int: +def _resolve_keep_count(keep: dict[str, Any]) -> int: keep_type = keep.get("type", "tool_uses") if keep_type != "tool_uses": return DEFAULT_KEEP_TOOL_USES @@ -84,10 +84,10 @@ def _resolve_keep_count(keep: Dict[str, Any]) -> int: def _last_completed_tool_use_id( - messages: List[Dict[str, Any]], -) -> Optional[str]: + messages: list[dict[str, Any]], +) -> str | None: """Latest completed tool_result id; never cleared.""" - last_id: Optional[str] = None + last_id: str | None = None for msg in messages: content = msg.get("content") if isinstance(content, list): @@ -99,17 +99,17 @@ def _last_completed_tool_use_id( return last_id -def _clear_tool_results(messages: List[Dict[str, Any]], ids_to_clear: set) -> Tuple[List[Dict[str, Any]], int]: +def _clear_tool_results(messages: list[dict[str, Any]], ids_to_clear: set) -> tuple[list[dict[str, Any]], int]: """Clear matching tool_result content; return (messages, cleared_count).""" cleared = 0 - new_messages: List[Dict[str, Any]] = [] + new_messages: list[dict[str, Any]] = [] for msg in messages: content = msg.get("content") if not isinstance(content, list): new_messages.append(msg) continue - new_blocks: List[Any] = [] + new_blocks: list[Any] = [] mutated = False for block in content: if ( @@ -138,11 +138,11 @@ def _clear_tool_results(messages: List[Dict[str, Any]], ids_to_clear: set) -> Tu def apply_clear_tool_uses_20250919( *, model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]], + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, system: Any, - edit_spec: Dict[str, Any], -) -> Tuple[List[Dict[str, Any]], Optional[AppliedEdit]]: + edit_spec: dict[str, Any], +) -> tuple[list[dict[str, Any]], AppliedEdit | None]: """Apply clear_tool_uses; return (messages, AppliedEdit or None).""" ignored_knobs = [knob for knob in ("clear_at_least", "exclude_tools", "clear_tool_inputs") if knob in edit_spec] for ignored_knob in ignored_knobs: diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py index c87014ffc1e..a28fc003e5d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py @@ -14,7 +14,7 @@ Mirrors Anthropic's native ``compact_20260112`` for non-Anthropic providers: import re from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, Union, cast import litellm from litellm._logging import verbose_logger @@ -83,7 +83,7 @@ _PROPAGATED_METADATA_KEYS = ( _SUMMARY_TAG_RE = re.compile(r"(.*?)", re.IGNORECASE | re.DOTALL) -def _read_summary_model_setting() -> Optional[str]: +def _read_summary_model_setting() -> str | None: """Look up the configured summarization model from proxy general_settings.""" try: from litellm.proxy.proxy_server import general_settings @@ -163,7 +163,7 @@ async def _check_summary_model_access( user_id = getattr(user_api_key_auth, "user_id", None) project_id = getattr(user_api_key_auth, "project_id", None) - checks: Tuple[Tuple[Literal["key", "team"], List[str]], ...] = ( + checks: tuple[tuple[Literal["key", "team"], list[str]], ...] = ( ("key", key_models), ("team", team_models), ) @@ -446,8 +446,8 @@ async def _check_summary_model_rate_limit( def _find_latest_compaction_index( - messages: List[Dict[str, object]], -) -> Tuple[Optional[int], Optional[int]]: + messages: list[dict[str, object]], +) -> tuple[int | None, int | None]: """Return (message_index, block_index) of the most recent compaction block. ``None, None`` if no compaction block is present. Iterates from the end so @@ -465,8 +465,8 @@ def _find_latest_compaction_index( def _slice_around_compaction_block( - messages: List[Dict[str, Any]], -) -> Tuple[List[Dict[str, object]], Optional[Dict[str, object]]]: + messages: list[dict[str, Any]], +) -> tuple[list[dict[str, object]], dict[str, object] | None]: """Apply Anthropic's "drop everything before the compaction block" rule. Returns ``(sliced_messages_with_compaction_block, compaction_block_dict)`` @@ -481,26 +481,26 @@ def _slice_around_compaction_block( original_msg = messages[msg_idx] original_content = original_msg["content"] - compaction_block = cast(Dict[str, object], original_content[blk_idx]) + compaction_block = cast(dict[str, object], original_content[blk_idx]) # Per Anthropic's contract everything before the compaction block is # dropped, including earlier blocks within the same assistant message. sliced_content = list(original_content[blk_idx:]) - sliced_messages: List[Dict[str, object]] = [{**original_msg, "content": sliced_content}] + sliced_messages: list[dict[str, object]] = [{**original_msg, "content": sliced_content}] sliced_messages.extend(messages[msg_idx + 1 :]) return sliced_messages, compaction_block def _strip_compaction_blocks( - messages: List[Dict[str, object]], -) -> List[Dict[str, object]]: + messages: list[dict[str, object]], +) -> list[dict[str, object]]: """Drop any ``compaction`` content blocks from messages. Used to build the downstream-bound message list — the adapter has no concept of a compaction block, so it must not see one. """ - cleaned: List[Dict[str, object]] = [] + cleaned: list[dict[str, object]] = [] for msg in messages: content = msg.get("content") if not isinstance(content, list): @@ -515,9 +515,9 @@ def _strip_compaction_blocks( def _augment_system_with_summary( - system: Optional[Union[str, List[Dict[str, object]]]], + system: str | list[dict[str, object]] | None, summary_text: str, -) -> Union[str, List[Dict[str, object]]]: +) -> str | list[dict[str, object]]: """Prepend a "Previous conversation summary: ..." block to ``system``.""" prefix = f"{COMPACT_SUMMARY_SYSTEM_PREFIX}{summary_text}\n\n" if system is None: @@ -534,14 +534,14 @@ def _augment_system_with_summary( return [{"type": "text", "text": prefix.rstrip()}, *system] -def _resolve_trigger_tokens(edit_spec: Dict[str, object]) -> Tuple[int, List[str]]: +def _resolve_trigger_tokens(edit_spec: dict[str, object]) -> tuple[int, list[str]]: """Validate and resolve ``trigger.value``. Raises ``AnthropicContextManagementError`` if the explicitly-supplied value is below the 50k minimum. Unknown ``trigger.type`` values fall back to ``input_tokens`` with a warning. """ - warnings: List[str] = [] + warnings: list[str] = [] trigger = edit_spec.get("trigger") or {} if not isinstance(trigger, dict): warnings.append("trigger_not_a_dict_using_default") @@ -568,7 +568,7 @@ def _resolve_trigger_tokens(edit_spec: Dict[str, object]) -> Tuple[int, List[str return value, warnings -def _build_summary_prompt(edit_spec: Dict[str, object], tools: Optional[List[Dict[str, object]]]) -> str: +def _build_summary_prompt(edit_spec: dict[str, object], tools: list[dict[str, object]] | None) -> str: custom = edit_spec.get("instructions") if isinstance(custom, str) and custom.strip(): return custom @@ -579,8 +579,8 @@ def _build_summary_prompt(edit_spec: Dict[str, object], tools: Optional[List[Dic def _propagate_metadata( - parent_litellm_metadata: Optional[Mapping[str, object]], -) -> Dict[str, object]: + parent_litellm_metadata: Mapping[str, object] | None, +) -> dict[str, object]: """Extract the parent request's auth/spend-attribution fields for the summary subcall. The proxy attaches ``user_api_key``, ``user_api_key_team_id`` etc. to @@ -591,7 +591,7 @@ def _propagate_metadata( """ if not parent_litellm_metadata: return {} - propagated: Dict[str, object] = {} + propagated: dict[str, object] = {} for key in _PROPAGATED_METADATA_KEYS: if key in parent_litellm_metadata: propagated[key] = parent_litellm_metadata[key] @@ -600,10 +600,10 @@ def _propagate_metadata( def _count_effective_tokens( model: str, - effective_messages: List[Dict[str, object]], - compaction_block: Optional[CompactionBlock], - tools: Optional[List[Dict[str, object]]], - system: Optional[Union[str, List[Dict[str, object]]]] = None, + effective_messages: list[dict[str, object]], + compaction_block: CompactionBlock | None, + tools: list[dict[str, object]] | None, + system: str | list[dict[str, object]] | None = None, ) -> int: """Token-count the conversation as it will appear downstream. @@ -623,7 +623,7 @@ def _count_effective_tokens( try: openai_shape = adapter.translate_anthropic_messages_to_openai( messages=cast( - "List[Union[AnthropicMessagesUserMessageParam, AnthopicMessagesAssistantMessageParam]]", + "list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]", messages_without_compaction, ) ) @@ -640,13 +640,13 @@ def _count_effective_tokens( # gets a consistent format regardless of which counting path it uses. # An inaccurate tool token count here could cause the polyfill to skip # needed compaction or trigger unnecessary summarization. - openai_tools: Optional[List[Dict[str, object]]] = None + openai_tools: list[dict[str, object]] | None = None if tools: try: translated_tools, _ = adapter.translate_anthropic_tools_to_openai( - tools=cast("List[AllAnthropicToolsValues]", tools) + tools=cast("list[AllAnthropicToolsValues]", tools) ) - openai_tools = cast(List[Dict[str, object]], translated_tools) + openai_tools = cast(list[dict[str, object]], translated_tools) except Exception as e: verbose_logger.debug( "compact_20260112: anthropic→openai tools translation failed " @@ -657,8 +657,8 @@ def _count_effective_tokens( total = litellm.token_counter( model=model, - messages=cast(List[Dict[str, object]], openai_shape), - tools=cast("Optional[List[ChatCompletionToolParam]]", openai_tools), + messages=cast(list[dict[str, object]], openai_shape), + tools=cast("list[ChatCompletionToolParam] | None", openai_tools), ) if compaction_block is not None: content = compaction_block.get("content") or "" @@ -671,7 +671,7 @@ def _count_effective_tokens( def _system_to_text( - system: Optional[Union[str, List[Dict[str, object]]]], + system: str | list[dict[str, object]] | None, ) -> str: """Flatten an Anthropic-style ``system`` value into a single string for token counting. Returns ``""`` when ``system`` carries no text.""" @@ -679,7 +679,7 @@ def _system_to_text( return "" if isinstance(system, str): return system - parts: List[str] = [] + parts: list[str] = [] for block in system: if isinstance(block, dict) and block.get("type") == "text": text = block.get("text") @@ -689,8 +689,8 @@ def _system_to_text( def _select_last_user_question( - messages: List[Dict[str, object]], -) -> List[Dict[str, object]]: + messages: list[dict[str, object]], +) -> list[dict[str, object]]: """Pick the most recent ``user`` turn that is a real question. Returns a one-element message list with any ``tool_result`` blocks @@ -724,7 +724,7 @@ def _select_last_user_question( ] -def _extract_summary_text(raw: Optional[str]) -> Optional[str]: +def _extract_summary_text(raw: str | None) -> str | None: if not raw: return None match = _SUMMARY_TAG_RE.search(raw) @@ -735,8 +735,8 @@ def _extract_summary_text(raw: Optional[str]) -> Optional[str]: def _system_to_openai_message( - system: Optional[Union[str, List[Dict[str, Any]]]], -) -> Optional[Dict[str, Any]]: + system: str | list[dict[str, Any]] | None, +) -> dict[str, Any] | None: """Translate Anthropic-shaped ``system`` to an OpenAI system message. Accepts a bare string or a list of Anthropic content blocks; returns @@ -754,10 +754,10 @@ def _system_to_openai_message( def _build_summary_messages( - effective_messages: List[Dict[str, object]], + effective_messages: list[dict[str, object]], prompt: str, - system: Optional[Union[str, List[Dict[str, object]]]] = None, -) -> List[Dict[str, object]]: + system: str | list[dict[str, object]] | None = None, +) -> list[dict[str, object]]: """Build the OpenAI-shape message list for the summary call. The caller's ``system`` prompt is prepended (the default summarization @@ -773,7 +773,7 @@ def _build_summary_messages( try: openai_messages = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( messages=cast( - "List[Union[AnthropicMessagesUserMessageParam, AnthopicMessagesAssistantMessageParam]]", + "list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]", stripped, ) ) @@ -785,7 +785,7 @@ def _build_summary_messages( ) openai_messages = stripped - summary_messages: List[Dict[str, object]] = [] + summary_messages: list[dict[str, object]] = [] system_message = _system_to_openai_message(system) if system_message is not None: summary_messages.append(system_message) @@ -827,10 +827,10 @@ def _append_text_to_content(content: Any, extra_text: str) -> Any: async def _call_summary_model( *, summary_model: str, - summary_messages: List[Dict[str, object]], + summary_messages: list[dict[str, object]], metadata: Mapping[str, object], llm_router: Any, - allowed_model_region: Optional[str] = None, + allowed_model_region: str | None = None, max_tokens: int = COMPACT_SUMMARY_MAX_TOKENS, ) -> Union["ModelResponse", "CustomStreamWrapper"]: """Invoke the configured summary model. @@ -860,7 +860,7 @@ async def _call_summary_model( # the parent ``/v1/messages`` request. On timeout the caller catches the # exception and surfaces ``applied_edits[0].error = "summary_call_failed"``, # forwarding the request without compaction rather than hanging. - call_kwargs: Dict[str, Any] = { + call_kwargs: dict[str, Any] = { "model": summary_model, "messages": summary_messages, "max_tokens": max_tokens, @@ -881,7 +881,7 @@ async def _call_summary_model( return await litellm.acompletion(**call_kwargs) -def _extract_response_text(response: Any) -> Optional[str]: +def _extract_response_text(response: Any) -> str | None: try: choice = response.choices[0] message = choice.message @@ -899,7 +899,7 @@ def _extract_response_text(response: Any) -> Optional[str]: return None -def _extract_usage(response: object) -> Tuple[int, int]: +def _extract_usage(response: object) -> tuple[int, int]: usage = getattr(response, "usage", None) if usage is None: return 0, 0 @@ -911,9 +911,9 @@ def _extract_usage(response: object) -> Tuple[int, int]: def apply_client_compaction_block_history( *, - messages: List[Dict[str, object]], - system: Optional[Union[str, List[Dict[str, object]]]], -) -> Optional[PolyfillResult]: + messages: list[dict[str, object]], + system: str | list[dict[str, object]] | None, +) -> PolyfillResult | None: """Honor client-sent compaction blocks without a ``compact_20260112`` edit. When the request omits ``context_management`` but the message history already @@ -933,7 +933,7 @@ def apply_client_compaction_block_history( ) prior_summary_text = prior_compaction_block.get("content") or "" - augmented_system: Union[str, List[Dict[str, object]], None] = system + augmented_system: str | list[dict[str, object]] | None = system if isinstance(prior_summary_text, str) and prior_summary_text: augmented_system = _augment_system_with_summary(system, prior_summary_text) verbose_logger.info( @@ -958,11 +958,11 @@ def apply_client_compaction_block_history( async def apply_compact_20260112( *, model: str, - messages: List[Dict[str, object]], - tools: Optional[List[Dict[str, object]]], - system: Optional[Union[str, List[Dict[str, object]]]], - edit_spec: Dict[str, object], - litellm_metadata: Optional[Mapping[str, object]] = None, + messages: list[dict[str, object]], + tools: list[dict[str, object]] | None, + system: str | list[dict[str, object]] | None, + edit_spec: dict[str, object], + litellm_metadata: Mapping[str, object] | None = None, llm_router: Optional["Router"] = None, user_api_key_auth: Optional["UserAPIKeyAuth"] = None, ) -> PolyfillResult: @@ -993,7 +993,7 @@ async def apply_compact_20260112( # non-Anthropic backends (which would reject them). effective_messages, prior_compaction_block = _slice_around_compaction_block(messages) prior_summary_text = prior_compaction_block.get("content") if prior_compaction_block else None - augmented_system: Union[str, List[Dict[str, object]], None] = system + augmented_system: str | list[dict[str, object]] | None = system if isinstance(prior_summary_text, str) and prior_summary_text: augmented_system = _augment_system_with_summary(system, prior_summary_text) verbose_logger.info( @@ -1143,7 +1143,7 @@ async def apply_compact_20260112( "type": "compaction", "content": summary_text, } - iterations_usage: List[UsageIteration] = [ + iterations_usage: list[UsageIteration] = [ { "type": "compaction", "input_tokens": summary_input_tokens, diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/placeholders.py b/litellm/llms/anthropic/experimental_pass_through/context_management/placeholders.py index f684d970df4..b3d3529f105 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/placeholders.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/placeholders.py @@ -1,13 +1,13 @@ """Placeholder content for cleared ``tool_result`` blocks (string or block list).""" -from typing import Any, List, Union +from typing import Any from .constants import CLEARED_TOOL_RESULT_PLACEHOLDER def build_cleared_tool_result_content( original_content: Any, -) -> Union[str, List[dict]]: +) -> str | list[dict]: """Return a string or single text block list, matching ``original_content`` shape.""" if isinstance(original_content, list): return [{"type": "text", "text": CLEARED_TOOL_RESULT_PLACEHOLDER}] diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/result.py b/litellm/llms/anthropic/experimental_pass_through/context_management/result.py index 14adeb9452a..33ad15e3885 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/result.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/result.py @@ -6,7 +6,7 @@ attach ``iterations`` to ``usage``. """ from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional, Union +from typing import Any from litellm.types.llms.anthropic import ( AppliedEdit, @@ -19,13 +19,13 @@ from .constants import COMPACT_EDIT_TYPE @dataclass class PolyfillResult: - messages: List[Dict[str, Any]] - system: Optional[Union[str, List[Dict[str, Any]]]] - applied_edits: List[AppliedEdit] = field(default_factory=list) - compaction_block: Optional[CompactionBlock] = None - iterations_usage: Optional[List[UsageIteration]] = None + messages: list[dict[str, Any]] + system: str | list[dict[str, Any]] | None + applied_edits: list[AppliedEdit] = field(default_factory=list) + compaction_block: CompactionBlock | None = None + iterations_usage: list[UsageIteration] | None = None - def applied_edits_for_response(self) -> Optional[List[AppliedEdit]]: + def applied_edits_for_response(self) -> list[AppliedEdit] | None: """``applied_edits`` to attach on the client-visible response. ``compact_20260112`` is included when a new compaction block was @@ -39,7 +39,7 @@ class PolyfillResult: (no block, no error, no warnings) are omitted. Other edit types are included when the editor returned an ``AppliedEdit``. """ - visible: List[AppliedEdit] = [] + visible: list[AppliedEdit] = [] for edit in self.applied_edits: if edit.get("type") == COMPACT_EDIT_TYPE: if self.compaction_block is not None or edit.get("error") or edit.get("warnings"): diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 3e3d9c19af9..95b1e9277b7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -10,7 +10,7 @@ follow-up response is chained as Phase 2 of the same iterator. import json from collections.abc import AsyncIterator -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast from litellm._logging import verbose_logger @@ -19,12 +19,12 @@ from litellm._logging import verbose_logger # --------------------------------------------------------------------------- -def _parse_sse_events(raw: bytes) -> List[tuple]: +def _parse_sse_events(raw: bytes) -> list[tuple]: """Return a list of (event_type, parsed_data_dict) from raw SSE bytes.""" text = raw.decode("utf-8", errors="replace") lines = text.split("\n") - events: List[tuple] = [] - current_event_type: Optional[str] = None + events: list[tuple] = [] + current_event_type: str | None = None for line in lines: stripped = line.strip() @@ -44,7 +44,7 @@ def _parse_sse_events(raw: bytes) -> List[tuple]: return events -def _handle_message_start(data: Dict, response: Dict) -> None: +def _handle_message_start(data: dict, response: dict) -> None: msg = data.get("message", {}) response["id"] = msg.get("id", response["id"]) response["model"] = msg.get("model", response["model"]) @@ -57,12 +57,12 @@ def _handle_message_start(data: Dict, response: Dict) -> None: response["usage"][key] = usage[key] -def _handle_content_block_start(data: Dict, content_blocks: Dict[int, Dict]) -> None: +def _handle_content_block_start(data: dict, content_blocks: dict[int, dict]) -> None: idx = data.get("index", len(content_blocks)) block = data.get("content_block", {}) block_type = block.get("type", "text") - _BLOCK_TEMPLATES: Dict[str, Dict] = { + _BLOCK_TEMPLATES: dict[str, dict] = { "text": {"type": "text", "text": ""}, "thinking": {"type": "thinking", "thinking": "", "signature": ""}, "redacted_thinking": { @@ -84,7 +84,7 @@ def _handle_content_block_start(data: Dict, content_blocks: Dict[int, Dict]) -> content_blocks[idx] = dict(block) -def _handle_content_block_delta(data: Dict, content_blocks: Dict[int, Dict]) -> None: +def _handle_content_block_delta(data: dict, content_blocks: dict[int, dict]) -> None: idx = data.get("index", 0) delta = data.get("delta", {}) delta_type = delta.get("type", "") @@ -102,7 +102,7 @@ def _handle_content_block_delta(data: Dict, content_blocks: Dict[int, Dict]) -> block["signature"] = delta.get("signature", block.get("signature", "")) -def _handle_content_block_stop(data: Dict, content_blocks: Dict[int, Dict]) -> None: +def _handle_content_block_stop(data: dict, content_blocks: dict[int, dict]) -> None: idx = data.get("index", 0) block = content_blocks.get(idx) if block and block.get("type") == "tool_use": @@ -114,7 +114,7 @@ def _handle_content_block_stop(data: Dict, content_blocks: Dict[int, Dict]) -> N block["input"] = {"_raw": partial} -def _handle_message_delta(data: Dict, response: Dict) -> None: +def _handle_message_delta(data: dict, response: dict) -> None: delta = data.get("delta", {}) if "stop_reason" in delta: response["stop_reason"] = delta["stop_reason"] @@ -150,12 +150,12 @@ class AgenticAnthropicStreamingIterator: completion_stream: AsyncIterator, http_handler: Any, model: str, - messages: List[Dict], + messages: list[dict], anthropic_messages_provider_config: Any, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: Any, custom_llm_provider: str, - kwargs: Dict, + kwargs: dict, ): self._inner = completion_stream.__aiter__() self._http_handler = http_handler @@ -167,10 +167,10 @@ class AgenticAnthropicStreamingIterator: self._custom_llm_provider = custom_llm_provider self._kwargs = kwargs - self._collected_bytes: List[bytes] = [] + self._collected_bytes: list[bytes] = [] self._stream_exhausted = False self._hook_processing_done = False - self._follow_up_iterator: Optional[AsyncIterator] = None + self._follow_up_iterator: AsyncIterator | None = None def __aiter__(self): return self @@ -265,8 +265,8 @@ class AgenticAnthropicStreamingIterator: @staticmethod def _rebuild_anthropic_response_from_sse( - raw_bytes: List[bytes], - ) -> Optional[Dict[str, Any]]: + raw_bytes: list[bytes], + ) -> dict[str, Any] | None: """ Parse collected SSE bytes into an Anthropic Messages response dict. @@ -280,7 +280,7 @@ class AgenticAnthropicStreamingIterator: """ events = _parse_sse_events(b"".join(raw_bytes)) - response: Dict[str, Any] = { + response: dict[str, Any] = { "id": "", "type": "message", "role": "assistant", @@ -290,7 +290,7 @@ class AgenticAnthropicStreamingIterator: "stop_sequence": None, "usage": {"input_tokens": 0, "output_tokens": 0}, } - content_blocks: Dict[int, Dict[str, Any]] = {} + content_blocks: dict[int, dict[str, Any]] = {} saw_message_start = False for event_type, data in events: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py index 184fede25e9..4238f1b1a20 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py @@ -9,7 +9,7 @@ the LLM doesn't make a tool call, and we need to return a stream to the user. """ import json -from typing import Any, Dict, List, cast +from typing import Any, cast from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -38,7 +38,7 @@ class FakeAnthropicMessagesStreamIterator: self.chunks = self._create_streaming_chunks() self.current_index = 0 - def _create_content_block_chunks(self, block_dict: Dict[str, Any], index: int) -> List[bytes]: + def _create_content_block_chunks(self, block_dict: dict[str, Any], index: int) -> list[bytes]: """Build SSE chunks for a single content block.""" chunks = [] block_type = block_dict.get("type") @@ -117,12 +117,12 @@ class FakeAnthropicMessagesStreamIterator: chunks.append(f"event: content_block_stop\ndata: {json.dumps(content_block_stop)}\n\n".encode()) return chunks - def _create_streaming_chunks(self) -> List[bytes]: + def _create_streaming_chunks(self) -> list[bytes]: """Convert the non-streaming response to streaming chunks""" chunks = [] # Cast response to dict for easier access - response_dict = cast(Dict[str, Any], self.response) + response_dict = cast(dict[str, Any], self.response) # 1. message_start event usage = response_dict.get("usage", {}) @@ -147,13 +147,13 @@ class FakeAnthropicMessagesStreamIterator: # 2-4. For each content block, send start/delta/stop events content_blocks = response_dict.get("content", []) for index, block in enumerate(content_blocks): - block_dict = cast(Dict[str, Any], block) + block_dict = cast(dict[str, Any], block) chunks.extend(self._create_content_block_chunks(block_dict, index)) # 5. message_delta event (with final usage and stop_reason) # Include cache usage fields so clients that only read message_delta # (like Claude Code's SDK) see the full input token breakdown. - delta_usage: Dict[str, Any] = { + delta_usage: dict[str, Any] = { "output_tokens": usage.get("output_tokens", 0) if usage else 0, } if usage: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 2230e8fc0db..5a7c85e41b4 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -11,10 +11,6 @@ from collections.abc import AsyncIterator, Coroutine, Iterator from functools import partial from typing import ( Any, - Dict, - List, - Optional, - Union, cast, ) @@ -48,7 +44,7 @@ from .utils import AnthropicMessagesRequestUtils, mock_response _RESPONSES_API_PROVIDERS = frozenset({"openai"}) -def _should_route_to_responses_api(custom_llm_provider: Optional[str]) -> bool: +def _should_route_to_responses_api(custom_llm_provider: str | None) -> bool: """Return True when the provider should use the Responses API path. Set ``litellm.use_chat_completions_url_for_anthropic_messages = True`` to @@ -80,12 +76,12 @@ base_llm_http_handler = BaseLLMHTTPHandler() async def _execute_pre_request_hooks( model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - stream: Optional[bool], - custom_llm_provider: Optional[str], + messages: list[dict], + tools: list[dict] | None, + stream: bool | None, + custom_llm_provider: str | None, **kwargs, -) -> Dict: +) -> dict: """ Execute pre-request hooks from CustomLogger callbacks. @@ -142,12 +138,12 @@ async def _execute_pre_request_hooks( async def _try_websearch_short_circuit( model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - custom_llm_provider: Optional[str], - stream: Optional[bool], - kwargs: Optional[dict] = None, -) -> Optional[Union[AnthropicMessagesResponse, AsyncIterator]]: + messages: list[dict], + tools: list[dict] | None, + custom_llm_provider: str | None, + stream: bool | None, + kwargs: dict | None = None, +) -> AnthropicMessagesResponse | AsyncIterator | None: """ Attempt to short-circuit a web-search-only request. @@ -194,24 +190,24 @@ async def _try_websearch_short_circuit( @client async def anthropic_messages( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[Union[str, list]] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, - custom_llm_provider: Optional[str] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | list | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + api_key: str | None = None, + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[AnthropicMessagesResponse, Iterator[bytes], AsyncIterator[Any]]: +) -> AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[Any]: """ Async: Make llm api request in Anthropic /messages API spec. @@ -365,7 +361,7 @@ async def anthropic_messages( return response -def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional[Dict]: +def validate_anthropic_api_metadata(metadata: dict | None = None) -> dict | None: """ Validate Anthropic API metadata - This is done to ensure only allowed `metadata` fields are passed to Anthropic API @@ -379,30 +375,30 @@ def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional def anthropic_messages_handler( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - metadata: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[Union[str, list]] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - container: Optional[Dict] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, - custom_llm_provider: Optional[str] = None, + metadata: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | list | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + container: dict | None = None, + api_key: str | None = None, + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ - AnthropicMessagesResponse, - Iterator[bytes], - AsyncIterator[Any], - Coroutine[Any, Any, Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]]], -]: +) -> ( + AnthropicMessagesResponse + | Iterator[bytes] + | AsyncIterator[Any] + | Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]] +): """ Makes Anthropic `/v1/messages` API calls In the Anthropic API Spec @@ -517,7 +513,7 @@ def anthropic_messages_handler( **kwargs, ) - anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None + anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None = None if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]: anthropic_messages_provider_config = ProviderConfigManager.get_provider_anthropic_messages_config( diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/__init__.py b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/__init__.py index 68f9f471809..d4c58bb6aea 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/__init__.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/__init__.py @@ -1,14 +1,12 @@ -from typing import List - from .advisor import AdvisorOrchestrationHandler from .base import MessagesInterceptor -_interceptors: List[MessagesInterceptor] = [ +_interceptors: list[MessagesInterceptor] = [ AdvisorOrchestrationHandler(), ] -def get_messages_interceptors() -> List[MessagesInterceptor]: +def get_messages_interceptors() -> list[MessagesInterceptor]: """Return the list of active MessagesInterceptors. Order matters: interceptors are tried in list order; the first one whose diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py index af1284b6c62..6e88e365b22 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py @@ -16,7 +16,7 @@ How it works: import uuid from collections.abc import AsyncIterator -from typing import Any, Dict, List, Optional, Union +from typing import Any import litellm import litellm.constants as _c @@ -44,8 +44,8 @@ class AdvisorOrchestrationHandler(MessagesInterceptor): def can_handle( self, - tools: Optional[List[Dict]], - custom_llm_provider: Optional[str], + tools: list[dict] | None, + custom_llm_provider: str | None, ) -> bool: if not tools: return False @@ -57,13 +57,13 @@ class AdvisorOrchestrationHandler(MessagesInterceptor): self, *, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - stream: Optional[bool], + messages: list[dict], + tools: list[dict] | None, + stream: bool | None, max_tokens: int, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, **kwargs, - ) -> Union[AnthropicMessagesResponse, AsyncIterator]: + ) -> AnthropicMessagesResponse | AsyncIterator: from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( FakeAnthropicMessagesStreamIterator, ) @@ -86,17 +86,17 @@ class AdvisorOrchestrationHandler(MessagesInterceptor): synthetic_advisor_tool = _make_synthetic_advisor_tool() # Executor tools = all original tools with advisor replaced by the synthetic one. - executor_tools: List[Dict] = [ + executor_tools: list[dict] = [ (synthetic_advisor_tool if t.get("type") == ANTHROPIC_ADVISOR_TOOL_TYPE else t) for t in (tools or []) ] # Strip prior advisor blocks from history, preserving advice text as context. - current_messages: List[Dict] = strip_advisor_blocks_from_messages( + current_messages: list[dict] = strip_advisor_blocks_from_messages( [dict(m) for m in messages], replace_with_text=True ) parent_request_id: str = str(kwargs.pop("litellm_call_id", None) or uuid.uuid4()) - metadata_base: Dict = dict(kwargs.pop("metadata", None) or {}) + metadata_base: dict = dict(kwargs.pop("metadata", None) or {}) iteration = 0 while True: @@ -187,7 +187,7 @@ def _allow_client_side_advisor_credentials() -> bool: return general_settings.get("allow_client_side_credentials") is True -def _resolve_advisor_credentials(advisor_tool: dict) -> tuple[Optional[str], Optional[str]]: +def _resolve_advisor_credentials(advisor_tool: dict) -> tuple[str | None, str | None]: """Resolve the (api_key, api_base) override for the advisor sub-call. A caller-supplied ``api_base`` is only honored alongside a caller-supplied @@ -206,8 +206,8 @@ def _resolve_advisor_credentials(advisor_tool: dict) -> tuple[Optional[str], Opt """ if not _allow_client_side_advisor_credentials(): return None, None - api_key: Optional[str] = advisor_tool.get("api_key") - api_base: Optional[str] = advisor_tool.get("api_base") + api_key: str | None = advisor_tool.get("api_key") + api_base: str | None = advisor_tool.get("api_base") if api_base is None: return api_key, None if not api_key: @@ -230,7 +230,7 @@ def _resolve_advisor_credentials(advisor_tool: dict) -> tuple[Optional[str], Opt return api_key, api_base -def _make_synthetic_advisor_tool() -> Dict: +def _make_synthetic_advisor_tool() -> dict: """Build a regular tool definition the executor provider can understand.""" return { "name": "advisor", @@ -248,7 +248,7 @@ def _make_synthetic_advisor_tool() -> Dict: } -def _find_advisor_tool_use(response: Any) -> Optional[Dict]: +def _find_advisor_tool_use(response: Any) -> dict | None: """Return the first tool_use block with name='advisor', or None.""" content = response.get("content") if isinstance(response, dict) else [] if not isinstance(content, list): @@ -272,10 +272,10 @@ _PROVIDER_SPECIFIC_KEYS = frozenset({"provider_specific_fields"}) def _build_advisor_context( - messages: List[Dict], + messages: list[dict], executor_response: Any, - advisor_use_block: Dict, -) -> List[Dict]: + advisor_use_block: dict, +) -> list[dict]: """ Build the message list for the advisor sub-call. @@ -303,11 +303,11 @@ def _build_advisor_context( def _inject_advisor_turn( - messages: List[Dict], + messages: list[dict], executor_response: Any, - advisor_use_block: Dict, + advisor_use_block: dict, advisor_text: str, -) -> List[Dict]: +) -> list[dict]: """ Append the executor's response (as an assistant turn) and the advisor result (as a user tool_result turn) so the executor can continue. @@ -331,10 +331,10 @@ def _inject_advisor_turn( def _inject_max_uses_error( - messages: List[Dict], + messages: list[dict], executor_response: Any, - advisor_use_block: Dict, -) -> List[Dict]: + advisor_use_block: dict, +) -> list[dict]: """ Inject a max_uses_exceeded error tool_result so the executor continues without further advisor calls (mirrors Anthropic's server-side behaviour). @@ -359,11 +359,11 @@ def _inject_max_uses_error( async def _call_messages_handler( model: str, - messages: List[Dict], - tools: Optional[List[Dict]], + messages: list[dict], + tools: list[dict] | None, stream: bool, max_tokens: int, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, **kwargs, ) -> Any: """ diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/base.py b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/base.py index 922c668e9c8..c32b21564a7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/base.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/base.py @@ -1,6 +1,5 @@ from abc import ABC, abstractmethod from collections.abc import AsyncIterator -from typing import Dict, List, Optional, Union from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -22,8 +21,8 @@ class MessagesInterceptor(ABC): @abstractmethod def can_handle( self, - tools: Optional[List[Dict]], - custom_llm_provider: Optional[str], + tools: list[dict] | None, + custom_llm_provider: str | None, ) -> bool: """Return True if this interceptor should handle the request.""" @@ -32,11 +31,11 @@ class MessagesInterceptor(ABC): self, *, model: str, - messages: List[Dict], - tools: Optional[List[Dict]], - stream: Optional[bool], + messages: list[dict], + tools: list[dict] | None, + stream: bool | None, max_tokens: int, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, **kwargs, - ) -> Union[AnthropicMessagesResponse, AsyncIterator]: + ) -> AnthropicMessagesResponse | AsyncIterator: """Execute the interception and return the response.""" diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py index 92dcfe9e5b1..887240bb1cc 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py @@ -8,7 +8,7 @@ tool through a ``tool_use`` content block, and results are fed back as """ from collections.abc import AsyncIterator, Mapping, Sequence -from typing import Any, Union +from typing import Any from litellm._logging import verbose_logger from litellm.responses.mcp.request_context import MCPRequestContext @@ -36,7 +36,7 @@ def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Ma return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use") -def _get_stop_reason(response: AnthropicMessagesResponse) -> Union[str, None]: +def _get_stop_reason(response: AnthropicMessagesResponse) -> str | None: stop_reason = response.get("stop_reason") return stop_reason if isinstance(stop_reason, str) else None @@ -60,9 +60,9 @@ async def anthropic_messages_with_mcp( max_tokens: int, messages: Sequence[Mapping[str, Any]], model: str, - tools: Union[Sequence[Mapping[str, Any]], None] = None, + tools: Sequence[Mapping[str, Any]] | None = None, **kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract -) -> Union[AnthropicMessagesResponse, AsyncIterator[Any]]: +) -> AnthropicMessagesResponse | AsyncIterator[Any]: """ Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop. diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 38c13e69ad7..a0cb1f75d3e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -2,7 +2,7 @@ import asyncio import json from collections.abc import AsyncIterator from datetime import datetime -from typing import Any, List, Protocol, Union, runtime_checkable +from typing import Any, Protocol, runtime_checkable import httpx from pydantic import TypeAdapter @@ -122,7 +122,7 @@ class BaseAnthropicMessagesStreamingIterator: self.start_time = datetime.now() self.completion_start_time: datetime | None = None - async def _handle_streaming_logging(self, collected_chunks: List[bytes]): + async def _handle_streaming_logging(self, collected_chunks: list[bytes]): """Handle the logging after all chunks have been collected.""" from litellm.proxy.pass_through_endpoints.streaming_handler import ( PassThroughStreamingHandler, @@ -169,7 +169,7 @@ class BaseAnthropicMessagesStreamingIterator: url_route="/v1/messages", ) - def _convert_chunk_to_sse_format(self, chunk: Union[dict, Any]) -> bytes: + def _convert_chunk_to_sse_format(self, chunk: dict | Any) -> bytes: """ Convert a chunk to Server-Sent Events format. @@ -186,7 +186,7 @@ class BaseAnthropicMessagesStreamingIterator: async def async_sse_wrapper( self, - completion_stream: AsyncIterator[Union[bytes, GenericStreamingChunk, ModelResponseStream, dict]], + completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict], ) -> AsyncIterator[bytes]: """ Generic async SSE wrapper that converts streaming chunks to SSE format diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index f82868bc1e1..29ac9133760 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator -from typing import Any, Dict, List, Optional, Tuple +from typing import Any import httpx @@ -42,7 +42,7 @@ DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING = ( class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "anthropic" @property @@ -72,7 +72,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): # "metadata", ] - def _remove_scope_from_cache_control(self, anthropic_messages_request: Dict) -> None: + def _remove_scope_from_cache_control(self, anthropic_messages_request: dict) -> None: """ Remove `scope` field from cache_control blocks. @@ -217,12 +217,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = AnthropicModelInfo.get_api_base(api_base) or "https://api.anthropic.com" if not api_base.endswith("/v1/messages"): @@ -233,12 +233,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: # Check for Anthropic OAuth token in Authorization header headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) @@ -259,7 +259,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): return headers, api_base @staticmethod - def _translate_reasoning_effort_to_anthropic(model: str, optional_params: Dict, custom_llm_provider: str) -> None: + def _translate_reasoning_effort_to_anthropic(model: str, optional_params: dict, custom_llm_provider: str) -> None: """Map OpenAI-style ``reasoning_effort`` to native Anthropic params. Caller-supplied ``thinking`` / ``output_config`` win over the alias. @@ -312,7 +312,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): @staticmethod def _translate_legacy_thinking_for_adaptive_model( - model: str, optional_params: Dict, custom_llm_provider: str + model: str, optional_params: dict, custom_llm_provider: str ) -> None: """Translate legacy ``thinking.type=enabled`` to adaptive for 4.6/4.7. Caller-provided ``output_config.effort`` is never overridden. @@ -346,7 +346,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): @staticmethod def _translate_adaptive_effort_for_non_adaptive_model( - model: str, optional_params: Dict, max_tokens: Optional[int], custom_llm_provider: str + model: str, optional_params: dict, max_tokens: int | None, custom_llm_provider: str ) -> None: """Translate the 4.6+ adaptive-thinking interface (``thinking.type=adaptive`` and/or ``output_config.effort``) down to what an older Anthropic model @@ -479,11 +479,11 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ No transformation is needed for Anthropic messages diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py index c8060d41fad..48d127af1c0 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py @@ -1,5 +1,5 @@ from functools import lru_cache -from typing import Any, Dict, FrozenSet, List, cast, get_type_hints +from typing import Any, cast, get_type_hints from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams from litellm.types.llms.anthropic_messages.anthropic_response import ( @@ -8,7 +8,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import ( @lru_cache(maxsize=1) -def _anthropic_messages_optional_param_keys() -> FrozenSet[str]: +def _anthropic_messages_optional_param_keys() -> frozenset[str]: """ Valid AnthropicMessagesRequestOptionalParams keys. @@ -22,7 +22,7 @@ def _anthropic_messages_optional_param_keys() -> FrozenSet[str]: class AnthropicMessagesRequestUtils: @staticmethod def get_requested_anthropic_messages_optional_param( - params: Dict[str, Any], + params: dict[str, Any], *, model: str | None = None, drop_params: bool = False, @@ -56,7 +56,7 @@ class AnthropicMessagesRequestUtils: def mock_response( model: str, - messages: List[Dict], + messages: list[dict], max_tokens: int, mock_response: str = "Hi! My name is Claude.", **kwargs, diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py index d98d984de4f..72f8ee524b8 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py @@ -5,7 +5,7 @@ Used when the target model is an OpenAI or Azure model. """ from collections.abc import AsyncIterator, Coroutine -from typing import Any, Dict, List, Optional, Union +from typing import Any import litellm from litellm.types.llms.anthropic import AnthropicMessagesRequest @@ -23,28 +23,28 @@ _ADAPTER = LiteLLMAnthropicToResponsesAPIAdapter() def _build_responses_kwargs( *, max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - context_management: Optional[Dict] = None, - metadata: Optional[Dict] = None, - output_config: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, - extra_kwargs: Optional[Dict[str, Any]] = None, -) -> Dict[str, Any]: + context_management: dict | None = None, + metadata: dict | None = None, + output_config: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, + extra_kwargs: dict[str, Any] | None = None, +) -> dict[str, Any]: """ Build the kwargs dict to pass directly to litellm.responses() / litellm.aresponses(). """ # Build a typed AnthropicMessagesRequest for the adapter - request_data: Dict[str, Any] = { + request_data: dict[str, Any] = { "model": model, "messages": messages, "max_tokens": max_tokens, @@ -124,23 +124,23 @@ class LiteLLMMessagesToResponsesAPIHandler: @staticmethod async def async_anthropic_messages_handler( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - context_management: Optional[Dict] = None, - metadata: Optional[Dict] = None, - output_config: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, + context_management: dict | None = None, + metadata: dict | None = None, + output_config: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, **kwargs, - ) -> Union[AnthropicMessagesResponse, AsyncIterator]: + ) -> AnthropicMessagesResponse | AsyncIterator: responses_kwargs = _build_responses_kwargs( max_tokens=max_tokens, messages=messages, @@ -175,28 +175,28 @@ class LiteLLMMessagesToResponsesAPIHandler: @staticmethod def anthropic_messages_handler( max_tokens: int, - messages: List[Dict], + messages: list[dict], model: str, - context_management: Optional[Dict] = None, - metadata: Optional[Dict] = None, - output_config: Optional[Dict] = None, - stop_sequences: Optional[List[str]] = None, - stream: Optional[bool] = False, - system: Optional[str] = None, - temperature: Optional[float] = None, - thinking: Optional[Dict] = None, - tool_choice: Optional[Dict] = None, - tools: Optional[List[Dict]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - output_format: Optional[Dict] = None, + context_management: dict | None = None, + metadata: dict | None = None, + output_config: dict | None = None, + stop_sequences: list[str] | None = None, + stream: bool | None = False, + system: str | None = None, + temperature: float | None = None, + thinking: dict | None = None, + tool_choice: dict | None = None, + tools: list[dict] | None = None, + top_k: int | None = None, + top_p: float | None = None, + output_format: dict | None = None, _is_async: bool = False, **kwargs, - ) -> Union[ - AnthropicMessagesResponse, - AsyncIterator[Any], - Coroutine[Any, Any, Union[AnthropicMessagesResponse, AsyncIterator[Any]]], - ]: + ) -> ( + AnthropicMessagesResponse + | AsyncIterator[Any] + | Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any]] + ): if _is_async: return LiteLLMMessagesToResponsesAPIHandler.async_anthropic_messages_handler( max_tokens=max_tokens, diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py index 29c9f3d059b..f698c78604d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py @@ -4,7 +4,7 @@ import json import traceback from collections import deque from collections.abc import AsyncIterator -from typing import Any, Dict +from typing import Any from litellm import verbose_logger from litellm._uuid import uuid @@ -34,14 +34,14 @@ class AnthropicResponsesStreamWrapper: self._message_id: str = f"msg_{uuid.uuid4()}" self._current_block_index: int = -1 # Map item_id -> content_block_index so we can stop the right block later - self._item_id_to_block_index: Dict[str, int] = {} + self._item_id_to_block_index: dict[str, int] = {} # Track open function_call items by item_id so we can emit tool_use start - self._pending_tool_ids: Dict[str, str] = {} # item_id -> call_id / name accumulator + self._pending_tool_ids: dict[str, str] = {} # item_id -> call_id / name accumulator self._sent_message_start = False self._sent_message_stop = False self._chunk_queue: deque = deque() - def _make_message_start(self) -> Dict[str, Any]: + def _make_message_start(self) -> dict[str, Any]: return { "type": "message_start", "message": { @@ -257,7 +257,7 @@ class AnthropicResponsesStreamWrapper: stop_reason = "tool_use" break - usage_delta: Dict[str, Any] = { + usage_delta: dict[str, Any] = { "input_tokens": input_tokens, "output_tokens": output_tokens, } @@ -280,7 +280,7 @@ class AnthropicResponsesStreamWrapper: def __aiter__(self) -> "AnthropicResponsesStreamWrapper": return self - async def __anext__(self) -> Dict[str, Any]: + async def __anext__(self) -> dict[str, Any]: # Return any queued chunks first if self._chunk_queue: return self._chunk_queue.popleft() diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 172e54de98e..0d707358d09 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -6,7 +6,7 @@ path used for OpenAI and Azure models. """ import json -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast from litellm.litellm_core_utils.reasoning_effort_utils import ( reasoning_effort_from_thinking_budget, @@ -43,7 +43,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: # ------------------------------------------------------------------ # @staticmethod - def _translate_anthropic_image_source_to_url(source: dict) -> Optional[str]: + def _translate_anthropic_image_source_to_url(source: dict) -> str | None: """Convert Anthropic image source to a URL string.""" source_type = source.get("type") if source_type == "base64": @@ -56,13 +56,8 @@ class LiteLLMAnthropicToResponsesAPIAdapter: def translate_messages_to_responses_input( self, - messages: List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], - ) -> List[Dict[str, Any]]: + messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam], + ) -> list[dict[str, Any]]: """ Convert Anthropic messages list to Responses API `input` items. @@ -73,7 +68,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: assistant text -> message(role=assistant, output_text) assistant tool_use -> function_call """ - input_items: List[Dict[str, Any]] = [] + input_items: list[dict[str, Any]] = [] for m in messages: role = m["role"] @@ -89,7 +84,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: } ) elif isinstance(content, list): - user_parts: List[Dict[str, Any]] = [] + user_parts: list[dict[str, Any]] = [] for block in content: if not isinstance(block, dict): continue @@ -141,7 +136,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: } ) elif isinstance(content, list): - asst_parts: List[Dict[str, Any]] = [] + asst_parts: list[dict[str, Any]] = [] for block in content: if not isinstance(block, dict): continue @@ -175,19 +170,19 @@ class LiteLLMAnthropicToResponsesAPIAdapter: def translate_tools_to_responses_api( self, - tools: List[AllAnthropicToolsValues], - ) -> List[Dict[str, Any]]: + tools: list[AllAnthropicToolsValues], + ) -> list[dict[str, Any]]: """Convert Anthropic tool definitions to Responses API function tools.""" - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] for tool in tools: - tool_dict = cast(Dict[str, Any], tool) + tool_dict = cast(dict[str, Any], tool) tool_type = tool_dict.get("type", "") tool_name = tool_dict.get("name", "") # web_search tool if (isinstance(tool_type, str) and tool_type.startswith("web_search")) or tool_name == "web_search": result.append({"type": "web_search_preview"}) continue - func_tool: Dict[str, Any] = {"type": "function", "name": tool_name} + func_tool: dict[str, Any] = {"type": "function", "name": tool_name} if "description" in tool_dict: func_tool["description"] = tool_dict["description"] if "input_schema" in tool_dict: @@ -198,7 +193,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: @staticmethod def translate_tool_choice_to_responses_api( tool_choice: AnthropicMessagesToolChoice, - ) -> Union[str, dict[str, Any]]: + ) -> str | dict[str, Any]: """Convert Anthropic tool_choice to Responses API tool_choice.""" tc_type = tool_choice.get("type") if tc_type == "any": @@ -211,8 +206,8 @@ class LiteLLMAnthropicToResponsesAPIAdapter: @staticmethod def translate_context_management_to_responses_api( - context_management: Dict[str, Any], - ) -> Optional[List[Dict[str, Any]]]: + context_management: dict[str, Any], + ) -> list[dict[str, Any]] | None: """ Convert Anthropic context_management dict to OpenAI Responses API array format. @@ -226,13 +221,13 @@ class LiteLLMAnthropicToResponsesAPIAdapter: if not isinstance(edits, list): return None - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] for edit in edits: if not isinstance(edit, dict): continue edit_type = edit.get("type", "") if edit_type == "compact_20260112": - entry: Dict[str, Any] = {"type": "compaction"} + entry: dict[str, Any] = {"type": "compaction"} trigger = edit.get("trigger") if isinstance(trigger, dict) and trigger.get("value") is not None: entry["compact_threshold"] = int(trigger["value"]) @@ -242,9 +237,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter: @staticmethod def translate_thinking_to_reasoning( - thinking: Dict[str, Any], - output_config: Optional[Dict[str, Any]] = None, - ) -> Optional[Dict[str, Any]]: + thinking: dict[str, Any], + output_config: dict[str, Any] | None = None, + ) -> dict[str, Any] | None: """ Convert Anthropic thinking param to Responses API reasoning param. @@ -269,7 +264,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: return None auto_summary = is_reasoning_auto_summary_enabled() - result: Dict[str, Any] = {"effort": effort} + result: dict[str, Any] = {"effort": effort} summary = thinking.get("summary") if summary: result["summary"] = summary @@ -280,23 +275,18 @@ class LiteLLMAnthropicToResponsesAPIAdapter: def translate_request( self, anthropic_request: AnthropicMessagesRequest, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Translate a full Anthropic /v1/messages request dict to litellm.responses() / litellm.aresponses() kwargs. """ model: str = anthropic_request["model"] messages_list = cast( - List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam], anthropic_request["messages"], ) - responses_kwargs: Dict[str, Any] = { + responses_kwargs: dict[str, Any] = { "model": model, "input": self.translate_messages_to_responses_input(messages_list), } @@ -325,7 +315,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: tools = anthropic_request.get("tools") if tools: responses_kwargs["tools"] = self.translate_tools_to_responses_api( - cast(List[AllAnthropicToolsValues], tools) + cast(list[AllAnthropicToolsValues], tools) ) # tool_choice @@ -341,7 +331,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: output_config = anthropic_request.get("output_config") reasoning = self.translate_thinking_to_reasoning( thinking, - output_config=cast(Optional[Dict[str, Any]], output_config), + output_config=cast(dict[str, Any] | None, output_config), ) if reasoning: responses_kwargs["reasoning"] = reasoning @@ -398,7 +388,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: from litellm.types.llms.openai import ResponseAPIUsage - content: List[Dict[str, Any]] = [] + content: list[dict[str, Any]] = [] stop_reason: AnthropicFinishReason = "end_turn" for item in response.output: @@ -464,7 +454,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: stop_reason = "max_tokens" # usage - raw_usage: Optional[ResponseAPIUsage] = response.usage + raw_usage: ResponseAPIUsage | None = response.usage input_tokens = int(getattr(raw_usage, "input_tokens", 0) or 0) output_tokens = int(getattr(raw_usage, "output_tokens", 0) or 0) diff --git a/litellm/llms/anthropic/experimental_pass_through/utils.py b/litellm/llms/anthropic/experimental_pass_through/utils.py index 827cce89dab..46091cd89a2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/utils.py +++ b/litellm/llms/anthropic/experimental_pass_through/utils.py @@ -1,5 +1,4 @@ import os -from typing import Optional import litellm from litellm.types.utils import ModelInfo @@ -13,7 +12,7 @@ def is_reasoning_auto_summary_enabled() -> bool: def normalize_reasoning_effort_value( effort: str, model: str, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> str: """ Normalize a reasoning effort value based on model capabilities. @@ -29,7 +28,7 @@ def normalize_reasoning_effort_value( from litellm.utils import get_model_info - model_info: Optional[ModelInfo] = None + model_info: ModelInfo | None = None try: model_info = get_model_info(model=model, custom_llm_provider=custom_llm_provider) except Exception: diff --git a/litellm/llms/anthropic/files/__init__.py b/litellm/llms/anthropic/files/__init__.py index 78c9dc89f70..fe93b39ce6c 100644 --- a/litellm/llms/anthropic/files/__init__.py +++ b/litellm/llms/anthropic/files/__init__.py @@ -1,4 +1,4 @@ from .handler import AnthropicFilesHandler from .transformation import AnthropicFilesConfig -__all__ = ["AnthropicFilesHandler", "AnthropicFilesConfig"] +__all__ = ["AnthropicFilesConfig", "AnthropicFilesHandler"] diff --git a/litellm/llms/anthropic/files/handler.py b/litellm/llms/anthropic/files/handler.py index ef86029de58..b911347b2ff 100644 --- a/litellm/llms/anthropic/files/handler.py +++ b/litellm/llms/anthropic/files/handler.py @@ -2,7 +2,7 @@ import asyncio import json import time from collections.abc import Coroutine -from typing import Any, Optional, Union +from typing import Any import httpx @@ -51,10 +51,10 @@ class AnthropicFilesHandler: async def afile_content( self, file_content_request: FileContentRequest, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Union[float, httpx.Timeout] = 600.0, - max_retries: Optional[int] = None, + api_base: str | None = None, + api_key: str | None = None, + timeout: float | httpx.Timeout = 600.0, + max_retries: int | None = None, ) -> HttpxBinaryResponseContent: """ Async: Retrieve file content from Anthropic. @@ -124,11 +124,11 @@ class AnthropicFilesHandler: self, _is_async: bool, file_content_request: FileContentRequest, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Union[float, httpx.Timeout] = 600.0, - max_retries: Optional[int] = None, - ) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: + api_base: str | None = None, + api_key: str | None = None, + timeout: float | httpx.Timeout = 600.0, + max_retries: int | None = None, + ) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: """ Retrieve file content from Anthropic. diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index 0fa01e09492..f2d292dd674 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -14,13 +14,13 @@ Anthropic Files API endpoints: import calendar import time -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast import httpx from openai.types.file_deleted import FileDeleted -from litellm.litellm_core_utils.url_utils import encode_url_path_segment 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 from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, @@ -62,12 +62,12 @@ class AnthropicFilesConfig(BaseFilesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = AnthropicModelInfo.get_api_base(api_base) or ANTHROPIC_FILES_API_BASE return f"{api_base.rstrip('/')}/v1/files" @@ -76,7 +76,7 @@ class AnthropicFilesConfig(BaseFilesConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: return AnthropicError( status_code=status_code, @@ -91,8 +91,8 @@ class AnthropicFilesConfig(BaseFilesConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_base is None and isinstance(litellm_params, dict): api_base = litellm_params.get("api_base") @@ -110,7 +110,7 @@ class AnthropicFilesConfig(BaseFilesConfig): ) return headers - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: return ["purpose"] def map_openai_params( @@ -154,7 +154,7 @@ class AnthropicFilesConfig(BaseFilesConfig): def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -220,13 +220,13 @@ class AnthropicFilesConfig(BaseFilesConfig): def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: api_base = AnthropicModelInfo.get_api_base(litellm_params.get("api_base")) or ANTHROPIC_FILES_API_BASE url = f"{api_base.rstrip('/')}/v1/files" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if purpose: params["purpose"] = purpose return url, params @@ -236,7 +236,7 @@ class AnthropicFilesConfig(BaseFilesConfig): raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: """ Anthropic list response: { diff --git a/litellm/llms/anthropic/skills/transformation.py b/litellm/llms/anthropic/skills/transformation.py index 896182b4763..a3aa694847c 100644 --- a/litellm/llms/anthropic/skills/transformation.py +++ b/litellm/llms/anthropic/skills/transformation.py @@ -2,7 +2,7 @@ Anthropic Skills API configuration and transformations """ -from typing import Any, Dict, Optional, Tuple +from typing import Any import httpx @@ -30,7 +30,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.ANTHROPIC - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """Add Anthropic-specific headers""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo @@ -69,9 +69,9 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, endpoint: str, - skill_id: Optional[str] = None, + skill_id: str | None = None, ) -> str: """Get complete URL for Anthropic Skills API""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo @@ -89,7 +89,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): create_request: CreateSkillRequest, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """Transform create skill request for Anthropic""" verbose_logger.debug("Transforming create skill request: %s", create_request) @@ -114,7 +114,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): list_params: ListSkillsParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform list skills request for Anthropic""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo @@ -122,7 +122,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): url = self.get_complete_url(api_base=api_base, endpoint="skills") # Build query parameters - query_params: Dict[str, Any] = {} + query_params: dict[str, Any] = {} if "limit" in list_params and list_params["limit"]: query_params["limit"] = list_params["limit"] if "page" in list_params and list_params["page"]: @@ -154,7 +154,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform get skill request for Anthropic""" url = self.get_complete_url(api_base=api_base, endpoint="skills", skill_id=skill_id) @@ -179,7 +179,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform delete skill request for Anthropic""" url = self.get_complete_url(api_base=api_base, endpoint="skills", skill_id=skill_id) diff --git a/litellm/llms/apiserpent/search/defaults.py b/litellm/llms/apiserpent/search/defaults.py index 3bd8e1f93f4..2851ab6ddbb 100644 --- a/litellm/llms/apiserpent/search/defaults.py +++ b/litellm/llms/apiserpent/search/defaults.py @@ -6,7 +6,7 @@ package-level defaults. See https://apiserpent.com/docs. """ from dataclasses import asdict, dataclass -from typing import Dict, Literal, Optional +from typing import Literal SearchEngine = Literal["google", "bing", "yahoo", "ddg"] SafeSearch = Literal["off", "moderate", "strict"] @@ -34,11 +34,11 @@ class APISerpentSearchParams: country: str = "us" num: int = 10 format: ResponseFormat = "full" - pages: Optional[int] = None - freshness: Optional[Freshness] = None - safe: Optional[SafeSearch] = None - language: Optional[str] = None - pixel_position: Optional[bool] = None + pages: int | None = None + freshness: Freshness | None = None + safe: SafeSearch | None = None + language: str | None = None + pixel_position: bool | None = None def __post_init__(self) -> None: # num's deep-search floor (NUM_MIN_DEEP) is endpoint-specific and enforced @@ -48,9 +48,9 @@ class APISerpentSearchParams: if self.pages is not None and not PAGES_MIN <= self.pages <= PAGES_MAX: raise ValueError(f"pages must be between {PAGES_MIN} and {PAGES_MAX}, got {self.pages}") - def to_request_params(self) -> Dict: + def to_request_params(self) -> dict: """Return non-None fields as request params, booleans lowercased.""" - params: Dict = {} + params: dict = {} for key, value in asdict(self).items(): if value is None: continue diff --git a/litellm/llms/apiserpent/search/transformation.py b/litellm/llms/apiserpent/search/transformation.py index 637b1472534..29efd47316c 100644 --- a/litellm/llms/apiserpent/search/transformation.py +++ b/litellm/llms/apiserpent/search/transformation.py @@ -8,7 +8,7 @@ Two endpoints under one provider, selected via the ``deep`` boolean param: APISerpent API Reference: https://apiserpent.com/docs """ -from typing import Dict, List, Literal, Optional, Union, cast +from typing import Literal, cast from urllib.parse import urlencode import httpx @@ -48,11 +48,11 @@ class APISerpentSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: api_key = self.resolve_server_api_key( caller_api_key=api_key, caller_api_base=api_base, @@ -68,9 +68,9 @@ class APISerpentSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -93,10 +93,10 @@ class APISerpentSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform a unified search request into APISerpent query params. @@ -114,7 +114,7 @@ class APISerpentSearchConfig(BaseSearchConfig): is_deep = self._is_deep_search(optional_params) - overrides: Dict = {} + overrides: dict = {} if "max_results" in optional_params: num_min = NUM_MIN_DEEP if is_deep else NUM_MIN overrides["num"] = max(num_min, min(optional_params["max_results"], NUM_MAX)) @@ -135,14 +135,14 @@ class APISerpentSearchConfig(BaseSearchConfig): return {APISERPENT_PARAMS_KEY: params} @staticmethod - def _append_domain_filters(query: str, domains: List[str]) -> str: + def _append_domain_filters(query: str, domains: list[str]) -> str: domain_clauses = " OR ".join(f"site:{domain}" for domain in domains) return f"({query}) ({domain_clauses})" def transform_search_response( self, raw_response: httpx.Response, - logging_obj: Optional[LiteLLMLoggingObj], + logging_obj: LiteLLMLoggingObj | None, **kwargs, ) -> SearchResponse: """ @@ -156,7 +156,7 @@ class APISerpentSearchConfig(BaseSearchConfig): raw_results = response_json.get("results") or {} organic = raw_results.get("organic", []) if isinstance(raw_results, dict) else raw_results - results: List[SearchResult] = [] + results: list[SearchResult] = [] for result in organic: results.append( SearchResult( diff --git a/litellm/llms/aws_polly/text_to_speech/transformation.py b/litellm/llms/aws_polly/text_to_speech/transformation.py index dcbed3d94df..1687afa23eb 100644 --- a/litellm/llms/aws_polly/text_to_speech/transformation.py +++ b/litellm/llms/aws_polly/text_to_speech/transformation.py @@ -7,7 +7,7 @@ Reference: https://docs.aws.amazon.com/polly/latest/dg/API_SynthesizeSpeech.html import json from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -69,16 +69,16 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): self, model: str, input: str, - voice: Optional[Union[str, Dict]], - optional_params: Dict, - litellm_params_dict: Dict, + voice: str | dict | None, + optional_params: dict, + litellm_params_dict: dict, logging_obj: "LiteLLMLoggingObj", - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, Any]], + timeout: float | httpx.Timeout, + extra_headers: dict[str, Any] | None, base_llm_http_handler: Any, aspeech: bool, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, **kwargs: Any, ) -> Union[ "HttpxBinaryResponseContent", @@ -98,7 +98,7 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): ) # Convert voice to string if it's a dict - voice_str: Optional[str] = None + voice_str: str | None = None if isinstance(voice, str): voice_str = voice elif isinstance(voice, dict): @@ -129,7 +129,7 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): return response - def _get_aws_region_name_for_polly(self, optional_params: Dict) -> str: + def _get_aws_region_name_for_polly(self, optional_params: dict) -> str: """Get AWS region name for Polly API calls.""" aws_region_name = optional_params.get("aws_region_name") if aws_region_name is None: @@ -145,18 +145,18 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Dict = {}, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict = {}, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to AWS Polly parameters """ mapped_params = {} # Map voice - support both native Polly voices and OpenAI voice mappings - mapped_voice: Optional[str] = None + mapped_voice: str | None = None if isinstance(voice, str): if voice in self.VOICE_MAPPINGS: # OpenAI voice -> Polly voice @@ -211,8 +211,8 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate AWS environment and set up headers. @@ -225,7 +225,7 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -250,10 +250,10 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): def _sign_polly_request( self, - request_body: Dict[str, Any], + request_body: dict[str, Any], endpoint_url: str, - litellm_params: Dict, - ) -> Tuple[Dict[str, str], str]: + litellm_params: dict, + ) -> tuple[dict[str, str], str]: """ Sign the AWS Polly request using SigV4. @@ -309,9 +309,9 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ @@ -336,7 +336,7 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): engine = optional_params.get("engine", self.DEFAULT_ENGINE) # Build request body - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "Engine": engine, "OutputFormat": output_format, "Text": input, diff --git a/litellm/llms/azure/assistants.py b/litellm/llms/azure/assistants.py index 2bd34ba6a3e..11da3130cb3 100644 --- a/litellm/llms/azure/assistants.py +++ b/litellm/llms/azure/assistants.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine, Iterable -from typing import Any, Dict, Literal, Optional, Union +from typing import Any, Literal import httpx from openai import AsyncAzureOpenAI, AzureOpenAI @@ -28,14 +28,14 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_azure_client( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI] = None, - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None = None, + litellm_params: dict | None = None, ) -> AzureOpenAI: if client is None: azure_client_params = self.initialize_azure_sdk_client( @@ -54,14 +54,14 @@ class AzureAssistantsAPI(BaseAzureLLM): def async_get_azure_client( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI] = None, - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None = None, + litellm_params: dict | None = None, ) -> AsyncAzureOpenAI: if client is None: azure_client_params = self.initialize_azure_sdk_client( @@ -84,14 +84,14 @@ class AzureAssistantsAPI(BaseAzureLLM): async def async_get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, + litellm_params: dict | None = None, ) -> AsyncCursorPage[Assistant]: azure_openai_client = self.async_get_azure_client( api_key=api_key, @@ -113,13 +113,13 @@ class AzureAssistantsAPI(BaseAzureLLM): @overload def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, aget_assistants: Literal[True], ) -> Coroutine[None, None, AsyncCursorPage[Assistant]]: ... @@ -127,14 +127,14 @@ class AzureAssistantsAPI(BaseAzureLLM): @overload def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI], - aget_assistants: Optional[Literal[False]], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None, + aget_assistants: Literal[False] | None, ) -> SyncCursorPage[Assistant]: ... @@ -142,15 +142,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, client=None, aget_assistants=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): if aget_assistants is not None and aget_assistants is True: return self.async_get_assistants( @@ -184,14 +184,14 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI] = None, - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None = None, + litellm_params: dict | None = None, ) -> OpenAIMessage: openai_client = self.async_get_azure_client( api_key=api_key, @@ -209,7 +209,7 @@ class AzureAssistantsAPI(BaseAzureLLM): **message_data, # type: ignore ) - response_obj: Optional[OpenAIMessage] = None + response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" response_obj = OpenAIMessage(**thread_message.dict()) @@ -224,15 +224,15 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, a_add_message: Literal[True], - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> Coroutine[None, None, OpenAIMessage]: ... @@ -241,15 +241,15 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI], - a_add_message: Optional[Literal[False]], - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None, + a_add_message: Literal[False] | None, + litellm_params: dict | None = None, ) -> OpenAIMessage: ... @@ -259,15 +259,15 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, client=None, - a_add_message: Optional[bool] = None, - litellm_params: Optional[dict] = None, + a_add_message: bool | None = None, + litellm_params: dict | None = None, ): if a_add_message is not None and a_add_message is True: return self.a_add_message( @@ -298,7 +298,7 @@ class AzureAssistantsAPI(BaseAzureLLM): **message_data, # type: ignore ) - response_obj: Optional[OpenAIMessage] = None + response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" response_obj = OpenAIMessage(**thread_message.dict()) @@ -309,14 +309,14 @@ class AzureAssistantsAPI(BaseAzureLLM): async def async_get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI] = None, - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None = None, + litellm_params: dict | None = None, ) -> AsyncCursorPage[OpenAIMessage]: openai_client = self.async_get_azure_client( api_key=api_key, @@ -339,15 +339,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, aget_messages: Literal[True], - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> Coroutine[None, None, AsyncCursorPage[OpenAIMessage]]: ... @@ -355,15 +355,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI], - aget_messages: Optional[Literal[False]], - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None, + aget_messages: Literal[False] | None, + litellm_params: dict | None = None, ) -> SyncCursorPage[OpenAIMessage]: ... @@ -372,15 +372,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, client=None, aget_messages=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): if aget_messages is not None and aget_messages is True: return self.async_get_messages( @@ -413,16 +413,16 @@ class AzureAssistantsAPI(BaseAzureLLM): async def async_create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], - litellm_params: Optional[dict] = None, + metadata: dict | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, + litellm_params: dict | None = None, ) -> Thread: openai_client = self.async_get_azure_client( api_key=api_key, @@ -450,34 +450,34 @@ class AzureAssistantsAPI(BaseAzureLLM): @overload def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], - client: Optional[AsyncAzureOpenAI], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, + client: AsyncAzureOpenAI | None, acreate_thread: Literal[True], - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> Coroutine[None, None, Thread]: ... @overload def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], - client: Optional[AzureOpenAI], - acreate_thread: Optional[Literal[False]], - litellm_params: Optional[dict] = None, + metadata: dict | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, + client: AzureOpenAI | None, + acreate_thread: Literal[False] | None, + litellm_params: dict | None = None, ) -> Thread: ... @@ -485,17 +485,17 @@ class AzureAssistantsAPI(BaseAzureLLM): def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, client=None, acreate_thread=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): """ Here's an example: @@ -544,14 +544,14 @@ class AzureAssistantsAPI(BaseAzureLLM): async def async_get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, + litellm_params: dict | None = None, ) -> Thread: openai_client = self.async_get_azure_client( api_key=api_key, @@ -574,15 +574,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, aget_thread: Literal[True], - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> Coroutine[None, None, Thread]: ... @@ -590,15 +590,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI], - aget_thread: Optional[Literal[False]], - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None, + aget_thread: Literal[False] | None, + litellm_params: dict | None = None, ) -> Thread: ... @@ -607,15 +607,15 @@ class AzureAssistantsAPI(BaseAzureLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, client=None, aget_thread=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): if aget_thread is not None and aget_thread is True: return self.async_get_thread( @@ -653,20 +653,20 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], - litellm_params: Optional[dict] = None, + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, + litellm_params: dict | None = None, ) -> Run: openai_client = self.async_get_azure_client( api_key=api_key, @@ -696,15 +696,15 @@ class AzureAssistantsAPI(BaseAzureLLM): client: AsyncAzureOpenAI, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - tools: Optional[Iterable[AssistantToolParam]], - event_handler: Optional[AssistantEventHandler], - litellm_params: Optional[dict] = None, + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + tools: Iterable[AssistantToolParam] | None, + event_handler: AssistantEventHandler | None, + litellm_params: dict | None = None, ) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: - data: Dict[str, Any] = { + data: dict[str, Any] = { "thread_id": thread_id, "assistant_id": assistant_id, "additional_instructions": additional_instructions, @@ -722,15 +722,15 @@ class AzureAssistantsAPI(BaseAzureLLM): client: AzureOpenAI, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - tools: Optional[Iterable[AssistantToolParam]], - event_handler: Optional[AssistantEventHandler], - litellm_params: Optional[dict] = None, + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + tools: Iterable[AssistantToolParam] | None, + event_handler: AssistantEventHandler | None, + litellm_params: dict | None = None, ) -> AssistantStreamManager[AssistantEventHandler]: - data: Dict[str, Any] = { + data: dict[str, Any] = { "thread_id": thread_id, "assistant_id": assistant_id, "additional_instructions": additional_instructions, @@ -750,19 +750,19 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, arun_thread: Literal[True], ) -> Coroutine[None, None, Run]: ... @@ -772,20 +772,20 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AzureOpenAI], - arun_thread: Optional[Literal[False]], + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | None, + arun_thread: Literal[False] | None, ) -> Run: ... @@ -795,22 +795,22 @@ class AzureAssistantsAPI(BaseAzureLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + additional_instructions: str | None, + instructions: str | None, + metadata: dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, client=None, arun_thread=None, - event_handler: Optional[AssistantEventHandler] = None, - litellm_params: Optional[dict] = None, + event_handler: AssistantEventHandler | None = None, + litellm_params: dict | None = None, ): if arun_thread is not None and arun_thread is True: if stream is not None and stream is True: @@ -894,15 +894,15 @@ class AzureAssistantsAPI(BaseAzureLLM): # Create Assistant async def async_create_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, create_assistant_data: dict, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> Assistant: azure_openai_client = self.async_get_azure_client( api_key=api_key, @@ -920,16 +920,16 @@ class AzureAssistantsAPI(BaseAzureLLM): def create_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, create_assistant_data: dict, client=None, async_create_assistants=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): if async_create_assistants is not None and async_create_assistants is True: return self.async_create_assistants( @@ -960,15 +960,15 @@ class AzureAssistantsAPI(BaseAzureLLM): # Delete Assistant async def async_delete_assistant( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[AsyncAzureOpenAI], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AsyncAzureOpenAI | None, assistant_id: str, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): azure_openai_client = self.async_get_azure_client( api_key=api_key, @@ -986,16 +986,16 @@ class AzureAssistantsAPI(BaseAzureLLM): def delete_assistant( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, assistant_id: str, - async_delete_assistants: Optional[bool] = None, + async_delete_assistants: bool | None = None, client=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ): if async_delete_assistants is not None and async_delete_assistants is True: return self.async_delete_assistant( diff --git a/litellm/llms/azure/audio_transcription/transformation.py b/litellm/llms/azure/audio_transcription/transformation.py index 77050ce6bca..fccbd81434e 100644 --- a/litellm/llms/azure/audio_transcription/transformation.py +++ b/litellm/llms/azure/audio_transcription/transformation.py @@ -5,7 +5,7 @@ Maps OpenAI-compatible audio transcription calls to Azure Speech REST recognition for short audio. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any from urllib.parse import urlencode, urlparse import httpx @@ -16,11 +16,11 @@ from litellm.llms.base_llm.audio_transcription.transformation import ( BaseAudioTranscriptionConfig, ) from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( AllMessageValues, OpenAIAudioTranscriptionOptionalParams, ) -from litellm.secret_managers.main import get_secret_str from litellm.types.utils import FileTypes, TranscriptionResponse @@ -41,7 +41,7 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): STT_ENDPOINT_PATH = "/speech/recognition/conversation/cognitiveservices/v1" DEFAULT_LANGUAGE = "en-US" - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: return ["language", "response_format"] def map_openai_params( @@ -61,11 +61,11 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("AZURE_SPEECH_API_KEY") if not api_key: @@ -82,12 +82,12 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = api_base or get_secret_str("AZURE_SPEECH_API_BASE") if api_base is None: @@ -140,9 +140,7 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response._hidden_params = response_json return response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return AzureSpeechAudioTranscriptionException( message=error_message, status_code=status_code, @@ -191,12 +189,12 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return f"https://{region}.{self.STT_SPEECH_DOMAIN}" return f"https://{self.STT_SPEECH_DOMAIN}" - def _get_azure_response_format(self, response_format: Optional[str]) -> str: + def _get_azure_response_format(self, response_format: str | None) -> str: if response_format == "verbose_json": return "detailed" return "simple" - def _extract_text(self, response_json: Dict[str, Any]) -> str: + def _extract_text(self, response_json: dict[str, Any]) -> str: if isinstance(response_json.get("DisplayText"), str): return response_json["DisplayText"] diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 14737714c97..25ab7534d6a 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine -from typing import Any, Optional, Union +from typing import Any from openai import AsyncAzureOpenAI, AzureOpenAI from pydantic import BaseModel @@ -27,14 +27,14 @@ class AzureAudioTranscription(AzureChatCompletion): model_response: TranscriptionResponse, timeout: float, max_retries: int, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, client=None, - azure_ad_token: Optional[str] = None, + azure_ad_token: str | None = None, atranscription: bool = False, - litellm_params: Optional[dict] = None, - ) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: + litellm_params: dict | None = None, + ) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]: data = {"model": model, "file": audio_file, **optional_params} if atranscription is True: @@ -113,12 +113,12 @@ class AzureAudioTranscription(AzureChatCompletion): model_response: TranscriptionResponse, timeout: float, logging_obj: Any, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_version: str | None = None, + api_key: str | None = None, + api_base: str | None = None, client=None, max_retries=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> TranscriptionResponse: response = None try: diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index a01ee93daee..373460c151d 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -2,7 +2,7 @@ import asyncio import json import time from collections.abc import Callable, Coroutine -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx # type: ignore from openai import ( @@ -84,7 +84,7 @@ class AzureOpenAIAssistantsAPIConfig: status_code=400, ) elif param == "attachments": # this is a v2 param. Azure currently supports the old 'file_id's param - file_ids: List[str] = [] + file_ids: list[str] = [] if isinstance(value, list): for item in value: if "file_id" in item: @@ -94,16 +94,12 @@ class AzureOpenAIAssistantsAPIConfig: pass else: raise litellm.utils.UnsupportedParamsError( - message="Azure doesn't support {}. To drop it from the call, set `litellm.drop_params = True.".format( - value - ), + message=f"Azure doesn't support {value}. To drop it from the call, set `litellm.drop_params = True.", status_code=400, ) else: raise litellm.utils.UnsupportedParamsError( - message="Invalid param. attachments should always be a list. Got={}, Expected=List. Raw value={}".format( - type(value), value - ), + message=f"Invalid param. attachments should always be a list. Got={type(value)}, Expected=List. Raw value={value}", status_code=400, ) return optional_params @@ -111,7 +107,7 @@ class AzureOpenAIAssistantsAPIConfig: def _check_dynamic_azure_params( azure_client_params: dict, - azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]], + azure_client: AzureOpenAI | AsyncAzureOpenAI | None, ) -> bool: """ Returns True if user passed in client params != initialized azure client @@ -136,9 +132,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): def make_sync_azure_openai_chat_completion_request( self, - azure_client: Union[AzureOpenAI, OpenAI], + azure_client: AzureOpenAI | OpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, ): """ Helper to: @@ -157,9 +153,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): @track_llm_api_timing() async def make_azure_openai_chat_completion_request( self, - azure_client: Union[AsyncAzureOpenAI, AsyncOpenAI], + azure_client: AsyncAzureOpenAI | AsyncOpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, ): """ @@ -187,21 +183,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): model: str, messages: list, model_response: ModelResponse, - api_key: Optional[str], + api_key: str | None, api_base: str, api_version: str, api_type: str, - azure_ad_token: Optional[str], - azure_ad_token_provider: Optional[Callable], + azure_ad_token: str | None, + azure_ad_token_provider: Callable | None, dynamic_params: bool, print_verbose: Callable, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, optional_params, litellm_params, logger_fn, acompletion: bool = False, - headers: Optional[dict] = None, + headers: dict | None = None, client=None, ): if headers: @@ -213,7 +209,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): max_retries = optional_params.pop("max_retries", None) if max_retries is None: max_retries = DEFAULT_MAX_RETRIES - json_mode: Optional[bool] = optional_params.pop("json_mode", False) + json_mode: bool | None = optional_params.pop("json_mode", False) ### CHECK IF CLOUDFLARE AI GATEWAY ### ### if so - set the model as part of the base url @@ -378,7 +374,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): async def acompletion( self, - api_key: Optional[str], + api_key: str | None, api_version: str, model: str, api_base: str, @@ -388,11 +384,11 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, max_retries: int, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, - convert_tool_call_to_json_mode: Optional[bool] = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, + convert_tool_call_to_json_mode: bool | None = None, client=None, # this is the AsyncAzureOpenAI - litellm_params: Optional[dict] = {}, + litellm_params: dict | None = {}, ): response = None try: @@ -488,17 +484,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): self, logging_obj, api_base: str, - api_key: Optional[str], + api_key: str | None, api_version: str, dynamic_params: bool, data: dict, model: str, timeout: Any, max_retries: int, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, client=None, - litellm_params: Optional[dict] = {}, + litellm_params: dict | None = {}, ): # init AzureOpenAI Client azure_client_params = { @@ -564,17 +560,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): self, logging_obj: LiteLLMLoggingObj, api_base: str, - api_key: Optional[str], + api_key: str | None, api_version: str, dynamic_params: bool, data: dict, model: str, timeout: Any, max_retries: int, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, client=None, - litellm_params: Optional[dict] = {}, + litellm_params: dict | None = {}, ): try: azure_client = self.get_azure_openai_client( @@ -645,14 +641,14 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): input: list, logging_obj: LiteLLMLoggingObj, api_base: str, - api_key: Optional[str] = None, - api_version: Optional[str] = None, - client: Optional[AsyncAzureOpenAI] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - max_retries: Optional[int] = None, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, - litellm_params: Optional[dict] = {}, + api_key: str | None = None, + api_version: str | None = None, + client: AsyncAzureOpenAI | None = None, + timeout: float | httpx.Timeout | None = None, + max_retries: int | None = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, + litellm_params: dict | None = {}, ) -> EmbeddingResponse: response = None try: @@ -688,7 +684,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): except json.JSONDecodeError as json_error: raise AzureOpenAIError( status_code=raw_response.status_code or 500, - message=f"Failed to parse raw Azure embedding response: {str(json_error)}", + message=f"Failed to parse raw Azure embedding response: {json_error!s}", ) from json_error if isinstance(response, str): raise AzureOpenAIError( @@ -737,15 +733,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): logging_obj: LiteLLMLoggingObj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str] = None, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, - max_retries: Optional[int] = None, + api_key: str | None = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, + max_retries: int | None = None, client=None, aembedding=None, - headers: Optional[dict] = None, - litellm_params: Optional[dict] = None, - ) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]: + headers: dict | None = None, + litellm_params: dict | None = None, + ) -> EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse]: if headers: optional_params["extra_headers"] = headers if self._client_session is None: @@ -830,8 +826,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): async def make_async_azure_httpx_request( self, - client: Optional[AsyncHTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, api_base: str, api_version: str, api_key: str, @@ -957,8 +953,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): def make_sync_azure_httpx_request( self, - client: Optional[HTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: HTTPHandler | None, + timeout: float | httpx.Timeout | None, api_base: str, api_version: str, api_key: str, @@ -1074,8 +1070,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): def create_azure_base_url( self, azure_client_params: dict, - model: Optional[str], - base_model: Optional[str] = None, + model: str | None, + base_model: str | None = None, ) -> str: from litellm.llms.azure_ai.image_generation import ( AzureFoundryFluxImageGenerationConfig, @@ -1117,7 +1113,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): async def aimage_generation( self, data: dict, - model_response: Optional[ImageResponse], + model_response: ImageResponse | None, azure_client_params: dict, api_key: str, input: list, @@ -1125,9 +1121,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): headers: dict, client=None, timeout=None, - model: Optional[str] = None, + model: str | None = None, ) -> ImageResponse: - response: Optional[dict] = None + response: dict | None = None try: # response = await azure_client.images.generate(**data, timeout=timeout) api_base: str = azure_client_params.get("api_base", "") # "https://example-endpoint.openai.azure.com" @@ -1206,16 +1202,16 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): optional_params: dict, logging_obj: LiteLLMLoggingObj, headers: dict, - model: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - model_response: Optional[ImageResponse] = None, - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + model_response: ImageResponse | None = None, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, client=None, aimg_generation=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> ImageResponse: try: if model and len(model) > 0: @@ -1247,7 +1243,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): headers["Authorization"] = f"Bearer {azure_ad_token}" # init AzureOpenAI Client - azure_client_params: Dict[str, Any] = self.initialize_azure_sdk_client( + azure_client_params: dict[str, Any] = self.initialize_azure_sdk_client( litellm_params=litellm_params or {}, api_key=api_key, model_name=model or "", @@ -1337,17 +1333,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): input: str, voice: str, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + api_version: str | None, + organization: str | None, max_retries: int, - timeout: Union[float, httpx.Timeout], - azure_ad_token: Optional[str] = None, - azure_ad_token_provider: Optional[Callable] = None, - aspeech: Optional[bool] = None, + timeout: float | httpx.Timeout, + azure_ad_token: str | None = None, + azure_ad_token_provider: Callable | None = None, + aspeech: bool | None = None, client=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> HttpxBinaryResponseContent: max_retries = optional_params.pop("max_retries", 2) @@ -1392,15 +1388,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): input: str, voice: str, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - azure_ad_token: Optional[str], - azure_ad_token_provider: Optional[Callable], + api_key: str | None, + api_base: str | None, + api_version: str | None, + azure_ad_token: str | None, + azure_ad_token_provider: Callable | None, max_retries: int, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, client=None, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> HttpxBinaryResponseContent: azure_client: AsyncAzureOpenAI = self.get_azure_openai_client( api_base=api_base, @@ -1423,15 +1419,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): def get_headers( self, - model: Optional[str], + model: str | None, api_key: str, api_base: str, api_version: str, timeout: float, mode: str, - messages: Optional[list] = None, - input: Optional[list] = None, - prompt: Optional[str] = None, + messages: list | None = None, + input: list | None = None, + prompt: str | None = None, ) -> dict: client_session = litellm.client_session or httpx.Client() if api_base is not None and "gateway.ai.cloudflare.com" in api_base: diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index f1bfd96de94..9fccc710e8f 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -1,7 +1,5 @@ """Support for Azure OpenAI gpt-5 model family.""" -from typing import List - import litellm from litellm.exceptions import UnsupportedParamsError from litellm.llms.openai.chat.gpt_5_transformation import ( @@ -56,7 +54,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): _normalized = model.split("/")[-1] # strip provider prefix, e.g. "azure/" return ("gpt-5" in model and not _normalized.startswith("gpt-5-chat")) or "gpt5_series" in model - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """Get supported parameters for Azure OpenAI GPT-5 models. Azure OpenAI GPT-5.2/5.4 models support logprobs, unlike OpenAI's GPT-5. @@ -142,7 +140,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/azure/chat/gpt_transformation.py b/litellm/llms/azure/chat/gpt_transformation.py index 50b3ba16326..948f88fd1c5 100644 --- a/litellm/llms/azure/chat/gpt_transformation.py +++ b/litellm/llms/azure/chat/gpt_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any from httpx._models import Headers, Response @@ -55,16 +55,16 @@ class AzureOpenAIConfig(BaseConfig): def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -75,7 +75,7 @@ class AzureOpenAIConfig(BaseConfig): def get_config(cls): return super().get_config() - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "temperature", "n", @@ -231,7 +231,7 @@ class AzureOpenAIConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -250,12 +250,12 @@ class AzureOpenAIConfig(BaseConfig): model_response: ModelResponse, logging_obj: LoggingClass, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "Azure OpenAI handler.py has custom logic for transforming response, as it uses the OpenAI SDK." @@ -270,13 +270,13 @@ class AzureOpenAIConfig(BaseConfig): optional_params["azure_ad_token"] = value return optional_params - def get_eu_regions(self) -> List[str]: + def get_eu_regions(self) -> list[str]: """ Source: https://learn.microsoft.com/en-us/azure/ai-services/openai/concepts/models#gpt-4-and-gpt-4-turbo-model-availability """ return ["europe", "sweden", "switzerland", "france", "uk"] - def get_us_regions(self) -> List[str]: + def get_us_regions(self) -> list[str]: """ Source: https://learn.microsoft.com/en-us/azure/ai-services/openai/concepts/models#gpt-4-and-gpt-4-turbo-model-availability """ @@ -293,18 +293,18 @@ class AzureOpenAIConfig(BaseConfig): "westus4", ] - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return AzureOpenAIError(message=error_message, status_code=status_code, headers=headers) def validate_environment( self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: raise NotImplementedError( "Azure OpenAI has custom logic for validating environment, as it uses the OpenAI SDK." diff --git a/litellm/llms/azure/chat/o_series_handler.py b/litellm/llms/azure/chat/o_series_handler.py index d6d902d67fb..64b6025f6ea 100644 --- a/litellm/llms/azure/chat/o_series_handler.py +++ b/litellm/llms/azure/chat/o_series_handler.py @@ -5,7 +5,7 @@ Written separately to handle faking streaming for o1 and o3 models. """ from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Optional import httpx @@ -22,26 +22,26 @@ class AzureOpenAIO1ChatCompletion(BaseAzureLLM, OpenAIChatCompletion): def completion( self, model_response: ModelResponse, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, optional_params: dict, litellm_params: dict, logging_obj: Any, - model: Optional[str] = None, - messages: Optional[list] = None, - print_verbose: Optional[Callable] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - dynamic_params: Optional[bool] = None, - azure_ad_token: Optional[str] = None, + model: str | None = None, + messages: list | None = None, + print_verbose: Callable | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + dynamic_params: bool | None = None, + azure_ad_token: str | None = None, acompletion: bool = False, logger_fn=None, - headers: Optional[dict] = None, + headers: dict | None = None, custom_prompt_dict: dict = {}, client=None, - organization: Optional[str] = None, - custom_llm_provider: Optional[str] = None, - drop_params: Optional[bool] = None, + organization: str | None = None, + custom_llm_provider: str | None = None, + drop_params: bool | None = None, shared_session: Optional["ClientSession"] = None, ): client = self.get_azure_openai_client( diff --git a/litellm/llms/azure/chat/o_series_transformation.py b/litellm/llms/azure/chat/o_series_transformation.py index b9cf77b89d8..2b8c2e88e51 100644 --- a/litellm/llms/azure/chat/o_series_transformation.py +++ b/litellm/llms/azure/chat/o_series_transformation.py @@ -12,8 +12,6 @@ Translations handled by LiteLLM: - Temperature => drop param (if user opts in to dropping param) """ -from typing import List, Optional - import litellm from litellm import verbose_logger from litellm.types.llms.openai import AllMessageValues @@ -68,9 +66,9 @@ class AzureOpenAIO1Config(OpenAIOSeriesConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Currently no Azure O Series models support native streaming. @@ -102,7 +100,7 @@ class AzureOpenAIO1Config(OpenAIOSeriesConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 36aab49c714..dcbd3985dfd 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -3,7 +3,7 @@ import hashlib import json import os from collections.abc import Callable -from typing import Any, Dict, Literal, NamedTuple, Optional, Union, cast +from typing import Any, Literal, NamedTuple, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -28,10 +28,10 @@ class AzureOpenAIError(BaseLLMException): self, status_code, message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, - body: Optional[dict] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, + body: dict | None = None, ): super().__init__( status_code=status_code, @@ -43,7 +43,7 @@ class AzureOpenAIError(BaseLLMException): ) -def process_azure_headers(headers: Union[httpx.Headers, dict]) -> dict: +def process_azure_headers(headers: httpx.Headers | dict) -> dict: openai_headers = {} if "x-ratelimit-limit-requests" in headers: openai_headers["x-ratelimit-limit-requests"] = headers["x-ratelimit-limit-requests"] @@ -153,9 +153,9 @@ def get_azure_ad_token_from_username_password( def get_azure_ad_token_from_oidc( azure_ad_token: str, - azure_client_id: Optional[str] = None, - azure_tenant_id: Optional[str] = None, - scope: Optional[str] = None, + azure_client_id: str | None = None, + azure_tenant_id: str | None = None, + scope: str | None = None, ) -> str: """ Get Azure AD token from OIDC token @@ -253,7 +253,7 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict): def get_azure_ad_token( litellm_params: GenericLiteLLMParams, -) -> Optional[str]: +) -> str | None: """ Get Azure AD token from various authentication methods. @@ -333,7 +333,7 @@ def get_azure_ad_token( verbose_logger.debug("Azure AD Token Provider could not be used.") except Exception as e: verbose_logger.error( - f"Error calling Azure AD token provider: {str(e)}. Follow docs - https://docs.litellm.ai/docs/providers/azure/#azure-ad-token-refresh---defaultazurecredential" + f"Error calling Azure AD token provider: {e!s}. Follow docs - https://docs.litellm.ai/docs/providers/azure/#azure-ad-token-refresh---defaultazurecredential" ) raise e @@ -359,8 +359,8 @@ def get_azure_ad_token( # Re-raise TypeError directly raise except Exception as e: - verbose_logger.error(f"Error calling Azure AD token provider: {str(e)}") - raise RuntimeError(f"Failed to get Azure AD token: {str(e)}") from e + verbose_logger.error(f"Error calling Azure AD token provider: {e!s}") + raise RuntimeError(f"Failed to get Azure AD token: {e!s}") from e return azure_ad_token @@ -369,7 +369,7 @@ class BaseAzureLLM(BaseOpenAILLM): @staticmethod def _try_get_default_azure_credential_provider( scope: str, - ) -> Optional[Callable[[], str]]: + ) -> Callable[[], str] | None: """ Try to get DefaultAzureCredential provider @@ -393,20 +393,20 @@ class BaseAzureLLM(BaseOpenAILLM): verbose_logger.debug("Successfully obtained Azure AD token provider using DefaultAzureCredential") return azure_ad_token_provider except Exception as e: - verbose_logger.debug(f"DefaultAzureCredential failed: {str(e)}") + verbose_logger.debug(f"DefaultAzureCredential failed: {e!s}") return None def get_azure_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str] = None, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None = None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, _is_async: bool = False, - model: Optional[str] = None, - ) -> Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]]: - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None + model: str | None = None, + ) -> AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None: + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None client_initialization_params: dict = locals() client_initialization_params["is_async"] = _is_async _lp = litellm_params or {} @@ -453,7 +453,7 @@ class BaseAzureLLM(BaseOpenAILLM): # on every request (via `_refresh_api_key`), so passing # `azure_ad_token_provider` directly preserves Azure AD token refresh # behavior that the regular AzureOpenAI client provides. - v1_api_key: Optional[Union[str, Callable[[], Any]]] = ( + v1_api_key: str | Callable[[], Any] | None = ( azure_client_params.get("api_key") or azure_client_params.get("azure_ad_token_provider") or azure_client_params.get("azure_ad_token") @@ -470,7 +470,7 @@ class BaseAzureLLM(BaseOpenAILLM): v1_api_key = _async_v1_api_key - v1_params: Dict[str, Any] = { + v1_params: dict[str, Any] = { "api_key": v1_api_key, "base_url": f"{api_base}/openai/v1/", } @@ -514,10 +514,10 @@ class BaseAzureLLM(BaseOpenAILLM): def initialize_azure_sdk_client( self, litellm_params: dict, - api_key: Optional[str], - api_base: Optional[str], - model_name: Optional[str], - api_version: Optional[str], + api_key: str | None, + api_base: str | None, + model_name: str | None, + api_version: str | None, is_async: bool, ) -> dict: azure_ad_token_provider = litellm_params.get("azure_ad_token_provider") @@ -580,7 +580,7 @@ class BaseAzureLLM(BaseOpenAILLM): # only show first 5 chars of api_key _api_key = _api_key[:8] + "*" * 15 verbose_logger.debug( - f"Initializing Azure OpenAI Client for {model_name}, Api Base: {str(api_base)}, Api Key:{_api_key}" + f"Initializing Azure OpenAI Client for {model_name}, Api Base: {api_base!s}, Api Key:{_api_key}" ) azure_client_params = { "api_key": api_key, @@ -615,14 +615,14 @@ class BaseAzureLLM(BaseOpenAILLM): model: str, api_version: str, max_retries: int, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, - api_key: Optional[str], - azure_ad_token: Optional[str], - azure_ad_token_provider: Optional[Callable[[], str]], + api_key: str | None, + azure_ad_token: str | None, + azure_ad_token_provider: Callable[[], str] | None, acompletion: bool, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[AzureOpenAI, AsyncAzureOpenAI]: + client: AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> AzureOpenAI | AsyncAzureOpenAI: ## build base url - assume api base includes resource name tenant_id = litellm_params.get("tenant_id", os.getenv("AZURE_TENANT_ID")) client_id = litellm_params.get("client_id", os.getenv("AZURE_CLIENT_ID")) @@ -635,7 +635,7 @@ class BaseAzureLLM(BaseOpenAILLM): api_base += "/" api_base += f"{model}" - azure_client_params: Dict[str, Any] = { + azure_client_params: dict[str, Any] = { "api_version": api_version, "base_url": f"{api_base}", "http_client": litellm.client_session, @@ -664,7 +664,7 @@ class BaseAzureLLM(BaseOpenAILLM): return client @staticmethod - def _base_validate_azure_environment(headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def _base_validate_azure_environment(headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() # Check if api-key is already in headers; if so, use it @@ -693,10 +693,10 @@ class BaseAzureLLM(BaseOpenAILLM): @staticmethod def _get_base_azure_url( - api_base: Optional[str], - litellm_params: Optional[Union[GenericLiteLLMParams, Dict[str, Any]]], - route: Union[Literal["/openai/responses", "/openai/vector_stores"], str], - default_api_version: Optional[Union[str, Literal["latest", "preview"]]] = None, + api_base: str | None, + litellm_params: GenericLiteLLMParams | dict[str, Any] | None, + route: Literal["/openai/responses", "/openai/vector_stores"] | str, + default_api_version: str | Literal["latest", "preview"] | None = None, ) -> str: """ Get the base Azure URL for the given route and API version. @@ -717,7 +717,7 @@ class BaseAzureLLM(BaseOpenAILLM): # Extract api_version or use default litellm_params = litellm_params or {} - api_version = cast(Optional[str], litellm_params.get("api_version")) or default_api_version + api_version = cast(str | None, litellm_params.get("api_version")) or default_api_version # Create a new dictionary with existing params query_params = dict(original_url.params) @@ -744,12 +744,12 @@ class BaseAzureLLM(BaseOpenAILLM): return str(final_url) @staticmethod - def _is_azure_v1_api_version(api_version: Optional[str]) -> bool: + def _is_azure_v1_api_version(api_version: str | None) -> bool: if api_version is None: return False return api_version in {"preview", "latest", "v1"} - def _resolve_env_var(self, litellm_params: Dict[str, Any], param_key: str, env_var_key: str) -> Optional[str]: + def _resolve_env_var(self, litellm_params: dict[str, Any], param_key: str, env_var_key: str) -> str | None: """Resolve the environment variable for a given parameter key. The logic here is different from `params.get(key, os.getenv(env_var))` because @@ -763,15 +763,15 @@ class BaseAzureLLM(BaseOpenAILLM): class AzureCredentials(NamedTuple): - api_base: Optional[str] - api_key: Optional[str] - api_version: Optional[str] + api_base: str | None + api_key: str | None + api_version: str | None def get_azure_credentials( - api_base: Optional[str] = None, - api_key: Optional[str] = None, - api_version: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, + api_version: str | None = None, ) -> AzureCredentials: """Resolve Azure credentials from params, litellm globals, and env vars.""" resolved_api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index a9c7b9f459c..5f910861dfb 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import Any, Optional +from typing import Any from openai import AsyncAzureOpenAI, AzureOpenAI @@ -31,12 +31,12 @@ class AzureTextCompletion(BaseAzureLLM): model: str, messages: list, model_response: ModelResponse, - api_key: Optional[str], + api_key: str | None, api_base: str, api_version: str, api_type: str, - azure_ad_token: Optional[str], - azure_ad_token_provider: Optional[Callable], + azure_ad_token: str | None, + azure_ad_token_provider: Callable | None, print_verbose: Callable, timeout, logging_obj, @@ -44,7 +44,7 @@ class AzureTextCompletion(BaseAzureLLM): litellm_params, logger_fn, acompletion: bool = False, - headers: Optional[dict] = None, + headers: dict | None = None, client=None, ): try: @@ -185,7 +185,7 @@ class AzureTextCompletion(BaseAzureLLM): async def acompletion( self, - api_key: Optional[str], + api_key: str | None, api_version: str, model: str, api_base: str, @@ -194,7 +194,7 @@ class AzureTextCompletion(BaseAzureLLM): model_response: ModelResponse, logging_obj: Any, max_retries: int, - azure_ad_token: Optional[str] = None, + azure_ad_token: str | None = None, client=None, # this is the AsyncAzureOpenAI litellm_params: dict = {}, ): @@ -248,12 +248,12 @@ class AzureTextCompletion(BaseAzureLLM): self, logging_obj, api_base: str, - api_key: Optional[str], + api_key: str | None, api_version: str, data: dict, model: str, timeout: Any, - azure_ad_token: Optional[str] = None, + azure_ad_token: str | None = None, client=None, litellm_params: dict = {}, ): @@ -301,12 +301,12 @@ class AzureTextCompletion(BaseAzureLLM): self, logging_obj, api_base: str, - api_key: Optional[str], + api_key: str | None, api_version: str, data: dict, model: str, timeout: Any, - azure_ad_token: Optional[str] = None, + azure_ad_token: str | None = None, client=None, litellm_params: dict = {}, ): diff --git a/litellm/llms/azure/completion/transformation.py b/litellm/llms/azure/completion/transformation.py index bc7b97c6ef2..99388803983 100644 --- a/litellm/llms/azure/completion/transformation.py +++ b/litellm/llms/azure/completion/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - from ...openai.completion.transformation import OpenAITextCompletionConfig @@ -32,14 +30,14 @@ class AzureOpenAITextConfig(OpenAITextCompletionConfig): def __init__( self, - frequency_penalty: Optional[int] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, + frequency_penalty: int | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, ) -> None: super().__init__( frequency_penalty=frequency_penalty, diff --git a/litellm/llms/azure/containers/transformation.py b/litellm/llms/azure/containers/transformation.py index 30cd3421d1b..e7a75a00446 100644 --- a/litellm/llms/azure/containers/transformation.py +++ b/litellm/llms/azure/containers/transformation.py @@ -1,4 +1,3 @@ -from typing import Optional from urllib.parse import parse_qs, urlparse, urlunparse from litellm.llms.azure.common_utils import BaseAzureLLM @@ -27,7 +26,7 @@ class AzureContainerConfig(OpenAIContainerConfig): def validate_environment( self, headers: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: return BaseAzureLLM._base_validate_azure_environment( headers=headers, @@ -35,7 +34,7 @@ class AzureContainerConfig(OpenAIContainerConfig): ) @staticmethod - def _normalize_api_base(api_base: Optional[str]) -> Optional[str]: + def _normalize_api_base(api_base: str | None) -> str | None: """Strip endpoint-specific path suffixes from api_base to get the resource root.""" if not api_base: return api_base @@ -47,7 +46,7 @@ class AzureContainerConfig(OpenAIContainerConfig): return api_base @staticmethod - def _extract_api_version(api_base: Optional[str]) -> Optional[str]: + def _extract_api_version(api_base: str | None) -> str | None: """Return the api-version query param from api_base if present.""" if not api_base: return None @@ -55,7 +54,7 @@ class AzureContainerConfig(OpenAIContainerConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/azure/cost_calculation.py b/litellm/llms/azure/cost_calculation.py index 2a20c55a6ce..6fddb8523e7 100644 --- a/litellm/llms/azure/cost_calculation.py +++ b/litellm/llms/azure/cost_calculation.py @@ -3,8 +3,6 @@ Helper util for handling azure openai-specific cost calculation - e.g.: prompt caching, audio tokens """ -from typing import Optional, Tuple - from litellm._logging import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.types.utils import Usage @@ -14,9 +12,9 @@ from litellm.utils import get_model_info def cost_per_token( model: str, usage: Usage, - response_time_ms: Optional[float] = 0.0, - service_tier: Optional[str] = None, -) -> Tuple[float, float]: + response_time_ms: float | None = 0.0, + service_tier: str | None = None, +) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/azure/exception_mapping.py b/litellm/llms/azure/exception_mapping.py index 07d589021a7..9e30d2f35e5 100644 --- a/litellm/llms/azure/exception_mapping.py +++ b/litellm/llms/azure/exception_mapping.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm.exceptions import ContentPolicyViolationError @@ -27,7 +27,7 @@ class AzureOpenAIExceptionMapping: # Keep the OpenAI-style body fields populated so downstream (proxy + SDK) # can surface `type` / `code` correctly. - openai_style_body: Dict[str, Any] = { + openai_style_body: dict[str, Any] = { "message": provider_message, "type": provider_type or "invalid_request_error", "code": provider_code or "content_policy_violation", @@ -54,7 +54,7 @@ class AzureOpenAIExceptionMapping: @staticmethod def _extract_azure_error( original_exception: Exception, - ) -> Tuple[Dict[str, Any], Optional[dict]]: + ) -> tuple[dict[str, Any], dict | None]: """Extract Azure OpenAI error payload and inner error details. Azure error formats can vary by endpoint/version. Common shapes: @@ -67,7 +67,7 @@ class AzureOpenAIExceptionMapping: return {}, None # Some SDKs place the payload under "error". - azure_error: Dict[str, Any] + azure_error: dict[str, Any] if isinstance(body_dict.get("error"), dict): azure_error = body_dict.get("error", {}) # type: ignore[assignment] else: diff --git a/litellm/llms/azure/files/handler.py b/litellm/llms/azure/files/handler.py index 7efffd5a102..ce4d35befc1 100644 --- a/litellm/llms/azure/files/handler.py +++ b/litellm/llms/azure/files/handler.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine -from typing import Any, Optional, Union, cast +from typing import Any, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -43,7 +43,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): async def acreate_file( self, create_file_data: CreateFileRequest, - openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI], + openai_client: AsyncAzureOpenAI | AsyncOpenAI, ) -> OpenAIFileObject: verbose_logger.debug("create_file_data=%s", create_file_data) response = await openai_client.files.create(**self._prepare_create_file_data(create_file_data)) # type: ignore[arg-type] @@ -54,23 +54,21 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): self, _is_async: bool, create_file_data: CreateFileRequest, - api_base: Optional[str], - api_key: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, - ) -> Union[OpenAIFileObject, Coroutine[Any, Any, OpenAIFileObject]]: - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = ( - self.get_azure_openai_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - api_version=api_version, - client=client, - _is_async=_is_async, - ) + api_base: str | None, + api_key: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, + ) -> OpenAIFileObject | Coroutine[Any, Any, OpenAIFileObject]: + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + api_version=api_version, + client=client, + _is_async=_is_async, ) if openai_client is None: raise ValueError( @@ -83,7 +81,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): "AzureOpenAI client is not an instance of AsyncAzureOpenAI. Make sure you passed an AsyncAzureOpenAI client." ) return self.acreate_file(create_file_data=create_file_data, openai_client=openai_client) - response = cast(Union[AzureOpenAI, OpenAI], openai_client).files.create( + response = cast(AzureOpenAI | OpenAI, openai_client).files.create( **self._prepare_create_file_data(create_file_data) ) # type: ignore[arg-type] return OpenAIFileObject(**response.model_dump()) @@ -91,7 +89,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): async def afile_content( self, file_content_request: FileContentRequest, - openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI], + openai_client: AsyncAzureOpenAI | AsyncOpenAI, ) -> HttpxBinaryResponseContent: response = await openai_client.files.content(**file_content_request) return HttpxBinaryResponseContent(response=response.response) @@ -100,23 +98,21 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): self, _is_async: bool, file_content_request: FileContentRequest, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - api_version: Optional[str] = None, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, - ) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = ( - self.get_azure_openai_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - api_version=api_version, - client=client, - _is_async=_is_async, - ) + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + api_version: str | None = None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, + ) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + api_version=api_version, + client=client, + _is_async=_is_async, ) if openai_client is None: raise ValueError( @@ -132,14 +128,14 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): file_content_request=file_content_request, openai_client=openai_client, ) - response = cast(Union[AzureOpenAI, OpenAI], openai_client).files.content(**file_content_request) + response = cast(AzureOpenAI | OpenAI, openai_client).files.content(**file_content_request) return HttpxBinaryResponseContent(response=response.response) async def aretrieve_file( self, file_id: str, - openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI], + openai_client: AsyncAzureOpenAI | AsyncOpenAI, ) -> FileObject: response = await openai_client.files.retrieve(file_id=file_id) return response @@ -148,23 +144,21 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): self, _is_async: bool, file_id: str, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - api_version: Optional[str] = None, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + api_version: str | None = None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, ): - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = ( - self.get_azure_openai_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - api_version=api_version, - client=client, - _is_async=_is_async, - ) + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + api_version=api_version, + client=client, + _is_async=_is_async, ) if openai_client is None: raise ValueError( @@ -187,7 +181,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): async def adelete_file( self, file_id: str, - openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI], + openai_client: AsyncAzureOpenAI | AsyncOpenAI, ) -> FileDeleted: response = await openai_client.files.delete(file_id=file_id) @@ -199,24 +193,22 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): self, _is_async: bool, file_id: str, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str] = None, - api_version: Optional[str] = None, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None = None, + api_version: str | None = None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, ): - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = ( - self.get_azure_openai_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - api_version=api_version, - client=client, - _is_async=_is_async, - ) + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + api_version=api_version, + client=client, + _is_async=_is_async, ) if openai_client is None: raise ValueError( @@ -241,8 +233,8 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): async def alist_files( self, - openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI], - purpose: Optional[str] = None, + openai_client: AsyncAzureOpenAI | AsyncOpenAI, + purpose: str | None = None, ): if isinstance(purpose, str): response = await openai_client.files.list(purpose=purpose) @@ -253,24 +245,22 @@ class AzureOpenAIFilesAPI(BaseAzureLLM): def list_files( self, _is_async: bool, - api_base: Optional[str], - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - purpose: Optional[str] = None, - api_version: Optional[str] = None, - client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None, - litellm_params: Optional[dict] = None, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + purpose: str | None = None, + api_version: str | None = None, + client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None, + litellm_params: dict | None = None, ): - openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = ( - self.get_azure_openai_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - api_version=api_version, - client=client, - _is_async=_is_async, - ) + openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + api_version=api_version, + client=client, + _is_async=_is_async, ) if openai_client is None: raise ValueError( diff --git a/litellm/llms/azure/fine_tuning/handler.py b/litellm/llms/azure/fine_tuning/handler.py index 100b5c43d90..3f6af16fcac 100644 --- a/litellm/llms/azure/fine_tuning/handler.py +++ b/litellm/llms/azure/fine_tuning/handler.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine -from typing import Any, Dict, Optional, Union, cast +from typing import Any, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -19,7 +19,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): """ @staticmethod - def _ensure_training_type(create_fine_tuning_job_data: Dict[str, Any]) -> None: + def _ensure_training_type(create_fine_tuning_job_data: dict[str, Any]) -> None: """ Azure requires trainingType in extra_body. Default to 1 (supervised) if omitted. """ @@ -34,7 +34,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): async def acreate_fine_tuning_job( self, create_fine_tuning_job_data: dict, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.create(**create_fine_tuning_job_data) return _litellm_fine_tuning_job_from_response(response, is_azure=True) @@ -42,7 +42,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): async def acancel_fine_tuning_job( self, fine_tuning_job_id: str, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) return _litellm_fine_tuning_job_from_response(response, is_azure=True) @@ -50,7 +50,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): async def aretrieve_fine_tuning_job( self, fine_tuning_job_id: str, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) return _litellm_fine_tuning_job_from_response(response, is_azure=True) @@ -59,17 +59,17 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): self, _is_async: bool, create_fine_tuning_job_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: self._ensure_training_type(create_fine_tuning_job_data) - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -102,15 +102,15 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): self, _is_async: bool, fine_tuning_job_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -142,15 +142,15 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): self, _is_async: bool, fine_tuning_job_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -180,23 +180,16 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): def get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, _is_async: bool = False, - api_version: Optional[str] = None, - litellm_params: Optional[dict] = None, - ) -> Optional[ - Union[ - OpenAI, - AsyncOpenAI, - AzureOpenAI, - AsyncAzureOpenAI, - ] - ]: + api_version: str | None = None, + litellm_params: dict | None = None, + ) -> OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None: # Override to use Azure-specific client initialization if isinstance(client, OpenAI) or isinstance(client, AsyncOpenAI): client = None diff --git a/litellm/llms/azure/image_edit/transformation.py b/litellm/llms/azure/image_edit/transformation.py index d28d92a0770..a784cfe55fd 100644 --- a/litellm/llms/azure/image_edit/transformation.py +++ b/litellm/llms/azure/image_edit/transformation.py @@ -1,4 +1,4 @@ -from typing import Optional, cast +from typing import cast import httpx @@ -28,9 +28,9 @@ class AzureImageEditConfig(OpenAIImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate Azure environment and set up authentication headers. @@ -70,7 +70,7 @@ class AzureImageEditConfig(OpenAIImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -99,7 +99,7 @@ class AzureImageEditConfig(OpenAIImageEditConfig): # Mirrors the fallback chain used by the Azure chat path in common_utils.py, # so callers that set a global / env api_version don't get an unversioned URL. api_version = ( - cast(Optional[str], litellm_params.get("api_version")) + cast(str | None, litellm_params.get("api_version")) or litellm.api_version or get_secret_str("AZURE_API_VERSION") or litellm.AZURE_DEFAULT_API_VERSION diff --git a/litellm/llms/azure/image_generation/dall_e_2_transformation.py b/litellm/llms/azure/image_generation/dall_e_2_transformation.py index 3fe702f57f0..788753e4ea9 100644 --- a/litellm/llms/azure/image_generation/dall_e_2_transformation.py +++ b/litellm/llms/azure/image_generation/dall_e_2_transformation.py @@ -5,5 +5,3 @@ class AzureDallE2ImageGenerationConfig(DallE2ImageGenerationConfig): """ Azure dall-e-2 image generation config """ - - pass diff --git a/litellm/llms/azure/image_generation/dall_e_3_transformation.py b/litellm/llms/azure/image_generation/dall_e_3_transformation.py index 5e0bfcd108f..e6974f9317e 100644 --- a/litellm/llms/azure/image_generation/dall_e_3_transformation.py +++ b/litellm/llms/azure/image_generation/dall_e_3_transformation.py @@ -5,5 +5,3 @@ class AzureDallE3ImageGenerationConfig(DallE3ImageGenerationConfig): """ Azure dall-e-3 image generation config """ - - pass diff --git a/litellm/llms/azure/image_generation/gpt_transformation.py b/litellm/llms/azure/image_generation/gpt_transformation.py index 2d46592e3fb..4ea3e9a592c 100644 --- a/litellm/llms/azure/image_generation/gpt_transformation.py +++ b/litellm/llms/azure/image_generation/gpt_transformation.py @@ -5,5 +5,3 @@ class AzureGPTImageGenerationConfig(GPTImageGenerationConfig): """ Azure gpt-image image generation config """ - - pass diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index dabcd4a1183..e9a57239445 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING, Optional import httpx from httpx import Response @@ -22,13 +22,13 @@ class AzurePassthroughConfig(BasePassthroughConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, endpoint: str, - request_query_params: Optional[dict], + request_query_params: dict | None, litellm_params: dict, - ) -> Tuple["URL", str]: + ) -> tuple["URL", str]: base_target_url = self.get_api_base(api_base) if base_target_url is None: @@ -54,11 +54,11 @@ class AzurePassthroughConfig(BasePassthroughConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return BaseAzureLLM._base_validate_azure_environment( headers=headers, @@ -67,21 +67,21 @@ class AzurePassthroughConfig(BasePassthroughConfig): @staticmethod def get_api_base( - api_base: Optional[str] = None, - ) -> Optional[str]: + api_base: str | None = None, + ) -> str | None: return api_base or get_secret_str("AZURE_API_BASE") @staticmethod def get_api_key( - api_key: Optional[str] = None, - ) -> Optional[str]: + api_key: str | None = None, + ) -> str | None: return api_key or get_secret_str("AZURE_API_KEY") @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: return model - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: return super().get_models(api_key, api_base) def logging_non_streaming_response( diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 86c1ed51b68..5349622560f 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -4,7 +4,7 @@ This file contains the calling Azure OpenAI's `/openai/realtime` endpoint. This requires websockets, and is currently only supported on LiteLLM Proxy. """ -from typing import Any, Optional, cast +from typing import Any, cast from litellm._logging import _redact_string, verbose_proxy_logger from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES @@ -34,9 +34,9 @@ class AzureOpenAIRealtime(AzureChatCompletion): self, api_base: str, model: str, - api_version: Optional[str], - realtime_protocol: Optional[str] = None, - query_params: Optional[RealtimeQueryParams] = None, + api_version: str | None, + realtime_protocol: str | None = None, + query_params: RealtimeQueryParams | None = None, ) -> str: """ Construct Azure realtime WebSocket URL. @@ -89,16 +89,16 @@ class AzureOpenAIRealtime(AzureChatCompletion): model: str, websocket: Any, logging_obj: LiteLLMLogging, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - api_version: Optional[str] = None, - azure_ad_token: Optional[str] = None, - client: Optional[Any] = None, - timeout: Optional[float] = None, - realtime_protocol: Optional[str] = None, - query_params: Optional[RealtimeQueryParams] = None, - user_api_key_dict: Optional[Any] = None, - litellm_metadata: Optional[dict] = None, + api_base: str | None = None, + api_key: str | None = None, + api_version: str | None = None, + azure_ad_token: str | None = None, + client: Any | None = None, + timeout: float | None = None, + realtime_protocol: str | None = None, + query_params: RealtimeQueryParams | None = None, + user_api_key_dict: Any | None = None, + litellm_metadata: dict | None = None, ): import websockets from websockets.asyncio.client import ClientConnection @@ -145,4 +145,3 @@ class AzureOpenAIRealtime(AzureChatCompletion): await websocket.close(code=e.status_code, reason=_redact_string(str(e))) except Exception: verbose_proxy_logger.exception("Error in AzureOpenAIRealtime.async_realtime") - pass diff --git a/litellm/llms/azure/realtime/http_transformation.py b/litellm/llms/azure/realtime/http_transformation.py index 55a86014423..f6376b875a2 100644 --- a/litellm/llms/azure/realtime/http_transformation.py +++ b/litellm/llms/azure/realtime/http_transformation.py @@ -1,20 +1,18 @@ """Azure OpenAI realtime HTTP transformation config (client_secrets + realtime_calls).""" -from typing import Optional - import litellm from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig from litellm.secret_managers.main import get_secret_str class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): - def get_api_base(self, api_base: Optional[str], **kwargs) -> str: + def get_api_base(self, api_base: str | None, **kwargs) -> str: return api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") or "" - def get_api_key(self, api_key: Optional[str], **kwargs) -> str: + def get_api_key(self, api_key: str | None, **kwargs) -> str: return api_key or litellm.api_key or get_secret_str("AZURE_API_KEY") or "" - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/client_secrets?api-version={version}" @@ -23,7 +21,7 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: return { **headers, @@ -31,14 +29,12 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): "Content-Type": "application/json", } - def get_realtime_calls_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/calls?api-version={version}" - def get_transcription_session_url( - self, api_base: Optional[str], model: str, api_version: Optional[str] = None - ) -> str: + def get_transcription_session_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/transcription_sessions?api-version={version}" diff --git a/litellm/llms/azure/responses/o_series_transformation.py b/litellm/llms/azure/responses/o_series_transformation.py index 2cc3e914307..7a88c42cb14 100644 --- a/litellm/llms/azure/responses/o_series_transformation.py +++ b/litellm/llms/azure/responses/o_series_transformation.py @@ -8,7 +8,7 @@ Translations handled by LiteLLM: - Other parameters follow base Azure OpenAI Responses API behavior """ -from typing import TYPE_CHECKING, Any, Dict +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams @@ -56,7 +56,7 @@ class AzureOpenAIOSeriesResponsesAPIConfig(AzureOpenAIResponsesAPIConfig): response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI parameters for Azure OpenAI O-series Responses API. diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index 1a860cca5a9..860b1b1dd5c 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -1,5 +1,5 @@ from copy import deepcopy -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Literal import httpx from openai.types.responses import ResponseReasoningItem @@ -36,7 +36,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): base_supported_params = super().get_supported_openai_params(model) return [param for param in base_supported_params if param not in self.AZURE_UNSUPPORTED_PARAMS] - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params) def get_stripped_model_name(self, model: str) -> str: @@ -47,7 +47,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): model = model.replace("o_series/", "") return model - def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]: + def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]: """ Handle reasoning items to filter out the status field. Issue: https://github.com/BerriAI/litellm/issues/13484 @@ -84,7 +84,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return filtered_item return item - def _validate_input_param(self, input: Union[str, ResponseInputParam]) -> Union[str, ResponseInputParam]: + def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam: """ Override parent method to also filter out 'status' field from message items. Azure OpenAI API does not accept 'status' field in input messages. @@ -96,7 +96,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): # Then filter out status from message items if isinstance(validated_input, list): - filtered_input: List[Any] = [] + filtered_input: list[Any] = [] for item in validated_input: if isinstance(item, dict) and item.get("type") == "message": # Filter out status field from message items @@ -111,11 +111,11 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """No transform applied since inputs are in OpenAI spec already""" stripped_model_name = self.get_stripped_model_name(model) @@ -123,10 +123,10 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): if "tools" in response_api_optional_request_params and isinstance( response_api_optional_request_params["tools"], list ): - new_tools: List[Dict[str, Any]] = [] + new_tools: list[dict[str, Any]] = [] for tool in response_api_optional_request_params["tools"]: if isinstance(tool, dict) and "function" in tool: - new_tool: Dict[str, Any] = deepcopy(tool) + new_tool: dict[str, Any] = deepcopy(tool) function_data = new_tool.pop("function") new_tool.update(function_data) new_tools.append(new_tool) @@ -144,7 +144,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -177,7 +177,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_websocket_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -239,7 +239,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the delete response API request into a URL and data @@ -251,7 +251,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): """ delete_url = self._construct_url_for_response_id_in_path(api_base=api_base, response_id=response_id) - data: Dict = {} + data: dict = {} verbose_logger.debug(f"delete response url={delete_url}") return delete_url, data @@ -264,7 +264,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the get response API request into a URL and data @@ -272,7 +272,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): - GET /v1/responses/{response_id} """ get_url = self._construct_url_for_response_id_in_path(api_base=api_base, response_id=response_id) - data: Dict = {} + data: dict = {} verbose_logger.debug(f"get response url={get_url}") return get_url, data @@ -282,16 +282,16 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: url = self._construct_url_for_response_id_in_path( api_base=api_base, response_id=response_id, path_suffix="/input_items" ) - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if after is not None: params["after"] = after if before is not None: @@ -314,7 +314,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the cancel response API request into a URL and data @@ -328,7 +328,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base=api_base, response_id=response_id, path_suffix="/cancel" ) - data: Dict = {} + data: dict = {} verbose_logger.debug(f"cancel response url={cancel_url}") return cancel_url, data diff --git a/litellm/llms/azure/text_to_speech/transformation.py b/litellm/llms/azure/text_to_speech/transformation.py index 31caf624ffc..2d47d85ebd4 100644 --- a/litellm/llms/azure/text_to_speech/transformation.py +++ b/litellm/llms/azure/text_to_speech/transformation.py @@ -5,7 +5,7 @@ Maps OpenAI TTS spec to Azure Cognitive Services TTS API """ from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union from urllib.parse import urlparse import httpx @@ -62,16 +62,16 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): self, model: str, input: str, - voice: Optional[Union[str, Dict]], - optional_params: Dict, - litellm_params_dict: Dict, + voice: str | dict | None, + optional_params: dict, + litellm_params_dict: dict, logging_obj: "LiteLLMLoggingObj", - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, Any]], + timeout: float | httpx.Timeout, + extra_headers: dict[str, Any] | None, base_llm_http_handler: Any, aspeech: bool, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, **kwargs: Any, ) -> Union[ "HttpxBinaryResponseContent", @@ -101,7 +101,7 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): ) # Convert voice to string if it's a dict (for Azure AVA, voice must be a string) - voice_str: Optional[str] = None + voice_str: str | None = None if isinstance(voice, str): voice_str = voice elif isinstance(voice, dict): @@ -162,9 +162,9 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): def _build_express_as_element( self, content: str, - style: Optional[str] = None, - styledegree: Optional[str] = None, - role: Optional[str] = None, + style: str | None = None, + styledegree: str | None = None, + role: str | None = None, ) -> str: """ Build mstts:express-as element with optional style, styledegree, and role attributes @@ -194,9 +194,9 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): def _get_voice_language( self, - voice_name: Optional[str], - explicit_lang: Optional[str] = None, - ) -> Optional[str]: + voice_name: str | None, + explicit_lang: str | None = None, + ) -> str | None: """ Get the language for the voice element's xml:lang attribute @@ -224,11 +224,11 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Dict = {}, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict = {}, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to Azure AVA TTS parameters """ @@ -238,7 +238,7 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): # OpenAI uses voice as a required param, hence not in optional_params ########################################################## # If it's already an Azure voice, use it directly - mapped_voice: Optional[str] = None + mapped_voice: str | None = None if isinstance(voice, str): if voice in self.VOICE_MAPPINGS: mapped_voice = self.VOICE_MAPPINGS[voice] @@ -282,8 +282,8 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate Azure environment and set up authentication headers @@ -312,7 +312,7 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -385,9 +385,9 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ diff --git a/litellm/llms/azure/vector_stores/transformation.py b/litellm/llms/azure/vector_stores/transformation.py index c340294c6b4..7e80ee46046 100644 --- a/litellm/llms/azure/vector_stores/transformation.py +++ b/litellm/llms/azure/vector_stores/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig from litellm.types.router import GenericLiteLLMParams @@ -8,7 +6,7 @@ from litellm.types.router import GenericLiteLLMParams class AzureOpenAIVectorStoreConfig(OpenAIVectorStoreConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: return BaseAzureLLM._get_base_azure_url( @@ -17,5 +15,5 @@ class AzureOpenAIVectorStoreConfig(OpenAIVectorStoreConfig): route="/openai/vector_stores", ) - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params) diff --git a/litellm/llms/azure/videos/transformation.py b/litellm/llms/azure/videos/transformation.py index 92e7c91fed3..daca765a1d2 100644 --- a/litellm/llms/azure/videos/transformation.py +++ b/litellm/llms/azure/videos/transformation.py @@ -1,15 +1,15 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any -from litellm.types.videos.main import VideoCreateOptionalRequestParams -from litellm.types.router import GenericLiteLLMParams from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.videos.transformation import OpenAIVideoConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.types.videos.main import VideoCreateOptionalRequestParams if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj - from ...base_llm.videos.transformation import BaseVideoConfig as _BaseVideoConfig from ...base_llm.chat.transformation import BaseLLMException as _BaseLLMException + from ...base_llm.videos.transformation import BaseVideoConfig as _BaseVideoConfig LiteLLMLoggingObj = _LiteLLMLoggingObj BaseVideoConfig = _BaseVideoConfig @@ -47,7 +47,7 @@ class AzureVideoConfig(OpenAIVideoConfig): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """No mapping applied since inputs are in OpenAI spec already""" return dict(video_create_optional_params) @@ -55,8 +55,8 @@ class AzureVideoConfig(OpenAIVideoConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: """ Validate Azure environment and set up authentication headers. @@ -77,7 +77,7 @@ class AzureVideoConfig(OpenAIVideoConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/azure_ai/agents/handler.py b/litellm/llms/azure_ai/agents/handler.py index 2061ba70a94..b12f2203e51 100644 --- a/litellm/llms/azure_ai/agents/handler.py +++ b/litellm/llms/azure_ai/agents/handler.py @@ -26,10 +26,6 @@ from collections.abc import AsyncIterator, Callable from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, ) import httpx @@ -96,7 +92,7 @@ class AzureAIAgentsHandler: # ------------------------------------------------------------------------- # Response Helpers # ------------------------------------------------------------------------- - def _extract_content_from_messages(self, messages_data: dict) -> Tuple[str, Optional[List[Dict[str, Any]]]]: + def _extract_content_from_messages(self, messages_data: dict) -> tuple[str, list[dict[str, Any]] | None]: """Extract assistant content and annotations from the messages response. Returns (content, annotations) where annotations is a list of @@ -115,8 +111,8 @@ class AzureAIAgentsHandler: def _transform_annotations( self, - raw_annotations: Optional[List[Dict[str, Any]]], - ) -> Optional[List[Dict[str, Any]]]: + raw_annotations: list[dict[str, Any]] | None, + ) -> list[dict[str, Any]] | None: """Transform Azure AI Foundry annotations to OpenAI-compatible format. Azure AI returns annotations like: @@ -130,7 +126,7 @@ class AzureAIAgentsHandler: if not raw_annotations: return None - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] for ann in raw_annotations: ann_type = ann.get("type") if ann_type == "url_citation": @@ -154,13 +150,13 @@ class AzureAIAgentsHandler: content: str, model_response: ModelResponse, thread_id: str, - messages: List[Dict[str, Any]], - annotations: Optional[List[Dict[str, Any]]] = None, + messages: list[dict[str, Any]], + annotations: list[dict[str, Any]] | None = None, ) -> ModelResponse: """Build the ModelResponse from agent output.""" from litellm.types.utils import Choices, Message, Usage - message_kwargs: Dict[str, Any] = { + message_kwargs: dict[str, Any] = { "content": content, "role": "assistant", } @@ -197,7 +193,7 @@ class AzureAIAgentsHandler: ), ) except Exception as e: - verbose_logger.warning(f"Failed to calculate token usage: {str(e)}") + verbose_logger.warning(f"Failed to calculate token usage: {e!s}") return model_response @@ -207,7 +203,7 @@ class AzureAIAgentsHandler: api_base: str, api_key: str, optional_params: dict, - headers: Optional[dict], + headers: dict | None, ) -> tuple: """Prepare common parameters for completion. @@ -234,7 +230,7 @@ class AzureAIAgentsHandler: return headers, api_version, agent_id, thread_id, api_base - def _check_response(self, response: httpx.Response, expected_codes: List[int], error_msg: str): + def _check_response(self, response: httpx.Response, expected_codes: list[int], error_msg: str): """Check response status and raise error if not expected.""" if response.status_code not in expected_codes: raise AzureAIAgentsError( @@ -248,7 +244,7 @@ class AzureAIAgentsHandler: def completion( self, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], api_base: str, api_key: str, model_response: ModelResponse, @@ -256,8 +252,8 @@ class AzureAIAgentsHandler: optional_params: dict, litellm_params: dict, timeout: float, - client: Optional[HTTPHandler] = None, - headers: Optional[dict] = None, + client: HTTPHandler | None = None, + headers: dict | None = None, ) -> ModelResponse: """Execute synchronous completion using Azure Agent Service.""" from litellm.llms.custom_httpx.http_handler import _get_httpx_client @@ -273,7 +269,7 @@ class AzureAIAgentsHandler: api_base, ) = self._prepare_completion_params(model, api_base, api_key, optional_params, headers) - def make_request(method: str, url: str, json_data: Optional[dict] = None) -> httpx.Response: + def make_request(method: str, url: str, json_data: dict | None = None) -> httpx.Response: if method == "GET": return client.get(url=url, headers=headers) return client.post( @@ -301,10 +297,10 @@ class AzureAIAgentsHandler: api_base: str, api_version: str, agent_id: str, - thread_id: Optional[str], - messages: List[Dict[str, Any]], + thread_id: str | None, + messages: list[dict[str, Any]], optional_params: dict, - ) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]: + ) -> tuple[str, str, list[dict[str, Any]] | None]: """Execute the agent flow synchronously. Returns (thread_id, content, annotations).""" # Step 1: Create thread if not provided @@ -367,7 +363,7 @@ class AzureAIAgentsHandler: async def acompletion( self, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], api_base: str, api_key: str, model_response: ModelResponse, @@ -375,8 +371,8 @@ class AzureAIAgentsHandler: optional_params: dict, litellm_params: dict, timeout: float, - client: Optional[AsyncHTTPHandler] = None, - headers: Optional[dict] = None, + client: AsyncHTTPHandler | None = None, + headers: dict | None = None, ) -> ModelResponse: """Execute asynchronous completion using Azure Agent Service.""" import litellm @@ -396,7 +392,7 @@ class AzureAIAgentsHandler: api_base, ) = self._prepare_completion_params(model, api_base, api_key, optional_params, headers) - async def make_request(method: str, url: str, json_data: Optional[dict] = None) -> httpx.Response: + async def make_request(method: str, url: str, json_data: dict | None = None) -> httpx.Response: if method == "GET": return await client.get(url=url, headers=headers) return await client.post( @@ -424,10 +420,10 @@ class AzureAIAgentsHandler: api_base: str, api_version: str, agent_id: str, - thread_id: Optional[str], - messages: List[Dict[str, Any]], + thread_id: str | None, + messages: list[dict[str, Any]], optional_params: dict, - ) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]: + ) -> tuple[str, str, list[dict[str, Any]] | None]: """Execute the agent flow asynchronously. Returns (thread_id, content, annotations).""" # Step 1: Create thread if not provided @@ -490,14 +486,14 @@ class AzureAIAgentsHandler: async def acompletion_stream( self, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], api_base: str, api_key: str, logging_obj: LiteLLMLoggingObj, optional_params: dict, litellm_params: dict, timeout: float, - headers: Optional[dict] = None, + headers: dict | None = None, ) -> AsyncIterator: """Execute async streaming completion using Azure Agent Service with native SSE.""" import litellm @@ -517,7 +513,7 @@ class AzureAIAgentsHandler: if msg.get("role") in ["user", "system"]: thread_messages.append({"role": "user", "content": msg.get("content", "")}) - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "assistant_id": agent_id, "stream": True, } @@ -566,7 +562,7 @@ class AzureAIAgentsHandler: response_id = f"chatcmpl-{uuid.uuid4().hex[:8]}" created = int(time.time()) thread_id = None - collected_annotations: Optional[List[Dict[str, Any]]] = None + collected_annotations: list[dict[str, Any]] | None = None current_event = None @@ -582,7 +578,7 @@ class AzureAIAgentsHandler: if data_str == "[DONE]": # Send final chunk with finish_reason - final_delta_kwargs: Dict[str, Any] = {"content": None} + final_delta_kwargs: dict[str, Any] = {"content": None} if collected_annotations: final_delta_kwargs["annotations"] = collected_annotations final_chunk = ModelResponseStream( diff --git a/litellm/llms/azure_ai/agents/transformation.py b/litellm/llms/azure_ai/agents/transformation.py index daf87b01579..fb18b0cb651 100644 --- a/litellm/llms/azure_ai/agents/transformation.py +++ b/litellm/llms/azure_ai/agents/transformation.py @@ -21,7 +21,7 @@ The API uses these endpoints: See: https://learn.microsoft.com/en-us/azure/ai-foundry/agents/quickstart """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -47,8 +47,6 @@ else: class AzureAIAgentsError(BaseLLMException): """Exception class for Azure AI Agent Service API errors.""" - pass - class AzureAIAgentsConfig(BaseConfig): """ @@ -104,9 +102,9 @@ class AzureAIAgentsConfig(BaseConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: """ Get Azure AI Agent Service API base and key from params or environment. @@ -120,7 +118,7 @@ class AzureAIAgentsConfig(BaseConfig): return api_base, api_key - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Azure Agents supports minimal OpenAI params since it's an agent runtime. """ @@ -144,12 +142,12 @@ class AzureAIAgentsConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the base URL for Azure AI Agent Service. @@ -188,7 +186,7 @@ class AzureAIAgentsConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -217,7 +215,7 @@ class AzureAIAgentsConfig(BaseConfig): converted_messages.append({"role": role, "content": content}) - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "agent_id": agent_id, "messages": converted_messages, "api_version": self._get_api_version(optional_params), @@ -238,11 +236,11 @@ class AzureAIAgentsConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate and set up environment for Azure Foundry Agents requests. @@ -261,16 +259,14 @@ class AzureAIAgentsConfig(BaseConfig): return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return AzureAIAgentsError(status_code=status_code, message=error_message) def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Azure Agents uses polling, so we fake stream by returning the final response. @@ -296,12 +292,12 @@ class AzureAIAgentsConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the Azure Agents response to LiteLLM ModelResponse format. @@ -312,17 +308,17 @@ class AzureAIAgentsConfig(BaseConfig): @staticmethod def completion( model: str, - messages: List, + messages: list, api_base: str, - api_key: Optional[str], + api_key: str | None, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, optional_params: dict, litellm_params: dict, - timeout: Union[float, int, Any], + timeout: float | Any, acompletion: bool, - stream: Optional[bool] = False, - headers: Optional[dict] = None, + stream: bool | None = False, + headers: dict | None = None, ) -> Any: """ Dispatch method for Azure Foundry Agents completion. diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/__init__.py b/litellm/llms/azure_ai/anthropic/count_tokens/__init__.py index 9605d401f8e..61ca3e11e86 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/__init__.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/__init__.py @@ -13,7 +13,7 @@ from litellm.llms.azure_ai.anthropic.count_tokens.transformation import ( ) __all__ = [ - "AzureAIAnthropicCountTokensHandler", "AzureAIAnthropicCountTokensConfig", + "AzureAIAnthropicCountTokensHandler", "AzureAIAnthropicTokenCounter", ] diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py index 0716e5ae988..3ac04729267 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py @@ -4,7 +4,7 @@ Azure AI Anthropic CountTokens API handler. Uses httpx for HTTP requests with Azure authentication. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -27,14 +27,14 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig): async def handle_count_tokens_request( self, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], api_key: str, api_base: str, - litellm_params: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Dict[str, Any]: + litellm_params: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> dict[str, Any]: """ Handle a CountTokens request using httpx with Azure authentication. @@ -114,14 +114,14 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig): raise except httpx.HTTPStatusError as e: # HTTP errors - preserve the actual status code - verbose_logger.error(f"HTTP error in CountTokens handler: {str(e)}") + verbose_logger.error(f"HTTP error in CountTokens handler: {e!s}") raise AnthropicError( status_code=e.response.status_code, message=e.response.text, ) except Exception as e: - verbose_logger.error(f"Error in CountTokens handler: {str(e)}") + verbose_logger.error(f"Error in CountTokens handler: {e!s}") raise AnthropicError( status_code=500, - message=f"CountTokens processing error: {str(e)}", + message=f"CountTokens processing error: {e!s}", ) diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/token_counter.py b/litellm/llms/azure_ai/anthropic/count_tokens/token_counter.py index 8e1ee73620f..129d7bb7aa9 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/token_counter.py @@ -3,7 +3,7 @@ Azure AI Anthropic Token Counter implementation using the CountTokens API. """ import os -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.llms.azure_ai.anthropic.count_tokens.handler import ( @@ -21,20 +21,20 @@ class AzureAIAnthropicTokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: return custom_llm_provider == LlmProviders.AZURE_AI.value async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: """ Count tokens using Azure AI Anthropic's CountTokens API. diff --git a/litellm/llms/azure_ai/anthropic/handler.py b/litellm/llms/azure_ai/anthropic/handler.py index b7a441169dd..8aa787fb69e 100644 --- a/litellm/llms/azure_ai/anthropic/handler.py +++ b/litellm/llms/azure_ai/anthropic/handler.py @@ -5,7 +5,6 @@ Azure Anthropic handler - reuses AnthropicChatCompletion logic with Azure authen import copy import json from collections.abc import Callable -from typing import TYPE_CHECKING, Union import httpx @@ -19,9 +18,6 @@ from litellm.utils import CustomStreamWrapper from .transformation import AzureAnthropicConfig -if TYPE_CHECKING: - pass - class AzureAnthropicChatCompletion(AnthropicChatCompletion): """ @@ -45,7 +41,7 @@ class AzureAnthropicChatCompletion(AnthropicChatCompletion): api_key, logging_obj, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, acompletion=None, logger_fn=None, diff --git a/litellm/llms/azure_ai/anthropic/messages_transformation.py b/litellm/llms/azure_ai/anthropic/messages_transformation.py index 9b05e754b7f..583870c5efe 100644 --- a/litellm/llms/azure_ai/anthropic/messages_transformation.py +++ b/litellm/llms/azure_ai/anthropic/messages_transformation.py @@ -2,7 +2,7 @@ Azure Anthropic messages transformation config - extends AnthropicMessagesConfig with Azure authentication """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import Any from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, @@ -10,9 +10,6 @@ from litellm.llms.anthropic.experimental_pass_through.messages.transformation im from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.types.router import GenericLiteLLMParams -if TYPE_CHECKING: - pass - class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): """ @@ -22,7 +19,7 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "azure_ai" def should_strip_billing_metadata(self) -> bool: @@ -32,12 +29,12 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: """ Validate environment and set up Azure authentication headers for /v1/messages endpoint. Azure Anthropic uses x-api-key header (not api-key). @@ -77,12 +74,12 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Azure Anthropic /v1/messages endpoint. @@ -99,10 +96,7 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): # Ensure the URL ends with /v1/messages api_base = api_base.rstrip("/") - if api_base.endswith("/v1/messages"): - # Already correct - pass - elif api_base.endswith("/anthropic/v1/messages"): + if api_base.endswith("/v1/messages") or api_base.endswith("/anthropic/v1/messages"): # Already correct pass else: @@ -120,7 +114,7 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): return api_base - def _remove_scope_from_cache_control(self, anthropic_messages_request: Dict) -> None: + def _remove_scope_from_cache_control(self, anthropic_messages_request: dict) -> None: """ Remove `scope` field from cache_control for Azure AI Foundry. @@ -154,11 +148,11 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: anthropic_messages_request = super().transform_anthropic_messages_request( model=model, messages=messages, diff --git a/litellm/llms/azure_ai/anthropic/transformation.py b/litellm/llms/azure_ai/anthropic/transformation.py index 26323ba707d..ba4bd4c45fa 100644 --- a/litellm/llms/azure_ai/anthropic/transformation.py +++ b/litellm/llms/azure_ai/anthropic/transformation.py @@ -2,15 +2,11 @@ Azure Anthropic transformation config - extends AnthropicConfig with Azure authentication """ -from typing import TYPE_CHECKING, Dict, List, Optional, Union from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.types.llms.openai import AllMessageValues from litellm.types.router import GenericLiteLLMParams -if TYPE_CHECKING: - pass - def _promote_extra_body_to_optional_params(optional_params: dict) -> None: """Promote anthropic-native passthrough keys out of ``extra_body``. @@ -37,7 +33,7 @@ class AzureAnthropicConfig(AnthropicConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "azure_ai" def should_strip_billing_metadata(self) -> bool: @@ -47,12 +43,12 @@ class AzureAnthropicConfig(AnthropicConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, - litellm_params: Union[dict, GenericLiteLLMParams], - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + litellm_params: dict | GenericLiteLLMParams, + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: """ Validate environment and set up Azure authentication headers. Azure supports: @@ -110,7 +106,7 @@ class AzureAnthropicConfig(AnthropicConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/azure_ai/azure_model_router/transformation.py b/litellm/llms/azure_ai/azure_model_router/transformation.py index f3045283840..9138b46839c 100644 --- a/litellm/llms/azure_ai/azure_model_router/transformation.py +++ b/litellm/llms/azure_ai/azure_model_router/transformation.py @@ -5,7 +5,7 @@ The Model Router is a special Azure AI deployment that automatically routes requ to the best available model. It has specific cost tracking requirements. """ -from typing import Any, List, Optional +from typing import Any from httpx import Response @@ -28,7 +28,7 @@ class AzureModelRouterConfig(AzureAIStudioConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -53,12 +53,12 @@ class AzureModelRouterConfig(AzureAIStudioConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform response for Model Router. @@ -88,7 +88,7 @@ class AzureModelRouterConfig(AzureAIStudioConfig): ) return model_response - def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> Optional[dict]: + def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> dict | None: """ Calculate additional costs for Azure Model Router. diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 27a98347087..707ddc9e12b 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -1,6 +1,6 @@ import enum import re -from typing import Any, List, Optional, Tuple, cast +from typing import Any, cast from urllib.parse import urlparse import httpx @@ -29,7 +29,7 @@ class AzureFoundryErrorStrings(str, enum.Enum): class AzureAIStudioConfig(OpenAIConfig): - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: model_supports_tool_choice = True # azure ai supports this by default if not supports_tool_choice(model=f"azure_ai/{model}"): model_supports_tool_choice = False @@ -61,11 +61,11 @@ class AzureAIStudioConfig(OpenAIConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key: if api_base and self._should_use_api_key_header(api_base): @@ -93,12 +93,12 @@ class AzureAIStudioConfig(OpenAIConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Constructs a complete URL for the API request. @@ -123,7 +123,7 @@ class AzureAIStudioConfig(OpenAIConfig): original_url = httpx.URL(api_base) # Extract api_version or use default - api_version = cast(Optional[str], litellm_params.get("api_version")) + api_version = cast(str | None, litellm_params.get("api_version")) # Create a new dictionary with existing params query_params = dict(original_url.params) @@ -143,7 +143,7 @@ class AzureAIStudioConfig(OpenAIConfig): return str(final_url) - def get_required_params(self) -> List[ProviderField]: + def get_required_params(self) -> list[ProviderField]: """For a given provider, return it's required fields with a description""" return [ ProviderField( @@ -162,9 +162,9 @@ class AzureAIStudioConfig(OpenAIConfig): def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, - ) -> List: + ) -> list: """ - Azure AI Studio doesn't support content as a list. This handles: 1. Transforms list content to a string. @@ -180,7 +180,7 @@ class AzureAIStudioConfig(OpenAIConfig): message["content"] = texts return messages - def _is_azure_openai_model(self, model: str, api_base: Optional[str]) -> bool: + def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool: try: if "/" in model: model = model.split("/", 1)[1] @@ -198,22 +198,22 @@ class AzureAIStudioConfig(OpenAIConfig): def _get_openai_compatible_provider_info( self, model: str, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, custom_llm_provider: str, - ) -> Tuple[Optional[str], Optional[str], str]: + ) -> tuple[str | None, str | None, str]: api_base = api_base or get_secret_str("AZURE_AI_API_BASE") dynamic_api_key = api_key or get_secret_str("AZURE_AI_API_KEY") if self._is_azure_openai_model(model=model, api_base=api_base): - verbose_logger.debug("Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(model)) + verbose_logger.debug(f"Model={model} is Azure OpenAI model. Setting custom_llm_provider='azure'.") custom_llm_provider = "azure" return api_base, dynamic_api_key, custom_llm_provider def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -231,12 +231,12 @@ class AzureAIStudioConfig(OpenAIConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: model_response.model = f"azure_ai/{model}" return super().transform_response( @@ -277,7 +277,7 @@ class AzureAIStudioConfig(OpenAIConfig): def transform_request_on_unprocessable_entity_error(self, e: httpx.HTTPStatusError, request_data: dict) -> dict: error_text = e.response.text - _messages = cast(Optional[List[AllMessageValues]], request_data.get("messages")) + _messages = cast(list[AllMessageValues] | None, request_data.get("messages")) if "unknown field: parameter index is not a valid field" in error_text and _messages is not None: litellm.remove_index_from_tool_calls( messages=_messages, @@ -307,7 +307,7 @@ class AzureAIStudioConfig(OpenAIConfig): request_data.pop(param, None) return request_data - def _extract_params_to_drop_from_error_text(self, error_text: str) -> Optional[List[str]]: + def _extract_params_to_drop_from_error_text(self, error_text: str) -> list[str] | None: """ Error text looks like this" "Extra parameters ['stream_options', 'extra-parameters'] are not allowed when extra-parameters is not set or set to be 'error'. diff --git a/litellm/llms/azure_ai/common_utils.py b/litellm/llms/azure_ai/common_utils.py index 9965aa693c3..e00c0f6e739 100644 --- a/litellm/llms/azure_ai/common_utils.py +++ b/litellm/llms/azure_ai/common_utils.py @@ -1,4 +1,4 @@ -from typing import List, Literal, Optional +from typing import Literal import litellm from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter @@ -9,7 +9,7 @@ from litellm.types.llms.openai import AllMessageValues class AzureFoundryModelInfo(BaseLLMModelInfo): """Model info for Azure AI / Azure Foundry models.""" - def __init__(self, model: Optional[str] = None): + def __init__(self, model: str | None = None): self._model = model @staticmethod @@ -38,19 +38,19 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): return "default" @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or litellm.api_base or get_secret_str("AZURE_AI_API_BASE") @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or litellm.api_key or litellm.openai_key or get_secret_str("AZURE_AI_API_KEY") @property - def api_version(self, api_version: Optional[str] = None) -> Optional[str]: + def api_version(self, api_version: str | None = None) -> str | None: api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") return api_version - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create a token counter for Azure AI. @@ -66,7 +66,7 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): return AzureAIAnthropicTokenCounter() return None - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: """ Returns a list of models supported by Azure AI. @@ -155,11 +155,11 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """Azure Foundry sends api key in query params""" raise NotImplementedError("Azure Foundry does not support environment validation") diff --git a/litellm/llms/azure_ai/cost_calculator.py b/litellm/llms/azure_ai/cost_calculator.py index e9c8cac0078..6cc0cb20e27 100644 --- a/litellm/llms/azure_ai/cost_calculator.py +++ b/litellm/llms/azure_ai/cost_calculator.py @@ -3,8 +3,6 @@ Azure AI cost calculation helper. Handles Azure AI Foundry Model Router flat cost and other Azure AI specific pricing. """ -from typing import Optional, Tuple - from litellm._logging import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.types.utils import Usage @@ -59,10 +57,10 @@ def calculate_azure_model_router_flat_cost(model: str, prompt_tokens: int) -> fl def cost_per_token( model: str, usage: Usage, - response_time_ms: Optional[float] = 0.0, - request_model: Optional[str] = None, - service_tier: Optional[str] = None, -) -> Tuple[float, float]: + response_time_ms: float | None = 0.0, + request_model: str | None = None, + service_tier: str | None = None, +) -> tuple[float, float]: """ Calculate the cost per token for Azure AI models. diff --git a/litellm/llms/azure_ai/embed/cohere_transformation.py b/litellm/llms/azure_ai/embed/cohere_transformation.py index 8a28d2f652a..0ef5d6b5543 100644 --- a/litellm/llms/azure_ai/embed/cohere_transformation.py +++ b/litellm/llms/azure_ai/embed/cohere_transformation.py @@ -9,8 +9,6 @@ Convers Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-titan-embed-text.html """ -from typing import List, Optional, Tuple - from litellm.types.llms.azure_ai import ImageEmbeddingInput, ImageEmbeddingRequest from litellm.types.llms.openai import EmbeddingCreateParams from litellm.types.utils import EmbeddingResponse, Usage @@ -29,24 +27,24 @@ class AzureAICohereConfig: return model - def _transform_request_image_embeddings(self, input: List[str], optional_params: dict) -> ImageEmbeddingRequest: + def _transform_request_image_embeddings(self, input: list[str], optional_params: dict) -> ImageEmbeddingRequest: """ Assume all str in list is base64 encoded string """ - image_input: List[ImageEmbeddingInput] = [] + image_input: list[ImageEmbeddingInput] = [] for i in input: embedding_input = ImageEmbeddingInput(image=i) image_input.append(embedding_input) return ImageEmbeddingRequest(input=image_input, **optional_params) def _transform_request( - self, input: List[str], optional_params: dict, model: str - ) -> Tuple[ImageEmbeddingRequest, EmbeddingCreateParams, List[int]]: + self, input: list[str], optional_params: dict, model: str + ) -> tuple[ImageEmbeddingRequest, EmbeddingCreateParams, list[int]]: """ Return the list of input to `/image/embeddings`, `/v1/embeddings`, list of image_embedding_idx for recombination """ - image_embeddings: List[str] = [] - image_embedding_idx: List[int] = [] + image_embeddings: list[str] = [] + image_embedding_idx: list[int] = [] for idx, i in enumerate(input): """ - is base64 -> route to image embeddings @@ -68,10 +66,10 @@ class AzureAICohereConfig: return image_embeddings_request, v1_embeddings_request, image_embedding_idx def _transform_response(self, response: EmbeddingResponse) -> EmbeddingResponse: - additional_headers: Optional[dict] = response._hidden_params.get("additional_headers") + additional_headers: dict | None = response._hidden_params.get("additional_headers") if additional_headers: # CALCULATE USAGE - input_tokens: Optional[str] = additional_headers.get("llm_provider-num_tokens") + input_tokens: str | None = additional_headers.get("llm_provider-num_tokens") if input_tokens: if response.usage: response.usage.prompt_tokens = int(input_tokens) @@ -79,7 +77,7 @@ class AzureAICohereConfig: response.usage = Usage(prompt_tokens=int(input_tokens)) # SET MODEL - base_model: Optional[str] = additional_headers.get("llm_provider-azureml-model-group") + base_model: str | None = additional_headers.get("llm_provider-azureml-model-group") if base_model: response.model = self._map_azure_model_group(base_model) diff --git a/litellm/llms/azure_ai/embed/handler.py b/litellm/llms/azure_ai/embed/handler.py index 62c80bd2568..970e6dfd965 100644 --- a/litellm/llms/azure_ai/embed/handler.py +++ b/litellm/llms/azure_ai/embed/handler.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - from openai import OpenAI import litellm @@ -19,11 +17,11 @@ from .cohere_transformation import AzureAICohereConfig class AzureAIEmbedding(OpenAIChatCompletion): def _process_response( self, - image_embedding_responses: Optional[List], - text_embedding_responses: Optional[List], - image_embeddings_idx: List[int], + image_embedding_responses: list | None, + text_embedding_responses: list | None, + image_embeddings_idx: list[int], model_response: EmbeddingResponse, - input: List, + input: list, ): combined_responses = [] if image_embedding_responses is not None and text_embedding_responses is not None: @@ -57,9 +55,9 @@ class AzureAIEmbedding(OpenAIChatCompletion): logging_obj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: str | None, + api_base: str | None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> EmbeddingResponse: if client is None or not isinstance(client, AsyncHTTPHandler): client = get_async_httpx_client( @@ -67,12 +65,12 @@ class AzureAIEmbedding(OpenAIChatCompletion): params={"timeout": timeout}, ) - url = "{}/images/embeddings".format(api_base) + url = f"{api_base}/images/embeddings" response = await client.post( url=url, json=data, # type: ignore - headers={"Authorization": "Bearer {}".format(api_key)}, + headers={"Authorization": f"Bearer {api_key}"}, ) embedding_response = response.json() @@ -94,9 +92,9 @@ class AzureAIEmbedding(OpenAIChatCompletion): logging_obj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: str | None, + api_base: str | None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ): if api_base is None: raise ValueError( @@ -110,12 +108,12 @@ class AzureAIEmbedding(OpenAIChatCompletion): if client is None or not isinstance(client, HTTPHandler): client = HTTPHandler(timeout=timeout, concurrent_limit=1) - url = "{}/images/embeddings".format(api_base) + url = f"{api_base}/images/embeddings" response = client.post( url=url, json=data, # type: ignore - headers={"Authorization": "Bearer {}".format(api_key)}, + headers={"Authorization": f"Bearer {api_key}"}, ) embedding_response = response.json() @@ -132,13 +130,13 @@ class AzureAIEmbedding(OpenAIChatCompletion): async def async_embedding( self, model: str, - input: List, + input: list, timeout: float, logging_obj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, client=None, ) -> EmbeddingResponse: ( @@ -147,8 +145,8 @@ class AzureAIEmbedding(OpenAIChatCompletion): image_embeddings_idx, ) = AzureAICohereConfig()._transform_request(input=input, optional_params=optional_params, model=model) - image_embedding_responses: Optional[List] = None - text_embedding_responses: Optional[List] = None + image_embedding_responses: list | None = None + text_embedding_responses: list | None = None if image_embeddings_request["input"]: image_response = await self.async_image_embedding( @@ -195,16 +193,16 @@ class AzureAIEmbedding(OpenAIChatCompletion): def embedding( self, model: str, - input: List, + input: list, timeout: float, logging_obj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, client=None, aembedding=None, - max_retries: Optional[int] = None, + max_retries: int | None = None, shared_session=None, ) -> EmbeddingResponse: """ @@ -233,8 +231,8 @@ class AzureAIEmbedding(OpenAIChatCompletion): image_embeddings_idx, ) = AzureAICohereConfig()._transform_request(input=input, optional_params=optional_params, model=model) - image_embedding_responses: Optional[List] = None - text_embedding_responses: Optional[List] = None + image_embedding_responses: list | None = None + text_embedding_responses: list | None = None if image_embeddings_request["input"]: image_response = self.image_embedding( diff --git a/litellm/llms/azure_ai/image_edit/__init__.py b/litellm/llms/azure_ai/image_edit/__init__.py index 42ece6d19ec..a03de5ecba3 100644 --- a/litellm/llms/azure_ai/image_edit/__init__.py +++ b/litellm/llms/azure_ai/image_edit/__init__.py @@ -11,8 +11,8 @@ from .mai_transformation import AzureFoundryMAIImageEditConfig from .transformation import AzureFoundryFluxImageEditConfig __all__ = [ - "AzureFoundryFluxImageEditConfig", "AzureFoundryFlux2ImageEditConfig", + "AzureFoundryFluxImageEditConfig", "AzureFoundryMAIImageEditConfig", ] diff --git a/litellm/llms/azure_ai/image_edit/flux2_transformation.py b/litellm/llms/azure_ai/image_edit/flux2_transformation.py index 1bc3bdcddc1..30305592e8f 100644 --- a/litellm/llms/azure_ai/image_edit/flux2_transformation.py +++ b/litellm/llms/azure_ai/image_edit/flux2_transformation.py @@ -1,6 +1,6 @@ import base64 from io import BufferedReader -from typing import Any, Dict, Optional, Tuple +from typing import Any from httpx._types import RequestFiles @@ -42,12 +42,12 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI params to FLUX 2 params. FLUX 2 uses the same param names as OpenAI for supported params. """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} supported_params = self.get_supported_openai_params(model) for key, value in dict(image_edit_optional_params).items(): @@ -64,9 +64,9 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate Azure AI Foundry environment and set up authentication @@ -89,12 +89,12 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform image edit request for FLUX 2. @@ -110,7 +110,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): image_b64 = self._convert_image_to_base64(image) # Build request body with required params - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "prompt": prompt, "image": image_b64, "model": model, @@ -145,7 +145,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/azure_ai/image_edit/mai_transformation.py b/litellm/llms/azure_ai/image_edit/mai_transformation.py index aa1092b0a53..b639fe49ff3 100644 --- a/litellm/llms/azure_ai/image_edit/mai_transformation.py +++ b/litellm/llms/azure_ai/image_edit/mai_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -33,8 +33,8 @@ class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: - optional_params: Dict[str, Any] = {} + ) -> dict: + optional_params: dict[str, Any] = {} supported_params = self.get_supported_openai_params(model) for key, value in dict(image_edit_optional_params).items(): @@ -87,9 +87,9 @@ class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: api_key = AzureFoundryModelInfo.get_api_key(api_key) @@ -105,7 +105,7 @@ class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = AzureFoundryModelInfo.get_api_base(api_base) @@ -125,12 +125,12 @@ class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: request_params = { "model": model, **image_edit_optional_request_params, @@ -139,7 +139,7 @@ class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): request_params["prompt"] = prompt data_without_files = {key: value for key, value in request_params.items() if key not in ["image", "mask"]} - files_list: List[Tuple[str, Any]] = [] + files_list: list[tuple[str, Any]] = [] if image is not None: image_list = [image] if not isinstance(image, list) else image diff --git a/litellm/llms/azure_ai/image_edit/transformation.py b/litellm/llms/azure_ai/image_edit/transformation.py index 5393a0ba55f..bd655232622 100644 --- a/litellm/llms/azure_ai/image_edit/transformation.py +++ b/litellm/llms/azure_ai/image_edit/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional - import httpx import litellm @@ -24,9 +22,9 @@ class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate Azure AI Foundry environment and set up authentication @@ -49,7 +47,7 @@ class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/azure_ai/image_generation/__init__.py b/litellm/llms/azure_ai/image_generation/__init__.py index 70821d5d764..fd511654665 100644 --- a/litellm/llms/azure_ai/image_generation/__init__.py +++ b/litellm/llms/azure_ai/image_generation/__init__.py @@ -10,10 +10,10 @@ from .gpt_transformation import AzureFoundryGPTImageGenerationConfig from .mai_transformation import AzureFoundryMAIImageGenerationConfig __all__ = [ - "AzureFoundryFluxImageGenerationConfig", - "AzureFoundryGPTImageGenerationConfig", "AzureFoundryDallE2ImageGenerationConfig", "AzureFoundryDallE3ImageGenerationConfig", + "AzureFoundryFluxImageGenerationConfig", + "AzureFoundryGPTImageGenerationConfig", "AzureFoundryMAIImageGenerationConfig", ] diff --git a/litellm/llms/azure_ai/image_generation/dall_e_2_transformation.py b/litellm/llms/azure_ai/image_generation/dall_e_2_transformation.py index 1ef93366f71..2c8a5426628 100644 --- a/litellm/llms/azure_ai/image_generation/dall_e_2_transformation.py +++ b/litellm/llms/azure_ai/image_generation/dall_e_2_transformation.py @@ -5,5 +5,3 @@ class AzureFoundryDallE2ImageGenerationConfig(DallE2ImageGenerationConfig): """ Azure dall-e-2 image generation config """ - - pass diff --git a/litellm/llms/azure_ai/image_generation/dall_e_3_transformation.py b/litellm/llms/azure_ai/image_generation/dall_e_3_transformation.py index 4688a5c3caa..3ceb9b405da 100644 --- a/litellm/llms/azure_ai/image_generation/dall_e_3_transformation.py +++ b/litellm/llms/azure_ai/image_generation/dall_e_3_transformation.py @@ -5,5 +5,3 @@ class AzureFoundryDallE3ImageGenerationConfig(DallE3ImageGenerationConfig): """ Azure dall-e-3 image generation config """ - - pass diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index a883893ceba..dcfc87da0f5 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.openai.image_generation import GPTImageGenerationConfig @@ -16,9 +14,9 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): @staticmethod def get_flux2_image_generation_url( - api_base: Optional[str], + api_base: str | None, model: str, - api_version: Optional[str], + api_version: str | None, ) -> str: """ Constructs the complete URL for Azure AI FLUX 2 image generation. diff --git a/litellm/llms/azure_ai/image_generation/gpt_transformation.py b/litellm/llms/azure_ai/image_generation/gpt_transformation.py index 3eead307463..02196d76c38 100644 --- a/litellm/llms/azure_ai/image_generation/gpt_transformation.py +++ b/litellm/llms/azure_ai/image_generation/gpt_transformation.py @@ -5,5 +5,3 @@ class AzureFoundryGPTImageGenerationConfig(GPTImageGenerationConfig): """ Azure gpt-image-1 image generation config """ - - pass diff --git a/litellm/llms/azure_ai/image_generation/mai_transformation.py b/litellm/llms/azure_ai/image_generation/mai_transformation.py index 7e79ea0b976..d63727381e8 100644 --- a/litellm/llms/azure_ai/image_generation/mai_transformation.py +++ b/litellm/llms/azure_ai/image_generation/mai_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -22,8 +22,8 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): @staticmethod def get_mai_image_generation_url( - api_base: Optional[str], - api_version: Optional[str], + api_base: str | None, + api_version: str | None, ) -> str: if api_base is None: raise ValueError("api_base is required for Azure AI MAI image generation") @@ -44,8 +44,8 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): @staticmethod def get_mai_image_edit_url( - api_base: Optional[str], - api_version: Optional[str], + api_base: str | None, + api_version: str | None, ) -> str: if api_base is None: raise ValueError("api_base is required for Azure AI MAI image editing") @@ -70,7 +70,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): return "maiimage" in model_normalized @staticmethod - def normalize_mai_image_usage(usage: Optional[Dict[str, Any]]) -> Dict[str, Any]: + def normalize_mai_image_usage(usage: dict[str, Any] | None) -> dict[str, Any]: """Map Azure MAI usage fields to OpenAI ImageUsage schema.""" if usage is None: return { @@ -126,7 +126,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): ) return normalized_usage - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "size"] def map_openai_params( @@ -200,8 +200,8 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: try: response = raw_response.json() diff --git a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py b/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py index 7d915892a28..66ce84cea0f 100644 --- a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py +++ b/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py @@ -11,20 +11,19 @@ The operation location must be polled until the analysis completes. import asyncio import re import time -from typing import Any, Dict +from typing import Any from urllib.parse import quote import httpx from pydantic import BaseModel from litellm._logging import verbose_logger -from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.constants import ( AZURE_DOCUMENT_INTELLIGENCE_API_VERSION, AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI, AZURE_OPERATION_POLLING_TIMEOUT, ) -from litellm.litellm_core_utils.url_utils import encode_url_path_segment +from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin, encode_url_path_segment from litellm.llms.base_llm.ocr.transformation import ( BaseOCRConfig, DocumentType, @@ -205,13 +204,13 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers for Azure Document Intelligence. @@ -374,7 +373,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): raise ValueError("Document URL is required") # Build Azure DI request - data: Dict[str, Any] = {} + data: dict[str, Any] = {} # Check if it's a data URI (base64) if document_url.startswith("data:"): @@ -498,7 +497,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): def _poll_operation_sync( self, operation_url: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: int, ) -> httpx.Response: """ @@ -541,7 +540,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): async def _poll_operation_async( self, operation_url: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: int, ) -> httpx.Response: """ @@ -579,7 +578,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig): retry_after = self._get_retry_after(response=response) await asyncio.sleep(retry_after) - def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, Dict[str, str]]: + def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, dict[str, str]]: operation_url = raw_response.headers.get("Operation-Location") if not operation_url: raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found") diff --git a/litellm/llms/azure_ai/ocr/transformation.py b/litellm/llms/azure_ai/ocr/transformation.py index a57e3e869cf..d757a7f1378 100644 --- a/litellm/llms/azure_ai/ocr/transformation.py +++ b/litellm/llms/azure_ai/ocr/transformation.py @@ -2,8 +2,6 @@ Azure AI OCR transformation implementation. """ -from typing import Dict - from litellm._logging import verbose_logger from litellm.litellm_core_utils.prompt_templates.image_handling import ( async_convert_url_to_base64, @@ -37,13 +35,13 @@ class AzureAIOCRConfig(MistralOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers for Azure AI OCR. diff --git a/litellm/llms/azure_ai/rerank/transformation.py b/litellm/llms/azure_ai/rerank/transformation.py index 928f53bd485..ebb995d9691 100644 --- a/litellm/llms/azure_ai/rerank/transformation.py +++ b/litellm/llms/azure_ai/rerank/transformation.py @@ -2,8 +2,6 @@ Translate between Cohere's `/rerank` format and Azure AI's `/rerank` format. """ -from typing import Optional - import httpx import litellm @@ -21,9 +19,9 @@ class AzureAIRerankConfig(CohereRerankConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, model: str, - optional_params: Optional[dict] = None, + optional_params: dict | None = None, ) -> str: if api_base is None: raise ValueError( @@ -62,8 +60,8 @@ class AzureAIRerankConfig(CohereRerankConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - optional_params: Optional[dict] = None, + api_key: str | None = None, + optional_params: dict | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("AZURE_AI_API_KEY") or litellm.azure_key @@ -90,7 +88,7 @@ class AzureAIRerankConfig(CohereRerankConfig): raw_response: httpx.Response, model_response: RerankResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -109,7 +107,7 @@ class AzureAIRerankConfig(CohereRerankConfig): rerank_response._hidden_params["model"] = base_model return rerank_response - def _get_base_model(self, azure_model_group: Optional[str]) -> Optional[str]: + def _get_base_model(self, azure_model_group: str | None) -> str | None: if azure_model_group is None: return None if azure_model_group == "offer-cohere-rerank-mul-paygo": diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index da6a4a93cd8..88a38fc1ec7 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -53,14 +53,14 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): } } - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: basic_headers = self._base_validate_azure_environment(headers, litellm_params) basic_headers.update({"Content-Type": "application/json"}) return basic_headers def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -86,13 +86,13 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict[str, Any]]: """ Transform search request for Azure AI Search API @@ -132,7 +132,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): ) query_vector = embedding_response.data[0]["embedding"] except Exception as e: - raise Exception(f"Failed to generate embedding for query: {str(e)}") + raise Exception(f"Failed to generate embedding for query: {e!s}") # Azure AI Search endpoint for search index_name = vector_store_id # vector_store_id is the index name @@ -187,7 +187,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): results = response_json.get("value", []) # Transform results to standard format - search_results: List[VectorStoreSearchResult] = [] + search_results: list[VectorStoreSearchResult] = [] for result in results: # Extract document ID document_id = result.get("id", "") @@ -245,7 +245,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: raise NotImplementedError def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: diff --git a/litellm/llms/base.py b/litellm/llms/base.py index 56d1643dd4e..e532db1f1e7 100644 --- a/litellm/llms/base.py +++ b/litellm/llms/base.py @@ -1,5 +1,5 @@ ## This is a template base class to be used for adding new LLM providers via API calls -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -11,7 +11,7 @@ if TYPE_CHECKING: class BaseLLM: - _client_session: Optional[httpx.Client] = None + _client_session: httpx.Client | None = None def process_response( self, @@ -22,7 +22,7 @@ class BaseLLM: logging_obj: Any, optional_params: dict, api_key: str, - data: Union[dict, str], + data: dict | str, messages: list, print_verbose, encoding, @@ -41,7 +41,7 @@ class BaseLLM: logging_obj: Any, optional_params: dict, api_key: str, - data: Union[dict, str], + data: dict | str, messages: list, print_verbose, encoding, @@ -75,9 +75,7 @@ class BaseLLM: if hasattr(self, "_aclient_session"): await self._aclient_session.aclose() # type: ignore - def validate_environment( - self, *args, **kwargs - ) -> Optional[Any]: # set up the environment required to run the model + def validate_environment(self, *args, **kwargs) -> Any | None: # set up the environment required to run the model return None def completion(self, *args, **kwargs) -> Any: # logic for parsing in - calling - parsing out model completion calls diff --git a/litellm/llms/base_llm/__init__.py b/litellm/llms/base_llm/__init__.py index 665e242969c..5040bf4351c 100644 --- a/litellm/llms/base_llm/__init__.py +++ b/litellm/llms/base_llm/__init__.py @@ -7,11 +7,11 @@ from .image_edit.transformation import BaseImageEditConfig from .image_generation.transformation import BaseImageGenerationConfig __all__ = [ - "BaseImageGenerationConfig", - "BaseConfig", - "BaseAudioTranscriptionConfig", "BaseAnthropicMessagesConfig", + "BaseAudioTranscriptionConfig", + "BaseBatchesConfig", + "BaseConfig", "BaseEmbeddingConfig", "BaseImageEditConfig", - "BaseBatchesConfig", + "BaseImageGenerationConfig", ] diff --git a/litellm/llms/base_llm/agents/transformation.py b/litellm/llms/base_llm/agents/transformation.py index 508e54cb7ab..970639939f1 100644 --- a/litellm/llms/base_llm/agents/transformation.py +++ b/litellm/llms/base_llm/agents/transformation.py @@ -10,7 +10,7 @@ InteractionsHTTPHandler). """ from abc import ABC, abstractmethod -from typing import Any, Dict, Optional, Tuple, Union +from typing import Any import httpx @@ -34,25 +34,25 @@ class BaseAgentsAPIConfig(ABC): @abstractmethod def get_complete_url( self, - api_base: Optional[str], - litellm_params: Dict[str, Any], + api_base: str | None, + litellm_params: dict[str, Any], ) -> str: """Return the full URL for POST /agents (create).""" @abstractmethod def validate_environment( self, - headers: Dict[str, str], - litellm_params: Dict[str, Any], - ) -> Dict[str, str]: + headers: dict[str, str], + litellm_params: dict[str, Any], + ) -> dict[str, str]: """Validate credentials and return auth headers.""" @abstractmethod def transform_create_request( self, name: str, - litellm_params: Dict[str, Any], - ) -> Dict[str, Any]: + litellm_params: dict[str, Any], + ) -> dict[str, Any]: """Map name + litellm_params to the provider's create-agent body.""" @abstractmethod @@ -70,9 +70,9 @@ class BaseAgentsAPIConfig(ABC): @abstractmethod def transform_list_request( self, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: """Return (url, query_params) for GET /agents.""" @abstractmethod @@ -90,9 +90,9 @@ class BaseAgentsAPIConfig(ABC): def transform_get_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: """Return (url, query_params) for GET /agents/{name}.""" @abstractmethod @@ -111,8 +111,8 @@ class BaseAgentsAPIConfig(ABC): def transform_delete_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], + api_base: str | None, + litellm_params: dict[str, Any], ) -> str: """Return the URL for DELETE /agents/{name}.""" @@ -132,9 +132,9 @@ class BaseAgentsAPIConfig(ABC): def transform_list_versions_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: """Return (url, query_params) for GET /agents/{name}/versions.""" @abstractmethod @@ -153,7 +153,7 @@ class BaseAgentsAPIConfig(ABC): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> Exception: """Map HTTP error status codes to provider-specific exceptions.""" from litellm.llms.base_llm.chat.transformation import BaseLLMException diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 146912347d8..6455bb010f4 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from collections.abc import AsyncIterator -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -24,12 +24,12 @@ class BaseAnthropicMessagesConfig(ABC): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: """ OPTIONAL @@ -44,12 +44,12 @@ class BaseAnthropicMessagesConfig(ABC): @abstractmethod def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -68,11 +68,11 @@ class BaseAnthropicMessagesConfig(ABC): def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: pass @abstractmethod @@ -90,11 +90,11 @@ class BaseAnthropicMessagesConfig(ABC): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: """ OPTIONAL @@ -138,7 +138,7 @@ class BaseAnthropicMessagesConfig(ABC): raise NotImplementedError("Subclasses must implement this method") def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + self, error_message: str, status_code: int, headers: dict | httpx.Headers ) -> "BaseLLMException": from litellm.llms.base_llm.chat.transformation import BaseLLMException diff --git a/litellm/llms/base_llm/audio_transcription/transformation.py b/litellm/llms/base_llm/audio_transcription/transformation.py index dc862b3dd92..af06fd0b33d 100644 --- a/litellm/llms/base_llm/audio_transcription/transformation.py +++ b/litellm/llms/base_llm/audio_transcription/transformation.py @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -30,24 +30,24 @@ class AudioTranscriptionRequestData: content_type: Optional content type override """ - data: Union[dict, bytes] - files: Optional[dict] = None - content_type: Optional[str] = None + data: dict | bytes + files: dict | None = None + content_type: str | None = None class BaseAudioTranscriptionConfig(BaseConfig, ABC): @abstractmethod - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: pass def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -81,7 +81,7 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -97,12 +97,12 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "AudioTranscriptionConfig does not need a response transformation for audio transcription models" @@ -112,7 +112,7 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC): self, model: str, optional_params: dict, - openai_params: List[OpenAIAudioTranscriptionOptionalParams], + openai_params: list[OpenAIAudioTranscriptionOptionalParams], ) -> dict: """ Get provider specific parameters that are not OpenAI compatible diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index 905a3ebda42..8a9b4935783 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -1,6 +1,6 @@ import json from abc import abstractmethod -from typing import TYPE_CHECKING, List, Optional, Union, cast +from typing import TYPE_CHECKING, cast import litellm @@ -35,7 +35,7 @@ def convert_model_response_to_streaming( ValueError: If the conversion fails """ try: - streaming_choices: List[StreamingChoices] = [] + streaming_choices: list[StreamingChoices] = [] for choice in model_response.choices: streaming_choices.append( StreamingChoices( @@ -64,11 +64,11 @@ def convert_model_response_to_streaming( class BaseModelResponseIterator: - def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): + def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): self.streaming_response = streaming_response self.response_iterator = self.streaming_response self.json_mode = json_mode - self.http_response: Optional["httpx.Response"] = None + self.http_response: httpx.Response | None = None async def aclose(self) -> None: """Close the upstream HTTP response so the provider connection is @@ -81,7 +81,7 @@ class BaseModelResponseIterator: if self.http_response is not None: await self.http_response.aclose() - def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]: + def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: return GenericStreamingChunk( text="", is_finished=False, @@ -96,8 +96,8 @@ class BaseModelResponseIterator: return self @staticmethod - def _string_to_dict_parser(str_line: str) -> Optional[dict]: - stripped_json_chunk: Optional[dict] = None + def _string_to_dict_parser(str_line: str) -> dict | None: + stripped_json_chunk: dict | None = None stripped_chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(str_line) try: if stripped_chunk is not None: @@ -108,7 +108,7 @@ class BaseModelResponseIterator: stripped_json_chunk = None return stripped_json_chunk - def _handle_string_chunk(self, str_line: str) -> Union[GenericStreamingChunk, ModelResponseStream]: + def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream: # chunk is a str at this point stripped_json_chunk = BaseModelResponseIterator._string_to_dict_parser(str_line=str_line) if "[DONE]" in str_line: @@ -202,7 +202,7 @@ class BaseModelResponseIterator: class MockResponseIterator: # for returning ai21 streaming responses - def __init__(self, model_response: ModelResponse, json_mode: Optional[bool] = False): + def __init__(self, model_response: ModelResponse, json_mode: bool | None = False): self.model_response = model_response self.json_mode = json_mode self.is_done = False @@ -232,7 +232,7 @@ class MockResponseIterator: # for returning ai21 streaming responses class FakeStreamResponseIterator: - def __init__(self, model_response, json_mode: Optional[bool] = False): + def __init__(self, model_response, json_mode: bool | None = False): self.model_response = model_response self.json_mode = json_mode self.is_done = False diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index 8eded37595b..789990e6c78 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -5,7 +5,7 @@ Utility functions for base LLM classes. import copy import json from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional, Type, Union +from typing import Any from openai.lib import _parsing, _pydantic from pydantic import BaseModel @@ -20,19 +20,19 @@ class BaseTokenCounter(ABC): async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: pass @abstractmethod def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: """ Returns True if we should the this API for token counting for the selected `custom_llm_provider` @@ -44,14 +44,14 @@ class BaseLLMModelInfo(ABC): def get_provider_info( self, model: str, - ) -> Optional[ProviderSpecificModelInfo]: + ) -> ProviderSpecificModelInfo | None: """ Default values all models of this provider support. """ return None @abstractmethod - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: """ Returns a list of models supported by this provider. """ @@ -59,14 +59,14 @@ class BaseLLMModelInfo(ABC): @staticmethod @abstractmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: pass @staticmethod @abstractmethod def get_api_base( - api_base: Optional[str] = None, - ) -> Optional[str]: + api_base: str | None = None, + ) -> str | None: pass @abstractmethod @@ -74,26 +74,25 @@ class BaseLLMModelInfo(ABC): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: pass @staticmethod @abstractmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: """ Returns the base model name from the given model name. Some providers like bedrock - can receive model=`invoke/anthropic.claude-3-opus-20240229-v1:0` or `converse/anthropic.claude-3-opus-20240229-v1:0` This function will return `anthropic.claude-3-opus-20240229-v1:0` """ - pass - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create a token counter for this provider. @@ -105,14 +104,14 @@ class BaseLLMModelInfo(ABC): def _convert_tool_response_to_message( - tool_calls: List[ChatCompletionToolCallChunk], -) -> Optional[Message]: + tool_calls: list[ChatCompletionToolCallChunk], +) -> Message | None: """ In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format """ ## HANDLE JSON MODE - anthropic returns single function call - json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments") + json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments") try: if json_mode_content_str is not None: args = json.loads(json_mode_content_str) @@ -130,7 +129,7 @@ def _convert_tool_response_to_message( return None -def _dict_to_response_format_helper(response_format: dict, ref_template: Optional[str] = None) -> dict: +def _dict_to_response_format_helper(response_format: dict, ref_template: str | None = None) -> dict: if ref_template is not None and response_format.get("type") == "json_schema": # Deep copy to avoid modifying original modified_format = copy.deepcopy(response_format) @@ -170,9 +169,9 @@ def _dict_to_response_format_helper(response_format: dict, ref_template: Optiona def type_to_response_format_param( - response_format: Optional[Union[Type[BaseModel], dict]], - ref_template: Optional[str] = None, -) -> Optional[dict]: + response_format: type[BaseModel] | dict | None, + ref_template: str | None = None, +) -> dict | None: """ Re-implementation of openai's 'type_to_response_format_param' function @@ -206,12 +205,12 @@ def type_to_response_format_param( def map_developer_role_to_system_role( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: """ Translate `developer` role to `system` role for non-OpenAI providers. """ - new_messages: List[AllMessageValues] = [] + new_messages: list[AllMessageValues] = [] for m in messages: if m["role"] == "developer": verbose_logger.debug( diff --git a/litellm/llms/base_llm/batches/transformation.py b/litellm/llms/base_llm/batches/transformation.py index aedaf0687cb..34c622d4cf6 100644 --- a/litellm/llms/base_llm/batches/transformation.py +++ b/litellm/llms/base_llm/batches/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx from httpx import Headers @@ -38,7 +38,6 @@ class BaseBatchesConfig(ABC): @abstractmethod def custom_llm_provider(self) -> LlmProviders: """Return the LLM provider type for this configuration.""" - pass @classmethod def get_config(cls): @@ -65,11 +64,11 @@ class BaseBatchesConfig(ABC): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate and prepare environment-specific headers and parameters. @@ -86,16 +85,15 @@ class BaseBatchesConfig(ABC): Returns: Updated headers dictionary """ - pass @abstractmethod def get_complete_batch_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, data: CreateBatchRequest, ) -> str: """ @@ -112,7 +110,6 @@ class BaseBatchesConfig(ABC): Returns: Complete URL for the batch request """ - pass @abstractmethod def transform_create_batch_request( @@ -121,7 +118,7 @@ class BaseBatchesConfig(ABC): create_batch_data: CreateBatchRequest, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, Dict[str, Any]]: + ) -> bytes | str | dict[str, Any]: """ Transform the batch creation request to provider-specific format. @@ -134,12 +131,11 @@ class BaseBatchesConfig(ABC): Returns: Transformed request data """ - pass @abstractmethod def transform_create_batch_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -156,7 +152,6 @@ class BaseBatchesConfig(ABC): Returns: LiteLLM batch object """ - pass @abstractmethod def transform_retrieve_batch_request( @@ -164,7 +159,7 @@ class BaseBatchesConfig(ABC): batch_id: str, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, Dict[str, Any]]: + ) -> bytes | str | dict[str, Any]: """ Transform the batch retrieval request to provider-specific format. @@ -176,12 +171,11 @@ class BaseBatchesConfig(ABC): Returns: Transformed request data """ - pass @abstractmethod def transform_retrieve_batch_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -198,12 +192,9 @@ class BaseBatchesConfig(ABC): Returns: LiteLLM batch object """ - pass @abstractmethod - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, Headers] - ) -> "BaseLLMException": + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> "BaseLLMException": """ Get the appropriate error class for this provider. @@ -215,4 +206,3 @@ class BaseBatchesConfig(ABC): Returns: Provider-specific exception class """ - pass diff --git a/litellm/llms/base_llm/bridges/completion_transformation.py b/litellm/llms/base_llm/bridges/completion_transformation.py index 7e6899b96e4..2d5879dc8e3 100644 --- a/litellm/llms/base_llm/bridges/completion_transformation.py +++ b/litellm/llms/base_llm/bridges/completion_transformation.py @@ -4,7 +4,7 @@ Bridge for transforming API requests to another API requests from abc import ABC, abstractmethod from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union if TYPE_CHECKING: from pydantic import BaseModel @@ -19,14 +19,13 @@ class CompletionTransformationBridge(ABC): def transform_request( self, model: str, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, headers: dict, litellm_logging_obj: "LiteLLMLoggingObj", ) -> dict: """Transform /chat/completions api request to another request""" - pass @abstractmethod def transform_response( @@ -36,21 +35,20 @@ class CompletionTransformationBridge(ABC): model_response: "ModelResponse", logging_obj: "LiteLLMLoggingObj", request_data: dict, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> "ModelResponse": """Transform another response to /chat/completions api response""" - pass @abstractmethod def get_model_response_iterator( self, streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse"], sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> "BaseModelResponseIterator": pass diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 01e133cd5a7..b0877ae6f04 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -8,10 +8,6 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - List, - Optional, - Tuple, - Type, Union, cast, ) @@ -51,10 +47,10 @@ class BaseLLMException(Exception): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - body: Optional[dict] = None, + headers: dict | httpx.Headers | None = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + body: dict | None = None, ): self.status_code = status_code self.message: str = message @@ -95,9 +91,7 @@ class BaseConfig(ABC): and not callable(v) # Filter out any callable objects including mocks } - def get_json_schema_from_pydantic_object( - self, response_format: Optional[Union[Type[BaseModel], dict]] - ) -> Optional[dict]: + def get_json_schema_from_pydantic_object(self, response_format: type[BaseModel] | dict | None) -> dict | None: return type_to_response_format_param(response_format=response_format) def is_thinking_enabled(self, non_default_params: dict) -> bool: @@ -129,16 +123,16 @@ class BaseConfig(ABC): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Returns True if the model/provider should fake stream """ return False - def _add_tools_to_optional_params(self, optional_params: dict, tools: List) -> dict: + def _add_tools_to_optional_params(self, optional_params: dict, tools: list) -> dict: """ Helper util to add tools to optional_params. """ @@ -153,8 +147,8 @@ class BaseConfig(ABC): def translate_developer_role_to_system_role( self, - messages: List[AllMessageValues], - ) -> List[AllMessageValues]: + messages: list[AllMessageValues], + ) -> list[AllMessageValues]: """ Translate `developer` role to `system` role for non-OpenAI providers. @@ -210,7 +204,7 @@ class BaseConfig(ABC): This is used to translate response_format to a tool call, for models/APIs that don't support response_format directly. """ - json_schema: Optional[dict] = None + json_schema: dict | None = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: @@ -252,11 +246,11 @@ class BaseConfig(ABC): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: pass @@ -266,11 +260,11 @@ class BaseConfig(ABC): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: """ Some providers like Bedrock require signing the request. The sign request funtion needs access to `request_data` and `complete_url` Args: @@ -287,12 +281,12 @@ class BaseConfig(ABC): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -309,7 +303,7 @@ class BaseConfig(ABC): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -319,7 +313,7 @@ class BaseConfig(ABC): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -345,12 +339,12 @@ class BaseConfig(ABC): model_response: "ModelResponse", logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> "ModelResponse": pass @@ -366,16 +360,14 @@ class BaseConfig(ABC): return parsed_response @abstractmethod - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: pass def get_model_response_iterator( self, streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse"], sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: pass @@ -388,9 +380,9 @@ class BaseConfig(ABC): headers: dict, data: dict, messages: list, - client: Optional[AsyncHTTPHandler] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": raise NotImplementedError @@ -403,14 +395,14 @@ class BaseConfig(ABC): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": raise NotImplementedError @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return None @property @@ -433,12 +425,12 @@ class BaseConfig(ABC): def apply_assembled_streaming_response_metadata( self, response: "ModelResponse", - chunks: List[Any], + chunks: list[Any], ) -> None: """Hook for providers to merge chunk metadata into assembled streaming responses.""" - return None + return - def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> Optional[dict]: + def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> dict | None: """ Calculate any additional costs beyond standard token costs. diff --git a/litellm/llms/base_llm/completion/transformation.py b/litellm/llms/base_llm/completion/transformation.py index 2309634f180..c38199b0966 100644 --- a/litellm/llms/base_llm/completion/transformation.py +++ b/litellm/llms/base_llm/completion/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -20,7 +20,7 @@ class BaseTextCompletionConfig(BaseConfig, ABC): def transform_text_completion_request( self, model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], optional_params: dict, headers: dict, ) -> dict: @@ -28,12 +28,12 @@ class BaseTextCompletionConfig(BaseConfig, ABC): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -47,7 +47,7 @@ class BaseTextCompletionConfig(BaseConfig, ABC): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -63,12 +63,12 @@ class BaseTextCompletionConfig(BaseConfig, ABC): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "AudioTranscriptionConfig does not need a response transformation for audio transcription models" diff --git a/litellm/llms/base_llm/embedding/transformation.py b/litellm/llms/base_llm/embedding/transformation.py index 07ffbb99626..0330c0118bd 100644 --- a/litellm/llms/base_llm/embedding/transformation.py +++ b/litellm/llms/base_llm/embedding/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -33,7 +33,7 @@ class BaseEmbeddingConfig(BaseConfig, ABC): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -42,12 +42,12 @@ class BaseEmbeddingConfig(BaseConfig, ABC): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -61,7 +61,7 @@ class BaseEmbeddingConfig(BaseConfig, ABC): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -75,11 +75,11 @@ class BaseEmbeddingConfig(BaseConfig, ABC): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError("EmbeddingConfig does not need a response transformation for chat models") diff --git a/litellm/llms/base_llm/evals/transformation.py b/litellm/llms/base_llm/evals/transformation.py index da8d7e12acb..49124e271a6 100644 --- a/litellm/llms/base_llm/evals/transformation.py +++ b/litellm/llms/base_llm/evals/transformation.py @@ -3,7 +3,7 @@ Base configuration class for Evals API """ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -46,7 +46,7 @@ class BaseEvalsAPIConfig(ABC): pass @abstractmethod - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate and update headers with provider-specific requirements @@ -62,9 +62,9 @@ class BaseEvalsAPIConfig(ABC): @abstractmethod def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, endpoint: str, - eval_id: Optional[str] = None, + eval_id: str | None = None, ) -> str: """ Get the complete URL for the API request @@ -87,7 +87,7 @@ class BaseEvalsAPIConfig(ABC): create_request: CreateEvalRequest, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Transform create eval request to provider-specific format @@ -99,7 +99,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Provider-specific request body """ - pass @abstractmethod def transform_create_eval_response( @@ -117,7 +116,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Eval object """ - pass @abstractmethod def transform_list_evals_request( @@ -125,7 +123,7 @@ class BaseEvalsAPIConfig(ABC): list_params: ListEvalsParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform list evals request parameters @@ -137,7 +135,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, query_params) """ - pass @abstractmethod def transform_list_evals_response( @@ -155,7 +152,6 @@ class BaseEvalsAPIConfig(ABC): Returns: ListEvalsResponse object """ - pass @abstractmethod def transform_get_eval_request( @@ -164,7 +160,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform get eval request @@ -177,7 +173,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers) """ - pass @abstractmethod def transform_get_eval_response( @@ -195,7 +190,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Eval object """ - pass @abstractmethod def transform_update_eval_request( @@ -205,7 +199,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict, Dict]: + ) -> tuple[str, dict, dict]: """ Transform update eval request @@ -219,7 +213,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers, body) """ - pass @abstractmethod def transform_update_eval_response( @@ -237,7 +230,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Eval object """ - pass @abstractmethod def transform_delete_eval_request( @@ -246,7 +238,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform delete eval request @@ -259,7 +251,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers) """ - pass @abstractmethod def transform_delete_eval_response( @@ -277,7 +268,6 @@ class BaseEvalsAPIConfig(ABC): Returns: DeleteEvalResponse object """ - pass @abstractmethod def transform_cancel_eval_request( @@ -286,7 +276,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict, Dict]: + ) -> tuple[str, dict, dict]: """ Transform cancel eval request @@ -299,7 +289,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers, body) """ - pass @abstractmethod def transform_cancel_eval_response( @@ -317,7 +306,6 @@ class BaseEvalsAPIConfig(ABC): Returns: CancelEvalResponse object """ - pass # Run API Transformations @abstractmethod @@ -327,7 +315,7 @@ class BaseEvalsAPIConfig(ABC): create_request: CreateRunRequest, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform create run request to provider-specific format @@ -340,7 +328,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, request_body) """ - pass @abstractmethod def transform_create_run_response( @@ -358,7 +345,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Run object """ - pass @abstractmethod def transform_list_runs_request( @@ -367,7 +353,7 @@ class BaseEvalsAPIConfig(ABC): list_params: ListRunsParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform list runs request parameters @@ -380,7 +366,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, query_params) """ - pass @abstractmethod def transform_list_runs_response( @@ -398,7 +383,6 @@ class BaseEvalsAPIConfig(ABC): Returns: ListRunsResponse object """ - pass @abstractmethod def transform_get_run_request( @@ -408,7 +392,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform get run request @@ -422,7 +406,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers) """ - pass @abstractmethod def transform_get_run_response( @@ -440,7 +423,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Run object """ - pass @abstractmethod def transform_cancel_run_request( @@ -450,7 +432,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict, Dict]: + ) -> tuple[str, dict, dict]: """ Transform cancel run request @@ -464,7 +446,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers, body) """ - pass @abstractmethod def transform_cancel_run_response( @@ -482,7 +463,6 @@ class BaseEvalsAPIConfig(ABC): Returns: CancelRunResponse object """ - pass @abstractmethod def transform_delete_run_request( @@ -492,7 +472,7 @@ class BaseEvalsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict, Dict]: + ) -> tuple[str, dict, dict]: """ Transform delete run request @@ -506,7 +486,6 @@ class BaseEvalsAPIConfig(ABC): Returns: Tuple of (url, headers, body) """ - pass @abstractmethod def transform_delete_run_response( @@ -524,7 +503,6 @@ class BaseEvalsAPIConfig(ABC): Returns: RunDeleteResponse object """ - pass def get_error_class( self, diff --git a/litellm/llms/base_llm/files/azure_blob_storage_backend.py b/litellm/llms/base_llm/files/azure_blob_storage_backend.py index 07dd339cac3..7c76003de3a 100644 --- a/litellm/llms/base_llm/files/azure_blob_storage_backend.py +++ b/litellm/llms/base_llm/files/azure_blob_storage_backend.py @@ -7,14 +7,13 @@ to reuse all authentication and Azure Storage operations. """ import time -from typing import Optional from urllib.parse import quote from litellm._logging import verbose_logger from litellm._uuid import uuid +from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger from .storage_backend import BaseFileStorageBackend -from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): @@ -69,14 +68,12 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): Override to do nothing - we're not using this as a logger. """ # Do nothing - this class is used for file storage, not logging - pass async def async_log_failure_event(self, *args, **kwargs): """ Override to do nothing - we're not using this as a logger. """ # Do nothing - this class is used for file storage, not logging - pass def _generate_file_name(self, original_filename: str, file_naming_strategy: str) -> str: """Generate file name based on naming strategy.""" @@ -99,7 +96,7 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): file_content: bytes, filename: str, content_type: str, - path_prefix: Optional[str] = None, + path_prefix: str | None = None, file_naming_strategy: str = "uuid", ) -> str: """ @@ -136,7 +133,7 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): return storage_url except Exception as e: - verbose_logger.exception(f"Error uploading file to Azure Blob Storage: {str(e)}") + verbose_logger.exception(f"Error uploading file to Azure Blob Storage: {e!s}") raise async def _upload_file_with_account_key(self, file_content: bytes, full_path: str) -> str: @@ -250,7 +247,7 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): return await self._download_file_with_azure_ad(file_path) except Exception as e: - verbose_logger.exception(f"Error downloading file from Azure Blob Storage: {str(e)}") + verbose_logger.exception(f"Error downloading file from Azure Blob Storage: {e!s}") raise async def _download_file_with_account_key(self, file_path: str) -> bytes: @@ -272,11 +269,11 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): # Reuse the logger's token management await self.set_valid_azure_ad_token() + from litellm.constants import AZURE_STORAGE_MSFT_VERSION from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) - from litellm.constants import AZURE_STORAGE_MSFT_VERSION async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) diff --git a/litellm/llms/base_llm/files/storage_backend.py b/litellm/llms/base_llm/files/storage_backend.py index 31e68a7002a..69b6a6f4a0f 100644 --- a/litellm/llms/base_llm/files/storage_backend.py +++ b/litellm/llms/base_llm/files/storage_backend.py @@ -6,7 +6,6 @@ This module defines the abstract base class that all file storage backends """ from abc import ABC, abstractmethod -from typing import Optional class BaseFileStorageBackend(ABC): @@ -23,7 +22,7 @@ class BaseFileStorageBackend(ABC): file_content: bytes, filename: str, content_type: str, - path_prefix: Optional[str] = None, + path_prefix: str | None = None, file_naming_strategy: str = "uuid", ) -> str: """ @@ -42,7 +41,6 @@ class BaseFileStorageBackend(ABC): Raises: Exception: If upload fails """ - pass @abstractmethod async def download_file(self, storage_url: str) -> bytes: @@ -58,7 +56,6 @@ class BaseFileStorageBackend(ABC): Raises: Exception: If download fails """ - pass async def delete_file(self, storage_url: str) -> None: """ @@ -75,4 +72,3 @@ class BaseFileStorageBackend(ABC): """ # Default implementation: no-op # Backends can override if they support deletion - pass diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index e5ae92d4b9c..174be93448b 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from collections.abc import Iterator -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union import httpx from openai.types.file_deleted import FileDeleted @@ -66,13 +66,13 @@ class BaseFilesConfig(BaseConfig): return "POST" @abstractmethod - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: pass def get_complete_file_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, @@ -102,12 +102,11 @@ class BaseFilesConfig(BaseConfig): - str/bytes: For traditional file uploads - TwoStepFileUploadConfig: For two-step upload process (e.g., Manus, GCS) """ - pass @abstractmethod def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -122,7 +121,6 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> tuple[str, dict]: """Transform file retrieve request into provider-specific format.""" - pass @abstractmethod def transform_retrieve_file_response( @@ -132,7 +130,6 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> OpenAIFileObject: """Transform file retrieve response into OpenAI format.""" - pass @abstractmethod def transform_delete_file_request( @@ -142,7 +139,6 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> tuple[str, dict]: """Transform file delete request into provider-specific format.""" - pass @abstractmethod def transform_delete_file_response( @@ -152,17 +148,15 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> "FileDeleted": """Transform file delete response into OpenAI format.""" - pass @abstractmethod def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: """Transform file list request into provider-specific format.""" - pass @abstractmethod def transform_list_files_response( @@ -170,9 +164,8 @@ class BaseFilesConfig(BaseConfig): raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: """Transform file list response into OpenAI format.""" - pass @abstractmethod def transform_file_content_request( @@ -182,7 +175,6 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> tuple[str, dict]: """Transform file content request into provider-specific format.""" - pass @abstractmethod def transform_file_content_response( @@ -192,12 +184,11 @@ class BaseFilesConfig(BaseConfig): litellm_params: dict, ) -> "HttpxBinaryResponseContent": """Transform file content response into OpenAI format.""" - pass def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -213,12 +204,12 @@ class BaseFilesConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "AudioTranscriptionConfig does not need a response transformation for audio transcription models" @@ -231,7 +222,7 @@ class BaseFileEndpoints(ABC): self, create_file_request: CreateFileRequest, llm_router: Router, - target_model_names_list: List[str], + target_model_names_list: list[str], litellm_parent_otel_span: Span, user_api_key_dict: UserAPIKeyAuth, ) -> OpenAIFileObject: @@ -241,27 +232,27 @@ class BaseFileEndpoints(ABC): async def afile_retrieve( self, file_id: str, - litellm_parent_otel_span: Optional[Span], - llm_router: Optional[Router] = None, + litellm_parent_otel_span: Span | None, + llm_router: Router | None = None, ) -> OpenAIFileObject: pass @abstractmethod async def afile_list( self, - purpose: Optional[OpenAIFilesPurpose], - litellm_parent_otel_span: Optional[Span], - **data: Dict, - ) -> List[OpenAIFileObject]: + purpose: OpenAIFilesPurpose | None, + litellm_parent_otel_span: Span | None, + **data: dict, + ) -> list[OpenAIFileObject]: pass @abstractmethod async def afile_delete( self, file_id: str, - litellm_parent_otel_span: Optional[Span], + litellm_parent_otel_span: Span | None, llm_router: Router, - **data: Dict, + **data: dict, ) -> OpenAIFileObject: pass @@ -269,8 +260,8 @@ class BaseFileEndpoints(ABC): async def afile_content( self, file_id: str, - litellm_parent_otel_span: Optional[Span], + litellm_parent_otel_span: Span | None, llm_router: Router, - **data: Dict, + **data: dict, ) -> "HttpxBinaryResponseContent": pass diff --git a/litellm/llms/base_llm/google_genai/transformation.py b/litellm/llms/base_llm/google_genai/transformation.py index 965c174df6e..bd0d29d5ea3 100644 --- a/litellm/llms/base_llm/google_genai/transformation.py +++ b/litellm/llms/base_llm/google_genai/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -48,7 +48,7 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): } @abstractmethod - def get_supported_generate_content_optional_params(self, model: str) -> List[str]: + def get_supported_generate_content_optional_params(self, model: str) -> list[str]: """ Get the list of supported Google GenAI parameters for the model. @@ -78,7 +78,7 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): self, generate_content_config_dict: GenerateContentConfigDict, model: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map Google GenAI parameters to provider-specific format. @@ -94,10 +94,10 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): @abstractmethod def validate_environment( self, - api_key: Optional[str], - headers: Optional[dict], + api_key: str | None, + headers: dict | None, model: str, - litellm_params: Optional[Union[GenericLiteLLMParams, dict]], + litellm_params: GenericLiteLLMParams | dict | None, ) -> dict: """ Validate the environment and return headers for the request. @@ -115,11 +115,11 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): def sync_get_auth_token_and_url( self, - api_base: Optional[str], + api_base: str | None, model: str, litellm_params: dict, stream: bool, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Sync version of get_auth_token_and_url. @@ -136,11 +136,11 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): async def get_auth_token_and_url( self, - api_base: Optional[str], + api_base: str | None, model: str, litellm_params: dict, stream: bool, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Get the complete URL for the request. @@ -159,9 +159,9 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): self, model: str, contents: GenerateContentContentListUnionDict, - tools: Optional[ToolConfigDict], - generate_content_config_dict: Dict, - system_instruction: Optional[Any] = None, + tools: ToolConfigDict | None, + generate_content_config_dict: dict, + system_instruction: Any | None = None, ) -> dict: """ Transform the request parameters for the generate content API. @@ -176,7 +176,6 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): Returns: Transformed request data """ - pass @abstractmethod def transform_generate_content_response( @@ -195,9 +194,8 @@ class BaseGoogleGenAIGenerateContentConfig(ABC): Returns: Transformed response data """ - pass - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]) -> Exception: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> Exception: """ Get the appropriate exception class for the error. diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index bed06832386..59a3704e53b 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Optional if TYPE_CHECKING: from litellm.integrations.custom_guardrail import ( @@ -36,7 +36,7 @@ class BaseTranslation(ABC): @staticmethod def transform_user_api_key_dict_to_metadata( user_api_key_dict: Any | None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform user_api_key_dict to a metadata dict with prefixed keys. @@ -85,7 +85,6 @@ class BaseTranslation(ABC): Note: user_api_key_dict metadata should be available in the data dict. """ - pass @abstractmethod async def process_output_response( @@ -105,11 +104,10 @@ class BaseTranslation(ABC): litellm_logging_obj: Optional logging object user_api_key_dict: User API key metadata (passed separately since response doesn't contain it) """ - pass async def process_output_streaming_response( self, - responses_so_far: List[Any], + responses_so_far: list[Any], guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, @@ -149,7 +147,7 @@ class BaseTranslation(ABC): """ return None - def get_structured_messages(self, data: dict) -> List["AllMessageValues"] | None: + def get_structured_messages(self, data: dict) -> list["AllMessageValues"] | None: """ Convert request data to OpenAI-spec structured messages. @@ -159,7 +157,7 @@ class BaseTranslation(ABC): """ return None - def extract_request_tool_names(self, data: dict) -> List[str]: + def extract_request_tool_names(self, data: dict) -> list[str]: """ Extract tool names from the request body for allowlist/policy checks. Override in tool-capable handlers; default returns []. diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index 8a06dd4ea52..1dacc056f50 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from typing import Any, List, Optional +from typing import Any from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage from litellm.types.llms.openai import AllMessageValues @@ -35,7 +35,7 @@ def _anthropic_stream_chunk_events(item: Any) -> list[dict]: return events -def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> Optional[AnthropicUsage]: +def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> AnthropicUsage | None: input_tokens = 0 output_tokens = 0 found_usage = False @@ -64,7 +64,7 @@ def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> Optiona return AnthropicUsage(input_tokens=input_tokens, output_tokens=output_tokens) -def blocked_response_usage(original_response: Optional[Any]) -> AnthropicUsage: +def blocked_response_usage(original_response: Any | None) -> AnthropicUsage: """ Token usage for a synthetic guardrail-blocked response. @@ -114,12 +114,12 @@ def effective_skip_tool_message_for_guardrail(guardrail_to_apply: Any) -> bool: def openai_messages_without_system( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: return [m for m in messages if str((m or {}).get("role") or "").lower() != "system"] def openai_messages_without_tool( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: + messages: list[AllMessageValues], +) -> list[AllMessageValues]: return [m for m in messages if str((m or {}).get("role") or "").lower() != "tool"] diff --git a/litellm/llms/base_llm/image_edit/transformation.py b/litellm/llms/base_llm/image_edit/transformation.py index 4c18702bc6c..9a25d3294e0 100644 --- a/litellm/llms/base_llm/image_edit/transformation.py +++ b/litellm/llms/base_llm/image_edit/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -58,7 +58,7 @@ class BaseImageEditConfig(ABC): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: pass @abstractmethod @@ -66,9 +66,9 @@ class BaseImageEditConfig(ABC): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: return {} @@ -76,7 +76,7 @@ class BaseImageEditConfig(ABC): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -94,12 +94,12 @@ class BaseImageEditConfig(ABC): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: pass def finalize_image_edit_request_data(self, data: dict, resolved_request_url: str) -> dict: diff --git a/litellm/llms/base_llm/image_generation/transformation.py b/litellm/llms/base_llm/image_generation/transformation.py index e80a970d806..4ce4add0432 100644 --- a/litellm/llms/base_llm/image_generation/transformation.py +++ b/litellm/llms/base_llm/image_generation/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -20,7 +20,7 @@ else: class BaseImageGenerationConfig(ABC): @abstractmethod - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: pass @abstractmethod @@ -35,12 +35,12 @@ class BaseImageGenerationConfig(ABC): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -55,17 +55,15 @@ class BaseImageGenerationConfig(ABC): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {} - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: raise BaseLLMException( status_code=status_code, message=error_message, @@ -94,8 +92,8 @@ class BaseImageGenerationConfig(ABC): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: raise NotImplementedError( "ImageVariationConfig implements 'transform_response_image_variation' for image variation models" diff --git a/litellm/llms/base_llm/image_variations/transformation.py b/litellm/llms/base_llm/image_variations/transformation.py index 23fc4dc88b9..beae828c301 100644 --- a/litellm/llms/base_llm/image_variations/transformation.py +++ b/litellm/llms/base_llm/image_variations/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx from aiohttp import ClientResponse @@ -26,17 +26,17 @@ else: class BaseImageVariationConfig(BaseConfig, ABC): @abstractmethod - def get_supported_openai_params(self, model: str) -> List[OpenAIImageVariationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageVariationOptionalParams]: pass def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -50,7 +50,7 @@ class BaseImageVariationConfig(BaseConfig, ABC): @abstractmethod def transform_request_image_variation( self, - model: Optional[str], + model: str | None, image: FileTypes, optional_params: dict, headers: dict, @@ -61,18 +61,18 @@ class BaseImageVariationConfig(BaseConfig, ABC): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {} @abstractmethod async def async_transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: ClientResponse, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -81,14 +81,14 @@ class BaseImageVariationConfig(BaseConfig, ABC): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: pass @abstractmethod def transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -97,14 +97,14 @@ class BaseImageVariationConfig(BaseConfig, ABC): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: pass def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -120,12 +120,12 @@ class BaseImageVariationConfig(BaseConfig, ABC): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "ImageVariationConfig implements 'transform_response_image_variation' for image variation models" diff --git a/litellm/llms/base_llm/interactions/transformation.py b/litellm/llms/base_llm/interactions/transformation.py index 3eba1858a23..b7bbf7bde73 100644 --- a/litellm/llms/base_llm/interactions/transformation.py +++ b/litellm/llms/base_llm/interactions/transformation.py @@ -11,7 +11,7 @@ Per OpenAPI spec (https://ai.google.dev/static/api/interactions.openapi.json): import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -57,7 +57,6 @@ class BaseInteractionsAPIConfig(ABC): @abstractmethod def custom_llm_provider(self) -> LlmProviders: """Return the LLM provider identifier.""" - pass @classmethod def get_config(cls): @@ -79,14 +78,13 @@ class BaseInteractionsAPIConfig(ABC): } @abstractmethod - def get_supported_params(self, model: str) -> List[str]: + def get_supported_params(self, model: str) -> list[str]: """ Return the list of supported parameters for the given model. """ - pass @abstractmethod - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate and prepare environment settings including headers. """ @@ -95,11 +93,11 @@ class BaseInteractionsAPIConfig(ABC): @abstractmethod def get_complete_url( self, - api_base: Optional[str], - model: Optional[str], - agent: Optional[str] = None, - litellm_params: Optional[dict] = None, - stream: Optional[bool] = None, + api_base: str | None, + model: str | None, + agent: str | None = None, + litellm_params: dict | None = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the interaction request. @@ -123,13 +121,13 @@ class BaseInteractionsAPIConfig(ABC): @abstractmethod def transform_request( self, - model: Optional[str], - agent: Optional[str], - input: Optional[InteractionInput], + model: str | None, + agent: str | None, + input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Transform the input request into the provider's expected format. @@ -148,12 +146,11 @@ class BaseInteractionsAPIConfig(ABC): Returns: The transformed request body as a dictionary """ - pass @abstractmethod def transform_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, ) -> InteractionsAPIResponse: @@ -162,12 +159,11 @@ class BaseInteractionsAPIConfig(ABC): Per OpenAPI spec, the response is an Interaction object. """ - pass @abstractmethod def transform_streaming_response( self, - model: Optional[str], + model: str | None, parsed_chunk: dict, logging_obj: LiteLLMLoggingObj, ) -> InteractionsAPIStreamingResponse: @@ -176,7 +172,6 @@ class BaseInteractionsAPIConfig(ABC): Per OpenAPI spec, streaming uses SSE with various event types. """ - pass # ========================================================= # GET INTERACTION TRANSFORMATION @@ -189,7 +184,7 @@ class BaseInteractionsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the get interaction request into URL and query params. @@ -198,7 +193,6 @@ class BaseInteractionsAPIConfig(ABC): Returns: Tuple of (URL, query_params) """ - pass @abstractmethod def transform_get_interaction_response( @@ -209,7 +203,6 @@ class BaseInteractionsAPIConfig(ABC): """ Transform the get interaction response. """ - pass # ========================================================= # DELETE INTERACTION TRANSFORMATION @@ -222,7 +215,7 @@ class BaseInteractionsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the delete interaction request into URL and body. @@ -231,7 +224,6 @@ class BaseInteractionsAPIConfig(ABC): Returns: Tuple of (URL, request_body) """ - pass @abstractmethod def transform_delete_interaction_response( @@ -243,7 +235,6 @@ class BaseInteractionsAPIConfig(ABC): """ Transform the delete interaction response. """ - pass # ========================================================= # CANCEL INTERACTION TRANSFORMATION @@ -256,14 +247,13 @@ class BaseInteractionsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the cancel interaction request into URL and body. Returns: Tuple of (URL, request_body) """ - pass @abstractmethod def transform_cancel_interaction_response( @@ -274,15 +264,12 @@ class BaseInteractionsAPIConfig(ABC): """ Transform the cancel interaction response. """ - pass # ========================================================= # ERROR HANDLING # ========================================================= - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get the appropriate exception class for an error. """ @@ -296,9 +283,9 @@ class BaseInteractionsAPIConfig(ABC): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Returns True if litellm should fake a stream for the given model. diff --git a/litellm/llms/base_llm/managed_resources/__init__.py b/litellm/llms/base_llm/managed_resources/__init__.py index a5543e631c0..4291cec835c 100644 --- a/litellm/llms/base_llm/managed_resources/__init__.py +++ b/litellm/llms/base_llm/managed_resources/__init__.py @@ -29,15 +29,15 @@ from .utils import ( __all__ = [ "BaseManagedResource", - "resolve_passthrough_managed_id_provider", - "is_base64_encoded_unified_id", - "extract_target_model_names_from_unified_id", - "extract_resource_type_from_unified_id", - "extract_unified_uuid_from_unified_id", + "decode_unified_id", + "encode_unified_id", "extract_model_id_from_unified_id", "extract_provider_resource_id_from_unified_id", + "extract_resource_type_from_unified_id", + "extract_target_model_names_from_unified_id", + "extract_unified_uuid_from_unified_id", "generate_unified_id_string", - "encode_unified_id", - "decode_unified_id", + "is_base64_encoded_unified_id", "parse_unified_id", + "resolve_passthrough_managed_id_provider", ] diff --git a/litellm/llms/base_llm/managed_resources/base_managed_resource.py b/litellm/llms/base_llm/managed_resources/base_managed_resource.py index 146a6aa6ae0..fd9eaf9801b 100644 --- a/litellm/llms/base_llm/managed_resources/base_managed_resource.py +++ b/litellm/llms/base_llm/managed_resources/base_managed_resource.py @@ -8,10 +8,7 @@ from abc import ABC, abstractmethod from typing import ( TYPE_CHECKING, Any, - Dict, Generic, - List, - Optional, TypeVar, Union, cast, @@ -84,7 +81,6 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): Return the resource type identifier (e.g., 'file', 'vector_store', 'vector_store_file'). Used for logging and unified ID generation. """ - pass @property @abstractmethod @@ -93,13 +89,12 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): Return the database table name for this resource type. Example: 'litellm_managedfiletable', 'litellm_managedvectorstoretable' """ - pass @abstractmethod def get_unified_resource_id_format( self, resource_object: ResourceObjectType, - target_model_names_list: List[str], + target_model_names_list: list[str], ) -> str: """ Generate the format string for the unified resource ID. @@ -115,14 +110,13 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): Returns: Format string to be base64 encoded """ - pass @abstractmethod async def create_resource_for_model( self, llm_router: Router, model: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], litellm_parent_otel_span: Span, ) -> ResourceObjectType: """ @@ -137,7 +131,6 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): Returns: Resource object from the provider """ - pass # ============================================================================ # COMMON STORAGE OPERATIONS @@ -146,11 +139,11 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def store_unified_resource_id( self, unified_resource_id: str, - resource_object: Optional[ResourceObjectType], - litellm_parent_otel_span: Optional[Span], - model_mappings: Dict[str, str], + resource_object: ResourceObjectType | None, + litellm_parent_otel_span: Span | None, + model_mappings: dict[str, str], user_api_key_dict: UserAPIKeyAuth, - additional_db_fields: Optional[Dict[str, Any]] = None, + additional_db_fields: dict[str, Any] | None = None, ) -> None: """ Store unified resource ID with model mappings in cache and database. @@ -228,8 +221,8 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def get_unified_resource_id( self, unified_resource_id: str, - litellm_parent_otel_span: Optional[Span] = None, - ) -> Optional[Dict[str, Any]]: + litellm_parent_otel_span: Span | None = None, + ) -> dict[str, Any] | None: """ Retrieve unified resource by ID from cache or database. @@ -242,7 +235,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): """ # Check cache first result = cast( - Optional[dict], + dict | None, await self.internal_usage_cache.async_get_cache( key=unified_resource_id, litellm_parent_otel_span=litellm_parent_otel_span, @@ -264,8 +257,8 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def delete_unified_resource_id( self, unified_resource_id: str, - litellm_parent_otel_span: Optional[Span] = None, - ) -> Optional[ResourceObjectType]: + litellm_parent_otel_span: Span | None = None, + ) -> ResourceObjectType | None: """ Delete unified resource from cache and database. @@ -299,7 +292,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): self, unified_resource_id: str, user_api_key_dict: UserAPIKeyAuth, - litellm_parent_otel_span: Optional[Span] = None, + litellm_parent_otel_span: Span | None = None, ) -> bool: """ Check if user has access to the unified resource ID. @@ -333,9 +326,9 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def get_model_resource_id_mapping( self, - resource_ids: List[str], + resource_ids: list[str], litellm_parent_otel_span: Span, - ) -> Dict[str, Dict[str, str]]: + ) -> dict[str, dict[str, str]]: """ Get model-specific resource IDs for a list of unified resource IDs. @@ -354,7 +347,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): } } """ - resource_id_mapping: Dict[str, Dict[str, str]] = {} + resource_id_mapping: dict[str, dict[str, str]] = {} for resource_id in resource_ids: # Get unified resource from cache/db @@ -378,10 +371,10 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def create_resource_for_each_model( self, llm_router: Router, - request_data: Dict[str, Any], - target_model_names_list: List[str], + request_data: dict[str, Any], + target_model_names_list: list[str], litellm_parent_otel_span: Span, - ) -> List[ResourceObjectType]: + ) -> list[ResourceObjectType]: """ Create a resource for each model in the target list. @@ -410,8 +403,8 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): def generate_unified_resource_id( self, - resource_objects: List[ResourceObjectType], - target_model_names_list: List[str], + resource_objects: list[ResourceObjectType], + target_model_names_list: list[str], ) -> str: """ Generate a unified resource ID from multiple resource objects. @@ -436,8 +429,8 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): def extract_model_mappings_from_responses( self, - resource_objects: List[ResourceObjectType], - ) -> Dict[str, str]: + resource_objects: list[ResourceObjectType], + ) -> dict[str, str]: """ Extract model mappings from resource objects. @@ -447,7 +440,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): Returns: Dictionary mapping model_id -> provider_resource_id """ - model_mappings: Dict[str, str] = {} + model_mappings: dict[str, str] = {} for resource_object in resource_objects: # Get hidden params if available @@ -466,11 +459,11 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - request_kwargs: Optional[Dict] = None, - parent_otel_span: Optional[Span] = None, + healthy_deployments: list, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, resource_id_key: str = "resource_id", - ) -> List[Dict]: + ) -> list[dict]: """ Filter deployments based on model mappings for a resource. @@ -490,9 +483,9 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): if request_kwargs is None: return healthy_deployments - resource_id = cast(Optional[str], request_kwargs.get(resource_id_key)) + resource_id = cast(str | None, request_kwargs.get(resource_id_key)) model_resource_id_mapping = cast( - Optional[Dict[str, Dict[str, str]]], + dict[str, dict[str, str]] | None, request_kwargs.get("model_resource_id_mapping"), ) @@ -526,10 +519,10 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): async def list_user_resources( self, user_api_key_dict: UserAPIKeyAuth, - limit: Optional[int] = None, - after: Optional[str] = None, - additional_filters: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + limit: int | None = None, + after: str | None = None, + additional_filters: dict[str, Any] | None = None, + ) -> dict[str, Any]: """ List resources created by a user. @@ -546,7 +539,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): if owner_filter is None: return build_list_page([]) - where_clause: Dict[str, Any] = {**owner_filter} + where_clause: dict[str, Any] = {**owner_filter} if after: where_clause["id"] = {"gt": after} @@ -564,7 +557,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): order={"created_at": "desc"}, ) - resource_objects: List[Any] = [] + resource_objects: list[Any] = [] for resource in resources: try: # Stop once we have enough diff --git a/litellm/llms/base_llm/managed_resources/isolation.py b/litellm/llms/base_llm/managed_resources/isolation.py index fd1e24f3e1d..c4de1624f61 100644 --- a/litellm/llms/base_llm/managed_resources/isolation.py +++ b/litellm/llms/base_llm/managed_resources/isolation.py @@ -9,15 +9,17 @@ identifying ids are denied so an empty user_id can never select an unscoped query. """ -from typing import Any, Dict, List, Optional +from typing import Any from litellm.proxy._types import ( UserAPIKeyAuth, +) +from litellm.proxy._types import ( user_api_key_has_admin_view as _user_has_admin_view, ) -def build_list_page(items: List[Any], has_more: bool = False) -> Dict[str, Any]: +def build_list_page(items: list[Any], has_more: bool = False) -> dict[str, Any]: """Build the OpenAI-style paginated list response shape used by managed file/batch/vector-store listings. ``first_id`` and ``last_id`` are sourced from each item's ``.id`` attribute.""" @@ -32,7 +34,7 @@ def build_list_page(items: List[Any], has_more: bool = False) -> Dict[str, Any]: def build_owner_filter( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: """Return a Prisma `where` fragment that scopes a managed-resource listing to records the caller is allowed to see. @@ -71,8 +73,8 @@ def build_owner_filter( def can_access_resource( user_api_key_dict: UserAPIKeyAuth, - created_by: Optional[str], - resource_team_id: Optional[str], + created_by: str | None, + resource_team_id: str | None, ) -> bool: """Return True iff the caller may read/modify a managed resource. diff --git a/litellm/llms/base_llm/managed_resources/utils.py b/litellm/llms/base_llm/managed_resources/utils.py index a93f62764f9..3296d019880 100644 --- a/litellm/llms/base_llm/managed_resources/utils.py +++ b/litellm/llms/base_llm/managed_resources/utils.py @@ -7,14 +7,14 @@ different managed resource types (files, vector stores, etc.). import base64 import re -from typing import Any, List, Literal, Optional, Union +from typing import Any, Literal PASSTHROUGH_MANAGED_ID_AZURE_PROVIDERS = ("azure", "azure_ai") def resolve_passthrough_managed_id_provider( custom_llm_provider: Any, -) -> Optional[str]: +) -> str | None: """Map a pass-through ``custom_llm_provider`` to the provider scope that namespaces passthrough managed object IDs, or ``None`` when the route is not an OpenAI/Azure pass-through and managed IDs must not apply. @@ -42,7 +42,7 @@ def resolve_passthrough_managed_id_provider( def is_base64_encoded_unified_id( resource_id: str, prefix: str = "litellm_proxy:", -) -> Union[str, Literal[False]]: +) -> str | Literal[False]: """ Check if a resource ID is a base64 encoded unified ID. @@ -73,7 +73,7 @@ def is_base64_encoded_unified_id( def extract_target_model_names_from_unified_id( unified_id: str, -) -> List[str]: +) -> list[str]: """ Extract target model names from a unified resource ID. @@ -110,7 +110,7 @@ def extract_target_model_names_from_unified_id( def extract_resource_type_from_unified_id( unified_id: str, -) -> Optional[str]: +) -> str | None: """ Extract resource type from a unified resource ID. @@ -146,7 +146,7 @@ def extract_resource_type_from_unified_id( def extract_unified_uuid_from_unified_id( unified_id: str, -) -> Optional[str]: +) -> str | None: """ Extract the UUID from a unified resource ID. @@ -182,7 +182,7 @@ def extract_unified_uuid_from_unified_id( def extract_model_id_from_unified_id( unified_id: str, -) -> Optional[str]: +) -> str | None: """ Extract model ID from a unified resource ID. @@ -224,7 +224,7 @@ def extract_model_id_from_unified_id( def extract_provider_resource_id_from_unified_id( unified_id: str, -) -> Optional[str]: +) -> str | None: """ Extract provider resource ID from a unified resource ID. @@ -268,10 +268,10 @@ def extract_provider_resource_id_from_unified_id( def generate_unified_id_string( resource_type: str, unified_uuid: str, - target_model_names: List[str], + target_model_names: list[str], provider_resource_id: str, model_id: str, - additional_fields: Optional[dict] = None, + additional_fields: dict | None = None, ) -> str: """ Generate a unified ID string (before base64 encoding). @@ -327,7 +327,7 @@ def encode_unified_id(unified_id_string: str) -> str: return base64.urlsafe_b64encode(unified_id_string.encode()).decode().rstrip("=") -def decode_unified_id(encoded_unified_id: str) -> Optional[str]: +def decode_unified_id(encoded_unified_id: str) -> str | None: """ Decode a base64 encoded unified ID. @@ -355,7 +355,7 @@ def decode_unified_id(encoded_unified_id: str) -> Optional[str]: def parse_unified_id( unified_id: str, -) -> Optional[dict]: +) -> dict | None: """ Parse a unified ID into its components. diff --git a/litellm/llms/base_llm/ocr/__init__.py b/litellm/llms/base_llm/ocr/__init__.py index 2aea2d67807..075c88a2ad0 100644 --- a/litellm/llms/base_llm/ocr/__init__.py +++ b/litellm/llms/base_llm/ocr/__init__.py @@ -14,10 +14,10 @@ from .transformation import ( __all__ = [ "BaseOCRConfig", "DocumentType", - "OCRResponse", "OCRPage", "OCRPageDimensions", "OCRPageImage", - "OCRUsageInfo", "OCRRequestData", + "OCRResponse", + "OCRUsageInfo", ] diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 0d878bd308c..96f86bc8dc0 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -2,7 +2,7 @@ Base OCR transformation configuration. """ -from typing import TYPE_CHECKING, Any, Dict, List, Union +from typing import TYPE_CHECKING, Any import httpx from pydantic import PrivateAttr @@ -19,7 +19,7 @@ else: # DocumentType for OCR - providers always receive a dict with # type="document_url" or type="image_url" (str values only). # File-type inputs are preprocessed to this format in litellm/ocr/main.py. -DocumentType = Dict[str, str] +DocumentType = dict[str, str] class OCRPageDimensions(LiteLLMPydanticObjectBase): @@ -34,7 +34,7 @@ class OCRPageImage(LiteLLMPydanticObjectBase): """Image extracted from OCR page.""" image_base64: str | None = None - bbox: Dict[str, Any] | None = None + bbox: dict[str, Any] | None = None model_config = {"extra": "allow"} @@ -44,7 +44,7 @@ class OCRPage(LiteLLMPydanticObjectBase): index: int markdown: str - images: List[OCRPageImage] | None = None + images: list[OCRPageImage] | None = None dimensions: OCRPageDimensions | None = None model_config = {"extra": "allow"} @@ -66,7 +66,7 @@ class OCRResponse(LiteLLMPydanticObjectBase): Standardized to Mistral OCR format - other providers should transform to this format. """ - pages: List[OCRPage] + pages: list[OCRPage] model: str document_annotation: Any | None = None usage_info: OCRUsageInfo | None = None @@ -84,8 +84,8 @@ class OCRResponse(LiteLLMPydanticObjectBase): class OCRRequestData(LiteLLMPydanticObjectBase): """OCR request data structure.""" - data: Union[Dict, bytes] | None = None - files: Dict[str, Any] | None = None + data: dict | bytes | None = None + files: dict[str, Any] | None = None class BaseOCRConfig: @@ -121,13 +121,13 @@ class BaseOCRConfig: def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. Override in provider-specific implementations. diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index e243d36a86a..61bfc867371 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -1,5 +1,5 @@ from abc import abstractmethod -from typing import TYPE_CHECKING, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Optional, Union from ..base_utils import BaseLLMModelInfo @@ -18,13 +18,12 @@ class BasePassthroughConfig(BaseLLMModelInfo): """ Check if the request is a streaming request """ - pass def format_url( self, endpoint: str, base_target_url: str, - request_query_params: Optional[dict], + request_query_params: dict | None, ) -> "URL": """ Helper function to add query params to the url @@ -53,29 +52,28 @@ class BasePassthroughConfig(BaseLLMModelInfo): @abstractmethod def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, endpoint: str, - request_query_params: Optional[dict], + request_query_params: dict | None, litellm_params: dict, - ) -> Tuple["URL", str]: + ) -> tuple["URL", str]: """ Get the complete url for the request Returns: - complete_url: URL - the complete url for the request - base_target_url: str - the base url to add the endpoint to. Useful for auth headers. """ - pass def sign_request( self, headers: dict, litellm_params: dict, - request_data: Optional[dict], + request_data: dict | None, api_base: str, - model: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + model: str | None = None, + ) -> tuple[dict, bytes | None]: """ Some providers like Bedrock require signing the request. The sign request funtion needs access to `request_data` and `complete_url` Args: @@ -110,7 +108,7 @@ class BasePassthroughConfig(BaseLLMModelInfo): def handle_logging_collected_chunks( self, - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: "LiteLLMLoggingObj", model: str, custom_llm_provider: str, @@ -118,7 +116,7 @@ class BasePassthroughConfig(BaseLLMModelInfo): ) -> Optional["CostResponseTypes"]: return None - def _convert_raw_bytes_to_str_lines(self, raw_bytes: List[bytes]) -> List[str]: + def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]: """ Converts a list of raw bytes into a list of string lines, similar to aiter_lines() diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 4c8cc30a8b3..44f3d001ce8 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -7,7 +7,6 @@ These are HTTP (not WebSocket) endpoints used by the WebRTC flow: """ from abc import ABC, abstractmethod -from typing import Optional, Union import httpx @@ -26,7 +25,7 @@ class BaseRealtimeHTTPConfig(ABC): @abstractmethod def get_api_base( self, - api_base: Optional[str], + api_base: str | None, **kwargs, ) -> str: """ @@ -39,7 +38,7 @@ class BaseRealtimeHTTPConfig(ABC): @abstractmethod def get_api_key( self, - api_key: Optional[str], + api_key: str | None, **kwargs, ) -> str: """ @@ -54,16 +53,13 @@ class BaseRealtimeHTTPConfig(ABC): # ------------------------------------------------------------------ # @abstractmethod - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: """Return the full URL for POST /realtime/client_secrets.""" - def get_transcription_session_url( - self, api_base: Optional[str], model: str, api_version: Optional[str] = None - ) -> str: + def get_transcription_session_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: """Return the full URL for POST /realtime/transcription_sessions.""" base = (api_base or "").rstrip("/") - if base.endswith("/v1"): - base = base[:-3] + base = base.removesuffix("/v1") return f"{base}/v1/realtime/transcription_sessions" @abstractmethod @@ -71,7 +67,7 @@ class BaseRealtimeHTTPConfig(ABC): self, headers: dict, model: str, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: """ Build and return the request headers for the client_secrets call. @@ -84,7 +80,7 @@ class BaseRealtimeHTTPConfig(ABC): # realtime_calls endpoint # # ------------------------------------------------------------------ # - def get_realtime_calls_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: """Return the full URL for POST /realtime/calls (SDP exchange).""" base = (api_base or "").rstrip("/") return f"{base}/v1/realtime/calls" @@ -104,7 +100,7 @@ class BaseRealtimeHTTPConfig(ABC): # Error handling # # ------------------------------------------------------------------ # - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]): + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers): """ Map HTTP errors to LiteLLM exception types. diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py index c24267ccc72..26c189504df 100644 --- a/litellm/llms/base_llm/realtime/transformation.py +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -25,12 +25,12 @@ class BaseRealtimeConfig(ABC): self, headers: dict, model: str, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: pass @abstractmethod - def get_complete_url(self, api_base: Optional[str], model: str, api_key: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_key: str | None = None) -> str: """ OPTIONAL @@ -40,9 +40,7 @@ class BaseRealtimeConfig(ABC): """ return api_base or "" - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: raise BaseLLMException( status_code=status_code, message=error_message, @@ -54,8 +52,8 @@ class BaseRealtimeConfig(ABC): self, message: str, model: str, - session_configuration_request: Optional[str] = None, - ) -> List[str]: + session_configuration_request: str | None = None, + ) -> list[str]: pass def is_setup_message(self, msg_obj: dict) -> bool: @@ -69,15 +67,15 @@ class BaseRealtimeConfig(ABC): ) -> bool: # initial configuration message sent to setup the realtime session return False - def session_configuration_request(self, model: str) -> Optional[str]: # message sent to setup the realtime session + def session_configuration_request(self, model: str) -> str | None: # message sent to setup the realtime session return None def transform_session_created_event( self, model: str, logging_session_id: str, - session_configuration_request: Optional[str] = None, - ) -> Optional[Union[dict, OpenAIRealtimeStreamSessionEvents]]: + session_configuration_request: str | None = None, + ) -> dict | OpenAIRealtimeStreamSessionEvents | None: """ Optional hook for providers that defer session setup until client `session.update`. @@ -89,7 +87,7 @@ class BaseRealtimeConfig(ABC): @abstractmethod def transform_realtime_response( self, - message: Union[str, bytes], + message: str | bytes, model: str, logging_obj: LiteLLMLoggingObj, realtime_response_transform_input: RealtimeResponseTransformInput, @@ -97,4 +95,3 @@ class BaseRealtimeConfig(ABC): """ Keep this state less - leave the state management (e.g. tracking current_output_item_id, current_response_id, current_conversation_id, current_delta_chunks) to the caller. """ - pass diff --git a/litellm/llms/base_llm/rerank/transformation.py b/litellm/llms/base_llm/rerank/transformation.py index eac44ba85c5..523b31b0902 100644 --- a/litellm/llms/base_llm/rerank/transformation.py +++ b/litellm/llms/base_llm/rerank/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -31,7 +31,7 @@ class BaseRerankConfig(ABC): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -78,20 +78,18 @@ class BaseRerankConfig(ABC): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: pass - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: raise BaseLLMException( status_code=status_code, message=error_message, @@ -104,7 +102,7 @@ class BaseRerankConfig(ABC): custom_llm_provider: str | None = None, billed_units: RerankBilledUnits | None = None, model_info: ModelInfo | None = None, - ) -> Tuple[float, float]: + ) -> tuple[float, float]: """ Calculates the cost per query for a given rerank model. diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index c6453745e5c..f55ae8f4692 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx @@ -68,11 +68,11 @@ class BaseResponsesAPIConfig(ABC): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: """Sign the request after the body is finalized. Default is a no-op (returns headers unchanged, no signed body). Providers @@ -92,17 +92,17 @@ class BaseResponsesAPIConfig(ABC): response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: pass @abstractmethod - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: return {} @abstractmethod def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -120,11 +120,11 @@ class BaseResponsesAPIConfig(ABC): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: pass @abstractmethod @@ -146,7 +146,6 @@ class BaseResponsesAPIConfig(ABC): """ Transform a parsed streaming response chunk into a ResponsesAPIStreamingResponse """ - pass ######################################################### ########## DELETE RESPONSE API TRANSFORMATION ############## @@ -158,7 +157,7 @@ class BaseResponsesAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -183,7 +182,7 @@ class BaseResponsesAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -204,12 +203,12 @@ class BaseResponsesAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -217,16 +216,14 @@ class BaseResponsesAPIConfig(ABC): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - ) -> Dict: + ) -> dict: pass ######################################################### ########## END GET RESPONSE API TRANSFORMATION ########## ######################################################### - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ..chat.transformation import BaseLLMException raise BaseLLMException( @@ -237,9 +234,9 @@ class BaseResponsesAPIConfig(ABC): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """Returns True if litellm should fake a stream for the given model and stream value""" return False @@ -258,7 +255,7 @@ class BaseResponsesAPIConfig(ABC): def get_websocket_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -289,7 +286,7 @@ class BaseResponsesAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -311,12 +308,12 @@ class BaseResponsesAPIConfig(ABC): def transform_compact_response_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -333,14 +330,14 @@ class BaseResponsesAPIConfig(ABC): @staticmethod def strip_custom_tool_call_namespace_from_responses_input( - input: Union[str, ResponseInputParam], - ) -> Union[str, ResponseInputParam]: + input: str | ResponseInputParam, + ) -> str | ResponseInputParam: """ Remove ``namespace`` from ``custom_tool_call`` input items. """ if not isinstance(input, list): return input - out: List[Any] = [] + out: list[Any] = [] for item in input: if isinstance(item, dict) and item.get("type") == "custom_tool_call": out.append({k: v for k, v in item.items() if k != "namespace"}) @@ -349,7 +346,7 @@ class BaseResponsesAPIConfig(ABC): return cast(ResponseInputParam, out) @staticmethod - def normalize_responses_api_request_dict(data: Dict[str, Any]) -> Dict[str, Any]: + def normalize_responses_api_request_dict(data: dict[str, Any]) -> dict[str, Any]: """Apply provider-agnostic fixes to an outbound Responses API request dict.""" if not isinstance(data, dict) or "input" not in data: return data diff --git a/litellm/llms/base_llm/sandbox/transformation.py b/litellm/llms/base_llm/sandbox/transformation.py index c807283ecd2..d94a4eb9ed4 100644 --- a/litellm/llms/base_llm/sandbox/transformation.py +++ b/litellm/llms/base_llm/sandbox/transformation.py @@ -6,10 +6,9 @@ returns whatever the sandbox produced. The lifecycle is create container -> run code -> delete container; `code_interpreter_tool` combines all three. """ -from typing import Any, Union +from typing import Any import httpx - from pydantic import Field, PrivateAttr from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -64,7 +63,7 @@ class BaseSandboxConfig: async def arun_code( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, code: str, api_key: str | None = None, **kwargs, @@ -74,7 +73,7 @@ class BaseSandboxConfig: async def adelete_sandbox( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, api_key: str | None = None, **kwargs, ) -> bool: diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index fdfac6f5f9f..422f36b73e3 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -2,7 +2,7 @@ Base Search transformation configuration. """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Literal from urllib.parse import urlsplit import httpx @@ -47,8 +47,8 @@ class SearchResult(LiteLLMPydanticObjectBase): title: str url: str snippet: str - date: Optional[str] = None - last_updated: Optional[str] = None + date: str | None = None + last_updated: str | None = None model_config = {"extra": "allow"} @@ -59,7 +59,7 @@ class SearchResponse(LiteLLMPydanticObjectBase): Standardized to Perplexity Search format - other providers should transform to this format. """ - results: List[SearchResult] + results: list[SearchResult] object: str = "search" model_config = {"extra": "allow"} @@ -167,11 +167,11 @@ class BaseSearchConfig: def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. Override in provider-specific implementations. @@ -180,9 +180,9 @@ class BaseSearchConfig: def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -207,10 +207,10 @@ class BaseSearchConfig: def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Union[Dict, List[Dict]]: + ) -> dict | list[dict]: """ Transform Search request to provider-specific format. Override in provider-specific implementations. diff --git a/litellm/llms/base_llm/skills/transformation.py b/litellm/llms/base_llm/skills/transformation.py index 5bb181f59fb..ba202e3db5d 100644 --- a/litellm/llms/base_llm/skills/transformation.py +++ b/litellm/llms/base_llm/skills/transformation.py @@ -3,7 +3,7 @@ Base configuration class for Skills API """ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -38,7 +38,7 @@ class BaseSkillsAPIConfig(ABC): pass @abstractmethod - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate and update headers with provider-specific requirements @@ -54,9 +54,9 @@ class BaseSkillsAPIConfig(ABC): @abstractmethod def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, endpoint: str, - skill_id: Optional[str] = None, + skill_id: str | None = None, ) -> str: """ Get the complete URL for the API request @@ -79,7 +79,7 @@ class BaseSkillsAPIConfig(ABC): create_request: CreateSkillRequest, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Transform create skill request to provider-specific format @@ -91,7 +91,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Provider-specific request body """ - pass @abstractmethod def transform_create_skill_response( @@ -109,7 +108,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Skill object """ - pass @abstractmethod def transform_list_skills_request( @@ -117,7 +115,7 @@ class BaseSkillsAPIConfig(ABC): list_params: ListSkillsParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform list skills request parameters @@ -129,7 +127,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Tuple of (url, query_params) """ - pass @abstractmethod def transform_list_skills_response( @@ -147,7 +144,6 @@ class BaseSkillsAPIConfig(ABC): Returns: ListSkillsResponse object """ - pass @abstractmethod def transform_get_skill_request( @@ -156,7 +152,7 @@ class BaseSkillsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform get skill request @@ -169,7 +165,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Tuple of (url, headers) """ - pass @abstractmethod def transform_get_skill_response( @@ -187,7 +182,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Skill object """ - pass @abstractmethod def transform_delete_skill_request( @@ -196,7 +190,7 @@ class BaseSkillsAPIConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform delete skill request @@ -209,7 +203,6 @@ class BaseSkillsAPIConfig(ABC): Returns: Tuple of (url, headers) """ - pass @abstractmethod def transform_delete_skill_response( @@ -227,7 +220,6 @@ class BaseSkillsAPIConfig(ABC): Returns: DeleteSkillResponse object """ - pass def get_error_class( self, diff --git a/litellm/llms/base_llm/text_to_speech/transformation.py b/litellm/llms/base_llm/text_to_speech/transformation.py index cbae6904ead..fb85cdb7687 100644 --- a/litellm/llms/base_llm/text_to_speech/transformation.py +++ b/litellm/llms/base_llm/text_to_speech/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, TypedDict, Union +from typing import TYPE_CHECKING, Any, TypedDict import httpx @@ -29,9 +29,9 @@ class TextToSpeechRequestData(TypedDict, total=False): Providers should set ONE of: dict_body, ssml_body, or text_body. """ - dict_body: Dict[str, Any] # JSON request body (e.g., OpenAI TTS) + dict_body: dict[str, Any] # JSON request body (e.g., OpenAI TTS) ssml_body: str # SSML/XML string body (e.g., Azure AVA TTS) - headers: Dict[str, str] # Provider-specific headers to merge with base headers + headers: dict[str, str] # Provider-specific headers to merge with base headers class BaseTextToSpeechConfig(ABC): @@ -62,29 +62,27 @@ class BaseTextToSpeechConfig(ABC): """ Get list of OpenAI TTS parameters supported by this provider """ - pass @abstractmethod def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Dict = {}, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict = {}, + ) -> tuple[str | None, dict]: """ Map OpenAI TTS parameters to provider-specific parameters """ - pass @abstractmethod def validate_environment( self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and return headers @@ -95,7 +93,7 @@ class BaseTextToSpeechConfig(ABC): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -110,9 +108,9 @@ class BaseTextToSpeechConfig(ABC): self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ @@ -123,7 +121,6 @@ class BaseTextToSpeechConfig(ABC): - body: The request body (JSON dict, XML string, or binary data) - headers: Provider-specific headers to merge with base headers """ - pass @abstractmethod def transform_text_to_speech_response( @@ -135,9 +132,8 @@ class BaseTextToSpeechConfig(ABC): """ Transform provider response to standard format """ - pass - def get_error_class(self, error_message: str, status_code: int, headers: Dict) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict) -> BaseLLMException: from ..chat.transformation import BaseLLMException raise BaseLLMException( diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index b222e3dd160..8083d2485ba 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -1,5 +1,5 @@ from abc import abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -27,7 +27,7 @@ else: class BaseVectorStoreConfig: - def get_supported_openai_params(self, model: str) -> List[VECTOR_STORE_OPENAI_PARAMS]: + def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]: return [] def map_openai_params( @@ -50,25 +50,25 @@ class BaseVectorStoreConfig: def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: pass async def atransform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Optional async version of transform_search_vector_store_request. If not implemented, the handler will fall back to the sync version. @@ -96,7 +96,7 @@ class BaseVectorStoreConfig: self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: pass @abstractmethod @@ -104,13 +104,13 @@ class BaseVectorStoreConfig: pass @abstractmethod - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: return {} @abstractmethod def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -124,9 +124,7 @@ class BaseVectorStoreConfig: raise ValueError("api_base is required") return api_base - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ..chat.transformation import BaseLLMException raise BaseLLMException( @@ -138,11 +136,11 @@ class BaseVectorStoreConfig: def sign_request( self, headers: dict, - optional_params: Dict, - request_data: Dict, + optional_params: dict, + request_data: dict, api_base: str, - api_key: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + ) -> tuple[dict, bytes | None]: """Optionally sign or modify the request before sending. Providers like AWS Bedrock require SigV4 signing. Providers that don't @@ -154,5 +152,5 @@ class BaseVectorStoreConfig: def calculate_vector_store_cost( self, response: VectorStoreSearchResponse, - ) -> Tuple[float, float]: + ) -> tuple[float, float]: return 0.0, 0.0 diff --git a/litellm/llms/base_llm/vector_store_files/transformation.py b/litellm/llms/base_llm/vector_store_files/transformation.py index e8799c56cae..74aa283113c 100644 --- a/litellm/llms/base_llm/vector_store_files/transformation.py +++ b/litellm/llms/base_llm/vector_store_files/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -34,7 +34,7 @@ class BaseVectorStoreFilesConfig(ABC): def get_supported_openai_params( self, operation: str, - ) -> Tuple[str, ...]: + ) -> tuple[str, ...]: """Return the set of OpenAI params supported for the given operation.""" return tuple() @@ -43,38 +43,38 @@ class BaseVectorStoreFilesConfig(ABC): self, *, operation: str, - non_default_params: Dict[str, Any], - optional_params: Dict[str, Any], + non_default_params: dict[str, Any], + optional_params: dict[str, Any], drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Map non-default OpenAI params to provider-specific params.""" return optional_params @abstractmethod - def get_auth_credentials(self, litellm_params: Dict[str, Any]) -> VectorStoreFileAuthCredentials: ... + def get_auth_credentials(self, litellm_params: dict[str, Any]) -> VectorStoreFileAuthCredentials: ... @abstractmethod def get_vector_store_file_endpoints_by_type( self, - ) -> Dict[str, Tuple[Tuple[str, str], ...]]: ... + ) -> dict[str, tuple[tuple[str, str], ...]]: ... @abstractmethod def validate_environment( self, *, - headers: Dict[str, str], - litellm_params: Optional[GenericLiteLLMParams], - ) -> Dict[str, str]: + headers: dict[str, str], + litellm_params: GenericLiteLLMParams | None, + ) -> dict[str, str]: return {} @abstractmethod def get_complete_url( self, *, - api_base: Optional[str], + api_base: str | None, vector_store_id: str, - litellm_params: Dict[str, Any], + litellm_params: dict[str, Any], ) -> str: if api_base is None: raise ValueError("api_base is required") @@ -87,7 +87,7 @@ class BaseVectorStoreFilesConfig(ABC): vector_store_id: str, create_request: VectorStoreFileCreateRequest, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_create_vector_store_file_response( @@ -103,7 +103,7 @@ class BaseVectorStoreFilesConfig(ABC): vector_store_id: str, query_params: VectorStoreFileListQueryParams, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_list_vector_store_files_response( @@ -119,7 +119,7 @@ class BaseVectorStoreFilesConfig(ABC): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_retrieve_vector_store_file_response( @@ -135,7 +135,7 @@ class BaseVectorStoreFilesConfig(ABC): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_retrieve_vector_store_file_content_response( @@ -152,7 +152,7 @@ class BaseVectorStoreFilesConfig(ABC): file_id: str, update_request: VectorStoreFileUpdateRequest, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_update_vector_store_file_response( @@ -168,7 +168,7 @@ class BaseVectorStoreFilesConfig(ABC): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: ... + ) -> tuple[str, dict[str, Any]]: ... @abstractmethod def transform_delete_vector_store_file_response( @@ -182,7 +182,7 @@ class BaseVectorStoreFilesConfig(ABC): *, error_message: str, status_code: int, - headers: Union[Dict[str, Any], httpx.Headers], + headers: dict[str, Any] | httpx.Headers, ) -> BaseLLMException: from ..chat.transformation import BaseLLMException @@ -195,16 +195,16 @@ class BaseVectorStoreFilesConfig(ABC): def sign_request( self, *, - headers: Dict[str, str], - optional_params: Dict[str, Any], - request_data: Dict[str, Any], + headers: dict[str, str], + optional_params: dict[str, Any], + request_data: dict[str, Any], api_base: str, - api_key: Optional[str] = None, - ) -> Tuple[Dict[str, str], Optional[bytes]]: + api_key: str | None = None, + ) -> tuple[dict[str, str], bytes | None]: return headers, None def prepare_chunking_strategy( self, - chunking_strategy: Optional[VectorStoreFileChunkingStrategy], - ) -> Optional[VectorStoreFileChunkingStrategy]: + chunking_strategy: VectorStoreFileChunkingStrategy | None, + ) -> VectorStoreFileChunkingStrategy | None: return chunking_strategy diff --git a/litellm/llms/base_llm/videos/transformation.py b/litellm/llms/base_llm/videos/transformation.py index e3a66af24a8..1aea3cafe33 100644 --- a/litellm/llms/base_llm/videos/transformation.py +++ b/litellm/llms/base_llm/videos/transformation.py @@ -1,6 +1,6 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -60,7 +60,7 @@ class BaseVideoConfig(ABC): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: pass @abstractmethod @@ -68,8 +68,8 @@ class BaseVideoConfig(ABC): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: return {} @@ -77,7 +77,7 @@ class BaseVideoConfig(ABC): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -97,10 +97,10 @@ class BaseVideoConfig(ABC): model: str, prompt: str, api_base: str, - video_create_optional_request_params: Dict, + video_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles, str]: + ) -> tuple[dict, RequestFiles, str]: pass @abstractmethod @@ -109,8 +109,8 @@ class BaseVideoConfig(ABC): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: pass @@ -121,15 +121,14 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - variant: Optional[str] = None, - ) -> Tuple[str, Dict]: + variant: str | None = None, + ) -> tuple[str, dict]: """ Transform the video content request into a URL and data/params Returns: Tuple[str, Dict]: (url, params) for the video content request """ - pass @abstractmethod def transform_video_content_response( @@ -172,22 +171,21 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video remix request into a URL and data Returns: Tuple[str, Dict]: (url, data) for the video remix request """ - pass @abstractmethod def transform_video_remix_response( self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: pass @@ -197,26 +195,25 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video list request into a URL and params Returns: Tuple[str, Dict]: (url, params) for the video list request """ - pass @abstractmethod def transform_video_list_response( self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - ) -> Dict[str, str]: + custom_llm_provider: str | None = None, + ) -> dict[str, str]: pass @abstractmethod @@ -226,14 +223,13 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video delete request into a URL and data Returns: Tuple[str, Dict]: (url, data) for the video delete request """ - pass @abstractmethod def transform_video_delete_response( @@ -250,21 +246,20 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video retrieve request into a URL and data/params Returns: Tuple[str, Dict]: (url, params) for the video retrieve request """ - pass @abstractmethod def transform_video_status_retrieve_response( self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: pass @@ -275,7 +270,7 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, list]: + ) -> tuple[str, list]: """ Transform the video create character request into a URL and files list (multipart). @@ -297,7 +292,7 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video get character request into a URL and params. @@ -319,7 +314,7 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Optional[Tuple[str, Dict]]: + ) -> tuple[str, dict] | None: """ Return (url, body) for a pre-fetch HTTP call that must be made before transform_video_edit_request, or None if no pre-fetch is required. @@ -337,9 +332,9 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - prefetched_source_data: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + prefetched_source_data: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video edit request into a URL and JSON data. @@ -352,8 +347,8 @@ class BaseVideoConfig(ABC): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: raise NotImplementedError("video edit is not supported for this provider") @@ -365,8 +360,8 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video extension request into a URL and JSON data. @@ -379,13 +374,11 @@ class BaseVideoConfig(ABC): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: raise NotImplementedError("video extension is not supported for this provider") - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ..chat.transformation import BaseLLMException raise BaseLLMException( diff --git a/litellm/llms/baseten/chat.py b/litellm/llms/baseten/chat.py index f5d52ef81ff..30b35e55e61 100644 --- a/litellm/llms/baseten/chat.py +++ b/litellm/llms/baseten/chat.py @@ -1,4 +1,3 @@ -from typing import Optional from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig @@ -9,33 +8,33 @@ class BasetenConfig(OpenAIGPTConfig): Below are the parameters: """ - max_tokens: Optional[int] = None - response_format: Optional[dict] = None - seed: Optional[int] = None - stream: Optional[bool] = None - top_p: Optional[int] = None - tool_choice: Optional[str] = None - tools: Optional[list] = None - user: Optional[str] = None - presence_penalty: Optional[int] = None - frequency_penalty: Optional[int] = None - stream_options: Optional[dict] = None + max_tokens: int | None = None + response_format: dict | None = None + seed: int | None = None + stream: bool | None = None + top_p: int | None = None + tool_choice: str | None = None + tools: list | None = None + user: str | None = None + presence_penalty: int | None = None + frequency_penalty: int | None = None + stream_options: dict | None = None def __init__( self, - max_tokens: Optional[int] = None, - response_format: Optional[dict] = None, - seed: Optional[int] = None, - stop: Optional[list] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, - top_p: Optional[int] = None, - tool_choice: Optional[str] = None, - tools: Optional[list] = None, - user: Optional[str] = None, - presence_penalty: Optional[int] = None, - frequency_penalty: Optional[int] = None, - stream_options: Optional[dict] = None, + max_tokens: int | None = None, + response_format: dict | None = None, + seed: int | None = None, + stop: list | None = None, + stream: bool | None = None, + temperature: float | None = None, + top_p: int | None = None, + tool_choice: str | None = None, + tools: list | None = None, + user: str | None = None, + presence_penalty: int | None = None, + frequency_penalty: int | None = None, + stream_options: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index f2e58df3015..e023e3a159b 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,5 +1,4 @@ import base64 -from typing import Union import httpx @@ -41,7 +40,7 @@ class BedrockAudioTranscriptionRustDispatch: custom_llm_provider: str, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: rust_response = rust_transcription_bridge.transcription( model=model, @@ -67,7 +66,7 @@ class BedrockAudioTranscriptionRustDispatch: custom_llm_provider: str, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: rust_response = await rust_transcription_bridge.atranscription( model=model, diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 03d443c5081..58478b2cb82 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -10,11 +10,7 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, - Dict, Literal, - Optional, - Tuple, - Union, cast, get_args, ) @@ -56,12 +52,12 @@ SIGV4_COMPUTED_HEADERS = frozenset({"authorization", "x-amz-date", "x-amz-securi class Boto3CredentialsInfo(BaseModel): credentials: Credentials aws_region_name: str - aws_bedrock_runtime_endpoint: Optional[str] + aws_bedrock_runtime_endpoint: str | None class _WebIdentityTokenClaims(BaseModel): - aud: Optional[Union[str, list[str]]] = None - iss: Optional[str] = None + aud: str | list[str] | None = None + iss: str | None = None class AwsAuthError(Exception): @@ -101,7 +97,7 @@ class BaseAWSLLM: "aws_external_id", ] - def _get_ssl_verify(self, ssl_verify: Optional[Union[bool, str]] = None): + def _get_ssl_verify(self, ssl_verify: bool | str | None = None): """ Get SSL verification setting for boto3 clients. @@ -116,7 +112,7 @@ class BaseAWSLLM: return get_ssl_verify(ssl_verify=ssl_verify) - def get_cache_key(self, credential_args: Dict[str, Optional[str]]) -> str: + def get_cache_key(self, credential_args: dict[str, str | None]) -> str: """ Generate a unique cache key based on the credential arguments. """ @@ -126,8 +122,8 @@ class BaseAWSLLM: def _get_or_set_cached_credentials( self, - credential_args: Dict[str, Optional[str]], - credential_fetcher: Callable[[], Tuple[Any, Optional[int]]], + credential_args: dict[str, str | None], + credential_fetcher: Callable[[], tuple[Any, int | None]], ) -> Any: """ Read-through IAM cache on the process-wide ``DualCache``. @@ -156,50 +152,50 @@ class BaseAWSLLM: @staticmethod def _is_auth_with_web_identity_token( - aws_web_identity_token: Optional[str], - aws_role_name: Optional[str], - aws_session_name: Optional[str], + aws_web_identity_token: str | None, + aws_role_name: str | None, + aws_session_name: str | None, ) -> bool: return aws_web_identity_token is not None and aws_role_name is not None and aws_session_name is not None @staticmethod - def _is_auth_with_aws_role(aws_role_name: Optional[str]) -> bool: + def _is_auth_with_aws_role(aws_role_name: str | None) -> bool: return aws_role_name is not None @staticmethod - def _is_auth_with_aws_profile(aws_profile_name: Optional[str]) -> bool: + def _is_auth_with_aws_profile(aws_profile_name: str | None) -> bool: return aws_profile_name is not None @staticmethod def _is_auth_with_aws_session_token_tuple( - aws_access_key_id: Optional[str], - aws_secret_access_key: Optional[str], - aws_session_token: Optional[str], + aws_access_key_id: str | None, + aws_secret_access_key: str | None, + aws_session_token: str | None, ) -> bool: return aws_access_key_id is not None and aws_secret_access_key is not None and aws_session_token is not None @staticmethod def _is_auth_with_access_key_and_secret_key( - aws_access_key_id: Optional[str], - aws_secret_access_key: Optional[str], - aws_region_name: Optional[str], + aws_access_key_id: str | None, + aws_secret_access_key: str | None, + aws_region_name: str | None, ) -> bool: return aws_access_key_id is not None and aws_secret_access_key is not None and aws_region_name is not None @tracer.wrap() def get_credentials( self, - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, - aws_region_name: Optional[str] = None, - aws_session_name: Optional[str] = None, - aws_profile_name: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_web_identity_token: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, - aws_external_id: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, + aws_access_key_id: str | None = None, + aws_secret_access_key: str | None = None, + aws_session_token: str | None = None, + aws_region_name: str | None = None, + aws_session_name: str | None = None, + aws_profile_name: str | None = None, + aws_role_name: str | None = None, + aws_web_identity_token: str | None = None, + aws_sts_endpoint: str | None = None, + aws_external_id: str | None = None, + ssl_verify: bool | str | None = None, ): """ Return a boto3.Credentials object @@ -345,7 +341,7 @@ class BaseAWSLLM: else: return self._get_or_set_cached_credentials(args, self._auth_with_env_vars) - def _get_aws_region_from_model_arn(self, model: Optional[str]) -> Optional[str]: + def _get_aws_region_from_model_arn(self, model: str | None) -> str | None: try: # First check if the string contains the expected prefix if not isinstance(model, str) or "arn:aws:bedrock" not in model: @@ -372,7 +368,7 @@ class BaseAWSLLM: @staticmethod def _get_provider_from_model_path( model_path: str, - ) -> Optional[BEDROCK_INVOKE_PROVIDERS_LITERAL]: + ) -> BEDROCK_INVOKE_PROVIDERS_LITERAL | None: """ Helper function to get the provider from a model path with format: provider/model-name @@ -392,7 +388,7 @@ class BaseAWSLLM: @staticmethod def get_bedrock_invoke_provider( model: str, - ) -> Optional[BEDROCK_INVOKE_PROVIDERS_LITERAL]: + ) -> BEDROCK_INVOKE_PROVIDERS_LITERAL | None: """ Helper function to get the bedrock provider from the model @@ -428,7 +424,7 @@ class BaseAWSLLM: @staticmethod def get_bedrock_model_id( optional_params: dict, - provider: Optional[BEDROCK_INVOKE_PROVIDERS_LITERAL], + provider: BEDROCK_INVOKE_PROVIDERS_LITERAL | None, model: str, ) -> str: model_id = optional_params.pop("model_id", None) @@ -501,7 +497,7 @@ class BaseAWSLLM: @staticmethod def get_bedrock_embedding_provider( model: str, - ) -> Optional[BEDROCK_EMBEDDING_PROVIDERS_LITERAL]: + ) -> BEDROCK_EMBEDDING_PROVIDERS_LITERAL | None: """ Helper function to get the bedrock embedding provider from the model @@ -542,8 +538,8 @@ class BaseAWSLLM: def _get_aws_region_name( self, optional_params: dict, - model: Optional[str] = None, - model_id: Optional[str] = None, + model: str | None = None, + model_id: str | None = None, ) -> str: """ Get the AWS region name from the environment variables. @@ -600,7 +596,7 @@ class BaseAWSLLM: return aws_region_name @staticmethod - def _validate_aws_region_name(aws_region_name: Optional[str]) -> None: + def _validate_aws_region_name(aws_region_name: str | None) -> None: """ Validate that an AWS region name conforms to the expected format (lowercase alphanumerics and hyphens). Raises ValueError otherwise. @@ -615,8 +611,8 @@ class BaseAWSLLM: @staticmethod def _parse_sts_region_from_endpoint( - aws_sts_endpoint: Optional[str], - ) -> Optional[str]: + aws_sts_endpoint: str | None, + ) -> str | None: """Extract region from sts.{region}.amazonaws.com or vpce-x.sts.{region}.vpce.amazonaws.com.""" if not aws_sts_endpoint: return None @@ -625,7 +621,7 @@ class BaseAWSLLM: return match.group(1) if match else None @staticmethod - def _resolve_sts_region(aws_sts_endpoint: Optional[str] = None) -> Optional[str]: + def _resolve_sts_region(aws_sts_endpoint: str | None = None) -> str | None: """STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION.""" return ( BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint) @@ -635,8 +631,8 @@ class BaseAWSLLM: def _build_sts_client_kwargs( self, - aws_sts_endpoint: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, + aws_sts_endpoint: str | None = None, + ssl_verify: bool | str | None = None, ) -> dict: """STS client kwargs with aligned endpoint_url and region_name (SigV4).""" kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)} @@ -649,7 +645,7 @@ class BaseAWSLLM: def get_aws_region_name_for_non_llm_api_calls( self, - aws_region_name: Optional[str] = None, + aws_region_name: str | None = None, ): """ Get the AWS region name for non-llm api calls. @@ -679,7 +675,7 @@ class BaseAWSLLM: @staticmethod def _parse_arn_account_and_role_name( arn: str, - ) -> Optional[Tuple[str, str, str]]: + ) -> tuple[str, str, str] | None: """ Parse an ARN and return (partition, account_id, role_name). @@ -717,7 +713,7 @@ class BaseAWSLLM: def _is_already_running_as_role( self, aws_role_name: str, - ssl_verify: Optional[Union[bool, str]] = None, + ssl_verify: bool | str | None = None, ) -> bool: """ Check if the current environment is already running as the target IAM role. @@ -774,7 +770,7 @@ class BaseAWSLLM: return False @staticmethod - def _unverified_web_identity_audience(oidc_token: str) -> Optional[str]: + def _unverified_web_identity_audience(oidc_token: str) -> str | None: """Return the public ``aud``/``iss`` claims of a web identity JWT without verifying its signature, so a rejected-token error can name the audience LiteLLM actually sent. The signature is never read, so no @@ -798,11 +794,11 @@ class BaseAWSLLM: aws_web_identity_token: str, aws_role_name: str, aws_session_name: str, - aws_region_name: Optional[str], - aws_sts_endpoint: Optional[str], - aws_external_id: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, - ) -> Tuple[Credentials, Optional[int]]: + aws_region_name: str | None, + aws_sts_endpoint: str | None, + aws_external_id: str | None = None, + ssl_verify: bool | str | None = None, + ) -> tuple[Credentials, int | None]: """ Authenticate with AWS Web Identity Token """ @@ -940,9 +936,9 @@ class BaseAWSLLM: aws_role_name: str, aws_session_name: str, web_identity_token_file: str, - aws_external_id: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, + aws_external_id: str | None = None, + aws_sts_endpoint: str | None = None, + ssl_verify: bool | str | None = None, ) -> dict: """Handle cross-account role assumption for IRSA.""" import boto3 @@ -1009,9 +1005,9 @@ class BaseAWSLLM: self, aws_role_name: str, aws_session_name: str, - aws_external_id: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, + aws_external_id: str | None = None, + aws_sts_endpoint: str | None = None, + ssl_verify: bool | str | None = None, ) -> dict: """Handle same-account role assumption for IRSA.""" import boto3 @@ -1045,7 +1041,7 @@ class BaseAWSLLM: return sts_client.assume_role(**assume_role_params) - def _extract_credentials_and_ttl(self, sts_response: dict) -> Tuple[Credentials, Optional[int]]: + def _extract_credentials_and_ttl(self, sts_response: dict) -> tuple[Credentials, int | None]: """Extract credentials and TTL from STS response.""" from botocore.credentials import Credentials @@ -1064,16 +1060,16 @@ class BaseAWSLLM: @tracer.wrap() def _auth_with_aws_role( self, - aws_access_key_id: Optional[str], - aws_secret_access_key: Optional[str], - aws_session_token: Optional[str], + aws_access_key_id: str | None, + aws_secret_access_key: str | None, + aws_session_token: str | None, aws_role_name: str, aws_session_name: str, - aws_region_name: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, - aws_external_id: Optional[str] = None, - ssl_verify: Optional[Union[bool, str]] = None, - ) -> Tuple[Credentials, Optional[int]]: + aws_region_name: str | None = None, + aws_sts_endpoint: str | None = None, + aws_external_id: str | None = None, + ssl_verify: bool | str | None = None, + ) -> tuple[Credentials, int | None]: """ Authenticate with AWS Role """ @@ -1196,7 +1192,7 @@ class BaseAWSLLM: return credentials, sts_ttl @tracer.wrap() - def _auth_with_aws_profile(self, aws_profile_name: str) -> Tuple[Credentials, Optional[int]]: + def _auth_with_aws_profile(self, aws_profile_name: str) -> tuple[Credentials, int | None]: """ Authenticate with AWS profile """ @@ -1213,7 +1209,7 @@ class BaseAWSLLM: aws_access_key_id: str, aws_secret_access_key: str, aws_session_token: str, - ) -> Tuple[Credentials, Optional[int]]: + ) -> tuple[Credentials, int | None]: """ Authenticate with AWS Session Token """ @@ -1233,8 +1229,8 @@ class BaseAWSLLM: self, aws_access_key_id: str, aws_secret_access_key: str, - aws_region_name: Optional[str], - ) -> Tuple[Credentials, Optional[int]]: + aws_region_name: str | None, + ) -> tuple[Credentials, int | None]: """ Authenticate with AWS Access Key and Secret Key """ @@ -1254,7 +1250,7 @@ class BaseAWSLLM: return credentials, self._get_default_ttl_for_boto3_credentials() @tracer.wrap() - def _auth_with_env_vars(self) -> Tuple[Credentials, Optional[int]]: + def _auth_with_env_vars(self) -> tuple[Credentials, int | None]: """ Authenticate with AWS Environment Variables """ @@ -1276,11 +1272,11 @@ class BaseAWSLLM: def get_runtime_endpoint( self, - api_base: Optional[str], - aws_bedrock_runtime_endpoint: Optional[str], + api_base: str | None, + aws_bedrock_runtime_endpoint: str | None, aws_region_name: str, - endpoint_type: Optional[Literal["runtime", "agent", "agentcore"]] = "runtime", - ) -> Tuple[str, str]: + endpoint_type: Literal["runtime", "agent", "agentcore"] | None = "runtime", + ) -> tuple[str, str]: env_aws_bedrock_runtime_endpoint = get_secret("AWS_BEDROCK_RUNTIME_ENDPOINT") if api_base is not None: endpoint_url = api_base @@ -1306,7 +1302,7 @@ class BaseAWSLLM: def _select_default_endpoint_url( self, - endpoint_type: Optional[Literal["runtime", "agent", "agentcore"]], + endpoint_type: Literal["runtime", "agent", "agentcore"] | None, aws_region_name: str, ) -> str: """ @@ -1322,7 +1318,7 @@ class BaseAWSLLM: return f"https://bedrock-runtime.{aws_region_name}.amazonaws.com" def _get_boto_credentials_from_optional_params( - self, optional_params: dict, model: Optional[str] = None + self, optional_params: dict, model: str | None = None ) -> Boto3CredentialsInfo: """ Get boto3 credentials from optional params @@ -1378,14 +1374,14 @@ class BaseAWSLLM: self, credentials: Credentials, aws_region_name: str, - extra_headers: Optional[dict], + extra_headers: dict | None, endpoint_url: str, - data: Union[str, bytes], + data: str | bytes, headers: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> AWSPreparedRequest: if api_key is not None: - aws_bearer_token: Optional[str] = api_key + aws_bearer_token: str | None = api_key else: aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") @@ -1471,11 +1467,11 @@ class BaseAWSLLM: optional_params: dict, request_data: dict, api_base: str, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - api_key: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + api_key: str | None = None, + ) -> tuple[dict, bytes | None]: """ Sign a request for Bedrock or Sagemaker @@ -1483,7 +1479,7 @@ class BaseAWSLLM: Tuple[dict, Optional[str]]: A tuple containing the headers and the json str body of the request """ if api_key is not None: - aws_bearer_token: Optional[str] = api_key + aws_bearer_token: str | None = api_key else: aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index b0c7f1a3695..395d1197037 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -1,5 +1,5 @@ from datetime import datetime -from typing import Any, Optional, cast +from typing import Any, cast from openai.types.batch import BatchRequestCounts from openai.types.batch import Metadata as OpenAIBatchMetadata @@ -23,7 +23,7 @@ _BEDROCK_MIJ_STATUS_TO_OPENAI = { } -def _extract_region_from_bedrock_arn(arn: str) -> Optional[str]: +def _extract_region_from_bedrock_arn(arn: str) -> str | None: """ARN shape: ``arn:aws:bedrock:::/``""" try: parts = arn.split(":") @@ -34,14 +34,14 @@ def _extract_region_from_bedrock_arn(arn: str) -> Optional[str]: return None -def _extract_job_id_from_arn(arn: str) -> Optional[str]: +def _extract_job_id_from_arn(arn: str) -> str | None: """``arn:aws:bedrock:::model-invocation-job/`` -> ````.""" if ":model-invocation-job/" not in arn: return None return arn.rsplit("/", 1)[-1] or None -def _predict_output_file_uri(output_prefix: str, input_uri: str, job_id: Optional[str]) -> Optional[str]: +def _predict_output_file_uri(output_prefix: str, input_uri: str, job_id: str | None) -> str | None: """ Compute the deterministic per-job result file URI Bedrock writes to. @@ -63,7 +63,7 @@ def _predict_output_file_uri(output_prefix: str, input_uri: str, job_id: Optiona return f"{output_prefix}{job_id}/{input_basename}.out" -def _to_epoch(value: Any) -> Optional[int]: +def _to_epoch(value: Any) -> int | None: if value is None: return None if isinstance(value, (int, float)): @@ -162,7 +162,7 @@ class BedrockBatchesHandler: @staticmethod def _handle_model_invocation_job_status( batch_id: str, - aws_region_name: Optional[str] = None, + aws_region_name: str | None = None, logging_obj=None, **kwargs, ) -> "LiteLLMBatch": diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index a4ff1c78467..8fdec6282e3 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -1,7 +1,7 @@ import os import re import time -from typing import Any, Dict, List, Literal, Optional, Union, cast +from typing import Any, Literal, cast from httpx import Headers, Response from pydantic import TypeAdapter, ValidationError @@ -67,7 +67,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): return LlmProviders.BEDROCK @classmethod - def _get_bare_model_name_from_s3_key(cls, object_key: str) -> Optional[str]: + def _get_bare_model_name_from_s3_key(cls, object_key: str) -> str | None: if not object_key.startswith(BEDROCK_MANAGED_S3_BATCH_PREFIX): return None model_part = object_key[len(BEDROCK_MANAGED_S3_BATCH_PREFIX) :] @@ -77,7 +77,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): return model_part[: match.start()] @classmethod - def is_unmanaged_s3_batch_input_file_id(cls, input_file_id: Optional[str]) -> bool: + def is_unmanaged_s3_batch_input_file_id(cls, input_file_id: str | None) -> bool: """ Returns True if `input_file_id` is a raw s3:// Bedrock batch input file (i.e. not a LiteLLM-managed unified file id) whose object key embeds the model name in the @@ -105,11 +105,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate and prepare environment for Bedrock batch requests. @@ -120,11 +120,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): def get_complete_batch_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, data: CreateBatchRequest, ) -> str: """ @@ -145,7 +145,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): create_batch_data: CreateBatchRequest, optional_params: dict, litellm_params: dict, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform the batch creation request to Bedrock format. @@ -251,7 +251,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): def transform_create_batch_response( self, - model: Optional[str], + model: str | None, raw_response: Response, logging_obj: Any, litellm_params: dict, @@ -269,7 +269,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): status_str: str = str(response_data.get("status", "Submitted")) # Map Bedrock status to OpenAI-compatible status - status_mapping: Dict[str, str] = { + status_mapping: dict[str, str] = { "Submitted": "validating", "Validating": "validating", "Scheduled": "in_progress", @@ -324,14 +324,14 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): ) @staticmethod - def _get_openai_compatible_batch_metadata(metadata: Any) -> Dict[str, str]: + def _get_openai_compatible_batch_metadata(metadata: Any) -> dict[str, str]: """ OpenAI Batch metadata only accepts string values. """ if not isinstance(metadata, dict): return {} - sanitized_metadata: Dict[str, str] = {} + sanitized_metadata: dict[str, str] = {} for key, value in metadata.items(): if key == "standard_logging_guardrail_information" or value is None: continue @@ -349,7 +349,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): batch_id: str, optional_params: dict, litellm_params: dict, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform batch retrieval request for Bedrock. @@ -405,7 +405,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): """Helper to parse timestamps based on status.""" import datetime - def parse_timestamp(ts_str: Optional[str]) -> Optional[int]: + def parse_timestamp(ts_str: str | None) -> int | None: if not ts_str: return None try: @@ -490,7 +490,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): ) # Enrich metadata with useful Bedrock fields - enriched_metadata_raw: Dict[str, Any] = { + enriched_metadata_raw: dict[str, Any] = { "jobName": response_data.get("jobName"), "clientRequestToken": response_data.get("clientRequestToken"), "modelId": response_data.get("modelId"), @@ -500,7 +500,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): } import json as _json - enriched_metadata: Dict[str, str] = {} + enriched_metadata: dict[str, str] = {} for _k, _v in enriched_metadata_raw.items(): if _v is None: continue @@ -516,7 +516,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): def transform_retrieve_batch_response( self, - model: Optional[str], + model: str | None, raw_response: Response, logging_obj: Any, litellm_params: dict, @@ -535,7 +535,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): status_str: str = str(response_data.get("status", "Submitted")) # Map Bedrock status to OpenAI-compatible status - status_mapping: Dict[str, str] = { + status_mapping: dict[str, str] = { "Submitted": "validating", "Validating": "validating", "Scheduled": "in_progress", @@ -600,7 +600,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): metadata=enriched_metadata, ) - def get_error_class(self, error_message: str, status_code: int, headers: Union[Dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: """ Get Bedrock-specific error class using common utility. """ diff --git a/litellm/llms/bedrock/chat/__init__.py b/litellm/llms/bedrock/chat/__init__.py index 37dcb270743..3c870226640 100644 --- a/litellm/llms/bedrock/chat/__init__.py +++ b/litellm/llms/bedrock/chat/__init__.py @@ -1,5 +1,3 @@ -from typing import Optional - from .converse_handler import BedrockConverseLLM from .invoke_handler import ( AmazonAnthropicClaudeStreamDecoder, @@ -8,7 +6,7 @@ from .invoke_handler import ( ) -def get_bedrock_event_stream_decoder(invoke_provider: Optional[str], model: str, sync_stream: bool, json_mode: bool): +def get_bedrock_event_stream_decoder(invoke_provider: str | None, model: str, sync_stream: bool, json_mode: bool): if invoke_provider and invoke_provider == "anthropic": decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder( model=model, diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 356bc829677..40b12e17e8a 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -6,7 +6,7 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agentcore_InvokeAgen import json from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Optional, Union, cast from urllib.parse import quote import httpx @@ -17,9 +17,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, ) from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.llms.a2a.common_utils import extract_text_from_a2a_response from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.a2a.common_utils import extract_text_from_a2a_response from litellm.llms.bedrock.common_utils import BedrockError from litellm.types.llms.bedrock_agentcore import ( AgentCoreMessage, @@ -53,7 +53,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): BaseConfig.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Bedrock AgentCore has 0 OpenAI compatible params """ @@ -73,12 +73,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -116,11 +116,11 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: # Set Accept header required by MCP servers on AgentCore # Per MCP spec (Streamable HTTP transport): client MUST include Accept header # listing both application/json and text/event-stream as supported content types @@ -186,11 +186,11 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): return session_id # Generate a session ID with 33+ characters - generated_id = f"litellm-session-{str(uuid.uuid4())}" + generated_id = f"litellm-session-{uuid.uuid4()!s}" verbose_logger.debug(f"Generated new session ID: {generated_id}") return generated_id - def _get_runtime_user_id(self, optional_params: dict) -> Optional[str]: + def _get_runtime_user_id(self, optional_params: dict) -> str | None: """ Get runtime user ID if provided """ @@ -202,7 +202,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -288,7 +288,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): return bool(value) return False - def _extract_sse_json(self, line: str) -> Optional[Dict]: + def _extract_sse_json(self, line: str) -> dict | None: """Extract and parse JSON from an SSE data line.""" if not line.startswith("data:"): return None @@ -305,7 +305,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): verbose_logger.debug(f"Skipping non-JSON line: {line[:100]}") return None - def _extract_usage_from_event(self, event_data: Dict) -> Optional[AgentCoreUsage]: + def _extract_usage_from_event(self, event_data: dict) -> AgentCoreUsage | None: """Extract usage information from event metadata.""" event_payload = event_data.get("event") if not event_payload: @@ -317,7 +317,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): return None - def _extract_content_delta(self, event_data: Dict) -> Optional[str]: + def _extract_content_delta(self, event_data: dict) -> str | None: """Extract text content from contentBlockDelta event.""" event_payload = event_data.get("event") if not event_payload: @@ -341,7 +341,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): return "".join(block["text"] for block in content_list if isinstance(block, dict) and "text" in block) - def _calculate_usage(self, model: str, messages: List[AllMessageValues], content: str) -> Optional[Usage]: + def _calculate_usage(self, model: str, messages: list[AllMessageValues], content: str) -> Usage | None: """ Calculate token usage using LiteLLM's token counter. @@ -370,7 +370,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): total_tokens=total_tokens, ) except Exception as e: - verbose_logger.warning(f"Failed to calculate token usage: {str(e)}") + verbose_logger.warning(f"Failed to calculate token usage: {e!s}") return None def _parse_json_response(self, response_json: dict) -> AgentCoreParsedResponse: @@ -483,9 +483,9 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): Returns: AgentCoreParsedResponse: Parsed response with content, usage, and message """ - final_message: Optional[AgentCoreMessage] = None - usage_data: Optional[AgentCoreUsage] = None - content_blocks: List[str] = [] + final_message: AgentCoreMessage | None = None + usage_data: AgentCoreUsage | None = None + content_blocks: list[str] = [] for line in response_text.strip().split("\n"): line = line.strip() @@ -636,9 +636,9 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, "AsyncHTTPHandler"]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": """ Simplified sync streaming - returns a generator that yields ModelResponse chunks. @@ -850,8 +850,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): data: dict, messages: list, client: Optional["AsyncHTTPHandler"] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": """ Simplified async streaming - returns an async generator that yields ModelResponse chunks. @@ -970,12 +970,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the AgentCore response to LiteLLM ModelResponse format. @@ -1023,9 +1023,9 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): return model_response except Exception as e: - verbose_logger.error(f"Error processing Bedrock AgentCore response: {str(e)}") + verbose_logger.error(f"Error processing Bedrock AgentCore response: {e!s}") raise BedrockError( - message=f"Error processing response: {str(e)}", + message=f"Error processing response: {e!s}", status_code=raw_response.status_code, ) @@ -1033,24 +1033,22 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BedrockError(status_code=status_code, message=error_message) def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: # AgentCore supports true streaming - don't buffer return False diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 292f570cc4e..2309965dbe5 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -1,5 +1,5 @@ import json -from typing import Any, Optional, Union +from typing import Any import httpx @@ -23,16 +23,16 @@ from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_ca def make_sync_call( - client: Optional[HTTPHandler], + client: HTTPHandler | None, api_base: str, headers: dict, data: str, model: str, messages: list, logging_obj: LiteLLMLoggingObject, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, fake_stream: bool = False, - stream_chunk_size: Optional[int] = None, + stream_chunk_size: int | None = None, ): if client is None: client = _get_httpx_client() # Create a new client if none provided @@ -87,7 +87,7 @@ class BedrockConverseLLM(BaseAWSLLM): messages: list, api_base: str, model_response: ModelResponse, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding, logging_obj, stream, @@ -96,11 +96,11 @@ class BedrockConverseLLM(BaseAWSLLM): credentials: Credentials, logger_fn=None, headers={}, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, fake_stream: bool = False, - json_mode: Optional[bool] = False, - api_key: Optional[str] = None, - stream_chunk_size: Optional[int] = None, + json_mode: bool | None = False, + api_key: str | None = None, + stream_chunk_size: int | None = None, ) -> CustomStreamWrapper: request_data = await litellm.AmazonConverseConfig()._async_transform_request( model=model, @@ -158,7 +158,7 @@ class BedrockConverseLLM(BaseAWSLLM): messages: list, api_base: str, model_response: ModelResponse, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding, logging_obj: LiteLLMLoggingObject, stream, @@ -167,9 +167,9 @@ class BedrockConverseLLM(BaseAWSLLM): credentials: Credentials, logger_fn=None, headers: dict = {}, - client: Optional[AsyncHTTPHandler] = None, - api_key: Optional[str] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: + client: AsyncHTTPHandler | None = None, + api_key: str | None = None, + ) -> ModelResponse | CustomStreamWrapper: request_data = await litellm.AmazonConverseConfig()._async_transform_request( model=model, messages=messages, @@ -242,19 +242,19 @@ class BedrockConverseLLM(BaseAWSLLM): self, model: str, messages: list, - api_base: Optional[str], + api_base: str | None, custom_prompt_dict: dict, model_response: ModelResponse, encoding, logging_obj: LiteLLMLoggingObject, optional_params: dict, acompletion: bool, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, litellm_params: dict, logger_fn=None, - extra_headers: Optional[dict] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - api_key: Optional[str] = None, + extra_headers: dict | None = None, + client: AsyncHTTPHandler | HTTPHandler | None = None, + api_key: str | None = None, ): ## SETUP ## stream = optional_params.pop("stream", None) @@ -274,7 +274,7 @@ class BedrockConverseLLM(BaseAWSLLM): break # Strip embedded region prefix (e.g. "bedrock/us-east-1/model" -> "model") # and capture it so it can be used as aws_region_name below. - _region_from_model: Optional[str] = None + _region_from_model: str | None = None _potential_region = _stripped.split("/", 1)[0] if _potential_region in _get_all_bedrock_regions() and "/" in _stripped: _region_from_model = _potential_region diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 8ce2b982955..5bd498a465e 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -6,7 +6,7 @@ import copy import json import time import types -from typing import List, Literal, Optional, Tuple, Union, cast, overload +from typing import Literal, cast, overload import httpx @@ -107,19 +107,19 @@ class AmazonConverseConfig(BaseConfig): #2 - https://docs.aws.amazon.com/bedrock/latest/userguide/conversation-inference.html#conversation-inference-supported-models-features """ - maxTokens: Optional[int] - stopSequences: Optional[List[str]] - temperature: Optional[int] - topP: Optional[int] - topK: Optional[int] + maxTokens: int | None + stopSequences: list[str] | None + temperature: int | None + topP: int | None + topK: int | None def __init__( self, - maxTokens: Optional[int] = None, - stopSequences: Optional[List[str]] = None, - temperature: Optional[int] = None, - topP: Optional[int] = None, - topK: Optional[int] = None, + maxTokens: int | None = None, + stopSequences: list[str] | None = None, + temperature: int | None = None, + topP: int | None = None, + topK: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -127,7 +127,7 @@ class AmazonConverseConfig(BaseConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock_converse" @classmethod @@ -140,8 +140,8 @@ class AmazonConverseConfig(BaseConfig): @staticmethod def _convert_consecutive_user_messages_to_guarded_text( - messages: List[AllMessageValues], optional_params: dict - ) -> List[AllMessageValues]: + messages: list[AllMessageValues], optional_params: dict + ) -> list[AllMessageValues]: """ Convert consecutive user messages at the end to guarded_text type if guardrailConfig is present and no guarded_text is already present in those messages. @@ -332,7 +332,7 @@ class AmazonConverseConfig(BaseConfig): # Also check for nova-2/ spec prefix for imported models return model_without_region.startswith("amazon.nova-2-") or model_without_region.startswith("nova-2/") - def _map_web_search_options(self, web_search_options: dict, model: str) -> Optional[BedrockToolBlock]: + def _map_web_search_options(self, web_search_options: dict, model: str) -> BedrockToolBlock | None: """ Map web_search_options to Nova grounding systemTool. @@ -493,7 +493,7 @@ class AmazonConverseConfig(BaseConfig): ) thinking["budget_tokens"] = BEDROCK_MIN_THINKING_BUDGET_TOKENS - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: from litellm.utils import supports_function_calling supported_params = [ @@ -572,16 +572,14 @@ class AmazonConverseConfig(BaseConfig): return supported_params def map_tool_choice_values( - self, model: str, tool_choice: Union[str, dict], drop_params: bool - ) -> Optional[ToolChoiceValuesBlock]: + self, model: str, tool_choice: str | dict, drop_params: bool + ) -> ToolChoiceValuesBlock | None: if tool_choice == "none": if litellm.drop_params is True or drop_params is True: return None else: raise litellm.utils.UnsupportedParamsError( - message="Bedrock doesn't support tool_choice={}. To drop it from the call, set `litellm.drop_params = True.".format( - tool_choice - ), + message=f"Bedrock doesn't support tool_choice={tool_choice}. To drop it from the call, set `litellm.drop_params = True.", status_code=400, ) elif tool_choice == "required": @@ -596,25 +594,23 @@ class AmazonConverseConfig(BaseConfig): return ToolChoiceValuesBlock(tool=specific_tool) else: raise litellm.utils.UnsupportedParamsError( - message="Bedrock doesn't support tool_choice={}. Supported tool_choice values=['auto', 'required', json object]. To drop it from the call, set `litellm.drop_params = True.".format( - tool_choice - ), + message=f"Bedrock doesn't support tool_choice={tool_choice}. Supported tool_choice values=['auto', 'required', json object]. To drop it from the call, set `litellm.drop_params = True.", status_code=400, ) - def get_supported_image_types(self) -> List[str]: + def get_supported_image_types(self) -> list[str]: return ["png", "jpeg", "gif", "webp"] - def get_supported_document_types(self) -> List[str]: + def get_supported_document_types(self) -> list[str]: return ["pdf", "csv", "doc", "docx", "xls", "xlsx", "html", "txt", "md"] - def get_supported_video_types(self) -> List[str]: + def get_supported_video_types(self) -> list[str]: return ["mp4", "mov", "mkv", "webm", "flv", "mpeg", "mpg", "wmv", "3gp"] - def get_all_supported_content_types(self) -> List[str]: + def get_all_supported_content_types(self) -> list[str]: return self.get_supported_image_types() + self.get_supported_document_types() + self.get_supported_video_types() - def is_computer_use_tool_used(self, tools: Optional[List[OpenAIChatCompletionToolParam]], model: str) -> bool: + def is_computer_use_tool_used(self, tools: list[OpenAIChatCompletionToolParam] | None, model: str) -> bool: """Check if computer use tools are being used in the request.""" if tools is None: return False @@ -627,9 +623,9 @@ class AmazonConverseConfig(BaseConfig): return True return False - def _transform_computer_use_tools(self, computer_use_tools: List[OpenAIChatCompletionToolParam]) -> List[dict]: + def _transform_computer_use_tools(self, computer_use_tools: list[OpenAIChatCompletionToolParam]) -> list[dict]: """Transform computer use tools to Bedrock format.""" - transformed_tools: List[dict] = [] + transformed_tools: list[dict] = [] for tool in computer_use_tools: tool_type = tool.get("type", "") @@ -668,8 +664,8 @@ class AmazonConverseConfig(BaseConfig): return transformed_tools def _separate_computer_use_tools( - self, tools: List[OpenAIChatCompletionToolParam], model: str - ) -> Tuple[List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam]]: + self, tools: list[OpenAIChatCompletionToolParam], model: str + ) -> tuple[list[OpenAIChatCompletionToolParam], list[OpenAIChatCompletionToolParam]]: """ Separate computer use tools from regular function tools. @@ -702,8 +698,8 @@ class AmazonConverseConfig(BaseConfig): def _create_json_tool_call_for_response_format( self, - json_schema: Optional[dict] = None, - description: Optional[str] = None, + json_schema: dict | None = None, + description: str | None = None, ) -> ChatCompletionToolParam: """ Handles creating a tool call for getting responses in JSON format. @@ -741,7 +737,7 @@ class AmazonConverseConfig(BaseConfig): return _tool @staticmethod - def _supports_native_structured_outputs(model: str, custom_llm_provider: Optional[str] = None) -> bool: + def _supports_native_structured_outputs(model: str, custom_llm_provider: str | None = None) -> bool: """Check if the Bedrock model supports native structured outputs (outputConfig.textFormat). Delegates to the standard ``supports_native_structured_output`` utility @@ -791,9 +787,9 @@ class AmazonConverseConfig(BaseConfig): @staticmethod def _create_output_config_for_response_format( - json_schema: Optional[dict] = None, - name: Optional[str] = None, - description: Optional[str] = None, + json_schema: dict | None = None, + name: str | None = None, + description: str | None = None, ) -> "OutputConfigBlock": """ Build an outputConfig block for Bedrock's native structured outputs API. @@ -832,7 +828,7 @@ class AmazonConverseConfig(BaseConfig): def _apply_tool_call_transformation( self, - tools: List[OpenAIChatCompletionToolParam], + tools: list[OpenAIChatCompletionToolParam], model: str, non_default_params: dict, optional_params: dict, @@ -881,7 +877,7 @@ class AmazonConverseConfig(BaseConfig): ) if param == "tools" and isinstance(value, list): self._apply_tool_call_transformation( - tools=cast(List[OpenAIChatCompletionToolParam], value), + tools=cast(list[OpenAIChatCompletionToolParam], value), model=model, non_default_params=non_default_params, optional_params=optional_params, @@ -966,14 +962,14 @@ class AmazonConverseConfig(BaseConfig): self._validate_request_metadata(value) # type: ignore optional_params["requestMetadata"] = value - def _map_context_management_param(self, value: Union[dict, list], optional_params: dict) -> None: + def _map_context_management_param(self, value: dict | list, optional_params: dict) -> None: # Match the dispatcher's ``_normalize_spec`` behavior: only run the # OpenAI→Anthropic mapper for list inputs. Dict inputs are already in # Anthropic-native shape (``{"edits": [...]}``) and should pass # through unchanged so an Anthropic-format ``context_management`` # value isn't silently dropped when the mapper can't classify it. if isinstance(value, list): - mapped = AnthropicConfig.map_openai_context_management_to_anthropic(cast(Union[dict, list], value)) + mapped = AnthropicConfig.map_openai_context_management_to_anthropic(cast(dict | list, value)) else: mapped = value # Skip when the mapper returned None for malformed input — leaving the @@ -1012,9 +1008,9 @@ class AmazonConverseConfig(BaseConfig): if value["type"] in ignore_response_format_types: # value is a no-op return optional_params - json_schema: Optional[dict] = None - name: Optional[str] = None - description: Optional[str] = None + json_schema: dict | None = None + name: str | None = None + description: str | None = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: @@ -1089,42 +1085,36 @@ class AmazonConverseConfig(BaseConfig): @overload def _get_cache_point_block( self, - message_block: Union[ - OpenAIMessageContentListBlock, - ChatCompletionUserMessage, - ChatCompletionSystemMessage, - ChatCompletionAssistantMessage, - ], + message_block: OpenAIMessageContentListBlock + | ChatCompletionUserMessage + | ChatCompletionSystemMessage + | ChatCompletionAssistantMessage, block_type: Literal["system"], - model: Optional[str] = None, - ) -> Optional[SystemContentBlock]: + model: str | None = None, + ) -> SystemContentBlock | None: pass @overload def _get_cache_point_block( self, - message_block: Union[ - OpenAIMessageContentListBlock, - ChatCompletionUserMessage, - ChatCompletionSystemMessage, - ChatCompletionAssistantMessage, - ], + message_block: OpenAIMessageContentListBlock + | ChatCompletionUserMessage + | ChatCompletionSystemMessage + | ChatCompletionAssistantMessage, block_type: Literal["content_block"], - model: Optional[str] = None, - ) -> Optional[ContentBlock]: + model: str | None = None, + ) -> ContentBlock | None: pass def _get_cache_point_block( self, - message_block: Union[ - OpenAIMessageContentListBlock, - ChatCompletionUserMessage, - ChatCompletionSystemMessage, - ChatCompletionAssistantMessage, - ], + message_block: OpenAIMessageContentListBlock + | ChatCompletionUserMessage + | ChatCompletionSystemMessage + | ChatCompletionAssistantMessage, block_type: Literal["system", "content_block"], - model: Optional[str] = None, - ) -> Optional[Union[SystemContentBlock, ContentBlock]]: + model: str | None = None, + ) -> SystemContentBlock | ContentBlock | None: cache_control = message_block.get("cache_control", None) if cache_control is None: return None @@ -1137,7 +1127,7 @@ class AmazonConverseConfig(BaseConfig): return ContentBlock(cachePoint=cache_point) @staticmethod - def _build_cache_point_block(control: Optional[dict], model: Optional[str] = None) -> CachePointBlock: + def _build_cache_point_block(control: dict | None, model: str | None = None) -> CachePointBlock: """Build a Bedrock ``cachePoint`` block from an OpenAI-style ``cache_control``/``control`` dict. ``type`` is always ``"default"`` (the only value Bedrock's Converse API @@ -1152,10 +1142,10 @@ class AmazonConverseConfig(BaseConfig): return cache_point def _transform_system_message( - self, messages: List[AllMessageValues], model: Optional[str] = None - ) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]: + self, messages: list[AllMessageValues], model: str | None = None + ) -> tuple[list[AllMessageValues], list[SystemContentBlock]]: system_prompt_indices = [] - system_content_blocks: List[SystemContentBlock] = [] + system_content_blocks: list[SystemContentBlock] = [] for idx, message in enumerate(messages): if message["role"] == "system": system_prompt_indices.append(idx) @@ -1225,7 +1215,7 @@ class AmazonConverseConfig(BaseConfig): def _prepare_request_params( self, optional_params: dict, model: str, drop_params: bool = False - ) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]: + ) -> tuple[dict, dict, dict, OutputConfigBlock | None]: """Prepare and separate request parameters.""" # Consume the internal ``_output_config_normalized`` marker set by # ``_handle_reasoning_effort_parameter`` so it does not linger on the @@ -1260,7 +1250,7 @@ class AmazonConverseConfig(BaseConfig): if request_metadata is not None: self._validate_request_metadata(request_metadata) - output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None) + output_config: OutputConfigBlock | None = inference_params.pop("outputConfig", None) base_model = BedrockModelInfo.get_base_model(model) if ( output_config is None @@ -1338,11 +1328,11 @@ class AmazonConverseConfig(BaseConfig): self, original_tools: list, model: str, - headers: Optional[dict], + headers: dict | None, additional_request_params: dict, - ) -> Tuple[List[ToolBlock], list]: + ) -> tuple[list[ToolBlock], list]: """Process tools and collect anthropic_beta values.""" - bedrock_tools: List[ToolBlock] = [] + bedrock_tools: list[ToolBlock] = [] # Collect anthropic_beta values from user headers anthropic_beta_list = [] @@ -1353,7 +1343,7 @@ class AmazonConverseConfig(BaseConfig): # Separate pre-formatted Bedrock tools (e.g. systemTool from web_search_options) # from OpenAI-format tools that need transformation via _bedrock_tools_pt filtered_tools = [] - pre_formatted_tools: List[ToolBlock] = [] + pre_formatted_tools: list[ToolBlock] = [] if original_tools: for tool in original_tools: # Already-formatted Bedrock tools (e.g. systemTool for Nova grounding) @@ -1397,9 +1387,7 @@ class AmazonConverseConfig(BaseConfig): or "sonnet_4.6" in model_lower or "sonnet-4-6" in model_lower or "sonnet_4_6" in model_lower - ): - computer_use_header = "computer-use-2025-11-24" - elif ( + ) or ( "opus-4.5" in model_lower or "opus_4.5" in model_lower or "opus-4-5" in model_lower @@ -1510,10 +1498,10 @@ class AmazonConverseConfig(BaseConfig): def _transform_request_helper( self, model: str, - system_content_blocks: List[SystemContentBlock], + system_content_blocks: list[SystemContentBlock], optional_params: dict, - messages: Optional[List[AllMessageValues]] = None, - headers: Optional[dict] = None, + messages: list[AllMessageValues] | None = None, + headers: dict | None = None, drop_params: bool = False, ) -> CommonRequestObject: ## VALIDATE REQUEST @@ -1573,7 +1561,7 @@ class AmazonConverseConfig(BaseConfig): bedrock_tools.append(ToolBlock(cachePoint=cache_point)) break - bedrock_tool_config: Optional[ToolConfigBlock] = None + bedrock_tool_config: ToolConfigBlock | None = None if len(bedrock_tools) > 0: tool_choice_values: ToolChoiceValuesBlock = inference_params.pop("tool_choice", None) bedrock_tool_config = ToolConfigBlock( @@ -1612,10 +1600,10 @@ class AmazonConverseConfig(BaseConfig): async def _async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - headers: Optional[dict] = None, + headers: dict | None = None, ) -> RequestObject: messages, system_content_blocks = self._transform_system_message(messages, model=model) @@ -1646,7 +1634,7 @@ class AmazonConverseConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -1665,10 +1653,10 @@ class AmazonConverseConfig(BaseConfig): def _transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - headers: Optional[dict] = None, + headers: dict | None = None, ) -> RequestObject: messages, system_content_blocks = self._transform_system_message(messages, model=model) @@ -1685,7 +1673,7 @@ class AmazonConverseConfig(BaseConfig): ) ## TRANSFORMATION ## - bedrock_messages: List[MessageBlock] = _bedrock_converse_messages_pt( + bedrock_messages: list[MessageBlock] = _bedrock_converse_messages_pt( messages=messages, model=model, llm_provider="bedrock_converse", @@ -1703,12 +1691,12 @@ class AmazonConverseConfig(BaseConfig): model_response: ModelResponse, logging_obj: Logging, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: return self._transform_response( model=model, @@ -1723,7 +1711,7 @@ class AmazonConverseConfig(BaseConfig): encoding=encoding, ) - def _transform_reasoning_content(self, reasoning_content_blocks: List[BedrockConverseReasoningContentBlock]) -> str: + def _transform_reasoning_content(self, reasoning_content_blocks: list[BedrockConverseReasoningContentBlock]) -> str: """ Extract the reasoning text from the reasoning content blocks @@ -1736,10 +1724,10 @@ class AmazonConverseConfig(BaseConfig): return reasoning_content_str def _transform_thinking_blocks( - self, thinking_blocks: List[BedrockConverseReasoningContentBlock] - ) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]: + self, thinking_blocks: list[BedrockConverseReasoningContentBlock] + ) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]: """Return a consistent format for thinking blocks between Anthropic and Bedrock.""" - thinking_blocks_list: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = [] + thinking_blocks_list: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = [] for block in thinking_blocks: if "reasoningText" in block: _thinking_block = ChatCompletionThinkingBlock(type="thinking") @@ -1760,7 +1748,7 @@ class AmazonConverseConfig(BaseConfig): def _transform_usage( self, usage: ConverseTokenUsageBlock, - reasoning_content: Optional[str] = None, + reasoning_content: str | None = None, ) -> Usage: input_tokens = usage["inputTokens"] output_tokens = usage["outputTokens"] @@ -1799,8 +1787,8 @@ class AmazonConverseConfig(BaseConfig): def get_tool_call_names( self, - tools: Optional[Union[List[ToolBlock], List[OpenAIChatCompletionToolParam]]] = None, - ) -> List[str]: + tools: list[ToolBlock] | list[OpenAIChatCompletionToolParam] | None = None, + ) -> list[str]: if tools is None: return [] tool_set: set[str] = set() @@ -1820,9 +1808,9 @@ class AmazonConverseConfig(BaseConfig): def apply_tool_call_transformation_if_needed( self, message: Message, - tools: Optional[List[ToolBlock]] = None, - initial_finish_reason: Optional[str] = None, - ) -> Tuple[Message, Optional[str]]: + tools: list[ToolBlock] | None = None, + initial_finish_reason: str | None = None, + ) -> tuple[Message, str | None]: """ Apply tool call transformation to a message. @@ -1850,12 +1838,12 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason def _translate_message_content( - self, content_blocks: List[ContentBlock] - ) -> Tuple[ + self, content_blocks: list[ContentBlock] + ) -> tuple[ str, - List[ChatCompletionToolCallChunk], - Optional[List[BedrockConverseReasoningContentBlock]], - Optional[List[CitationsContentBlock]], + list[ChatCompletionToolCallChunk], + list[BedrockConverseReasoningContentBlock] | None, + list[CitationsContentBlock] | None, ]: """ Translate the message content to a string and a list of tool calls, reasoning content blocks, and citations. @@ -1867,14 +1855,14 @@ class AmazonConverseConfig(BaseConfig): citationsContentBlocks: Optional[List[CitationsContentBlock]] - Citations from Nova grounding """ content_str = "" - tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None - citationsContentBlocks: Optional[List[CitationsContentBlock]] = None + tools: list[ChatCompletionToolCallChunk] = [] + reasoningContentBlocks: list[BedrockConverseReasoningContentBlock] | None = None + citationsContentBlocks: list[CitationsContentBlock] | None = None for idx, content in enumerate(content_blocks): """ - Content is either a tool response or text """ - extracted_reasoning_content_str: Optional[str] = None + extracted_reasoning_content_str: str | None = None if "text" in content: ( extracted_reasoning_content_str, @@ -1922,8 +1910,8 @@ class AmazonConverseConfig(BaseConfig): @staticmethod def _transform_citations_to_annotations( - citations_content_blocks: Optional[List[CitationsContentBlock]], - ) -> Tuple[Optional[str], Optional[List[ChatCompletionAnnotation]]]: + citations_content_blocks: list[CitationsContentBlock] | None, + ) -> tuple[str | None, list[ChatCompletionAnnotation] | None]: """ Convert Bedrock citationsContent blocks into OpenAI-style annotations. @@ -1934,8 +1922,8 @@ class AmazonConverseConfig(BaseConfig): if not citations_content_blocks: return None, None - annotations: List[ChatCompletionAnnotation] = [] - citations_text_parts: List[str] = [] + annotations: list[ChatCompletionAnnotation] = [] + citations_text_parts: list[str] = [] content_offset = 0 for citations_block in citations_content_blocks: @@ -2014,10 +2002,10 @@ class AmazonConverseConfig(BaseConfig): @staticmethod def _filter_json_mode_tools( - json_mode: Optional[bool], - tools: List[ChatCompletionToolCallChunk], + json_mode: bool | None, + tools: list[ChatCompletionToolCallChunk], chat_completion_message: ChatCompletionResponseMessage, - ) -> Optional[List[ChatCompletionToolCallChunk]]: + ) -> list[ChatCompletionToolCallChunk] | None: """ When json_mode is True, Bedrock may return the internal `json_tool_call` tool alongside real user-defined tools. This method handles 3 scenarios: @@ -2038,7 +2026,7 @@ class AmazonConverseConfig(BaseConfig): if len(json_tool_indices) == len(tools): # All tools are json_tool_call — convert first one to content verbose_logger.debug("Processing JSON tool call response for response_format") - json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments") + json_mode_content_str: str | None = tools[0]["function"].get("arguments") if json_mode_content_str is not None: json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties(json_mode_content_str) chat_completion_message["content"] = json_mode_content_str @@ -2063,11 +2051,11 @@ class AmazonConverseConfig(BaseConfig): response: httpx.Response, model_response: ModelResponse, stream: bool, - logging_obj: Optional[Logging], + logging_obj: Logging | None, optional_params: dict, - api_key: Optional[str], - data: Union[dict, str], - messages: List, + api_key: str | None, + data: dict | str, + messages: list, encoding, ) -> ModelResponse: ## LOGGING @@ -2079,15 +2067,13 @@ class AmazonConverseConfig(BaseConfig): additional_args={"complete_input_dict": data}, ) - json_mode: Optional[bool] = optional_params.get("json_mode", None) + json_mode: bool | None = optional_params.get("json_mode", None) ## RESPONSE OBJECT try: completion_response = ConverseResponseBlock(**response.json()) # type: ignore except Exception as e: raise BedrockError( - message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format( - str(e) - ), + message=f"Error converting to valid response block={e!s}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues", status_code=422, ) @@ -2126,12 +2112,12 @@ class AmazonConverseConfig(BaseConfig): } """ - message: Optional[MessageBlock] = completion_response["output"]["message"] + message: MessageBlock | None = completion_response["output"]["message"] chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" - tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None - citationsContentBlocks: Optional[List[CitationsContentBlock]] = None + tools: list[ChatCompletionToolCallChunk] = [] + reasoningContentBlocks: list[BedrockConverseReasoningContentBlock] | None = None + citationsContentBlocks: list[CitationsContentBlock] | None = None if message is not None: ( @@ -2228,9 +2214,7 @@ class AmazonConverseConfig(BaseConfig): return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BedrockError( message=error_message, status_code=status_code, @@ -2241,11 +2225,11 @@ class AmazonConverseConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key: headers["Authorization"] = f"Bearer {api_key}" @@ -2253,10 +2237,10 @@ class AmazonConverseConfig(BaseConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, - fake_stream: Optional[bool] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, + fake_stream: bool | None = None, ) -> bool: """ Returns True if the model/provider should fake stream diff --git a/litellm/llms/bedrock/chat/invoke_agent/transformation.py b/litellm/llms/bedrock/chat/invoke_agent/transformation.py index 413cdad45e0..d877ca81244 100644 --- a/litellm/llms/bedrock/chat/invoke_agent/transformation.py +++ b/litellm/llms/bedrock/chat/invoke_agent/transformation.py @@ -6,16 +6,16 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agent-runtime_Invoke import base64 import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from litellm._logging import verbose_logger from litellm._uuid import uuid -from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.bedrock.common_utils import BedrockError @@ -49,7 +49,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): BaseConfig.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ This is a base invoke agent model mapping. For Invoke Agent - define a bedrock provider specific config that extends this class. @@ -73,12 +73,12 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -110,11 +110,11 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: return self._sign_request( service_name="bedrock", headers=headers, @@ -147,7 +147,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -222,7 +222,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): return events - def _parse_message_from_event(self, event, parser) -> Optional[str]: + def _parse_message_from_event(self, event, parser) -> str | None: """Extract message content from an AWS event, adapted from AWSEventStreamDecoder.""" try: response_dict = event.to_response_dict() @@ -314,7 +314,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): model=None, ) - response_model: Optional[str] = None + response_model: str | None = None for event in events: if not self._is_trace_event(event): @@ -346,7 +346,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): payload = event.get("payload") return event_type == "trace" and payload is not None - def _get_trace_data(self, event: InvokeAgentEvent) -> Optional[InvokeAgentTrace]: + def _get_trace_data(self, event: InvokeAgentEvent) -> InvokeAgentTrace | None: """Extract trace data from a trace event.""" payload = event.get("payload") if not payload: @@ -359,34 +359,34 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): self, trace_data: InvokeAgentTrace, usage_info: InvokeAgentUsage ) -> None: """Extract usage information from preprocessing trace.""" - pre_processing: Optional[InvokeAgentPreProcessingTrace] = trace_data.get("preProcessingTrace") + pre_processing: InvokeAgentPreProcessingTrace | None = trace_data.get("preProcessingTrace") if not pre_processing: return - model_output: Optional[InvokeAgentModelInvocationOutput] = ( + model_output: InvokeAgentModelInvocationOutput | None = ( pre_processing.get("modelInvocationOutput") or InvokeAgentModelInvocationOutput() ) if not model_output: return - metadata: Optional[InvokeAgentMetadata] = model_output.get("metadata") or InvokeAgentMetadata() + metadata: InvokeAgentMetadata | None = model_output.get("metadata") or InvokeAgentMetadata() if not metadata: return - usage: Optional[Union[InvokeAgentUsage, Dict]] = metadata.get("usage", {}) + usage: InvokeAgentUsage | dict | None = metadata.get("usage", {}) if not usage: return usage_info["inputTokens"] += usage.get("inputTokens", 0) usage_info["outputTokens"] += usage.get("outputTokens", 0) - def _extract_orchestration_model(self, trace_data: InvokeAgentTrace) -> Optional[str]: + def _extract_orchestration_model(self, trace_data: InvokeAgentTrace) -> str | None: """Extract model information from orchestration trace.""" - orchestration_trace: Optional[InvokeAgentOrchestrationTrace] = trace_data.get("orchestrationTrace") + orchestration_trace: InvokeAgentOrchestrationTrace | None = trace_data.get("orchestrationTrace") if not orchestration_trace: return None - model_invocation: Optional[InvokeAgentModelInvocationInput] = ( + model_invocation: InvokeAgentModelInvocationInput | None = ( orchestration_trace.get("modelInvocationInput") or InvokeAgentModelInvocationInput() ) if not model_invocation: @@ -433,12 +433,12 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: # Get the raw binary content @@ -464,9 +464,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): ) except Exception as e: - verbose_logger.error(f"Error processing Bedrock Invoke Agent response: {str(e)}") + verbose_logger.error(f"Error processing Bedrock Invoke Agent response: {e!s}") raise BedrockError( - message=f"Error processing response: {str(e)}", + message=f"Error processing response: {e!s}", status_code=raw_response.status_code, ) @@ -474,23 +474,21 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BedrockError(status_code=status_code, message=error_message) def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: return True diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 9a2576b9dd4..d069929df92 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1,8 +1,6 @@ import types from collections.abc import AsyncIterator, Iterator from typing import ( - Optional, - Tuple, cast, ) @@ -61,37 +59,37 @@ class AmazonCohereChatConfig: Reference - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-command-r-plus.html """ - documents: Optional[List[Document]] = None - search_queries_only: Optional[bool] = None - preamble: Optional[str] = None - max_tokens: Optional[int] = None - temperature: Optional[float] = None - p: Optional[float] = None - k: Optional[float] = None - prompt_truncation: Optional[str] = None - frequency_penalty: Optional[float] = None - presence_penalty: Optional[float] = None - seed: Optional[int] = None - return_prompt: Optional[bool] = None - stop_sequences: Optional[List[str]] = None - raw_prompting: Optional[bool] = None + documents: List[Document] | None = None + search_queries_only: bool | None = None + preamble: str | None = None + max_tokens: int | None = None + temperature: float | None = None + p: float | None = None + k: float | None = None + prompt_truncation: str | None = None + frequency_penalty: float | None = None + presence_penalty: float | None = None + seed: int | None = None + return_prompt: bool | None = None + stop_sequences: List[str] | None = None + raw_prompting: bool | None = None def __init__( self, - documents: Optional[List[Document]] = None, - search_queries_only: Optional[bool] = None, - preamble: Optional[str] = None, - max_tokens: Optional[int] = None, - temperature: Optional[float] = None, - p: Optional[float] = None, - k: Optional[float] = None, - prompt_truncation: Optional[str] = None, - frequency_penalty: Optional[float] = None, - presence_penalty: Optional[float] = None, - seed: Optional[int] = None, - return_prompt: Optional[bool] = None, - stop_sequences: Optional[str] = None, - raw_prompting: Optional[bool] = None, + documents: List[Document] | None = None, + search_queries_only: bool | None = None, + preamble: str | None = None, + max_tokens: int | None = None, + temperature: float | None = None, + p: float | None = None, + k: float | None = None, + prompt_truncation: str | None = None, + frequency_penalty: float | None = None, + presence_penalty: float | None = None, + seed: int | None = None, + return_prompt: bool | None = None, + stop_sequences: str | None = None, + raw_prompting: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -156,7 +154,7 @@ class AmazonCohereChatConfig: async def make_call( - client: Optional[AsyncHTTPHandler], + client: AsyncHTTPHandler | None, api_base: str, headers: dict, data: str, @@ -164,9 +162,9 @@ async def make_call( messages: list, logging_obj: Logging, fake_stream: bool = False, - json_mode: Optional[bool] = False, - bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, - stream_chunk_size: Optional[int] = None, + json_mode: bool | None = False, + bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None, + stream_chunk_size: int | None = None, ): try: if client is None: @@ -240,18 +238,18 @@ async def make_call( def make_sync_call( - client: Optional[HTTPHandler], + client: HTTPHandler | None, api_base: str, headers: dict, data: str, - signed_json_body: Optional[bytes], + signed_json_body: bytes | None, model: str, messages: list, logging_obj: Logging, fake_stream: bool = False, - json_mode: Optional[bool] = False, - bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, - stream_chunk_size: Optional[int] = None, + json_mode: bool | None = False, + bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None, + stream_chunk_size: int | None = None, ): try: if client is None: @@ -324,16 +322,16 @@ def make_sync_call( class AWSEventStreamDecoder: - def __init__(self, model: str, json_mode: Optional[bool] = False) -> None: + def __init__(self, model: str, json_mode: bool | None = False) -> None: from botocore.parsers import EventStreamJSONParser self.model = model self.parser = EventStreamJSONParser() self.content_blocks: List[ContentBlockDeltaEvent] = [] - self.tool_calls_index: Optional[int] = None - self.response_id: Optional[str] = None + self.tool_calls_index: int | None = None + self.response_id: str | None = None self.json_mode = json_mode - self._current_tool_name: Optional[str] = None + self._current_tool_name: str | None = None def check_empty_tool_call_args(self) -> bool: """ @@ -359,20 +357,20 @@ class AWSEventStreamDecoder: def extract_reasoning_content_str( self, reasoning_content_block: BedrockConverseReasoningContentBlockDelta - ) -> Optional[str]: + ) -> str | None: if "text" in reasoning_content_block: return reasoning_content_block["text"] return None def translate_thinking_blocks( self, thinking_block: BedrockConverseReasoningContentBlockDelta - ) -> Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]]: + ) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None: """ Translate the thinking blocks to a string """ thinking_blocks_list: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = [] - _thinking_block: Optional[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = None + _thinking_block: Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] | None = None if "text" in thinking_block: _thinking_block = ChatCompletionThinkingBlock(type="thinking") @@ -403,15 +401,15 @@ class AWSEventStreamDecoder: def _handle_converse_start_event( self, start_obj: ContentBlockStartEvent, - ) -> Tuple[ - Optional[ChatCompletionToolCallChunk], + ) -> tuple[ + ChatCompletionToolCallChunk | None, dict, - Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], + List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None, ]: """Handle 'start' event in converse chunk parsing.""" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None provider_specific_fields: dict = {} - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None + thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None self.content_blocks = [] # reset if start_obj is not None: @@ -449,19 +447,19 @@ class AWSEventStreamDecoder: self, delta_obj: ContentBlockDeltaEvent, index: int, - ) -> Tuple[ + ) -> tuple[ str, - Optional[ChatCompletionToolCallChunk], + ChatCompletionToolCallChunk | None, dict, - Optional[str], - Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], + str | None, + List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None, ]: """Handle 'delta' event in converse chunk parsing.""" text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None provider_specific_fields: dict = {} - reasoning_content: Optional[str] = None - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None + reasoning_content: str | None = None + thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None self.content_blocks.append(delta_obj) if "text" in delta_obj: @@ -502,9 +500,9 @@ class AWSEventStreamDecoder: thinking_blocks, ) - def _handle_converse_stop_event(self, index: int) -> Optional[ChatCompletionToolCallChunk]: + def _handle_converse_stop_event(self, index: int) -> ChatCompletionToolCallChunk | None: """Handle stop/contentBlockIndex event in converse chunk parsing.""" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None # If the ending block was the internal json_tool_call, skip emitting # the empty-args tool chunk and reset tracking state @@ -532,16 +530,14 @@ class AWSEventStreamDecoder: # and use it as the consistent ID for all subsequent chunks. self._initialize_converse_response_id(chunk_data) - verbose_logger.debug("\n\nRaw Chunk: {}\n\n".format(chunk_data)) + verbose_logger.debug(f"\n\nRaw Chunk: {chunk_data}\n\n") text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None finish_reason = "" - usage: Optional[Usage] = None + usage: Usage | None = None provider_specific_fields: dict = {} - reasoning_content: Optional[str] = None - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = ( - None - ) + reasoning_content: str | None = None + thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None content_block_index = int(chunk_data.get("contentBlockIndex", 0)) if "start" in chunk_data: @@ -594,7 +590,7 @@ class AWSEventStreamDecoder: return response except Exception as e: - raise Exception("Received streaming error - {}".format(str(e))) + raise Exception(f"Received streaming error - {e!s}") def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]: text = "" @@ -680,7 +676,7 @@ class AWSEventStreamDecoder: _data = json.loads(message) yield self._chunk_parser(chunk_data=_data) - def _parse_message_from_event(self, event) -> Optional[str]: + def _parse_message_from_event(self, event) -> str | None: response_stream_shape = get_bedrock_response_stream_shape() if response_stream_shape is None: raise BedrockError( @@ -713,7 +709,7 @@ class AmazonAnthropicClaudeStreamDecoder(AWSEventStreamDecoder): self, model: str, sync_stream: bool, - json_mode: Optional[bool] = None, + json_mode: bool | None = None, ) -> None: """ Child class of AWSEventStreamDecoder that handles the streaming response from the Anthropic family of models @@ -752,7 +748,7 @@ class AmazonDeepSeekR1StreamDecoder(AWSEventStreamDecoder): class MockResponseIterator: # for returning ai21 streaming responses - def __init__(self, model_response, json_mode: Optional[bool] = False): + def __init__(self, model_response, json_mode: bool | None = False): self.model_response = model_response self.json_mode = json_mode self.is_done = False @@ -762,8 +758,8 @@ class MockResponseIterator: # for returning ai21 streaming responses return self def _handle_json_mode_chunk( - self, text: str, tool_calls: Optional[List[ChatCompletionToolCallChunk]] - ) -> Tuple[str, Optional[ChatCompletionToolCallChunk]]: + self, text: str, tool_calls: List[ChatCompletionToolCallChunk] | None + ) -> tuple[str, ChatCompletionToolCallChunk | None]: """ If JSON mode is enabled, convert the tool call to a message. @@ -779,7 +775,7 @@ class MockResponseIterator: # for returning ai21 streaming responses text: The text to use in the content tool_use: The ChatCompletionToolCallChunk to use in the chunk response """ - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None if self.json_mode is True and tool_calls is not None: message = litellm.AnthropicConfig()._convert_tool_response_to_message(tool_calls=tool_calls) if message is not None: @@ -795,7 +791,7 @@ class MockResponseIterator: # for returning ai21 streaming responses text = chunk_data.choices[0].message.content or "" # type: ignore tool_use = None _model_response_tool_call = cast( - Optional[List[ChatCompletionMessageToolCall]], + List[ChatCompletionMessageToolCall] | None, cast(Choices, chunk_data.choices[0]).message.tool_calls, ) if self.json_mode is True: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py index 50fa6f170b3..5179940cd4b 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_ai21_transformation.py @@ -1,5 +1,4 @@ import types -from typing import List, Optional from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( @@ -28,23 +27,23 @@ class AmazonAI21Config(AmazonInvokeConfig, BaseConfig): - `countPenalty` (object): Placeholder for count penalty object. """ - maxTokens: Optional[int] = None - temperature: Optional[float] = None - topP: Optional[float] = None - stopSequences: Optional[list] = None - frequencePenalty: Optional[dict] = None - presencePenalty: Optional[dict] = None - countPenalty: Optional[dict] = None + maxTokens: int | None = None + temperature: float | None = None + topP: float | None = None + stopSequences: list | None = None + frequencePenalty: dict | None = None + presencePenalty: dict | None = None + countPenalty: dict | None = None def __init__( self, - maxTokens: Optional[int] = None, - temperature: Optional[float] = None, - topP: Optional[float] = None, - stopSequences: Optional[list] = None, - frequencePenalty: Optional[dict] = None, - presencePenalty: Optional[dict] = None, - countPenalty: Optional[dict] = None, + maxTokens: int | None = None, + temperature: float | None = None, + topP: float | None = None, + stopSequences: list | None = None, + frequencePenalty: dict | None = None, + presencePenalty: dict | None = None, + countPenalty: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -72,7 +71,7 @@ class AmazonAI21Config(AmazonInvokeConfig, BaseConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "max_tokens", "temperature", diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py index 8b411b7b576..04f10fae984 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_cohere_transformation.py @@ -1,5 +1,4 @@ import types -from typing import List, Optional from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, @@ -18,14 +17,14 @@ class AmazonCohereConfig(AmazonInvokeConfig, CohereChatConfig): - `return_likelihood` (string) n/a """ - max_tokens: Optional[int] = None - return_likelihood: Optional[str] = None + max_tokens: int | None = None + return_likelihood: str | None = None def __init__( self, - max_tokens: Optional[int] = None, - temperature: Optional[float] = None, - return_likelihood: Optional[str] = None, + max_tokens: int | None = None, + temperature: float | None = None, + return_likelihood: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -53,7 +52,7 @@ class AmazonCohereConfig(AmazonInvokeConfig, CohereChatConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: supported_params = CohereChatConfig.get_supported_openai_params(self, model=model) return supported_params diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py index d3025e13a99..491e6a4839c 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py @@ -1,4 +1,4 @@ -from typing import Any, List, Optional, cast +from typing import Any, cast from httpx import Response @@ -33,12 +33,12 @@ class AmazonDeepSeekR1Config(AmazonLlamaConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Extract the reasoning content, and return it as a separate field in the response. @@ -56,8 +56,8 @@ class AmazonDeepSeekR1Config(AmazonLlamaConfig): api_key, json_mode, ) - prompt = cast(Optional[str], request_data.get("prompt")) - message_content = cast(Optional[str], cast(Choices, response.choices[0]).message.get("content")) + prompt = cast(str | None, request_data.get("prompt")) + message_content = cast(str | None, cast(Choices, response.choices[0]).message.get("content")) if prompt and prompt.strip().endswith("") and message_content: message_content_with_reasoning_token = "" + message_content reasoning, content = _parse_content_for_reasoning(message_content_with_reasoning_token) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py index 9f84844fcb6..389a2633eb9 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_llama_transformation.py @@ -1,5 +1,4 @@ import types -from typing import List, Optional from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( @@ -18,15 +17,15 @@ class AmazonLlamaConfig(AmazonInvokeConfig, BaseConfig): - `top_p` (float) top p for model """ - max_gen_len: Optional[int] = None - temperature: Optional[float] = None - topP: Optional[float] = None + max_gen_len: int | None = None + temperature: float | None = None + topP: float | None = None def __init__( self, - maxTokenCount: Optional[int] = None, - temperature: Optional[float] = None, - topP: Optional[int] = None, + maxTokenCount: int | None = None, + temperature: float | None = None, + topP: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -53,7 +52,7 @@ class AmazonLlamaConfig(AmazonInvokeConfig, BaseConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "max_tokens", "temperature", diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py index 58dfa17a722..d48abe1c395 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_mistral_transformation.py @@ -1,5 +1,5 @@ import types -from typing import List, Optional, TYPE_CHECKING +from typing import TYPE_CHECKING from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( @@ -23,19 +23,19 @@ class AmazonMistralConfig(AmazonInvokeConfig, BaseConfig): - `top_k` (float) top k for model """ - max_tokens: Optional[int] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - top_k: Optional[float] = None - stop: Optional[List[str]] = None + max_tokens: int | None = None + temperature: float | None = None + top_p: float | None = None + top_k: float | None = None + stop: list[str] | None = None def __init__( self, - max_tokens: Optional[int] = None, - temperature: Optional[float] = None, - top_p: Optional[int] = None, - top_k: Optional[float] = None, - stop: Optional[List[str]] = None, + max_tokens: int | None = None, + temperature: float | None = None, + top_p: int | None = None, + top_k: float | None = None, + stop: list[str] | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -63,7 +63,7 @@ class AmazonMistralConfig(AmazonInvokeConfig, BaseConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["max_tokens", "temperature", "top_p", "stop", "stream"] def map_openai_params( diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py index 0532d677e5a..0ef52b1e1c3 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py @@ -7,8 +7,8 @@ Model format: bedrock/moonshot.kimi-k2-thinking-v1:0 Reference: https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/ """ -from typing import TYPE_CHECKING, Any, List, Optional, Union import re +from typing import TYPE_CHECKING, Any import httpx @@ -56,7 +56,7 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): MoonshotChatConfig.__init__(self, **kwargs) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock" def _get_model_id(self, model: str) -> str: @@ -69,12 +69,10 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): - moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking """ # Remove bedrock/ prefix if present - if model.startswith("bedrock/"): - model = model[8:] + model = model.removeprefix("bedrock/") # Remove invoke/ prefix if present - if model.startswith("invoke/"): - model = model[7:] + model = model.removeprefix("invoke/") # Remove any provider prefix (e.g., moonshot/) if "/" in model and not model.startswith("arn:"): @@ -84,7 +82,7 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): return model - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Get the supported OpenAI params for Moonshot AI models on Bedrock. @@ -96,13 +94,13 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): Note: kimi-k2-thinking DOES support tool calls (unlike kimi-thinking-preview) The parent MoonshotChatConfig class handles the kimi-thinking-preview exclusion. """ - excluded_params: List[str] = [ + excluded_params: list[str] = [ "functions", "stop", ] # Bedrock doesn't support stopSequences base_openai_params = super(MoonshotChatConfig, self).get_supported_openai_params(model=model) - final_params: List[str] = [] + final_params: list[str] = [] for param in base_openai_params: if param not in excluded_params: final_params.append(param) @@ -135,7 +133,7 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -166,7 +164,7 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): headers=headers, ) - def _extract_reasoning_from_content(self, content: str) -> tuple[Optional[str], str]: + def _extract_reasoning_from_content(self, content: str) -> tuple[str | None, str]: """ Extract reasoning content from tags in the response. @@ -199,12 +197,12 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): model_response: "ModelResponse", logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> "ModelResponse": """ Transform the response from Bedrock Moonshot AI models. @@ -249,8 +247,6 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BedrockError: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BedrockError: """Return the appropriate error class for Bedrock.""" return BedrockError(status_code=status_code, message=error_message) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_nova_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_nova_transformation.py index acfa5021507..522a473407c 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_nova_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_nova_transformation.py @@ -6,7 +6,7 @@ Inherits from `AmazonConverseConfig` Nova + Invoke API Tutorial: https://docs.aws.amazon.com/nova/latest/userguide/using-invoke-api.html """ -from typing import Any, List, Optional +from typing import Any import httpx @@ -42,7 +42,7 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -67,12 +67,12 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig): model_response: ModelResponse, logging_obj: Logging, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: return AmazonConverseConfig.transform_response( self, @@ -105,4 +105,3 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig): _system_message = bedrock_invoke_nova_request.get("system", None) if isinstance(_system_message, list) and len(_system_message) == 0: bedrock_invoke_nova_request.pop("system", None) - return diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_openai_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_openai_transformation.py index d3f9d8bffb8..8132fdefbc6 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_openai_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_openai_transformation.py @@ -7,7 +7,7 @@ Model format: bedrock/openai/ Example: bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123 """ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -45,7 +45,7 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): BaseAWSLLM.__init__(self, **kwargs) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock" def _get_openai_model_id(self, model: str) -> str: @@ -56,23 +56,21 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): Returns: """ # Remove bedrock/ prefix if present - if model.startswith("bedrock/"): - model = model[8:] + model = model.removeprefix("bedrock/") # Remove openai/ prefix - if model.startswith("openai/"): - model = model[7:] + model = model.removeprefix("openai/") return model def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the Bedrock invoke endpoint. @@ -109,11 +107,11 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: """ Sign the request using AWS Signature Version 4. """ @@ -132,7 +130,7 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -162,11 +160,11 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate the environment and return headers. @@ -175,8 +173,6 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM): """ return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BedrockError: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BedrockError: """Return the appropriate error class for Bedrock.""" return BedrockError(status_code=status_code, message=error_message) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py index a2aa98d6676..7270d987095 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py @@ -7,7 +7,7 @@ The main difference is in the response format: Qwen2 uses "text" field while Qwe Qwen2 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html """ -from typing import Any, List, Optional +from typing import Any import httpx @@ -38,12 +38,12 @@ class AmazonQwen2Config(AmazonQwen3Config): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform Qwen2 Bedrock response to OpenAI format @@ -57,10 +57,8 @@ class AmazonQwen2Config(AmazonQwen3Config): generated_text = response_data.get("generation", "") or response_data.get("text", "") # Clean up the response (remove assistant start token if present) - if generated_text.startswith("<|im_start|>assistant\n"): - generated_text = generated_text[len("<|im_start|>assistant\n") :] - if generated_text.endswith("<|im_end|>"): - generated_text = generated_text[: -len("<|im_end|>")] + generated_text = generated_text.removeprefix("<|im_start|>assistant\n") + generated_text = generated_text.removesuffix("<|im_end|>") # Set the content in the existing model_response structure if hasattr(model_response, "choices") and len(model_response.choices) > 0: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py index 4f496df084e..c4e2bfc93f6 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py @@ -6,7 +6,7 @@ Inherits from `AmazonInvokeConfig` Qwen3 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html """ -from typing import Any, List, Optional +from typing import Any import httpx @@ -26,19 +26,19 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html """ - max_tokens: Optional[int] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - top_k: Optional[int] = None - stop: Optional[List[str]] = None + max_tokens: int | None = None + temperature: float | None = None + top_p: float | None = None + top_k: int | None = None + stop: list[str] | None = None def __init__( self, - max_tokens: Optional[int] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - stop: Optional[List[str]] = None, + max_tokens: int | None = None, + temperature: float | None = None, + top_p: float | None = None, + top_k: int | None = None, + stop: list[str] | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -46,7 +46,7 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): setattr(self.__class__, key, value) AmazonInvokeConfig.__init__(self) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "max_tokens", "temperature", @@ -81,7 +81,7 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -111,7 +111,7 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): return request_body - def _convert_messages_to_prompt(self, messages: List[AllMessageValues]) -> str: + def _convert_messages_to_prompt(self, messages: list[AllMessageValues]) -> str: """ Convert OpenAI messages format to Qwen3 prompt format Supports tool calls, multimodal content, and various message types @@ -164,12 +164,12 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform Qwen3 Bedrock response to OpenAI format @@ -181,10 +181,8 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): generated_text = response_data.get("generation", "") # Clean up the response (remove assistant start token if present) - if generated_text.startswith("<|im_start|>assistant\n"): - generated_text = generated_text[len("<|im_start|>assistant\n") :] - if generated_text.endswith("<|im_end|>"): - generated_text = generated_text[: -len("<|im_end|>")] + generated_text = generated_text.removeprefix("<|im_start|>assistant\n") + generated_text = generated_text.removesuffix("<|im_end|>") # Set the content in the existing model_response structure if hasattr(model_response, "choices") and len(model_response.choices) > 0: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py index ff9a2ee0c6d..585596dc4cf 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_titan_transformation.py @@ -1,6 +1,5 @@ import re import types -from typing import List, Optional, Union import litellm from litellm.llms.base_llm.chat.transformation import BaseConfig @@ -21,17 +20,17 @@ class AmazonTitanConfig(AmazonInvokeConfig, BaseConfig): - `topP` (int) top p for model """ - maxTokenCount: Optional[int] = None - stopSequences: Optional[list] = None - temperature: Optional[float] = None - topP: Optional[int] = None + maxTokenCount: int | None = None + stopSequences: list | None = None + temperature: float | None = None + topP: int | None = None def __init__( self, - maxTokenCount: Optional[int] = None, - stopSequences: Optional[list] = None, - temperature: Optional[float] = None, - topP: Optional[int] = None, + maxTokenCount: int | None = None, + stopSequences: list | None = None, + temperature: float | None = None, + topP: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -64,7 +63,7 @@ class AmazonTitanConfig(AmazonInvokeConfig, BaseConfig): supported_params: dict, provider: str, model: str, - stop: Union[List[str], str], + stop: list[str] | str, ): """ filter params to fit the required provider format, drop those that don't fit if user sets `litellm.drop_params = True`. @@ -82,7 +81,7 @@ class AmazonTitanConfig(AmazonInvokeConfig, BaseConfig): return supported_params - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "max_tokens", "max_completion_tokens", diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_twelvelabs_pegasus_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_twelvelabs_pegasus_transformation.py index 6d25bb32309..b96756f1e4e 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_twelvelabs_pegasus_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_twelvelabs_pegasus_transformation.py @@ -7,7 +7,7 @@ https://docs.twelvelabs.io/docs/models/pegasus import json import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -42,7 +42,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): response_format, max_tokens) are translated to the TwelveLabs schema. """ - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "max_tokens", "max_completion_tokens", @@ -102,13 +102,13 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, ) -> dict: input_prompt = self._convert_messages_to_prompt(messages=messages) - request_data: Dict[str, Any] = {"inputPrompt": input_prompt} + request_data: dict[str, Any] = {"inputPrompt": input_prompt} media_source = self._build_media_source(optional_params) if media_source is not None: @@ -128,7 +128,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): return request_data - def _build_media_source(self, optional_params: dict) -> Optional[dict]: + def _build_media_source(self, optional_params: dict) -> dict | None: direct_source = optional_params.get("mediaSource") or optional_params.get("media_source") if isinstance(direct_source, dict): return direct_source @@ -154,8 +154,8 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): return {"s3Location": s3_location} return None - def _convert_messages_to_prompt(self, messages: List[AllMessageValues]) -> str: - prompt_parts: List[str] = [] + def _convert_messages_to_prompt(self, messages: list[AllMessageValues]) -> str: + prompt_parts: list[str] = [] for message in messages: role = message.get("role", "user") content = message.get("content", "") @@ -185,12 +185,12 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform TwelveLabs Pegasus response to LiteLLM format. @@ -208,7 +208,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): completion_response = raw_response.json() except Exception as e: raise BedrockError( - message=f"Error parsing response: {raw_response.text}, error: {str(e)}", + message=f"Error parsing response: {raw_response.text}, error: {e!s}", status_code=raw_response.status_code, ) @@ -237,7 +237,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig): raise Exception("Unable to set message content") except Exception as e: raise BedrockError( - message=f"Error setting response content: {str(e)}. Response: {completion_response}", + message=f"Error setting response content: {e!s}. Response: {completion_response}", status_code=raw_response.status_code, ) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py index 9cc6195cfbb..b516571b111 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py @@ -1,5 +1,4 @@ import types -from typing import Optional import litellm @@ -20,21 +19,21 @@ class AmazonAnthropicConfig(AmazonInvokeConfig): - `anthropic_version` (string) version of anthropic for bedrock - e.g. "bedrock-2023-05-31" """ - max_tokens_to_sample: Optional[int] = litellm.max_tokens - stop_sequences: Optional[list] = None - temperature: Optional[float] = None - top_k: Optional[int] = None - top_p: Optional[int] = None - anthropic_version: Optional[str] = None + max_tokens_to_sample: int | None = litellm.max_tokens + stop_sequences: list | None = None + temperature: float | None = None + top_k: int | None = None + top_p: int | None = None + anthropic_version: str | None = None def __init__( self, - max_tokens_to_sample: Optional[int] = None, - stop_sequences: Optional[list] = None, - temperature: Optional[float] = None, - top_k: Optional[int] = None, - top_p: Optional[int] = None, - anthropic_version: Optional[str] = None, + max_tokens_to_sample: int | None = None, + stop_sequences: list | None = None, + temperature: float | None = None, + top_k: int | None = None, + top_p: int | None = None, + anthropic_version: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index 6b5cb304bec..b4061227143 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -57,13 +57,13 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): anthropic_version: str = "bedrock-2023-05-31" @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock" def should_strip_billing_metadata(self) -> bool: return True - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return AnthropicConfig.get_supported_openai_params(self, model) def map_openai_params( @@ -127,7 +127,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -155,7 +155,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -183,7 +183,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): def _build_bedrock_anthropic_request_base( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -251,10 +251,10 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): def _compute_bedrock_invoke_beta_headers( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, headers: dict, - ) -> List[str]: + ) -> list[str]: tools = optional_params.get("tools") tool_search_used = self.is_tool_search_used(tools) programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(tools) @@ -309,7 +309,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): if not isinstance(source_url, str): continue - inferred_format: Optional[str] = None + inferred_format: str | None = None if source_url.lower().endswith(".pdf"): inferred_format = "application/pdf" base64_url = convert_url_to_base64(url=source_url) @@ -348,7 +348,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): if not isinstance(source_url, str): continue - inferred_format: Optional[str] = None + inferred_format: str | None = None if source_url.lower().endswith(".pdf"): inferred_format = "application/pdf" base64_url = await async_convert_url_to_base64(url=source_url) @@ -394,12 +394,12 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: return AnthropicConfig.transform_response( self, diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index dd7cf12604d..0c6436030af 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -2,7 +2,7 @@ import copy import json import time from functools import partial -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast, get_args +from typing import TYPE_CHECKING, Any, cast, get_args import httpx from pydantic import TypeAdapter, ValidationError @@ -52,16 +52,12 @@ def _bedrock_invoke_guardrail_headers(raw_guardrail_config: object) -> "dict[str except ValidationError as e: raise BedrockError( status_code=400, - message="Invalid guardrailConfig={}. Expected format: {}. Error: {}".format( - raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT, e - ), + message=f"Invalid guardrailConfig={raw_guardrail_config}. Expected format: {_GUARDRAIL_CONFIG_EXPECTED_FORMAT}. Error: {e}", ) if "guardrailIdentifier" not in guardrail_config: raise BedrockError( status_code=400, - message="guardrailConfig={} is missing 'guardrailIdentifier'. Expected format: {}".format( - raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT - ), + message=f"guardrailConfig={raw_guardrail_config} is missing 'guardrailIdentifier'. Expected format: {_GUARDRAIL_CONFIG_EXPECTED_FORMAT}", ) trace = guardrail_config.get("trace") candidate_headers = { @@ -77,7 +73,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): BaseConfig.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ This is a base invoke model mapping. For Invoke - define a bedrock provider specific config that extends this class. """ @@ -106,12 +102,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -147,11 +143,11 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: return self._sign_request( service_name="bedrock", headers=headers, @@ -173,7 +169,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -272,9 +268,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): else: raise BedrockError( status_code=404, - message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/`.".format( - provider, model - ), + message=f"Bedrock Invoke HTTPX: Unknown provider={provider}, model={model}. Try calling via converse route - `bedrock/converse/`.", ) return request_data @@ -286,12 +280,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: completion_response = raw_response.json() @@ -302,7 +296,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): json.dumps(completion_response, indent=4, default=str), ) provider = self.get_bedrock_invoke_provider(model) - outputText: Optional[str] = None + outputText: str | None = None try: if provider == "cohere": if "text" in completion_response: @@ -362,7 +356,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): outputText = completion_response.get("results")[0].get("outputText") except Exception as e: raise BedrockError( - message="Error processing={}, Received error={}".format(raw_response.text, str(e)), + message=f"Error processing={raw_response.text}, Received error={e!s}", status_code=422, ) @@ -385,7 +379,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): raise Exception() except Exception as e: raise BedrockError( - message="Error parsing received text={}.\nError-{}".format(outputText, str(e)), + message=f"Error parsing received text={outputText}.\nError-{e!s}", status_code=raw_response.status_code, ) @@ -418,11 +412,11 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: raw_guardrail_config = optional_params.pop("guardrailConfig", None) if raw_guardrail_config is None: @@ -435,9 +429,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): } return {**headers, **guardrail_headers} - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BedrockError(status_code=status_code, message=error_message) @track_llm_api_timing() @@ -450,9 +442,9 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): headers: dict, data: dict, messages: list, - client: Optional[AsyncHTTPHandler] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: streaming_response = CustomStreamWrapper( completion_stream=None, @@ -485,9 +477,9 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -527,7 +519,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): @staticmethod def get_bedrock_invoke_provider( model: str, - ) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]: + ) -> litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None: """ Helper function to get the bedrock provider from the model @@ -563,7 +555,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): @staticmethod def _get_provider_from_model_path( model_path: str, - ) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]: + ) -> litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None: """ Helper function to get the provider from a model path with format: provider/model-name @@ -580,10 +572,10 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider) return None - def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]: + def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> tuple[str, list | None]: # handle anthropic prompts and amazon titan prompts prompt = "" - chat_history: Optional[list] = None + chat_history: list | None = None ## CUSTOM PROMPT if model in custom_prompt_dict: # check if the model has a registered custom prompt @@ -596,11 +588,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): ) return prompt, None ## ELSE - if provider == "anthropic" or provider == "amazon": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "mistral": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "meta" or provider == "llama": + if ( + provider == "anthropic" + or provider == "amazon" + or provider == "mistral" + or provider == "meta" + or provider == "llama" + ): prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") elif provider == "cohere": prompt, chat_history = cohere_message_pt(messages=messages) diff --git a/litellm/llms/bedrock/chat/mantle/transformation.py b/litellm/llms/bedrock/chat/mantle/transformation.py index d7deade8df9..b293d6ddf7f 100644 --- a/litellm/llms/bedrock/chat/mantle/transformation.py +++ b/litellm/llms/bedrock/chat/mantle/transformation.py @@ -8,7 +8,7 @@ at a different endpoint (bedrock-mantle.{region}.api.aws) with AWS SigV4 auth. """ from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( AmazonAnthropicClaudeConfig, @@ -37,12 +37,12 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: region = self._get_aws_region_name(optional_params=optional_params, model=model) return build_mantle_messages_url( @@ -55,11 +55,11 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = super().validate_environment( headers=headers, @@ -78,7 +78,7 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -105,7 +105,7 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -139,7 +139,7 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig): self, streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: from litellm.llms.anthropic.chat.handler import ModelResponseIterator diff --git a/litellm/llms/bedrock/claude_platform/__init__.py b/litellm/llms/bedrock/claude_platform/__init__.py index 88d4e9783c7..e05334a0350 100644 --- a/litellm/llms/bedrock/claude_platform/__init__.py +++ b/litellm/llms/bedrock/claude_platform/__init__.py @@ -1,8 +1,8 @@ -from .transformation import ( - BedrockClaudePlatformConfig, -) from .messages_transformation import ( BedrockClaudePlatformMessagesConfig, ) +from .transformation import ( + BedrockClaudePlatformConfig, +) __all__ = ["BedrockClaudePlatformConfig", "BedrockClaudePlatformMessagesConfig"] diff --git a/litellm/llms/bedrock/claude_platform/common_utils.py b/litellm/llms/bedrock/claude_platform/common_utils.py index b93577e2bca..e6056ce8216 100644 --- a/litellm/llms/bedrock/claude_platform/common_utils.py +++ b/litellm/llms/bedrock/claude_platform/common_utils.py @@ -1,4 +1,4 @@ -from typing import Literal, Optional, Tuple +from typing import Literal import litellm from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM @@ -16,7 +16,7 @@ def strip_claude_platform_route(model: str) -> str: class BedrockClaudePlatformMixin(BaseAWSLLM): @staticmethod - def _get_workspace_id(optional_params: dict, litellm_params: dict) -> Optional[str]: + def _get_workspace_id(optional_params: dict, litellm_params: dict) -> str | None: workspace_id = ( optional_params.get("workspace_id") or litellm_params.get("workspace_id") @@ -52,12 +52,12 @@ class BedrockClaudePlatformMixin(BaseAWSLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = ( api_base @@ -78,11 +78,11 @@ class BedrockClaudePlatformMixin(BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: if api_key or get_secret_str("ANTHROPIC_AWS_API_KEY"): return headers, None diff --git a/litellm/llms/bedrock/claude_platform/messages_transformation.py b/litellm/llms/bedrock/claude_platform/messages_transformation.py index 1b0d21a724c..4cfda162cee 100644 --- a/litellm/llms/bedrock/claude_platform/messages_transformation.py +++ b/litellm/llms/bedrock/claude_platform/messages_transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, Tuple +from typing import Any import litellm from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -16,12 +16,12 @@ class BedrockClaudePlatformMessagesConfig(BedrockClaudePlatformMixin, AnthropicM self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: workspace_id = self._get_workspace_id(optional_params, litellm_params) if workspace_id is None: raise litellm.AuthenticationError( @@ -53,11 +53,11 @@ class BedrockClaudePlatformMessagesConfig(BedrockClaudePlatformMixin, AnthropicM def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: return super().transform_anthropic_messages_request( model=strip_claude_platform_route(model), messages=messages, diff --git a/litellm/llms/bedrock/claude_platform/transformation.py b/litellm/llms/bedrock/claude_platform/transformation.py index 6f5ccececc7..e308e547360 100644 --- a/litellm/llms/bedrock/claude_platform/transformation.py +++ b/litellm/llms/bedrock/claude_platform/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional +from typing import Any import litellm from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -14,7 +14,7 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock" def should_strip_billing_metadata(self) -> bool: @@ -24,12 +24,12 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: workspace_id = self._get_workspace_id(optional_params, litellm_params) if workspace_id is None: raise litellm.AuthenticationError( @@ -70,7 +70,7 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig): self, streaming_response: Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: from litellm.llms.anthropic.chat.handler import ModelResponseIterator diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index a6615a23866..03bb6fe1fbf 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -13,12 +13,8 @@ from collections.abc import Mapping from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, TypedDict, - Union, ) if TYPE_CHECKING: @@ -49,7 +45,7 @@ class BedrockError(BaseLLMException): _get_model_info = None BedrockOutputConfigEffort = Literal["low", "medium", "high", "max", "xhigh"] -_BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER: Dict[BedrockOutputConfigEffort, int] = { +_BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER: dict[BedrockOutputConfigEffort, int] = { "low": 0, "medium": 1, "high": 2, @@ -75,13 +71,13 @@ def get_cached_model_info(): @functools.lru_cache(maxsize=1) -def _get_local_model_cost_map() -> Dict: +def _get_local_model_cost_map() -> dict: from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap return GetModelCostMap.load_local_model_cost_map() -def pop_bedrock_invoke_output_config_format(request_body: Dict) -> Optional[Dict]: +def pop_bedrock_invoke_output_config_format(request_body: dict) -> dict | None: """ Remove and return Anthropic's nested ``output_config.format`` field. @@ -102,8 +98,8 @@ def pop_bedrock_invoke_output_config_format(request_body: Dict) -> Optional[Dict def convert_bedrock_invoke_output_format_to_inline_schema( - output_format: Dict, - request_body: Dict, + output_format: dict, + request_body: dict, ) -> None: """ Embed an Anthropic structured-output schema into the last user message. @@ -177,7 +173,7 @@ def normalize_json_schema_custom_types_to_object(schema: dict) -> None: Uses an explicit stack (not recursion) to satisfy recursive-function guards in CI. """ - stack: List[Any] = [schema] + stack: list[Any] = [schema] seen: set[int] = set() while stack: node = stack.pop() @@ -267,7 +263,7 @@ class AmazonBedrockGlobalConfig: optional_params[mapped_params[param]] = value return optional_params - def get_all_regions(self) -> List[str]: + def get_all_regions(self) -> list[str]: return ( self.get_us_regions() + self.get_eu_regions() @@ -276,7 +272,7 @@ class AmazonBedrockGlobalConfig: + self.get_sa_regions() ) - def get_ap_regions(self) -> List[str]: + def get_ap_regions(self) -> list[str]: """ Source: https://www.aws-services.info/bedrock.html """ @@ -290,10 +286,10 @@ class AmazonBedrockGlobalConfig: "ap-southeast-2", # Asia Pacific (Sydney) ] - def get_sa_regions(self) -> List[str]: + def get_sa_regions(self) -> list[str]: return ["sa-east-1"] - def get_eu_regions(self) -> List[str]: + def get_eu_regions(self) -> list[str]: """ Source: https://www.aws-services.info/bedrock.html """ @@ -308,10 +304,10 @@ class AmazonBedrockGlobalConfig: "eu-north-1", # Europe (Stockholm) ] - def get_ca_regions(self) -> List[str]: + def get_ca_regions(self) -> list[str]: return ["ca-central-1"] - def get_us_regions(self) -> List[str]: + def get_us_regions(self) -> list[str]: """ Source: https://www.aws-services.info/bedrock.html """ @@ -336,7 +332,7 @@ def add_custom_header(headers): return callback -def _get_bedrock_client_ssl_verify() -> Union[bool, str]: +def _get_bedrock_client_ssl_verify() -> bool | str: """ Get SSL verification setting for Bedrock client. @@ -352,16 +348,16 @@ def _get_bedrock_client_ssl_verify() -> Union[bool, str]: def init_bedrock_client( region_name=None, - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_region_name: Optional[str] = None, - aws_bedrock_runtime_endpoint: Optional[str] = None, - aws_session_name: Optional[str] = None, - aws_profile_name: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_web_identity_token: Optional[str] = None, - extra_headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + aws_access_key_id: str | None = None, + aws_secret_access_key: str | None = None, + aws_region_name: str | None = None, + aws_bedrock_runtime_endpoint: str | None = None, + aws_session_name: str | None = None, + aws_profile_name: str | None = None, + aws_role_name: str | None = None, + aws_web_identity_token: str | None = None, + extra_headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, ): # check for custom AWS_REGION_NAME and use it if not passed to init_bedrock_client litellm_aws_region_name = get_secret("AWS_REGION_NAME", None) @@ -567,10 +563,10 @@ def get_bedrock_tool_name(response_tool_name: str) -> str: # Cache the global regions list at module level -_BEDROCK_GLOBAL_REGIONS: Optional[List[str]] = None +_BEDROCK_GLOBAL_REGIONS: list[str] | None = None -def _get_all_bedrock_regions() -> List[str]: +def _get_all_bedrock_regions() -> list[str]: """Get all Bedrock regions, cached at module level.""" global _BEDROCK_GLOBAL_REGIONS if _BEDROCK_GLOBAL_REGIONS is None: @@ -578,7 +574,7 @@ def _get_all_bedrock_regions() -> List[str]: return _BEDROCK_GLOBAL_REGIONS -def get_bedrock_cross_region_inference_regions() -> List[str]: +def get_bedrock_cross_region_inference_regions() -> list[str]: """Abbreviations of regions AWS Bedrock supports for cross region inference.""" return ["global", "us", "eu", "apac", "jp", "au", "us-gov"] @@ -627,8 +623,8 @@ MANTLE_MESSAGES_PATH = "/anthropic/v1/messages" def build_mantle_messages_url( - api_base: Optional[str], - aws_bedrock_runtime_endpoint: Optional[str], + api_base: str | None, + aws_bedrock_runtime_endpoint: str | None, region: str, ) -> str: """Build the bedrock-mantle Anthropic /messages URL. @@ -730,7 +726,7 @@ def bedrock_converse_supports_strict_tools(model: str) -> bool: return flag if flag is not None else True -def _get_bedrock_converse_strict_tools_flag(base_model: str) -> Optional[bool]: +def _get_bedrock_converse_strict_tools_flag(base_model: str) -> bool | None: candidates = dict.fromkeys((base_model, _BEDROCK_MODEL_VERSION_SUFFIX_RE.sub("", base_model))) for candidate in candidates: with contextlib.suppress(Exception): @@ -782,7 +778,7 @@ def normalize_bedrock_opus_output_config_effort(model: str, output_config: Any) def _get_bedrock_output_config_effort_ceiling( model: str, -) -> Optional[BedrockOutputConfigEffort]: +) -> BedrockOutputConfigEffort | None: try: model_info = get_cached_model_info()( model=model, @@ -815,14 +811,14 @@ class BedrockModelInfo(BaseLLMModelInfo): all_global_regions = global_config.get_all_regions() @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: """ Get the API base for the given model. """ return api_base @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: """ Get the API key for the given model. """ @@ -832,15 +828,15 @@ class BedrockModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List["AllMessageValues"], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return headers - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: return [] # def get_provider_info(self, model: str) -> Optional[ProviderSpecificModelInfo]: @@ -859,7 +855,7 @@ class BedrockModelInfo(BaseLLMModelInfo): # return overrides if overrides else None - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create a Bedrock token counter. @@ -884,7 +880,7 @@ class BedrockModelInfo(BaseLLMModelInfo): return get_bedrock_base_model(model) @staticmethod - def _supported_cross_region_inference_region() -> List[str]: + def _supported_cross_region_inference_region() -> list[str]: """Wrapper for standalone function. See get_bedrock_cross_region_inference_regions().""" return get_bedrock_cross_region_inference_regions() @@ -905,7 +901,7 @@ class BedrockModelInfo(BaseLLMModelInfo): """ Get the bedrock route for the given model. """ - route_mappings: Dict[ + route_mappings: dict[ str, Literal[ "invoke", @@ -1056,7 +1052,7 @@ class BedrockModelInfo(BaseLLMModelInfo): @staticmethod def get_bedrock_provider_config_for_messages_api( model: str, - ) -> Optional[BaseAnthropicMessagesConfig]: + ) -> BaseAnthropicMessagesConfig | None: """ Get the bedrock provider config for the given model. @@ -1248,7 +1244,7 @@ class BedrockEventStreamDecoderBase: self.parser = EventStreamJSONParser() - def _parse_message_from_event(self, event) -> Optional[str]: + def _parse_message_from_event(self, event) -> str | None: response_stream_shape = get_bedrock_response_stream_shape() if response_stream_shape is None: raise BedrockError( @@ -1276,7 +1272,7 @@ class BedrockEventStreamDecoderBase: return chunk.decode() # type: ignore[no-any-return] -def get_anthropic_beta_from_headers(headers: dict) -> List[str]: +def get_anthropic_beta_from_headers(headers: dict) -> list[str]: """ Extract anthropic-beta header values and convert them to a list. Supports both JSON array format and comma-separated values from user headers. @@ -1388,8 +1384,7 @@ class CommonBatchFilesUtils: if len(parts) > 1: # Reconstruct model name (everything except the last UUID part and .jsonl) model_name = "-".join(parts[:-1]) - if model_name.endswith(".jsonl"): - model_name = model_name[:-6] # Remove .jsonl + model_name = model_name.removesuffix(".jsonl") # Remove .jsonl return model_name except Exception: pass @@ -1400,7 +1395,7 @@ class CommonBatchFilesUtils: def sign_aws_request( self, service_name: str, - data: Union[str, dict, "BedrockCreateBatchRequest"], + data: str | dict | BedrockCreateBatchRequest, endpoint_url: str, optional_params: dict, method: str = "POST", @@ -1523,9 +1518,7 @@ class CommonBatchFilesUtils: return bucket_name, object_key - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get Bedrock-specific error class. """ diff --git a/litellm/llms/bedrock/cost_calculation.py b/litellm/llms/bedrock/cost_calculation.py index 9a164d02eeb..0a2e39ab972 100644 --- a/litellm/llms/bedrock/cost_calculation.py +++ b/litellm/llms/bedrock/cost_calculation.py @@ -3,7 +3,7 @@ Helper util for handling bedrock-specific cost calculation - e.g.: prompt caching """ -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token @@ -11,7 +11,7 @@ if TYPE_CHECKING: from litellm.types.utils import Usage -def cost_per_token(model: str, usage: "Usage", service_tier: Optional[str] = None) -> Tuple[float, float]: +def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py index 1ea870a1d32..934d416d256 100644 --- a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -2,7 +2,7 @@ Bedrock Token Counter implementation using the CountTokens API. """ -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.llms.base_llm.base_utils import BaseTokenCounter @@ -16,7 +16,7 @@ class BedrockTokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: """ Returns True if we should use the Bedrock CountTokens API for token counting. @@ -26,13 +26,13 @@ class BedrockTokenCounter(BaseTokenCounter): async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: """ Count tokens using AWS Bedrock's CountTokens API. @@ -56,7 +56,7 @@ class BedrockTokenCounter(BaseTokenCounter): litellm_params = deployment.get("litellm_params", {}) # Build request data in the format expected by BedrockCountTokensHandler - request_data: Dict[str, Any] = { + request_data: dict[str, Any] = { "model": model_to_use, "messages": messages, } diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index 2c40e14129d..8e993c6f8b2 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -4,7 +4,7 @@ AWS Bedrock CountTokens API handler. Simplified handler leveraging existing LiteLLM Bedrock infrastructure. """ -from typing import Any, Dict +from typing import Any import httpx @@ -24,10 +24,10 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): async def handle_count_tokens_request( self, - request_data: Dict[str, Any], - litellm_params: Dict[str, Any], + request_data: dict[str, Any], + litellm_params: dict[str, Any], resolved_model: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Handle a CountTokens request using existing LiteLLM patterns. @@ -120,14 +120,14 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): raise except httpx.HTTPStatusError as e: # HTTP errors - preserve the actual status code - verbose_logger.error(f"HTTP error in CountTokens handler: {str(e)}") + verbose_logger.error(f"HTTP error in CountTokens handler: {e!s}") raise BedrockError( status_code=e.response.status_code, message=e.response.text, ) except Exception as e: - verbose_logger.error(f"Error in CountTokens handler: {str(e)}") + verbose_logger.error(f"Error in CountTokens handler: {e!s}") raise BedrockError( status_code=500, - message=f"CountTokens processing error: {str(e)}", + message=f"CountTokens processing error: {e!s}", ) diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index 38eaf13893d..645c4845ebb 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -6,7 +6,7 @@ to AWS Bedrock's CountTokens API format and vice versa. """ import re -from typing import Any, Dict, List, Optional +from typing import Any from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.bedrock.common_utils import get_bedrock_base_model @@ -27,7 +27,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): - Response: {"inputTokens": } """ - def _detect_input_type(self, request_data: Dict[str, Any]) -> str: + def _detect_input_type(self, request_data: dict[str, Any]) -> str: """ Detect whether to use 'converse' or 'invokeModel' input format. @@ -57,8 +57,8 @@ class BedrockCountTokensConfig(BaseAWSLLM): def transform_anthropic_to_bedrock_count_tokens( self, - request_data: Dict[str, Any], - ) -> Dict[str, Any]: + request_data: dict[str, Any], + ) -> dict[str, Any]: """ Transform request to Bedrock CountTokens format. Supports both Converse and InvokeModel input types. @@ -95,7 +95,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): else: return self._transform_to_invoke_model_format(request_data) - def _transform_to_converse_format(self, request_data: Dict[str, Any]) -> Dict[str, Any]: + def _transform_to_converse_format(self, request_data: dict[str, Any]) -> dict[str, Any]: """Transform to Converse input format, including system and tools.""" messages = request_data.get("messages", []) system = request_data.get("system") @@ -104,7 +104,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): # Transform messages user_messages = [] for message in messages: - transformed_message: Dict[str, Any] = { + transformed_message: dict[str, Any] = { "role": message.get("role"), "content": [], } @@ -115,7 +115,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): transformed_message["content"] = content user_messages.append(transformed_message) - converse_input: Dict[str, Any] = {"messages": user_messages} + converse_input: dict[str, Any] = {"messages": user_messages} # Transform system prompt (string or list of blocks → Bedrock format) system_blocks = self._transform_system(system) @@ -129,7 +129,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): return {"input": {"converse": converse_input}} - def _transform_system(self, system: Optional[Any]) -> List[Dict[str, Any]]: + def _transform_system(self, system: Any | None) -> list[dict[str, Any]]: """Transform Anthropic system prompt to Bedrock system blocks.""" if system is None: return [] @@ -140,7 +140,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): return [{"text": block.get("text", "")} for block in system if isinstance(block, dict)] return [] - def _transform_tools(self, tools: Optional[List[Dict[str, Any]]]) -> Optional[Dict[str, Any]]: + def _transform_tools(self, tools: list[dict[str, Any]] | None) -> dict[str, Any] | None: """Transform Anthropic tools to Bedrock toolConfig format.""" if not tools: return None @@ -169,7 +169,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): return {"tools": bedrock_tools} - def _transform_to_invoke_model_format(self, request_data: Dict[str, Any]) -> Dict[str, Any]: + def _transform_to_invoke_model_format(self, request_data: dict[str, Any]) -> dict[str, Any]: """Transform to InvokeModel input format.""" import base64 import json @@ -192,8 +192,8 @@ class BedrockCountTokensConfig(BaseAWSLLM): self, model: str, aws_region_name: str, - api_base: Optional[str] = None, - aws_bedrock_runtime_endpoint: Optional[str] = None, + api_base: str | None = None, + aws_bedrock_runtime_endpoint: str | None = None, ) -> str: """ Construct the AWS Bedrock CountTokens API endpoint using existing LiteLLM functions. @@ -211,8 +211,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): model_id = get_bedrock_base_model(model) # Remove bedrock/ prefix if present - if model_id.startswith("bedrock/"): - model_id = model_id[8:] # Remove "bedrock/" prefix + model_id = model_id.removeprefix("bedrock/") # Remove "bedrock/" prefix encoded_model_id = self.encode_model_id(model_id=model_id) base_url, _ = self.get_runtime_endpoint( @@ -224,7 +223,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): return endpoint - def transform_bedrock_response_to_anthropic(self, bedrock_response: Dict[str, Any]) -> Dict[str, Any]: + def transform_bedrock_response_to_anthropic(self, bedrock_response: dict[str, Any]) -> dict[str, Any]: """ Transform Bedrock CountTokens response to Anthropic format. @@ -242,7 +241,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): return {"input_tokens": input_tokens} - def validate_count_tokens_request(self, request_data: Dict[str, Any]) -> None: + def validate_count_tokens_request(self, request_data: dict[str, Any]) -> None: """ Validate the incoming count tokens request. Supports both Converse and InvokeModel input formats. diff --git a/litellm/llms/bedrock/embed/amazon_nova_transformation.py b/litellm/llms/bedrock/embed/amazon_nova_transformation.py index 58519d0d061..86b52fb5d22 100644 --- a/litellm/llms/bedrock/embed/amazon_nova_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_nova_transformation.py @@ -12,8 +12,6 @@ Supports: Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/nova-embed.html """ -from typing import List, Optional - from litellm.types.utils import ( Embedding, EmbeddingResponse, @@ -35,7 +33,7 @@ class AmazonNovaEmbeddingConfig: def __init__(self) -> None: pass - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return [ "dimensions", ] @@ -88,8 +86,8 @@ class AmazonNovaEmbeddingConfig: input: str, inference_params: dict, async_invoke_route: bool = False, - model_id: Optional[str] = None, - output_s3_uri: Optional[str] = None, + model_id: str | None = None, + output_s3_uri: str | None = None, ) -> dict: """ Transform OpenAI-style input to Nova format. @@ -208,7 +206,7 @@ class AmazonNovaEmbeddingConfig: self, model_input: dict, model_id: str, - output_s3_uri: Optional[str] = None, + output_s3_uri: str | None = None, ) -> dict: """ Wrap the transformed request in the AWS Bedrock async invoke format. @@ -240,9 +238,9 @@ class AmazonNovaEmbeddingConfig: def _transform_response( self, - response_list: List[dict], + response_list: list[dict], model: str, - batch_data: Optional[List[dict]] = None, + batch_data: list[dict] | None = None, ) -> EmbeddingResponse: """ Transform Nova response to OpenAI format. @@ -258,7 +256,7 @@ class AmazonNovaEmbeddingConfig: ] } """ - embeddings: List[Embedding] = [] + embeddings: list[Embedding] = [] total_tokens = 0 for response in response_list: @@ -302,7 +300,7 @@ class AmazonNovaEmbeddingConfig: if "image" in params: image_count += 1 - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if image_count > 0: prompt_tokens_details = PromptTokensDetailsWrapper( image_count=image_count, diff --git a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py index 57cbb3263de..6c97e69f635 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py @@ -10,7 +10,6 @@ Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-tit """ import types -from typing import List from litellm.types.llms.bedrock import ( AmazonTitanG1EmbeddingRequest, @@ -50,7 +49,7 @@ class AmazonTitanG1Config: and v is not None } - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return [] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: @@ -59,10 +58,10 @@ class AmazonTitanG1Config: def _transform_request(self, input: str, inference_params: dict) -> AmazonTitanG1EmbeddingRequest: return AmazonTitanG1EmbeddingRequest(inputText=input) - def _transform_response(self, response_list: List[dict], model: str) -> EmbeddingResponse: + def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse: total_prompt_tokens = 0 - transformed_responses: List[Embedding] = [] + transformed_responses: list[Embedding] = [] for index, response in enumerate(response_list): _parsed_response = AmazonTitanG1EmbeddingResponse(**response) # type: ignore transformed_responses.append( diff --git a/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py index 878d5f7e850..e3ba7b856cd 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py @@ -6,8 +6,6 @@ Why separate file? Make it easy to see how transformation works Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-titan-embed-mm.html """ -from typing import List, Optional - from litellm.types.llms.bedrock import ( AmazonTitanMultimodalEmbeddingConfig, AmazonTitanMultimodalEmbeddingRequest, @@ -30,7 +28,7 @@ class AmazonTitanMultimodalEmbeddingG1Config: def __init__(self) -> None: pass - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return ["dimensions"] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: @@ -54,12 +52,12 @@ class AmazonTitanMultimodalEmbeddingG1Config: def _transform_response( self, - response_list: List[dict], + response_list: list[dict], model: str, - batch_data: Optional[List[dict]] = None, + batch_data: list[dict] | None = None, ) -> EmbeddingResponse: total_prompt_tokens = 0 - transformed_responses: List[Embedding] = [] + transformed_responses: list[Embedding] = [] for index, response in enumerate(response_list): _parsed_response = AmazonTitanMultimodalEmbeddingResponse(**response) # type: ignore transformed_responses.append( @@ -78,7 +76,7 @@ class AmazonTitanMultimodalEmbeddingG1Config: if "inputImage" in request_data: image_count += 1 - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if image_count > 0: prompt_tokens_details = PromptTokensDetailsWrapper( image_count=image_count, diff --git a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py index 2c7b0ba465a..72734f963d0 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py @@ -10,7 +10,6 @@ Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-tit """ import types -from typing import List, Optional, Union from litellm.types.llms.bedrock import ( AmazonTitanV2EmbeddingRequest, @@ -27,10 +26,10 @@ class AmazonTitanV2Config: dimensions: int - The number of dimensions the output embeddings should have. The following values are accepted: 1024 (default), 512, 256. """ - normalize: Optional[bool] = None - dimensions: Optional[int] = None + normalize: bool | None = None + dimensions: int | None = None - def __init__(self, normalize: Optional[bool] = None, dimensions: Optional[int] = None) -> None: + def __init__(self, normalize: bool | None = None, dimensions: int | None = None) -> None: locals_ = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: @@ -54,7 +53,7 @@ class AmazonTitanV2Config: and v is not None } - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return ["dimensions", "encoding_format"] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: @@ -76,17 +75,17 @@ class AmazonTitanV2Config: def _transform_request(self, input: str, inference_params: dict) -> AmazonTitanV2EmbeddingRequest: return AmazonTitanV2EmbeddingRequest(inputText=input, **inference_params) # type: ignore - def _transform_response(self, response_list: List[dict], model: str) -> EmbeddingResponse: + def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse: total_prompt_tokens = 0 - transformed_responses: List[Embedding] = [] + transformed_responses: list[Embedding] = [] for index, response in enumerate(response_list): _parsed_response = AmazonTitanV2EmbeddingResponse(**response) # type: ignore # According to AWS docs, embeddingsByType is always present # If binary was requested (encoding_format="base64"), use binary data # Otherwise, use float data from embeddingsByType or fallback to embedding field - embedding_data: Union[List[float], List[int]] + embedding_data: list[float] | list[int] if "embeddingsByType" in _parsed_response and "binary" in _parsed_response["embeddingsByType"]: # Use binary data if available (for encoding_format="base64") diff --git a/litellm/llms/bedrock/embed/cohere_transformation.py b/litellm/llms/bedrock/embed/cohere_transformation.py index ac3130ea434..a5211baa5e8 100644 --- a/litellm/llms/bedrock/embed/cohere_transformation.py +++ b/litellm/llms/bedrock/embed/cohere_transformation.py @@ -4,8 +4,6 @@ Transformation logic from OpenAI /v1/embeddings format to Bedrock Cohere /invoke Why separate file? Make it easy to see how transformation works """ -from typing import List - from litellm.llms.cohere.embed.transformation import CohereEmbeddingConfig from litellm.types.llms.bedrock import CohereEmbeddingRequest @@ -14,7 +12,7 @@ class BedrockCohereEmbeddingConfig: def __init__(self) -> None: pass - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return ["encoding_format", "dimensions"] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: @@ -28,7 +26,7 @@ class BedrockCohereEmbeddingConfig: def _is_v3_model(self, model: str) -> bool: return "3" in model - def _transform_request(self, model: str, input: List[str], inference_params: dict) -> CohereEmbeddingRequest: + def _transform_request(self, model: str, input: list[str], inference_params: dict) -> CohereEmbeddingRequest: transformed_request = CohereEmbeddingConfig()._transform_request(model, input, inference_params) new_transformed_request = CohereEmbeddingRequest( diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index 0e82baad74b..d408692ff95 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -6,7 +6,7 @@ import copy import json import urllib.parse from collections.abc import Callable -from typing import Any, List, Optional, Tuple, Union, get_args +from typing import Any, get_args import httpx @@ -42,7 +42,7 @@ class BedrockEmbedding(BaseAWSLLM): def _load_credentials( self, optional_params: dict, - ) -> Tuple[Any, str]: + ) -> tuple[Any, str]: try: from botocore.credentials import Credentials except ImportError: @@ -92,8 +92,8 @@ class BedrockEmbedding(BaseAWSLLM): def _make_sync_call( self, - client: Optional[HTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: HTTPHandler | None, + timeout: float | httpx.Timeout | None, api_base: str, headers: dict, data: dict, @@ -120,8 +120,8 @@ class BedrockEmbedding(BaseAWSLLM): async def _make_async_call( self, - client: Optional[AsyncHTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, api_base: str, headers: dict, data: dict, @@ -149,16 +149,16 @@ class BedrockEmbedding(BaseAWSLLM): def _transform_response( self, - response_list: List[dict], + response_list: list[dict], model: str, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, - is_async_invoke: Optional[bool] = False, - batch_data: Optional[List[dict]] = None, - ) -> Optional[EmbeddingResponse]: + is_async_invoke: bool | None = False, + batch_data: list[dict] | None = None, + ) -> EmbeddingResponse | None: """ Transforms the response from the Bedrock embedding provider to the OpenAI format. """ - returned_response: Optional[EmbeddingResponse] = None + returned_response: EmbeddingResponse | None = None # Handle async invoke responses (single response with invocationArn) if is_async_invoke and len(response_list) == 1 and "invocationArn" in response_list[0]: @@ -220,25 +220,25 @@ class BedrockEmbedding(BaseAWSLLM): # Validate returned response ########################################################## if returned_response is None: - raise Exception("Unable to map model response to known provider format. model={}".format(model)) + raise Exception(f"Unable to map model response to known provider format. model={model}") return returned_response def _single_func_embeddings( self, - client: Optional[HTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], - batch_data: List[dict], + client: HTTPHandler | None, + timeout: float | httpx.Timeout | None, + batch_data: list[dict], credentials: Any, - extra_headers: Optional[dict], + extra_headers: dict | None, endpoint_url: str, aws_region_name: str, model: str, logging_obj: Any, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, - api_key: Optional[str] = None, - is_async_invoke: Optional[bool] = False, + api_key: str | None = None, + is_async_invoke: bool | None = False, ): - responses: List[dict] = [] + responses: list[dict] = [] for data in batch_data: headers = {"Content-Type": "application/json"} if extra_headers is not None: @@ -293,20 +293,20 @@ class BedrockEmbedding(BaseAWSLLM): async def _async_single_func_embeddings( self, - client: Optional[AsyncHTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], - batch_data: List[dict], + client: AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + batch_data: list[dict], credentials: Any, - extra_headers: Optional[dict], + extra_headers: dict | None, endpoint_url: str, aws_region_name: str, model: str, logging_obj: Any, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, - api_key: Optional[str] = None, - is_async_invoke: Optional[bool] = False, + api_key: str | None = None, + is_async_invoke: bool | None = False, ): - responses: List[dict] = [] + responses: list[dict] = [] for data in batch_data: headers = {"Content-Type": "application/json"} if extra_headers is not None: @@ -364,19 +364,19 @@ class BedrockEmbedding(BaseAWSLLM): def embeddings( self, model: str, - input: List[str], - api_base: Optional[str], + input: list[str], + api_base: str | None, model_response: EmbeddingResponse, print_verbose: Callable, encoding, logging_obj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]], - timeout: Optional[Union[float, httpx.Timeout]], - aembedding: Optional[bool], - extra_headers: Optional[dict], + client: HTTPHandler | AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + aembedding: bool | None, + extra_headers: dict | None, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> EmbeddingResponse: credentials, aws_region_name = self._load_credentials(optional_params) @@ -404,8 +404,8 @@ class BedrockEmbedding(BaseAWSLLM): } inference_params.pop("user", None) # make sure user is not passed in for bedrock call - data: Optional[CohereEmbeddingRequest] = None - batch_data: Optional[List] = None + data: CohereEmbeddingRequest | None = None + batch_data: list | None = None if provider == "cohere": data = BedrockCohereEmbeddingConfig()._transform_request( model=model, input=input, inference_params=inference_params diff --git a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py index 56ac2c00560..90001d2ef58 100644 --- a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py +++ b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py @@ -6,7 +6,7 @@ Why separate file? Make it easy to see how transformation works Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html """ -from typing import List, Optional, Union, cast +from typing import cast from litellm.types.llms.bedrock import ( TWELVELABS_EMBEDDING_INPUT_TYPES, @@ -31,7 +31,7 @@ class TwelveLabsMarengoEmbeddingConfig: def __init__(self) -> None: pass - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return [ "encoding_format", "textTruncate", @@ -75,9 +75,9 @@ class TwelveLabsMarengoEmbeddingConfig: input: str, inference_params: dict, async_invoke_route: bool = False, - model_id: Optional[str] = None, - output_s3_uri: Optional[str] = None, - ) -> Union[TwelveLabsMarengoEmbeddingRequest, TwelveLabsAsyncInvokeRequest]: + model_id: str | None = None, + output_s3_uri: str | None = None, + ) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsAsyncInvokeRequest: """ Transform OpenAI-style input to TwelveLabs Marengo format/async-invoke format. @@ -156,7 +156,7 @@ class TwelveLabsMarengoEmbeddingConfig: self, model_input: TwelveLabsMarengoEmbeddingRequest, model_id: str, - output_s3_uri: Optional[str] = None, + output_s3_uri: str | None = None, ) -> TwelveLabsAsyncInvokeRequest: """ Wrap the transformed request in the correct AWS Bedrock async invoke format. @@ -188,12 +188,12 @@ class TwelveLabsMarengoEmbeddingConfig: ), ) - def _transform_response(self, response_list: List[dict], model: str) -> EmbeddingResponse: + def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse: """ Transform TwelveLabs response to OpenAI format. Handles the actual TwelveLabs response format: {"data": [{"embedding": [...]}]} """ - embeddings: List[Embedding] = [] + embeddings: list[Embedding] = [] total_tokens = 0 for response in response_list: diff --git a/litellm/llms/bedrock/files/handler.py b/litellm/llms/bedrock/files/handler.py index 5015ebe5774..12ebc52dff3 100644 --- a/litellm/llms/bedrock/files/handler.py +++ b/litellm/llms/bedrock/files/handler.py @@ -1,6 +1,6 @@ import asyncio from collections.abc import Coroutine, Mapping -from typing import Any, Optional, Tuple, Union +from typing import Any import httpx @@ -43,7 +43,7 @@ class BedrockFilesHandler(BaseAWSLLM): s3_uri: str, configured_bucket_name: str, allow_legacy_cloud_file_ids: bool = False, - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Parse S3 URI to extract bucket name and object key. @@ -70,8 +70,8 @@ class BedrockFilesHandler(BaseAWSLLM): self, file_content_request: FileContentRequest, optional_params: dict, - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + timeout: float | httpx.Timeout, + max_retries: int | None, ) -> HttpxBinaryResponseContent: """ Download file content from S3 bucket for Bedrock files. @@ -130,7 +130,7 @@ class BedrockFilesHandler(BaseAWSLLM): response = s3_client.get_object(Bucket=bucket_name, Key=object_key) file_content = response["Body"].read() except Exception as e: - raise ValueError(f"Failed to download file from S3: {s3_uri}. Error: {str(e)}") + raise ValueError(f"Failed to download file from S3: {s3_uri}. Error: {e!s}") # Create mock HTTP response mock_response = httpx.Response( @@ -146,11 +146,11 @@ class BedrockFilesHandler(BaseAWSLLM): self, _is_async: bool, file_content_request: FileContentRequest, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - ) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: + timeout: float | httpx.Timeout, + max_retries: int | None, + ) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: """ Download file content from S3 bucket for Bedrock files. Supports both sync and async operations. diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index d4865a1c87a..3656088cb9d 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -6,11 +6,6 @@ from collections.abc import Mapping, MutableMapping from types import MappingProxyType from typing import ( Any, - Dict, - List, - Optional, - Tuple, - Union, ) from urllib.parse import unquote @@ -157,7 +152,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): self, headers: MutableMapping[str, object], model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: MutableMapping[str, object], api_key: str | None = None, @@ -180,7 +175,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): - Tuple formats: (filename, content, [content_type], [headers]) - PathLike objects """ - content: Union[str, bytes] = b"" + content: str | bytes = b"" # Extract file content from tuple if necessary if isinstance(openai_file_content, tuple): # Take the second element which is always the file content @@ -210,7 +205,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def _get_s3_object_name_from_batch_jsonl( self, - openai_jsonl_content: List[Dict[str, Any]], + openai_jsonl_content: list[dict[str, Any]], ) -> str: """ Gets a unique S3 object name for the Bedrock batch processing job @@ -219,8 +214,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): """ _model = openai_jsonl_content[0].get("body", {}).get("model", "") # Remove bedrock/ prefix if present - if _model.startswith("bedrock/"): - _model = _model[8:] + _model = _model.removeprefix("bedrock/") safe_model = sanitize_cloud_object_component(_model.replace(":", "-"), fallback="model") @@ -255,11 +249,11 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def get_complete_file_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, data: CreateFileRequest, ) -> str: """ @@ -294,7 +288,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}" - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: return [] def map_openai_params( @@ -318,7 +312,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): OPENAI_EMBEDDINGS_URL = "/v1/embeddings" @staticmethod - def _is_embedding_record(openai_jsonl_record: Dict[str, Any]) -> bool: + def _is_embedding_record(openai_jsonl_record: dict[str, Any]) -> bool: """ Decide whether an OpenAI batch JSONL line is an embedding request. @@ -406,8 +400,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): # Registry silence -> substring fallback for unmapped ids only. normalized = model.lower() - if normalized.startswith("bedrock/"): - normalized = normalized[len("bedrock/") :] + normalized = normalized.removeprefix("bedrock/") marker = BedrockFilesConfig._TITAN_V2_EMBED_MODEL_MARKER idx = normalized.find(marker) if idx < 0: @@ -416,7 +409,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): return end == len(normalized) or normalized[end] in (":", "/") @staticmethod - def _lookup_provider_specific_field(model_id: str, field: str) -> Optional[str]: + def _lookup_provider_specific_field(model_id: str, field: str) -> str | None: """ Read a nested string field from the registry entry's `provider_specific_entry` dict via `litellm.get_model_info`. @@ -507,8 +500,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def _map_openai_embedding_to_bedrock_params( self, - openai_request_body: Dict[str, Any], - ) -> Dict[str, Any]: + openai_request_body: dict[str, Any], + ) -> dict[str, Any]: """ Transform an OpenAI /v1/embeddings request body into the Bedrock InvokeModel `modelInput` for embedding models that AWS @@ -555,9 +548,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def _map_openai_to_bedrock_params( self, - openai_request_body: Dict[str, Any], - provider: Optional[str] = None, - ) -> Dict[str, Any]: + openai_request_body: dict[str, Any], + provider: str | None = None, + ) -> dict[str, Any]: """ Transform OpenAI request body to Bedrock-compatible modelInput parameters using existing transformation logic. @@ -624,8 +617,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): } def _transform_openai_jsonl_content_to_bedrock_jsonl_content( - self, openai_jsonl_content: List[Dict[str, Any]] - ) -> List[Dict[str, Any]]: + self, openai_jsonl_content: list[dict[str, Any]] + ) -> list[dict[str, Any]]: """ Transforms OpenAI JSONL content to Bedrock batch format @@ -659,7 +652,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): ) except Exception as e: verbose_logger.exception( - f"litellm.llms.bedrock.files.transformation.py::_transform_openai_jsonl_content_to_bedrock_jsonl_content() - Error inferring custom_llm_provider - {str(e)}" + f"litellm.llms.bedrock.files.transformation.py::_transform_openai_jsonl_content_to_bedrock_jsonl_content() - Error inferring custom_llm_provider - {e!s}" ) # Determine provider from model name @@ -688,7 +681,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): create_file_data: CreateFileRequest, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, dict]: + ) -> bytes | str | dict: """ Transform file request and return a pre-signed request for S3. This keeps the HTTP handler clean by doing all the signing here. @@ -758,7 +751,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): content: str, api_base: str, optional_params: dict, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Sign S3 PUT request using the same proven logic as S3Logger. Reuses the exact pattern from litellm/integrations/s3_v2.py @@ -879,7 +872,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -911,7 +904,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): object="file", ) - def get_error_class(self, error_message: str, status_code: int, headers: Union[Dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return BedrockError(status_code=status_code, message=error_message, headers=headers) def transform_retrieve_file_request( @@ -948,7 +941,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: @@ -959,7 +952,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: raise NotImplementedError("BedrockFilesConfig does not support file listing") def transform_file_content_request( @@ -1071,8 +1064,8 @@ class BedrockJsonlFilesTransformation: """ def transform_openai_file_content_to_bedrock_file_content( - self, openai_file_content: Optional[FileTypes] = None - ) -> Tuple[str, str]: + self, openai_file_content: FileTypes | None = None + ) -> tuple[str, str]: """ Transforms OpenAI FileContentRequest to Bedrock S3 file format """ @@ -1089,7 +1082,7 @@ class BedrockJsonlFilesTransformation: object_name = self._get_s3_object_name(openai_jsonl_content=openai_jsonl_content) return bedrock_jsonl_string, object_name - def _transform_openai_jsonl_content_to_bedrock_jsonl_content(self, openai_jsonl_content: List[Dict[str, Any]]): + def _transform_openai_jsonl_content_to_bedrock_jsonl_content(self, openai_jsonl_content: list[dict[str, Any]]): """ Delegate to the main BedrockFilesConfig transformation method """ @@ -1098,7 +1091,7 @@ class BedrockJsonlFilesTransformation: def _get_s3_object_name( self, - openai_jsonl_content: List[Dict[str, Any]], + openai_jsonl_content: list[dict[str, Any]], ) -> str: """ Gets a unique S3 object name for the Bedrock batch processing job @@ -1107,8 +1100,7 @@ class BedrockJsonlFilesTransformation: """ _model = openai_jsonl_content[0].get("body", {}).get("model", "") # Remove bedrock/ prefix if present - if _model.startswith("bedrock/"): - _model = _model[8:] + _model = _model.removeprefix("bedrock/") safe_model = sanitize_cloud_object_component(_model.replace(":", "-"), fallback="model") object_name = f"{BEDROCK_MANAGED_S3_BATCH_PREFIX}{safe_model}-{uuid.uuid4()}.jsonl" return object_name @@ -1122,7 +1114,7 @@ class BedrockJsonlFilesTransformation: - Tuple formats: (filename, content, [content_type], [headers]) - PathLike objects """ - content: Union[str, bytes] = b"" + content: str | bytes = b"" # Extract file content from tuple if necessary if isinstance(openai_file_content, tuple): # Take the second element which is always the file content @@ -1151,7 +1143,7 @@ class BedrockJsonlFilesTransformation: return content def transform_s3_bucket_response_to_openai_file_object( - self, create_file_data: CreateFileRequest, s3_upload_response: Dict[str, Any] + self, create_file_data: CreateFileRequest, s3_upload_response: dict[str, Any] ) -> OpenAIFileObject: """ Transforms S3 Bucket upload file response to OpenAI FileObject diff --git a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py index 1008924ab0e..7b1f621d3c7 100644 --- a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py +++ b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py @@ -14,7 +14,7 @@ from __future__ import annotations import base64 import os -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -40,14 +40,14 @@ else: def _nova_canvas_task_body( *, image_b64: str, - mask_b64: Optional[str], + mask_b64: str | None, text: str, - negative_text: Optional[str], - similarity_strength: Optional[float], - task_type: Optional[str], - mask_prompt: Optional[str], - out_painting_mode: Optional[str], -) -> Dict[str, Any]: + negative_text: str | None, + similarity_strength: float | None, + task_type: str | None, + mask_prompt: str | None, + out_painting_mode: str | None, +) -> dict[str, Any]: """Build InvokeModel body task section (without imageGenerationConfig).""" if task_type == "BACKGROUND_REMOVAL": return { @@ -60,7 +60,7 @@ def _nova_canvas_task_body( "OUTPAINTING requires either a mask image or a mask prompt. " "Pass mask= or maskPrompt= in the request." ) - out_params: Dict[str, Any] = { + out_params: dict[str, Any] = { "image": image_b64, "text": text, } @@ -79,7 +79,7 @@ def _nova_canvas_task_body( # Honour explicit IMAGE_VARIATION even when a mask is present (mask is ignored # for this task type; callers use INPAINTING when they want mask semantics). if task_type == "IMAGE_VARIATION": - var_params_explicit: Dict[str, Any] = { + var_params_explicit: dict[str, Any] = { "images": [image_b64], "text": text, } @@ -100,7 +100,7 @@ def _nova_canvas_task_body( "or omit taskType for automatic routing (mask → INPAINTING, else IMAGE_VARIATION)." ) if mask_b64 is not None or mask_prompt is not None or task_type == "INPAINTING": - in_params: Dict[str, Any] = {"image": image_b64, "text": text} + in_params: dict[str, Any] = {"image": image_b64, "text": text} if mask_prompt is not None: in_params["maskPrompt"] = mask_prompt elif mask_b64 is not None: @@ -114,7 +114,7 @@ def _nova_canvas_task_body( "See https://docs.aws.amazon.com/nova/latest/userguide/image-gen-req-resp-structure.html" ) return {"taskType": "INPAINTING", "inPaintingParams": in_params} - var_params: Dict[str, Any] = { + var_params: dict[str, Any] = { "images": [image_b64], "text": text, } @@ -128,7 +128,7 @@ def _nova_canvas_task_body( } -def _file_types_to_b64(image: Optional[FileTypes]) -> str: +def _file_types_to_b64(image: FileTypes | None) -> str: """Encode OpenAI image input to base64 string for Nova Canvas.""" if image is None: raise ValueError("Nova Canvas image edit requires an image input") @@ -165,9 +165,9 @@ def _supports_nova_canvas_image_edit_from_model_cost(model: str) -> bool: return False seen: set[str] = set() - candidates: List[str] = [] + candidates: list[str] = [] - def _add(name: Optional[str]) -> None: + def _add(name: str | None) -> None: if name and name not in seen: seen.add(name) candidates.append(name) @@ -220,7 +220,7 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): """ @classmethod - def _is_nova_canvas_image_edit_model(cls, model: Optional[str] = None) -> bool: + def _is_nova_canvas_image_edit_model(cls, model: str | None = None) -> bool: """ Use model_cost.supports_nova_canvas_image_edit so new Nova Canvas inference IDs are added via model_prices_and_context_window.json only (not get_model_info, which @@ -250,9 +250,9 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: supported = set(self.get_supported_openai_params(model)) - mapped: Dict[str, Any] = dict(image_edit_optional_params) + mapped: dict[str, Any] = dict(image_edit_optional_params) _size = mapped.pop("size", None) if _size is not None and isinstance(_size, str) and "x" in _size: w, h = _size.split("x", 1) @@ -298,17 +298,17 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, Any]: + ) -> tuple[dict, Any]: op = dict(image_edit_optional_request_params) image_b64 = _file_types_to_b64(image) mask_raw = op.pop("mask", None) - mask_b64: Optional[str] = None + mask_b64: str | None = None if mask_raw is not None: mask_b64 = _file_types_to_b64(mask_raw) # type: ignore[arg-type] @@ -327,7 +327,7 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): cfg_scale = op.pop("cfgScale", None) seed = op.pop("seed", None) - image_generation_config: Dict[str, Any] = {} + image_generation_config: dict[str, Any] = {} nested_igc = op.pop("imageGenerationConfig", None) if isinstance(nested_igc, dict): image_generation_config.update(nested_igc) @@ -380,8 +380,8 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: try: response_data = raw_response.json() @@ -399,7 +399,7 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): headers=raw_response.headers, ) - images: List[str] = response_data.get("images") or [] + images: list[str] = response_data.get("images") or [] if "errors" in response_data and not images: raise self.get_error_class( @@ -462,7 +462,7 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: raise NotImplementedError( @@ -475,9 +475,9 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: if headers is None: headers = {} diff --git a/litellm/llms/bedrock/image_edit/handler.py b/litellm/llms/bedrock/image_edit/handler.py index 01a40c0e475..e82b1251fa6 100644 --- a/litellm/llms/bedrock/image_edit/handler.py +++ b/litellm/llms/bedrock/image_edit/handler.py @@ -7,7 +7,7 @@ Handles image edit requests for Bedrock stability models. from __future__ import annotations import json -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any import httpx from pydantic import BaseModel @@ -70,16 +70,16 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: Optional[str], + prompt: str | None, model_response: ImageResponse, optional_params: dict, logging_obj: LitellmLogging, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, aimage_edit: bool = False, - api_base: Optional[str] = None, - extra_headers: Optional[dict] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_key: Optional[str] = None, + api_base: str | None = None, + extra_headers: dict | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_key: str | None = None, ): prepared_request = self._prepare_request( model=model, @@ -132,12 +132,12 @@ class BedrockImageEdit(BaseAWSLLM): async def async_image_edit( self, prepared_request: BedrockImageEditPreparedRequest, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, model: str, logging_obj: LitellmLogging, - prompt: Optional[str], + prompt: str | None, model_response: ImageResponse, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: """ Asynchronous handler for bedrock image edit @@ -175,12 +175,12 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: Optional[str], + prompt: str | None, optional_params: dict, - api_base: Optional[str], - extra_headers: Optional[dict], + api_base: str | None, + extra_headers: dict | None, logging_obj: LitellmLogging, - api_key: Optional[str], + api_key: str | None, ) -> BedrockImageEditPreparedRequest: """ Prepare the request body, headers, and endpoint URL for the Bedrock Image Edit API @@ -258,7 +258,7 @@ class BedrockImageEdit(BaseAWSLLM): self, model: str, image: list, - prompt: Optional[str], + prompt: str | None, optional_params: dict, ) -> dict: """ @@ -286,7 +286,7 @@ class BedrockImageEdit(BaseAWSLLM): model_response: ImageResponse, model: str, logging_obj: LitellmLogging, - prompt: Optional[str], + prompt: str | None, response: httpx.Response, data: dict, ) -> ImageResponse: diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index 0b45aba219f..b9cd4624d24 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -22,7 +22,7 @@ API Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parame """ import base64 -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -51,7 +51,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): """ @classmethod - def _is_stability_edit_model(cls, model: Optional[str] = None) -> bool: + def _is_stability_edit_model(cls, model: str | None = None) -> bool: """ Returns True if the model is a Bedrock Stability edit model. @@ -100,7 +100,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI parameters to Bedrock Stability parameters. @@ -116,7 +116,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): } # Create a copy to not mutate original - convert TypedDict to regular dict - mapped_params: Dict[str, Any] = dict(image_edit_optional_params) + mapped_params: dict[str, Any] = dict(image_edit_optional_params) for k, v in image_edit_optional_params.items(): if k in param_mapping: @@ -144,27 +144,26 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): # Remove OpenAI params that have been mapped unless they're in stability for mapped in ["size", "n", "response_format"]: - if mapped in mapped_params: - del mapped_params[mapped] + mapped_params.pop(mapped, None) return mapped_params def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, Any]: + ) -> tuple[dict, Any]: """ Transform OpenAI-style request to Bedrock Stability request format. Returns the request body dict that will be JSON-encoded by the handler. """ # Build Bedrock Stability request - data: Dict[str, Any] = { + data: dict[str, Any] = { "output_format": "png", # Default to PNG } @@ -273,8 +272,8 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Bedrock Stability response to OpenAI-compatible ImageResponse. @@ -349,7 +348,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -369,9 +368,9 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment for Bedrock Stability image edit. diff --git a/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py b/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py index 626baf707a5..674182feda2 100644 --- a/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py @@ -1,8 +1,9 @@ import types -from typing import Any, Dict, List, Optional +from typing import Any from openai.types.image import Image +from litellm.llms.bedrock.common_utils import get_cached_model_info from litellm.types.llms.bedrock import ( AmazonNovaCanvasColorGuidedGenerationParams, AmazonNovaCanvasColorGuidedRequest, @@ -14,7 +15,6 @@ from litellm.types.llms.bedrock import ( AmazonNovaCanvasTextToImageRequest, AmazonNovaCanvasTextToImageResponse, ) -from litellm.llms.bedrock.common_utils import get_cached_model_info from litellm.types.utils import ImageResponse @@ -43,12 +43,12 @@ class AmazonNovaCanvasConfig: } @classmethod - def get_supported_openai_params(cls, model: Optional[str] = None) -> List: + def get_supported_openai_params(cls, model: str | None = None) -> list: """ """ return ["n", "size", "quality"] @classmethod - def _is_nova_model(cls, model: Optional[str] = None) -> bool: + def _is_nova_model(cls, model: str | None = None) -> bool: """ Returns True if the model is a Nova Canvas model @@ -73,7 +73,7 @@ class AmazonNovaCanvasConfig: image_generation_config = {**image_generation_config, **optional_params} if task_type == "TEXT_IMAGE": - text_to_image_params: Dict[str, Any] = image_generation_config.pop("textToImageParams", {}) + text_to_image_params: dict[str, Any] = image_generation_config.pop("textToImageParams", {}) text_to_image_params = {"text": text, **text_to_image_params} try: text_to_image_params_typed = AmazonNovaCanvasTextToImageParams( @@ -97,7 +97,7 @@ class AmazonNovaCanvasConfig: imageGenerationConfig=image_generation_config_typed, ) if task_type == "COLOR_GUIDED_GENERATION": - color_guided_generation_params: Dict[str, Any] = image_generation_config.pop( + color_guided_generation_params: dict[str, Any] = image_generation_config.pop( "colorGuidedGenerationParams", {} ) color_guided_generation_params = { @@ -126,7 +126,7 @@ class AmazonNovaCanvasConfig: imageGenerationConfig=image_generation_config_typed, ) if task_type == "INPAINTING": - inpainting_params: Dict[str, Any] = image_generation_config.pop("inpaintingParams", {}) + inpainting_params: dict[str, Any] = image_generation_config.pop("inpaintingParams", {}) inpainting_params = {"text": text, **inpainting_params} try: inpainting_params_typed = AmazonNovaCanvasInpaintingParams( @@ -181,7 +181,7 @@ class AmazonNovaCanvasConfig: """ nova_response = AmazonNovaCanvasTextToImageResponse(**response_dict) - openai_images: List[Image] = [] + openai_images: list[Image] = [] for _img in nova_response.get("images", []): openai_images.append(Image(b64_json=_img)) @@ -193,8 +193,8 @@ class AmazonNovaCanvasConfig: cls, model: str, image_response: ImageResponse, - size: Optional[str] = None, - optional_params: Optional[dict] = None, + size: str | None = None, + optional_params: dict | None = None, ) -> float: get_model_info = get_cached_model_info() model_info = get_model_info( diff --git a/litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py b/litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py index 0e8214fd81f..a1b46e7706b 100644 --- a/litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py @@ -1,7 +1,6 @@ import copy import os import types -from typing import List, Optional from openai.types.image import Image @@ -38,19 +37,19 @@ class AmazonStabilityConfig: - SD v1.6: must be between 320x320 and 1536x1536 """ - cfg_scale: Optional[int] = None - seed: Optional[float] = None - steps: Optional[List[str]] = None - width: Optional[int] = None - height: Optional[int] = None + cfg_scale: int | None = None + seed: float | None = None + steps: list[str] | None = None + width: int | None = None + height: int | None = None def __init__( self, - cfg_scale: Optional[int] = None, - seed: Optional[float] = None, - steps: Optional[List[str]] = None, - width: Optional[int] = None, - height: Optional[int] = None, + cfg_scale: int | None = None, + seed: float | None = None, + steps: list[str] | None = None, + width: int | None = None, + height: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -76,7 +75,7 @@ class AmazonStabilityConfig: } @classmethod - def get_supported_openai_params(cls, model: Optional[str] = None) -> List: + def get_supported_openai_params(cls, model: str | None = None) -> list: return ["size"] @classmethod @@ -120,7 +119,7 @@ class AmazonStabilityConfig: def transform_response_dict_to_openai_response( cls, model_response: ImageResponse, response_dict: dict ) -> ImageResponse: - image_list: List[Image] = [] + image_list: list[Image] = [] for artifact in response_dict["artifacts"]: _image = Image(b64_json=artifact["base64"]) image_list.append(_image) @@ -134,8 +133,8 @@ class AmazonStabilityConfig: cls, model: str, image_response: ImageResponse, - size: Optional[str] = None, - optional_params: Optional[dict] = None, + size: str | None = None, + optional_params: dict | None = None, ) -> float: optional_params = optional_params or {} diff --git a/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py b/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py index a5449679941..25393c0bda9 100644 --- a/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py @@ -1,14 +1,12 @@ import types -from typing import List, Optional from openai.types.image import Image -from litellm.llms.bedrock.common_utils import BedrockError +from litellm.llms.bedrock.common_utils import BedrockError, get_cached_model_info from litellm.types.llms.bedrock import ( AmazonStability3TextToImageRequest, AmazonStability3TextToImageResponse, ) -from litellm.llms.bedrock.common_utils import get_cached_model_info from litellm.types.utils import ImageResponse @@ -38,14 +36,14 @@ class AmazonStability3Config: } @classmethod - def get_supported_openai_params(cls, model: Optional[str] = None) -> List: + def get_supported_openai_params(cls, model: str | None = None) -> list: """ No additional OpenAI params are mapped for stability 3 """ return [] @classmethod - def _is_stability_3_model(cls, model: Optional[str] = None) -> bool: + def _is_stability_3_model(cls, model: str | None = None) -> bool: """ Returns True if the model is a Stability 3 model @@ -98,7 +96,7 @@ class AmazonStability3Config: if len(finish_reasons) > 0: raise BedrockError(status_code=400, message="; ".join(finish_reasons)) - openai_images: List[Image] = [] + openai_images: list[Image] = [] for _img in stability_3_response.get("images", []): openai_images.append(Image(b64_json=_img)) @@ -110,8 +108,8 @@ class AmazonStability3Config: cls, model: str, image_response: ImageResponse, - size: Optional[str] = None, - optional_params: Optional[dict] = None, + size: str | None = None, + optional_params: dict | None = None, ) -> float: get_model_info = get_cached_model_info() model_info = get_model_info( diff --git a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py index 5a975b6ab11..03550c1bc9f 100644 --- a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py @@ -3,17 +3,16 @@ Transformation logic for Amazon Titan Image Generation. """ import types -from typing import List, Optional from openai.types.image import Image -from litellm.utils import get_model_info from litellm.types.llms.bedrock import ( AmazonNovaCanvasImageGenerationConfig, AmazonTitanImageGenerationRequestBody, AmazonTitanTextToImageParams, ) from litellm.types.utils import ImageResponse +from litellm.utils import get_model_info class AmazonTitanImageGenerationConfig: @@ -21,19 +20,19 @@ class AmazonTitanImageGenerationConfig: Reference: https://us-west-2.console.aws.amazon.com/bedrock/home?region=us-west-2#/providers?model=stability.stable-diffusion-xl-v0 """ - cfg_scale: Optional[int] = None - seed: Optional[float] = None - steps: Optional[List[str]] = None - width: Optional[int] = None - height: Optional[int] = None + cfg_scale: int | None = None + seed: float | None = None + steps: list[str] | None = None + width: int | None = None + height: int | None = None def __init__( self, - cfg_scale: Optional[int] = None, - seed: Optional[float] = None, - steps: Optional[List[str]] = None, - width: Optional[int] = None, - height: Optional[int] = None, + cfg_scale: int | None = None, + seed: float | None = None, + steps: list[str] | None = None, + width: int | None = None, + height: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -59,7 +58,7 @@ class AmazonTitanImageGenerationConfig: } @classmethod - def _is_titan_model(cls, model: Optional[str] = None) -> bool: + def _is_titan_model(cls, model: str | None = None) -> bool: """ Returns True if the model is a Titan model @@ -71,7 +70,7 @@ class AmazonTitanImageGenerationConfig: return False @classmethod - def get_supported_openai_params(cls, model: Optional[str] = None) -> List: + def get_supported_openai_params(cls, model: str | None = None) -> list: return ["size", "n", "quality"] @classmethod @@ -80,9 +79,9 @@ class AmazonTitanImageGenerationConfig: non_default_params: dict, optional_params: dict, ): - from typing import Any, Dict + from typing import Any - image_generation_config: Dict[str, Any] = {} + image_generation_config: dict[str, Any] = {} for k, v in non_default_params.items(): if k == "size" and v is not None: width, height = v.split("x") @@ -106,11 +105,11 @@ class AmazonTitanImageGenerationConfig: text: str, optional_params: dict, ) -> AmazonTitanImageGenerationRequestBody: - from typing import Any, Dict + from typing import Any image_generation_config = optional_params.pop("imageGenerationConfig", {}) negative_text = optional_params.pop("negativeText", None) - text_to_image_params: Dict[str, Any] = {"text": text} + text_to_image_params: dict[str, Any] = {"text": text} if negative_text: text_to_image_params["negativeText"] = negative_text task_type = optional_params.pop("taskType", "TEXT_IMAGE") @@ -129,7 +128,7 @@ class AmazonTitanImageGenerationConfig: def transform_response_dict_to_openai_response( cls, model_response: ImageResponse, response_dict: dict ) -> ImageResponse: - image_list: List[Image] = [] + image_list: list[Image] = [] for image in response_dict["images"]: _image = Image(b64_json=image) image_list.append(_image) @@ -143,8 +142,8 @@ class AmazonTitanImageGenerationConfig: cls, model: str, image_response: ImageResponse, - size: Optional[str] = None, - optional_params: Optional[dict] = None, + size: str | None = None, + optional_params: dict | None = None, ) -> float: model_info = get_model_info(model=model) output_cost_per_image = model_info.get("output_cost_per_image") or 0.0 diff --git a/litellm/llms/bedrock/image_generation/cost_calculator.py b/litellm/llms/bedrock/image_generation/cost_calculator.py index b04acc3e809..2455f88cb0a 100644 --- a/litellm/llms/bedrock/image_generation/cost_calculator.py +++ b/litellm/llms/bedrock/image_generation/cost_calculator.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.bedrock.image_generation.image_handler import BedrockImageGeneration from litellm.types.utils import ImageResponse @@ -7,8 +5,8 @@ from litellm.types.utils import ImageResponse def cost_calculator( model: str, image_response: ImageResponse, - size: Optional[str] = None, - optional_params: Optional[dict] = None, + size: str | None = None, + optional_params: dict | None = None, ) -> float: """ Bedrock image generation cost calculator diff --git a/litellm/llms/bedrock/image_generation/image_handler.py b/litellm/llms/bedrock/image_generation/image_handler.py index 03e40565d95..024f26e60eb 100644 --- a/litellm/llms/bedrock/image_generation/image_handler.py +++ b/litellm/llms/bedrock/image_generation/image_handler.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union import httpx from pydantic import BaseModel @@ -80,12 +80,12 @@ class BedrockImageGeneration(BaseAWSLLM): model_response: ImageResponse, optional_params: dict, logging_obj: LitellmLogging, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, aimg_generation: bool = False, - api_base: Optional[str] = None, - extra_headers: Optional[dict] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_key: Optional[str] = None, + api_base: str | None = None, + extra_headers: dict | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_key: str | None = None, ): prepared_request = self._prepare_request( model=model, @@ -136,12 +136,12 @@ class BedrockImageGeneration(BaseAWSLLM): async def async_image_generation( self, prepared_request: BedrockImagePreparedRequest, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, model: str, logging_obj: LitellmLogging, prompt: str, model_response: ImageResponse, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: """ Asynchronous handler for bedrock image generation @@ -196,11 +196,11 @@ class BedrockImageGeneration(BaseAWSLLM): self, model: str, optional_params: dict, - api_base: Optional[str], - extra_headers: Optional[dict], + api_base: str | None, + extra_headers: dict | None, logging_obj: LitellmLogging, prompt: str, - api_key: Optional[str], + api_key: str | None, ) -> BedrockImagePreparedRequest: """ Prepare the request body, headers, and endpoint URL for the Bedrock Image Generation API diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 7b6084c03fc..e4eef800847 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -2,11 +2,6 @@ from collections.abc import AsyncIterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Union, cast, ) @@ -77,7 +72,7 @@ class AmazonAnthropicClaudeMessagesConfig( DEFAULT_BEDROCK_ANTHROPIC_API_VERSION = "bedrock-2023-05-31" @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock" BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS = frozenset(BedrockInvokeAnthropicMessagesRequest.__annotations__.keys()) @@ -90,12 +85,12 @@ class AmazonAnthropicClaudeMessagesConfig( self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: return headers, api_base def sign_request( @@ -104,11 +99,11 @@ class AmazonAnthropicClaudeMessagesConfig( optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: return AmazonInvokeConfig.sign_request( self=self, headers=headers, @@ -123,12 +118,12 @@ class AmazonAnthropicClaudeMessagesConfig( def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return AmazonInvokeConfig.get_complete_url( self=self, @@ -140,7 +135,7 @@ class AmazonAnthropicClaudeMessagesConfig( stream=stream, ) - def _remove_ttl_from_cache_control(self, anthropic_messages_request: Dict, model: Optional[str] = None) -> None: + def _remove_ttl_from_cache_control(self, anthropic_messages_request: dict, model: str | None = None) -> None: """ Remove unsupported fields from cache_control for Bedrock. @@ -234,7 +229,7 @@ class AmazonAnthropicClaudeMessagesConfig( def _ensure_thinking_for_clear_thinking_context_management( self, - anthropic_messages_request: Dict, + anthropic_messages_request: dict, model: str, ) -> bool: """ @@ -463,14 +458,14 @@ class AmazonAnthropicClaudeMessagesConfig( # Bedrock InvokeModel DOES support ``clear_tool_uses_20250919`` under the # ``context-management-2025-06-27`` beta. AWS docs: # https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-tool-use.md - _BEDROCK_INVOKE_SUPPORTED_CONTEXT_MANAGEMENT_EDITS: Dict[str, str] = { + _BEDROCK_INVOKE_SUPPORTED_CONTEXT_MANAGEMENT_EDITS: dict[str, str] = { "compact_20260112": ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value, "clear_tool_uses_20250919": ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value, } @staticmethod def _filter_context_management_for_bedrock_invoke( - anthropic_messages_request: Dict, + anthropic_messages_request: dict, beta_set: set, ) -> None: """ @@ -514,15 +509,15 @@ class AmazonAnthropicClaudeMessagesConfig( def _get_bedrock_invoke_anthropic_beta_headers( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, headers: dict, - anthropic_messages_request: Dict, + anthropic_messages_request: dict, injected_thinking_for_clear_thinking: bool, - ) -> List[str]: + ) -> list[str]: anthropic_model_info = AnthropicModelInfo() tools = anthropic_messages_optional_request_params.get("tools") - messages_typed = cast(List[AllMessageValues], messages) + messages_typed = cast(list[AllMessageValues], messages) tool_search_used = anthropic_model_info.is_tool_search_used(tools) programmatic_tool_calling_used = anthropic_model_info.is_programmatic_tool_calling_used(tools) input_examples_used = anthropic_model_info.is_input_examples_used(tools) @@ -583,8 +578,8 @@ class AmazonAnthropicClaudeMessagesConfig( def _strip_unsupported_bedrock_invoke_fields( self, - anthropic_messages_request: Dict, - ) -> Dict: + anthropic_messages_request: dict, + ) -> dict: allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS stripped = sorted(k for k in anthropic_messages_request if k not in allowed) if stripped: @@ -595,7 +590,7 @@ class AmazonAnthropicClaudeMessagesConfig( return {k: v for k, v in anthropic_messages_request.items() if k in allowed} @staticmethod - def _clamp_adaptive_reasoning_effort_for_bedrock(model: str, optional_params: Dict) -> None: + def _clamp_adaptive_reasoning_effort_for_bedrock(model: str, optional_params: dict) -> None: """Lower ``reasoning_effort`` to the Bedrock effort ceiling before validation. The shared ``/v1/messages`` effort gate rejects tiers a model does not @@ -617,11 +612,11 @@ class AmazonAnthropicClaudeMessagesConfig( def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: self._clamp_adaptive_reasoning_effort_for_bedrock( model=model, optional_params=anthropic_messages_optional_request_params, @@ -768,7 +763,7 @@ class AmazonAnthropicClaudeMessagesConfig( async def bedrock_sse_wrapper( self, - completion_stream: AsyncIterator[Union[bytes, GenericStreamingChunk, ModelResponseStream, dict]], + completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict], litellm_logging_obj: LiteLLMLoggingObj, request_body: dict, ): @@ -800,8 +795,8 @@ class AmazonAnthropicClaudeMessagesConfig( @staticmethod def _merge_message_start_cache_into_delta_usage( - delta_usage: Dict[str, Any], - start_usage: Optional[Dict[str, Any]], + delta_usage: dict[str, Any], + start_usage: dict[str, Any] | None, ) -> None: """ Copy cache breakdown from message_start onto message_delta usage when @@ -821,16 +816,16 @@ class AmazonAnthropicClaudeMessagesConfig( @staticmethod async def _promote_message_stop_usage( - completion_stream: AsyncIterator[Union[bytes, GenericStreamingChunk, ModelResponseStream, dict]], - ) -> AsyncIterator[Union[bytes, GenericStreamingChunk, ModelResponseStream, dict]]: + completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict], + ) -> AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict]: """ Promote cache usage fields onto message_delta from message_stop (and, when stop lacks them, from message_start). Ensures the final usage chunk that logging/cost sees is always self-consistent. """ _CACHE_FIELDS = ("cache_creation_input_tokens", "cache_read_input_tokens") - pending_delta: Optional[Dict[str, Any]] = None - start_usage_snapshot: Optional[Dict[str, Any]] = None + pending_delta: dict[str, Any] | None = None + start_usage_snapshot: dict[str, Any] | None = None async for chunk in completion_stream: if not isinstance(chunk, dict): @@ -843,7 +838,7 @@ class AmazonAnthropicClaudeMessagesConfig( chunk_type = chunk.get("type") if chunk_type == "message_start": - msg: Dict[str, Any] = cast(Dict[str, Any], chunk.get("message") or {}) + msg: dict[str, Any] = cast(dict[str, Any], chunk.get("message") or {}) u = msg.get("usage") if isinstance(u, dict): start_usage_snapshot = dict(u) @@ -854,7 +849,7 @@ class AmazonAnthropicClaudeMessagesConfig( continue if chunk_type == "message_delta": - pending_delta = cast(Dict[str, Any], chunk) + pending_delta = cast(dict[str, Any], chunk) continue if chunk_type == "message_stop" and pending_delta is not None: @@ -908,7 +903,7 @@ class AmazonAnthropicClaudeMessagesStreamDecoder(AWSEventStreamDecoder): super().__init__(model=model) self.DEFAULT_CHUNK_SIZE = 1024 - def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]: + def _chunk_parser(self, chunk_data: dict) -> GChunk | ModelResponseStream | dict: """ Parse the chunk data into anthropic /messages format diff --git a/litellm/llms/bedrock/messages/mantle_transformation.py b/litellm/llms/bedrock/messages/mantle_transformation.py index 2be6ba736c2..f7714ac2352 100644 --- a/litellm/llms/bedrock/messages/mantle_transformation.py +++ b/litellm/llms/bedrock/messages/mantle_transformation.py @@ -7,7 +7,7 @@ stripping that are specific to the bedrock-mantle endpoint. """ from collections.abc import AsyncIterator -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx @@ -42,12 +42,12 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: region = self._get_aws_region_name(optional_params=optional_params, model=model) return build_mantle_messages_url( @@ -60,12 +60,12 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: headers, api_base = super().validate_anthropic_messages_environment( headers=headers, model=model, @@ -83,11 +83,11 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig): def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: # Strip "mantle/" routing prefix to get the real model ID model_id = model.replace("mantle/", "", 1) diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index 137f1e333eb..f660f4c74fe 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -37,10 +37,10 @@ def _generic_passthrough_handler() -> BaseTranslation: return PassThroughEndpointHandler() -_StringHolder = Tuple[Any, Union[str, int]] +_StringHolder = tuple[Any, str | int] -def _collect_strings(node: Any, holders: List[_StringHolder]) -> None: +def _collect_strings(node: Any, holders: list[_StringHolder]) -> None: """ Record a (container, key) holder for every non-empty string value nested under an arbitrary JSON node, so prompt content a caller hides in fields @@ -48,7 +48,7 @@ def _collect_strings(node: Any, holders: List[_StringHolder]) -> None: and can be written back in place. Iterative to avoid unbounded recursion on deeply nested payloads. """ - stack: List[Any] = [node] + stack: list[Any] = [node] while stack: current = stack.pop() if isinstance(current, dict): @@ -67,7 +67,7 @@ def _collect_strings(node: Any, holders: List[_StringHolder]) -> None: stack.append(value) -def _collect_block_text(block: dict, holders: List[_StringHolder]) -> None: +def _collect_block_text(block: dict, holders: list[_StringHolder]) -> None: text = block.get("text") if isinstance(text, str) and text: holders.append((block, "text")) @@ -77,7 +77,7 @@ def _extract_converse_texts( body: dict, skip_system: bool, skip_tool: bool, -) -> Tuple[List[str], List[_StringHolder]]: +) -> tuple[list[str], list[_StringHolder]]: """ Walk a Bedrock Converse request body and collect text content. @@ -92,7 +92,7 @@ def _extract_converse_texts( message blocks are skipped when tool messages are excluded, but tool definitions are always scanned to match the chat-completions guardrail path. """ - holders: List[_StringHolder] = [] + holders: list[_StringHolder] = [] if not skip_system: for block in body.get("system") or []: @@ -129,8 +129,8 @@ def _extract_converse_texts( def _extract_converse_output_texts( - content_blocks: List[Any], -) -> Tuple[List[str], List[_StringHolder]]: + content_blocks: list[Any], +) -> tuple[list[str], list[_StringHolder]]: """ Collect user-visible text from Bedrock Converse output content blocks. @@ -139,7 +139,7 @@ def _extract_converse_output_texts( ``citationsContent.content[].text`` -- while leaving structural values such as reasoning signatures and citation sources untouched. """ - holders: List[_StringHolder] = [] + holders: list[_StringHolder] = [] for block in content_blocks: if not isinstance(block, dict): continue @@ -162,8 +162,8 @@ def _extract_converse_output_texts( def _write_back_texts( - guardrailed_texts: List[str], - holders: List[_StringHolder], + guardrailed_texts: list[str], + holders: list[_StringHolder], ) -> None: if len(guardrailed_texts) < len(holders): verbose_proxy_logger.warning( @@ -178,10 +178,10 @@ def _write_back_texts( container[key] = guardrailed_texts[idx] -_DeltaHolder = Tuple[Any, Any, Union[str, int]] +_DeltaHolder = tuple[Any, Any, str | int] -def _collect_stream_delta_text_holders(delta: Any) -> List[_DeltaHolder]: +def _collect_stream_delta_text_holders(delta: Any) -> list[_DeltaHolder]: """ Collect the user-visible text strings a Bedrock Converse ``contentBlockDelta`` can carry, matching the coverage of the non-streaming output handler. @@ -193,7 +193,7 @@ def _collect_stream_delta_text_holders(delta: Any) -> List[_DeltaHolder]: values such as reasoning signatures, redacted reasoning and citation sources are left out so they are never rewritten. """ - holders: List[_DeltaHolder] = [] + holders: list[_DeltaHolder] = [] if not isinstance(delta, dict): return holders if isinstance(delta.get("text"), str): @@ -263,7 +263,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): frames.append({"raw": frame_raw, "texts": []}) continue - texts: List[Tuple[Any, str]] = [] + texts: list[tuple[Any, str]] = [] if event_type == "contentBlockDelta": try: payload_dict = _json.loads(payload_bytes) @@ -282,8 +282,8 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): trailing_bytes = body_bytes[offset:] - group_order: List[Any] = [] - group_members: dict[Any, list[Tuple[int, int]]] = {} + group_order: list[Any] = [] + group_members: dict[Any, list[tuple[int, int]]] = {} group_texts: dict[Any, list[str]] = {} for frame_idx, frame in enumerate(frames): for local_idx, (group_key, text) in enumerate(frame["texts"]): @@ -328,7 +328,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): except (KeyError, IndexError, TypeError): return body_bytes - new_text_map: dict[Tuple[int, int], str] = {} + new_text_map: dict[tuple[int, int], str] = {} for group_key, de_anonymized_text in zip(active_groups, de_anonymized_texts): members = group_members[group_key] orig_texts = group_texts[group_key] @@ -431,8 +431,8 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): response: Any, guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: endpoint = (request_data or {}).get("endpoint", "") if endpoint and not _is_converse_endpoint(endpoint): diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index cc8840526f0..6733e887930 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -1,5 +1,5 @@ import json -from typing import TYPE_CHECKING, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Optional, cast from httpx import Response @@ -36,9 +36,10 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD Returns: The encoded model_id suitable for use in endpoint URLs """ - from litellm.passthrough.utils import CommonUtils import re + from litellm.passthrough.utils import CommonUtils + # Create a temporary endpoint with the model_id to check if encoding is needed temp_endpoint = f"/model/{model_id}/converse" encoded_temp_endpoint = CommonUtils.encode_bedrock_runtime_modelid_arn(temp_endpoint) @@ -53,13 +54,13 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, endpoint: str, - request_query_params: Optional[dict], + request_query_params: dict | None, litellm_params: dict, - ) -> Tuple["URL", str]: + ) -> tuple["URL", str]: optional_params = litellm_params.copy() model_id = optional_params.get("model_id", None) @@ -96,10 +97,10 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD self, headers: dict, litellm_params: dict, - request_data: Optional[dict], + request_data: dict | None, api_base: str, - model: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + model: str | None = None, + ) -> tuple[dict, bytes | None]: optional_params = litellm_params.copy() return self._sign_request( service_name="bedrock", @@ -153,7 +154,7 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD return litellm_model_response - def _convert_raw_bytes_to_str_lines(self, raw_bytes: List[bytes]) -> List[str]: + def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]: from botocore.eventstream import EventStreamBuffer all_chunks = [] @@ -169,7 +170,7 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD def handle_logging_collected_chunks( self, - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: "LiteLLMLoggingObj", model: str, custom_llm_provider: str, diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index b7237d288ec..17007f48fb0 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -7,7 +7,7 @@ This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic. import asyncio import contextlib import json -from typing import Any, Optional +from typing import Any from pydantic import TypeAdapter @@ -32,20 +32,20 @@ class BedrockRealtime(BaseAWSLLM): model: str, websocket: Any, logging_obj: LiteLLMLogging, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Optional[float] = None, - aws_region_name: Optional[str] = None, - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_session_name: Optional[str] = None, - aws_profile_name: Optional[str] = None, - aws_web_identity_token: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, - aws_bedrock_runtime_endpoint: Optional[str] = None, - aws_external_id: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, + timeout: float | None = None, + aws_region_name: str | None = None, + aws_access_key_id: str | None = None, + aws_secret_access_key: str | None = None, + aws_session_token: str | None = None, + aws_role_name: str | None = None, + aws_session_name: str | None = None, + aws_profile_name: str | None = None, + aws_web_identity_token: str | None = None, + aws_sts_endpoint: str | None = None, + aws_bedrock_runtime_endpoint: str | None = None, + aws_external_id: str | None = None, **kwargs, ): """ @@ -175,7 +175,7 @@ class BedrockRealtime(BaseAWSLLM): except Exception as e: verbose_proxy_logger.exception(f"Error in BedrockRealtime.async_realtime: {e}") try: - await websocket.close(code=1011, reason=_redact_string(f"Internal error: {str(e)}")) + await websocket.close(code=1011, reason=_redact_string(f"Internal error: {e!s}")) except Exception: pass raise diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py index 24a40ebea1b..39f5d25cf89 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -7,7 +7,7 @@ Transforms between OpenAI Realtime API format and Bedrock Nova Sonic format. import base64 import json import uuid as uuid_lib -from typing import Any, List, Optional, Union +from typing import Any from pydantic import BaseModel @@ -40,7 +40,7 @@ from litellm.utils import get_empty_usage class BedrockContentEnd(BaseModel): - stopReason: Optional[str] = None + stopReason: str | None = None TRIGGER_AUDIO_SAMPLE_RATE_HERTZ = 16000 @@ -87,11 +87,11 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): # Text configuration self.text_media_type = "text/plain" - def validate_environment(self, headers: dict, model: str, api_key: Optional[str] = None) -> dict: + def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict: """Validate environment - no special validation needed for Bedrock.""" return headers - def get_complete_url(self, api_base: Optional[str], model: str, api_key: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_key: str | None = None) -> str: """Get complete URL - handled by aws_sdk_bedrock_runtime.""" return api_base or "" @@ -99,7 +99,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): """Bedrock requires session configuration.""" return True - def session_configuration_request(self, model: str, tools: Optional[List[dict]] = None) -> str: + def session_configuration_request(self, model: str, tools: list[dict] | None = None) -> str: """ Create initial session configuration for Bedrock Nova Sonic. @@ -145,7 +145,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): # Return as a marker that we've sent the configuration return json.dumps({"session_start": session_start, "prompt_start": prompt_start}) - def _transform_tools_to_bedrock_format(self, tools: List[dict]) -> List[dict]: + def _transform_tools_to_bedrock_format(self, tools: list[dict]) -> list[dict]: """ Transform OpenAI tool format to Bedrock tool format. @@ -188,7 +188,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): return 8000 # G.711 typically uses 8kHz return 24000 if is_output else 16000 - def transform_session_update_event(self, json_message: dict) -> List[str]: + def transform_session_update_event(self, json_message: dict) -> list[str]: """ Transform session.update event to Bedrock session configuration. @@ -199,7 +199,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): List of Bedrock format messages (JSON strings) """ verbose_logger.debug("Handling session.update") - messages: List[str] = [] + messages: list[str] = [] session_config = json_message.get("session", {}) @@ -311,7 +311,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): return messages - def transform_input_audio_buffer_append_event(self, json_message: dict) -> List[str]: + def transform_input_audio_buffer_append_event(self, json_message: dict) -> list[str]: """ Transform input_audio_buffer.append event to Bedrock audio input. @@ -323,7 +323,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): """ verbose_logger.debug("Handling input_audio_buffer.append") self.client_audio_streamed = True - messages: List[str] = [] + messages: list[str] = [] if hasattr(self, "_audio_content_started") and self._audio_content_sample_rate != self.input_sample_rate_hertz: mismatched_content_end = { @@ -378,7 +378,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): return messages - def transform_input_audio_buffer_commit_event(self, json_message: dict) -> List[str]: + def transform_input_audio_buffer_commit_event(self, json_message: dict) -> list[str]: """ Transform input_audio_buffer.commit event to Bedrock audio content end. @@ -389,7 +389,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): List of Bedrock format messages (JSON strings) """ verbose_logger.debug("Handling input_audio_buffer.commit") - messages: List[str] = [] + messages: list[str] = [] if hasattr(self, "_audio_content_started"): audio_content_end = { @@ -405,7 +405,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): return messages - def transform_conversation_item_create_event(self, json_message: dict) -> List[str]: + def transform_conversation_item_create_event(self, json_message: dict) -> list[str]: """ Transform conversation.item.create event to Bedrock text input or tool result. @@ -473,7 +473,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): return messages - def transform_response_create_event(self, json_message: dict) -> List[str]: + def transform_response_create_event(self, json_message: dict) -> list[str]: """ Transform response.create event to Bedrock format. @@ -538,7 +538,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): for offset in range(0, len(pcm), TRIGGER_AUDIO_CHUNK_SIZE) ] - def transform_response_cancel_event(self, json_message: dict) -> List[str]: + def transform_response_cancel_event(self, json_message: dict) -> list[str]: """ Transform response.cancel event to Bedrock format. @@ -585,8 +585,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): self, message: str, model: str, - session_configuration_request: Optional[str] = None, - ) -> List[str]: + session_configuration_request: str | None = None, + ) -> list[str]: """ Transform OpenAI realtime request to Bedrock Nova Sonic format. @@ -665,15 +665,15 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_content_start_event( self, event: dict, - current_response_id: Optional[str], - current_output_item_id: Optional[str], - current_conversation_id: Optional[str], + current_response_id: str | None, + current_output_item_id: str | None, + current_conversation_id: str | None, ) -> tuple[ - List[OpenAIRealtimeEvents], - Optional[str], - Optional[str], - Optional[str], - Optional[ALL_DELTA_TYPES], + list[OpenAIRealtimeEvents], + str | None, + str | None, + str | None, + ALL_DELTA_TYPES | None, ]: """ Transform Bedrock contentStart event to OpenAI response events. @@ -713,7 +713,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): content_type = content_start.get("type", "TEXT") current_delta_type: ALL_DELTA_TYPES = "text" if content_type == "TEXT" else "audio" - returned_messages: List[OpenAIRealtimeEvents] = [] + returned_messages: list[OpenAIRealtimeEvents] = [] # Send response.created response_created = OpenAIRealtimeStreamResponseBaseObject( @@ -770,10 +770,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_text_output_event( self, event: dict, - current_output_item_id: Optional[str], - current_response_id: Optional[str], - current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], - ) -> tuple[List[OpenAIRealtimeEvents], Optional[List[OpenAIRealtimeResponseDelta]]]: + current_output_item_id: str | None, + current_response_id: str | None, + current_delta_chunks: list[OpenAIRealtimeResponseDelta] | None, + ) -> tuple[list[OpenAIRealtimeEvents], list[OpenAIRealtimeResponseDelta] | None]: """ Transform Bedrock textOutput event to OpenAI response.text.delta. @@ -812,9 +812,9 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_audio_output_event( self, event: dict, - current_output_item_id: Optional[str], - current_response_id: Optional[str], - ) -> List[OpenAIRealtimeEvents]: + current_output_item_id: str | None, + current_response_id: str | None, + ) -> list[OpenAIRealtimeEvents]: """ Transform Bedrock audioOutput event to OpenAI response.audio.delta. @@ -847,11 +847,11 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_content_end_event( self, event: dict, - current_output_item_id: Optional[str], - current_response_id: Optional[str], - current_delta_type: Optional[str], - current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], - ) -> tuple[List[OpenAIRealtimeEvents], Optional[List[OpenAIRealtimeResponseDelta]]]: + current_output_item_id: str | None, + current_response_id: str | None, + current_delta_type: str | None, + current_delta_chunks: list[OpenAIRealtimeResponseDelta] | None, + ) -> tuple[list[OpenAIRealtimeEvents], list[OpenAIRealtimeResponseDelta] | None]: """ Transform Bedrock contentEnd event to OpenAI response done events. @@ -871,7 +871,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): if not current_output_item_id or not current_response_id: return [], current_delta_chunks - returned_messages: List[OpenAIRealtimeEvents] = [] + returned_messages: list[OpenAIRealtimeEvents] = [] # Send appropriate done event based on type if current_delta_type == "text": @@ -949,13 +949,13 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_prompt_end_event( self, event: dict, - current_response_id: Optional[str], - current_conversation_id: Optional[str], + current_response_id: str | None, + current_conversation_id: str | None, ) -> tuple[ - List[OpenAIRealtimeEvents], - Optional[str], - Optional[str], - Optional[ALL_DELTA_TYPES], + list[OpenAIRealtimeEvents], + str | None, + str | None, + ALL_DELTA_TYPES | None, ]: """ Transform a Bedrock end-of-response event (promptEnd, completionEnd, or an @@ -974,13 +974,13 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def _response_done_events( self, - current_response_id: Optional[str], - current_conversation_id: Optional[str], + current_response_id: str | None, + current_conversation_id: str | None, ) -> tuple[ - List[OpenAIRealtimeEvents], - Optional[str], - Optional[str], - Optional[ALL_DELTA_TYPES], + list[OpenAIRealtimeEvents], + str | None, + str | None, + ALL_DELTA_TYPES | None, ]: if not current_response_id or not current_conversation_id: return [], None, None, None @@ -1009,9 +1009,9 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_tool_use_event( self, event: dict, - current_output_item_id: Optional[str], - current_response_id: Optional[str], - ) -> tuple[List[OpenAIRealtimeEvents], str, str]: + current_output_item_id: str | None, + current_response_id: str | None, + ) -> tuple[list[OpenAIRealtimeEvents], str, str]: """ Transform Bedrock toolUse event to OpenAI format. @@ -1061,7 +1061,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): tool_name, ) - def transform_conversation_item_create_tool_result_event(self, json_message: dict) -> List[str]: + def transform_conversation_item_create_tool_result_event(self, json_message: dict) -> list[str]: """ Transform conversation.item.create with tool result to Bedrock format. @@ -1072,7 +1072,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): List of Bedrock format messages (JSON strings) """ verbose_logger.debug("Handling conversation.item.create for tool result") - messages: List[str] = [] + messages: list[str] = [] item = json_message.get("item", {}) if item.get("type") == "function_call_output": @@ -1126,7 +1126,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def transform_realtime_response( self, - message: Union[str, bytes], + message: str | bytes, model: str, logging_obj: LiteLLMLoggingObj, realtime_response_transform_input: RealtimeResponseTransformInput, @@ -1169,7 +1169,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): current_delta_type = realtime_response_transform_input.get("current_delta_type") session_configuration_request = realtime_response_transform_input.get("session_configuration_request") - returned_messages: List[OpenAIRealtimeEvents] = [] + returned_messages: list[OpenAIRealtimeEvents] = [] # Parse Bedrock event event = json_message.get("event", {}) diff --git a/litellm/llms/bedrock/rerank/handler.py b/litellm/llms/bedrock/rerank/handler.py index 1728f52a413..fb359bc65e5 100644 --- a/litellm/llms/bedrock/rerank/handler.py +++ b/litellm/llms/bedrock/rerank/handler.py @@ -1,5 +1,5 @@ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx @@ -29,8 +29,8 @@ class BedrockRerankHandler(BaseAWSLLM): async def arerank( self, prepared_request: BedrockPreparedRequest, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ): if client is None: client = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) @@ -54,18 +54,18 @@ class BedrockRerankHandler(BaseAWSLLM): self, model: str, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], optional_params: dict, logging_obj: LitellmLogging, - top_n: Optional[int] = None, - rank_fields: Optional[List[str]] = None, - return_documents: Optional[bool] = True, - max_chunks_per_doc: Optional[int] = None, - _is_async: Optional[bool] = False, - timeout: Optional[Union[float, httpx.Timeout]] = None, - api_base: Optional[str] = None, - extra_headers: Optional[dict] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + top_n: int | None = None, + rank_fields: list[str] | None = None, + return_documents: bool | None = True, + max_chunks_per_doc: int | None = None, + _is_async: bool | None = False, + timeout: float | httpx.Timeout | None = None, + api_base: str | None = None, + extra_headers: dict | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> RerankResponse: request_data = RerankRequest( model=model, @@ -130,8 +130,8 @@ class BedrockRerankHandler(BaseAWSLLM): def _prepare_request( self, model: str, - api_base: Optional[str], - extra_headers: Optional[dict], + api_base: str | None, + extra_headers: dict | None, data: dict, optional_params: dict, ) -> BedrockPreparedRequest: diff --git a/litellm/llms/bedrock/rerank/transformation.py b/litellm/llms/bedrock/rerank/transformation.py index 38625a26939..dd060f681a4 100644 --- a/litellm/llms/bedrock/rerank/transformation.py +++ b/litellm/llms/bedrock/rerank/transformation.py @@ -5,8 +5,6 @@ Why separate file? Make it easy to see how transformation works """ from litellm._uuid import uuid -from typing import List, Optional, Union - from litellm.types.llms.bedrock import ( BedrockRerankBedrockRerankingConfiguration, BedrockRerankConfiguration, @@ -29,7 +27,7 @@ from litellm.types.rerank import ( class BedrockRerankConfig: - def _transform_sources(self, documents: List[Union[str, dict]]) -> List[BedrockRerankSource]: + def _transform_sources(self, documents: list[str | dict]) -> list[BedrockRerankSource]: """ Transform the sources from RerankRequest format to Bedrock format. """ @@ -88,7 +86,7 @@ class BedrockRerankConfig: _tokens = RerankTokens(**response.get("usage", {})) rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) - _results: Optional[List[RerankResponseResult]] = None + _results: list[RerankResponseResult] | None = None bedrock_results = response.get("results") if bedrock_results: diff --git a/litellm/llms/bedrock/vector_stores/transformation.py b/litellm/llms/bedrock/vector_stores/transformation.py index c1b124caec1..eec14a3aeb2 100644 --- a/litellm/llms/bedrock/vector_stores/transformation.py +++ b/litellm/llms/bedrock/vector_stores/transformation.py @@ -1,5 +1,5 @@ from copy import deepcopy -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast from urllib.parse import urlparse import httpx @@ -10,15 +10,15 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.types.integrations.rag.bedrock_knowledgebase import ( BedrockKBContent, - BedrockKBRetrievalConfiguration, BedrockKBResponse, + BedrockKBRetrievalConfiguration, BedrockKBRetrievalQuery, ) from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + VECTOR_STORE_OPENAI_PARAMS, BaseVectorStoreAuthCredentials, VectorStoreIndexEndpoints, - VECTOR_STORE_OPENAI_PARAMS, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, @@ -47,7 +47,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): "write": [], } - def get_supported_openai_params(self, model: str) -> List[VECTOR_STORE_OPENAI_PARAMS]: + def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]: return ["filters", "max_num_results", "ranking_options"] def _map_operator_to_aws(self, operator: str) -> str: @@ -157,7 +157,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): # 1. check if filter is in openai format # 2. if it is, map it to the aws kb filters format # 3. if it is not, assume it is in aws kb filters format and add it to the optional_params - aws_filters: Optional[Dict] = None + aws_filters: dict | None = None if isinstance(value, dict): if "operator" in value.keys(): @@ -172,12 +172,12 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): return optional_params - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: headers = headers or {} headers.setdefault("Content-Type", "application/json") return headers - def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str: + def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str: aws_region_name = litellm_params.get("aws_region_name") endpoint_url, _ = self.get_runtime_endpoint( api_base=api_base, @@ -190,24 +190,24 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: if isinstance(query, list): query = " ".join(query) encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url = f"{api_base}/{encoded_vector_store_id}/retrieve" - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "retrievalQuery": BedrockKBRetrievalQuery(text=query), } - retrieval_config: Dict[str, Any] = {} + retrieval_config: dict[str, Any] = {} if isinstance(extra_body, dict): retrieval_config = deepcopy( @@ -240,11 +240,11 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): def sign_request( self, headers: dict, - optional_params: Dict, - request_data: Dict, + optional_params: dict, + request_data: dict, api_base: str, - api_key: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + ) -> tuple[dict, bytes | None]: return self._sign_request( service_name="bedrock", headers=headers, @@ -254,7 +254,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): api_key=api_key, ) - def _get_file_id_from_metadata(self, metadata: Dict[str, Any]) -> str: + def _get_file_id_from_metadata(self, metadata: dict[str, Any]) -> str: """ Extract file_id from Bedrock KB metadata. Uses source URI if available, otherwise generates a fallback ID. @@ -266,7 +266,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): chunk_id = metadata.get("x-amz-bedrock-kb-chunk-id", "unknown") if metadata else "unknown" return f"bedrock-kb-{chunk_id}" - def _get_filename_from_metadata(self, metadata: Dict[str, Any]) -> str: + def _get_filename_from_metadata(self, metadata: dict[str, Any]) -> str: """ Extract filename from Bedrock KB metadata. Tries to extract filename from source URI, falls back to domain name or data source ID. @@ -288,7 +288,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): data_source_id = metadata.get("x-amz-bedrock-kb-data-source-id", "unknown") if metadata else "unknown" return f"bedrock-kb-document-{data_source_id}" - def _get_attributes_from_metadata(self, metadata: Dict[str, Any]) -> Dict[str, Any]: + def _get_attributes_from_metadata(self, metadata: dict[str, Any]) -> dict[str, Any]: """ Extract all attributes from Bedrock KB metadata. Returns a copy of the metadata dictionary. @@ -302,9 +302,9 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): ) -> VectorStoreSearchResponse: try: response_data = BedrockKBResponse(**response.json()) - results: List[VectorStoreSearchResult] = [] + results: list[VectorStoreSearchResult] = [] for item in response_data.get("retrievalResults", []) or []: - content: Optional[BedrockKBContent] = item.get("content") + content: BedrockKBContent | None = item.get("content") text = content.get("text") if content else None if text is None: continue @@ -341,7 +341,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): self, vector_store_create_optional_params, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: raise NotImplementedError def transform_create_vector_store_response(self, response: httpx.Response): diff --git a/litellm/llms/bedrock_mantle/chat/transformation.py b/litellm/llms/bedrock_mantle/chat/transformation.py index 007258b50f6..d85edbd0c86 100644 --- a/litellm/llms/bedrock_mantle/chat/transformation.py +++ b/litellm/llms/bedrock_mantle/chat/transformation.py @@ -11,7 +11,7 @@ Auth: Bearer token (litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the """ from collections.abc import AsyncIterator, Iterator -from typing import Any, List, Optional, Tuple, Union +from typing import Any import litellm from litellm._logging import verbose_logger @@ -38,7 +38,7 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): self._aws_signer = aws_signer or BaseAWSLLM() @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "bedrock_mantle" @classmethod @@ -47,11 +47,11 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - litellm_params: Optional[GenericLiteLLMParams] = None, + api_base: str | None, + api_key: str | None, + litellm_params: GenericLiteLLMParams | None = None, model: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: region = ( (litellm_params.aws_region_name if litellm_params else None) or get_secret_str("BEDROCK_MANTLE_REGION") @@ -75,11 +75,11 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = super().validate_environment( headers=headers, @@ -107,9 +107,9 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], Any], + streaming_response: Iterator[str] | AsyncIterator[str] | Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: from litellm.llms.openai.chat.gpt_transformation import ( OpenAIChatCompletionStreamingHandler, diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index eedb57ea386..d3c0af7f932 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -13,7 +13,6 @@ global state. """ import re -from typing import Tuple from botocore.exceptions import ( CredentialRetrievalError, @@ -66,7 +65,7 @@ class BedrockMantleAuthMixin: model: str | None = None, stream: bool | None = None, fake_stream: bool | None = None, - ) -> Tuple[dict, bytes | None]: + ) -> tuple[dict, bytes | None]: bearer = self._resolve_bearer_token(api_key) if not bearer: # Pin the credential-scope region to the region of the actual signing URL diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index 08579b6bf0d..a5b143a2679 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -15,7 +15,7 @@ role / access key / profile / web identity), signed via the shared BaseAWSLLM._sign_request after the request body is finalized. """ -from typing import Any, Dict, List, Optional +from typing import Any import litellm from litellm._logging import verbose_logger @@ -54,7 +54,7 @@ _CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE = "additional_tools" class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPIConfig): def __init__( self, - aws_signer: Optional[BaseAWSLLM] = None, + aws_signer: BaseAWSLLM | None = None, use_openai_path: bool = True, ): super().__init__() @@ -67,7 +67,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: region = self._resolve_region({**litellm_params, "api_base": api_base}) @@ -85,7 +85,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI path = "/openai/v1/responses" if self.use_openai_path else "/v1/responses" return f"{base}{path}" - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() bearer = self._resolve_bearer_token(litellm_params.api_key) if bearer: @@ -101,10 +101,10 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI return False @staticmethod - def _filter_unsupported_tools(tools: List[Any]) -> List[Any]: + def _filter_unsupported_tools(tools: list[Any]) -> list[Any]: """Keep only tool types Mantle's Responses API accepts.""" - kept: List[Any] = [] - dropped_types: List[str] = [] + kept: list[Any] = [] + dropped_types: list[str] = [] for tool in tools: if not isinstance(tool, dict): kept.append(tool) @@ -215,7 +215,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: params = self._handle_unsupported_service_tier( super().map_openai_params( response_api_optional_params=response_api_optional_params, diff --git a/litellm/llms/black_forest_labs/__init__.py b/litellm/llms/black_forest_labs/__init__.py index 7a78638c8c7..a7cb7ff52bf 100644 --- a/litellm/llms/black_forest_labs/__init__.py +++ b/litellm/llms/black_forest_labs/__init__.py @@ -10,12 +10,12 @@ from .image_edit import BlackForestLabsImageEditConfig from .image_generation import BlackForestLabsImageGenerationConfig __all__ = [ - "BlackForestLabsError", - "BlackForestLabsImageEditConfig", - "BlackForestLabsImageGenerationConfig", "DEFAULT_API_BASE", "DEFAULT_MAX_POLLING_TIME", "DEFAULT_POLLING_INTERVAL", "IMAGE_EDIT_MODELS", "IMAGE_GENERATION_MODELS", + "BlackForestLabsError", + "BlackForestLabsImageEditConfig", + "BlackForestLabsImageGenerationConfig", ] diff --git a/litellm/llms/black_forest_labs/common_utils.py b/litellm/llms/black_forest_labs/common_utils.py index 71c09093679..818c0c76914 100644 --- a/litellm/llms/black_forest_labs/common_utils.py +++ b/litellm/llms/black_forest_labs/common_utils.py @@ -4,7 +4,6 @@ Black Forest Labs Common Utilities Common utilities, constants, and error handling for Black Forest Labs API. """ -from typing import Dict from urllib.parse import urlparse from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -13,8 +12,6 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException class BlackForestLabsError(BaseLLMException): """Exception class for Black Forest Labs API errors.""" - pass - # API Constants DEFAULT_API_BASE = "https://api.bfl.ai" @@ -58,7 +55,7 @@ DEFAULT_POLLING_INTERVAL = 1.5 # seconds DEFAULT_MAX_POLLING_TIME = 300 # 5 minutes # Model to endpoint mapping for image edit -IMAGE_EDIT_MODELS: Dict[str, str] = { +IMAGE_EDIT_MODELS: dict[str, str] = { "flux-kontext-pro": "/v1/flux-kontext-pro", "flux-kontext-max": "/v1/flux-kontext-max", "flux-pro-1.0-fill": "/v1/flux-pro-1.0-fill", @@ -66,7 +63,7 @@ IMAGE_EDIT_MODELS: Dict[str, str] = { } # Model to endpoint mapping for image generation -IMAGE_GENERATION_MODELS: Dict[str, str] = { +IMAGE_GENERATION_MODELS: dict[str, str] = { "flux-pro-1.1": "/v1/flux-pro-1.1", "flux-pro-1.1-ultra": "/v1/flux-pro-1.1-ultra", "flux-dev": "/v1/flux-dev", diff --git a/litellm/llms/black_forest_labs/image_edit/__init__.py b/litellm/llms/black_forest_labs/image_edit/__init__.py index 73af716e062..efbd3e8b26a 100644 --- a/litellm/llms/black_forest_labs/image_edit/__init__.py +++ b/litellm/llms/black_forest_labs/image_edit/__init__.py @@ -2,7 +2,7 @@ from .handler import BlackForestLabsImageEdit, bfl_image_edit from .transformation import BlackForestLabsImageEditConfig __all__ = [ - "BlackForestLabsImageEditConfig", "BlackForestLabsImageEdit", + "BlackForestLabsImageEditConfig", "bfl_image_edit", ] diff --git a/litellm/llms/black_forest_labs/image_edit/handler.py b/litellm/llms/black_forest_labs/image_edit/handler.py index ab191c165fd..62aaa6da77a 100644 --- a/litellm/llms/black_forest_labs/image_edit/handler.py +++ b/litellm/llms/black_forest_labs/image_edit/handler.py @@ -8,7 +8,7 @@ then we poll until the result is ready. import asyncio import time -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -47,16 +47,16 @@ class BlackForestLabsImageEdit: def image_edit( self, model: str, - image: Union[FileTypes, List[FileTypes]], - prompt: Optional[str], - image_edit_optional_request_params: Dict, - litellm_params: Union[GenericLiteLLMParams, Dict], + image: FileTypes | list[FileTypes], + prompt: str | None, + image_edit_optional_request_params: dict, + litellm_params: GenericLiteLLMParams | dict, logging_obj: LiteLLMLoggingObj, - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[Dict[str, Any]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout | None, + extra_headers: dict[str, Any] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, aimage_edit: bool = False, - ) -> Union[ImageResponse, Any]: + ) -> ImageResponse | Any: """ Main entry point for image edit requests. @@ -159,7 +159,7 @@ class BlackForestLabsImageEdit: except Exception as e: raise BlackForestLabsError( status_code=500, - message=f"Request failed: {str(e)}", + message=f"Request failed: {e!s}", ) # Poll for result @@ -179,14 +179,14 @@ class BlackForestLabsImageEdit: async def async_image_edit( self, model: str, - image: Union[FileTypes, List[FileTypes]], - prompt: Optional[str], - image_edit_optional_request_params: Dict, - litellm_params: Union[GenericLiteLLMParams, Dict], + image: FileTypes | list[FileTypes], + prompt: str | None, + image_edit_optional_request_params: dict, + litellm_params: GenericLiteLLMParams | dict, logging_obj: LiteLLMLoggingObj, - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[Dict[str, Any]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None, + extra_headers: dict[str, Any] | None = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: """ Async version of image edit. @@ -262,7 +262,7 @@ class BlackForestLabsImageEdit: except Exception as e: raise BlackForestLabsError( status_code=500, - message=f"Request failed: {str(e)}", + message=f"Request failed: {e!s}", ) # Poll for result @@ -286,7 +286,7 @@ class BlackForestLabsImageEdit: sync_client: HTTPHandler, max_wait: float = DEFAULT_MAX_POLLING_TIME, interval: float = DEFAULT_POLLING_INTERVAL, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ) -> httpx.Response: """ Poll BFL API until result is ready (sync version). @@ -388,7 +388,7 @@ class BlackForestLabsImageEdit: async_client: AsyncHTTPHandler, max_wait: float = DEFAULT_MAX_POLLING_TIME, interval: float = DEFAULT_POLLING_INTERVAL, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ) -> httpx.Response: """ Poll BFL API until result is ready (async version). diff --git a/litellm/llms/black_forest_labs/image_edit/transformation.py b/litellm/llms/black_forest_labs/image_edit/transformation.py index a80ca491d74..7bf57819b47 100644 --- a/litellm/llms/black_forest_labs/image_edit/transformation.py +++ b/litellm/llms/black_forest_labs/image_edit/transformation.py @@ -9,7 +9,7 @@ API Reference: https://docs.bfl.ai/ import base64 import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -51,7 +51,7 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): This class only handles data transformation. """ - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Return list of OpenAI params supported by Black Forest Labs. @@ -78,13 +78,13 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI parameters to Black Forest Labs parameters. BFL-specific params are passed through directly. """ - optional_params: Dict[str, Any] = {} + optional_params: dict[str, Any] = {} # Pass through BFL-specific params bfl_params = [ @@ -124,16 +124,16 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Black Forest Labs. BFL uses x-key header for authentication. """ - final_api_key: Optional[str] = ( + final_api_key: str | None = ( api_key or get_secret_str("BFL_API_KEY") or get_secret_str("BLACK_FOREST_LABS_API_KEY") ) @@ -175,7 +175,7 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -230,12 +230,12 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform OpenAI-style request to Black Forest Labs request format. @@ -246,7 +246,7 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): b64_image = base64.b64encode(image_bytes).decode("utf-8") # Build request body - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "prompt": prompt, "input_image": b64_image, } @@ -314,7 +314,7 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): ) def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + self, error_message: str, status_code: int, headers: dict | httpx.Headers ) -> BlackForestLabsError: """Return the appropriate error class for Black Forest Labs.""" return BlackForestLabsError( diff --git a/litellm/llms/black_forest_labs/image_generation/__init__.py b/litellm/llms/black_forest_labs/image_generation/__init__.py index 2ccee2069ef..95940afb272 100644 --- a/litellm/llms/black_forest_labs/image_generation/__init__.py +++ b/litellm/llms/black_forest_labs/image_generation/__init__.py @@ -5,8 +5,8 @@ from .transformation import ( ) __all__ = [ - "BlackForestLabsImageGenerationConfig", - "get_black_forest_labs_image_generation_config", "BlackForestLabsImageGeneration", + "BlackForestLabsImageGenerationConfig", "bfl_image_generation", + "get_black_forest_labs_image_generation_config", ] diff --git a/litellm/llms/black_forest_labs/image_generation/handler.py b/litellm/llms/black_forest_labs/image_generation/handler.py index f797fac4193..af321fad580 100644 --- a/litellm/llms/black_forest_labs/image_generation/handler.py +++ b/litellm/llms/black_forest_labs/image_generation/handler.py @@ -8,7 +8,7 @@ then we poll until the result is ready. import asyncio import time -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -49,14 +49,14 @@ class BlackForestLabsImageGeneration: model: str, prompt: str, model_response: ImageResponse, - optional_params: Dict, - litellm_params: Union[GenericLiteLLMParams, Dict], + optional_params: dict, + litellm_params: GenericLiteLLMParams | dict, logging_obj: LiteLLMLoggingObj, - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[Dict[str, Any]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout | None, + extra_headers: dict[str, Any] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, aimg_generation: bool = False, - ) -> Union[ImageResponse, Any]: + ) -> ImageResponse | Any: """ Main entry point for image generation requests. @@ -156,7 +156,7 @@ class BlackForestLabsImageGeneration: except Exception as e: raise BlackForestLabsError( status_code=500, - message=f"Request failed: {str(e)}", + message=f"Request failed: {e!s}", ) # Poll for result @@ -183,12 +183,12 @@ class BlackForestLabsImageGeneration: model: str, prompt: str, model_response: ImageResponse, - optional_params: Dict, - litellm_params: Union[GenericLiteLLMParams, Dict], + optional_params: dict, + litellm_params: GenericLiteLLMParams | dict, logging_obj: LiteLLMLoggingObj, - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[Dict[str, Any]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None, + extra_headers: dict[str, Any] | None = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: """ Async version of image generation. @@ -262,7 +262,7 @@ class BlackForestLabsImageGeneration: except Exception as e: raise BlackForestLabsError( status_code=500, - message=f"Request failed: {str(e)}", + message=f"Request failed: {e!s}", ) # Poll for result @@ -291,7 +291,7 @@ class BlackForestLabsImageGeneration: sync_client: HTTPHandler, max_wait: float = DEFAULT_MAX_POLLING_TIME, interval: float = DEFAULT_POLLING_INTERVAL, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ) -> httpx.Response: """ Poll BFL API until result is ready (sync version). @@ -382,7 +382,7 @@ class BlackForestLabsImageGeneration: async_client: AsyncHTTPHandler, max_wait: float = DEFAULT_MAX_POLLING_TIME, interval: float = DEFAULT_POLLING_INTERVAL, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ) -> httpx.Response: """ Poll BFL API until result is ready (async version). diff --git a/litellm/llms/black_forest_labs/image_generation/transformation.py b/litellm/llms/black_forest_labs/image_generation/transformation.py index 7176247b4be..535c290bf5b 100644 --- a/litellm/llms/black_forest_labs/image_generation/transformation.py +++ b/litellm/llms/black_forest_labs/image_generation/transformation.py @@ -8,7 +8,7 @@ API Reference: https://docs.bfl.ai/ """ import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -50,7 +50,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): This class only handles data transformation. """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Return list of OpenAI params supported by Black Forest Labs. @@ -140,18 +140,18 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Black Forest Labs. BFL uses x-key header for authentication. """ - final_api_key: Optional[str] = ( + final_api_key: str | None = ( api_key or get_secret_str("BFL_API_KEY") or get_secret_str("BLACK_FOREST_LABS_API_KEY") ) @@ -187,12 +187,12 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the Black Forest Labs API request. @@ -217,7 +217,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): https://docs.bfl.ai/flux_models/flux_1_1_pro """ # Build request body with prompt - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "prompt": prompt, } @@ -257,8 +257,8 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Black Forest Labs response to OpenAI-compatible ImageResponse. @@ -300,7 +300,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): return model_response def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + self, error_message: str, status_code: int, headers: dict | httpx.Headers ) -> BlackForestLabsError: """Return the appropriate error class for Black Forest Labs.""" return BlackForestLabsError( diff --git a/litellm/llms/brave/search/transformation.py b/litellm/llms/brave/search/transformation.py index 54fb574087c..3f89020bab1 100644 --- a/litellm/llms/brave/search/transformation.py +++ b/litellm/llms/brave/search/transformation.py @@ -4,11 +4,13 @@ Documentation: https://api-dashboard.search.brave.com/app/documentation/web-sear """ from __future__ import annotations -from datetime import datetime, timezone -from dateutil import parser # type: ignore[import-untyped] -from typing import Dict, List, Literal, Optional, TypedDict, Union -import httpx + import re +from datetime import datetime, timezone +from typing import Literal, TypedDict + +import httpx +from dateutil import parser # type: ignore[import-untyped] _ISO_YMD = re.compile(r"^\s*\d{4}[-/]\d{1,2}[-/]\d{1,2}\s*$") _UNIX_TIMESTAMP = re.compile(r"^\s*-?\d+(\.\d+)?\s*$") @@ -20,16 +22,15 @@ from litellm.llms.base_llm.search.transformation import ( SearchResponse, SearchResult, ) - from litellm.secret_managers.main import get_secret_str def to_yyyy_mm_dd( - s: Union[str, int, float, None], + s: str | float | None, *, dayfirst: bool = False, yearfirst: bool = False, -) -> Optional[str]: +) -> str | None: """ Convert a string/int/float to YYYY-MM-DD; return None if parsing fails. """ @@ -107,11 +108,11 @@ class BraveSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -135,9 +136,9 @@ class BraveSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -160,12 +161,12 @@ class BraveSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - api_key: Optional[str] = None, - search_engine_id: Optional[str] = None, + api_key: str | None = None, + search_engine_id: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Brave Search API format. @@ -227,7 +228,7 @@ class BraveSearchConfig(BaseSearchConfig): } @staticmethod - def _append_domain_filters(query: str, domains: List[str]) -> str: + def _append_domain_filters(query: str, domains: list[str]) -> str: """ Add site: filters to emulate domain restriction in Brave. """ @@ -239,7 +240,7 @@ class BraveSearchConfig(BaseSearchConfig): def transform_search_response( self, raw_response: httpx.Response, - logging_obj: Optional[LiteLLMLoggingObj], + logging_obj: LiteLLMLoggingObj | None, **kwargs, ) -> SearchResponse: """ @@ -248,7 +249,7 @@ class BraveSearchConfig(BaseSearchConfig): response_json = raw_response.json() # Transform results to SearchResult objects - results: List[SearchResult] = [] + results: list[SearchResult] = [] query_params = raw_response.request.url.params if raw_response.request else {} sections_to_process = self._sections_from_params(dict(query_params)) @@ -285,14 +286,14 @@ class BraveSearchConfig(BaseSearchConfig): ) @staticmethod - def _sections_from_params(query_params: dict) -> List[str]: + def _sections_from_params(query_params: dict) -> list[str]: """ Returns a list of sections the user has requested via the Brave Search API's `result_filter` parameter. If no `result_filter` parameter is provided, returns all sections. """ raw_filter = query_params.get("result_filter") - requested_filters: List[str] = [] + requested_filters: list[str] = [] if raw_filter and isinstance(raw_filter, str): requested_filters = [part.strip() for part in raw_filter.split(",")] diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index e5d91c6533f..22dc39040c4 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -1,13 +1,13 @@ import json import time import traceback -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx -from litellm.litellm_core_utils.url_utils import encode_url_path_segments from litellm.litellm_core_utils.exception_mapping_utils import exception_type from litellm.litellm_core_utils.logging_utils import track_llm_api_timing +from litellm.litellm_core_utils.url_utils import encode_url_path_segments from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -77,7 +77,7 @@ class BytezChatConfig(BaseConfig): "web_search_options": False, } - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: supported_params = [] for key, value in self.openai_to_bytez_param_map.items(): if value: @@ -117,11 +117,11 @@ class BytezChatConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers.update( { @@ -141,12 +141,12 @@ class BytezChatConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: encoded_model = encode_url_path_segments(model, field_name="model") return f"{API_BASE}/{encoded_model}" @@ -154,7 +154,7 @@ class BytezChatConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -182,12 +182,12 @@ class BytezChatConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: json = raw_response.json() @@ -254,9 +254,9 @@ class BytezChatConfig(BaseConfig): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -296,9 +296,9 @@ class BytezChatConfig(BaseConfig): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={}) @@ -328,9 +328,7 @@ class BytezChatConfig(BaseConfig): ) return streaming_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BytezError(status_code=status_code, message=error_message) @@ -338,7 +336,7 @@ class BytezCustomStreamWrapper(CustomStreamWrapper): def chunk_creator(self, chunk: Any): try: model_response = self.model_response_creator() - response_obj: Dict[str, Any] = {} + response_obj: dict[str, Any] = {} response_obj = { "text": chunk, @@ -346,7 +344,7 @@ class BytezCustomStreamWrapper(CustomStreamWrapper): "finish_reason": "", } - completion_obj: Dict[str, Any] = {"content": chunk} + completion_obj: dict[str, Any] = {"content": chunk} return self.return_processed_chunk_logic( completion_obj=completion_obj, @@ -377,7 +375,7 @@ open_ai_to_bytez_content_item_map = { } -def adapt_messages_to_bytez_standard(messages: List[Dict]): +def adapt_messages_to_bytez_standard(messages: list[dict]): messages = _adapt_string_only_content_to_lists(messages) new_messages = [] @@ -389,7 +387,7 @@ def adapt_messages_to_bytez_standard(messages: List[Dict]): new_content = [] for content_item in content: - type: Union[str, None] = content_item.get("type") + type: str | None = content_item.get("type") if not type: raise Exception("Prop `type` is not a string") @@ -403,7 +401,7 @@ def adapt_messages_to_bytez_standard(messages: List[Dict]): value_name = content_item_map["value_name"] - value: Union[str, None] = content_item.get(value_name) + value: str | None = content_item.get(value_name) if not value: raise Exception(f"Prop `{value_name}` is not a string") @@ -418,7 +416,7 @@ def adapt_messages_to_bytez_standard(messages: List[Dict]): # "content": "The cat ran so fast" # becomes # "content": [{"type": "text", "text": "The cat ran so fast"}] -def _adapt_string_only_content_to_lists(messages: List[Dict]): +def _adapt_string_only_content_to_lists(messages: list[dict]): new_messages = [] for message in messages: @@ -453,11 +451,11 @@ def _adapt_string_only_content_to_lists(messages: List[Dict]): # TODO get this from the api instead of doing it here, will require backend work -def get_tokens_from_messages(messages: List[dict]): +def get_tokens_from_messages(messages: list[dict]): total = 0 for message in messages: - content: List[dict] = message["content"] + content: list[dict] = message["content"] for content_item in content: type = content_item["type"] diff --git a/litellm/llms/bytez/common_utils.py b/litellm/llms/bytez/common_utils.py index d6593a06b71..65742b84297 100644 --- a/litellm/llms/bytez/common_utils.py +++ b/litellm/llms/bytez/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -12,7 +10,7 @@ class BytezError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[httpx.Headers] = None, + headers: httpx.Headers | None = None, ): self.status_code = status_code self.message = message diff --git a/litellm/llms/cerebras/chat.py b/litellm/llms/cerebras/chat.py index 9929e2ab9a2..4ae8a74c5de 100644 --- a/litellm/llms/cerebras/chat.py +++ b/litellm/llms/cerebras/chat.py @@ -4,8 +4,6 @@ Cerebras Chat Completions API this is OpenAI compatible - no translation needed / occurs """ -from typing import Optional - from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.utils import supports_reasoning @@ -17,29 +15,29 @@ class CerebrasConfig(OpenAIGPTConfig): Below are the parameters: """ - max_tokens: Optional[int] = None - response_format: Optional[dict] = None - seed: Optional[int] = None - stream: Optional[bool] = None - top_p: Optional[int] = None - tool_choice: Optional[str] = None - tools: Optional[list] = None - user: Optional[str] = None - reasoning_effort: Optional[str] = None + max_tokens: int | None = None + response_format: dict | None = None + seed: int | None = None + stream: bool | None = None + top_p: int | None = None + tool_choice: str | None = None + tools: list | None = None + user: str | None = None + reasoning_effort: str | None = None def __init__( self, - max_tokens: Optional[int] = None, - response_format: Optional[dict] = None, - seed: Optional[int] = None, - stop: Optional[str] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, - top_p: Optional[int] = None, - tool_choice: Optional[str] = None, - tools: Optional[list] = None, - user: Optional[str] = None, - reasoning_effort: Optional[str] = None, + max_tokens: int | None = None, + response_format: dict | None = None, + seed: int | None = None, + stop: str | None = None, + stream: bool | None = None, + temperature: float | None = None, + top_p: int | None = None, + tool_choice: str | None = None, + tools: list | None = None, + user: str | None = None, + reasoning_effort: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 277bcfa18d0..15452d0f864 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -2,7 +2,7 @@ import base64 import json import os import time -from typing import Any, Dict, Optional +from typing import Any import httpx @@ -63,7 +63,7 @@ class Authenticator: tokens = self._login_device_code() return tokens["access_token"] - def get_account_id(self) -> Optional[str]: + def get_account_id(self) -> str | None: auth_data = self._read_auth_file() if not auth_data: return None @@ -82,24 +82,24 @@ class Authenticator: if not os.path.exists(self.token_dir): os.makedirs(self.token_dir, exist_ok=True) - def _read_auth_file(self) -> Optional[Dict[str, Any]]: + def _read_auth_file(self) -> dict[str, Any] | None: try: with open(self.auth_file, "r") as f: return json.load(f) - except IOError: + except OSError: return None except json.JSONDecodeError as exc: verbose_logger.warning("Invalid ChatGPT auth file: %s", exc) return None - def _write_auth_file(self, data: Dict[str, Any]) -> None: + def _write_auth_file(self, data: dict[str, Any]) -> None: try: with open(self.auth_file, "w") as f: json.dump(data, f) - except IOError as exc: + except OSError as exc: verbose_logger.error("Failed to write ChatGPT auth file: %s", exc) - def _is_token_expired(self, auth_data: Dict[str, Any], access_token: str) -> bool: + def _is_token_expired(self, auth_data: dict[str, Any], access_token: str) -> bool: expires_at = auth_data.get("expires_at") if expires_at is None: expires_at = self._get_expires_at(access_token) @@ -110,14 +110,14 @@ class Authenticator: return True return time.time() >= float(expires_at) - TOKEN_EXPIRY_SKEW_SECONDS - def _get_expires_at(self, token: str) -> Optional[int]: + def _get_expires_at(self, token: str) -> int | None: claims = self._decode_jwt_claims(token) exp = claims.get("exp") if isinstance(exp, (int, float)): return int(exp) return None - def _decode_jwt_claims(self, token: str) -> Dict[str, Any]: + def _decode_jwt_claims(self, token: str) -> dict[str, Any]: try: parts = token.split(".") if len(parts) < 2: @@ -129,7 +129,7 @@ class Authenticator: except Exception: return {} - def _extract_account_id(self, token: Optional[str]) -> Optional[str]: + def _extract_account_id(self, token: str | None) -> str | None: if not token: return None claims = self._decode_jwt_claims(token) @@ -140,7 +140,7 @@ class Authenticator: return account_id return None - def _login_device_code(self) -> Dict[str, str]: + def _login_device_code(self) -> dict[str, str]: cooldown_remaining = self._get_device_code_cooldown_remaining(self._read_auth_file()) if cooldown_remaining > 0: token = self._wait_for_access_token(cooldown_remaining) @@ -162,7 +162,7 @@ class Authenticator: self._write_auth_file(auth_data) return tokens - def _request_device_code(self) -> Dict[str, str]: + def _request_device_code(self) -> dict[str, str]: try: client = _get_httpx_client() resp = client.post( @@ -196,7 +196,7 @@ class Authenticator: "interval": str(interval or "5"), } - def _poll_for_authorization_code(self, device_code: Dict[str, str]) -> Dict[str, str]: + def _poll_for_authorization_code(self, device_code: dict[str, str]) -> dict[str, str]: client = _get_httpx_client() interval = int(device_code.get("interval", "5")) start_time = time.time() @@ -245,7 +245,7 @@ class Authenticator: status_code=408, ) - def _exchange_code_for_tokens(self, code_data: Dict[str, str]) -> Dict[str, str]: + def _exchange_code_for_tokens(self, code_data: dict[str, str]) -> dict[str, str]: try: client = _get_httpx_client() redirect_uri = f"{CHATGPT_AUTH_BASE}/deviceauth/callback" @@ -285,7 +285,7 @@ class Authenticator: "id_token": data["id_token"], } - def _refresh_tokens(self, refresh_token: str) -> Dict[str, str]: + def _refresh_tokens(self, refresh_token: str) -> dict[str, str]: try: client = _get_httpx_client() resp = client.post( @@ -327,7 +327,7 @@ class Authenticator: self._write_auth_file(auth_data) return refreshed - def _build_auth_record(self, tokens: Dict[str, str]) -> Dict[str, Any]: + def _build_auth_record(self, tokens: dict[str, str]) -> dict[str, Any]: access_token = tokens.get("access_token") id_token = tokens.get("id_token") expires_at = self._get_expires_at(access_token) if access_token else None @@ -340,7 +340,7 @@ class Authenticator: "account_id": account_id, } - def _get_device_code_cooldown_remaining(self, auth_data: Optional[Dict[str, Any]]) -> float: + def _get_device_code_cooldown_remaining(self, auth_data: dict[str, Any] | None) -> float: if not auth_data: return 0.0 requested_at = auth_data.get("device_code_requested_at") @@ -359,7 +359,7 @@ class Authenticator: auth_data["device_code_requested_at"] = time.time() self._write_auth_file(auth_data) - def _wait_for_access_token(self, timeout_seconds: float) -> Optional[str]: + def _wait_for_access_token(self, timeout_seconds: float) -> str | None: deadline = time.time() + timeout_seconds while time.time() < deadline: auth_data = self._read_auth_file() diff --git a/litellm/llms/chatgpt/chat/streaming_utils.py b/litellm/llms/chatgpt/chat/streaming_utils.py index 3232b452a37..953309266e6 100644 --- a/litellm/llms/chatgpt/chat/streaming_utils.py +++ b/litellm/llms/chatgpt/chat/streaming_utils.py @@ -4,7 +4,7 @@ Streaming utilities for ChatGPT provider. Normalizes non-spec-compliant tool_call chunks from the ChatGPT backend API. """ -from typing import Any, Dict, Optional +from typing import Any class ChatGPTToolCallNormalizer: @@ -22,9 +22,9 @@ class ChatGPTToolCallNormalizer: def __init__(self, stream: Any): self._stream = stream - self._seen_ids: Dict[str, int] = {} # tool_call_id -> assigned_index + self._seen_ids: dict[str, int] = {} # tool_call_id -> assigned_index self._next_index: int = 0 - self._last_id: Optional[str] = None # tracks which tool call the next delta belongs to + self._last_id: str | None = None # tracks which tool call the next delta belongs to def __getattr__(self, name: str) -> Any: return getattr(self._stream, name) diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index 9b0d8dc2e65..4433d42e5d4 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, List, Optional, Tuple +from typing import Any from litellm.exceptions import AuthenticationError from litellm.llms.openai.openai import OpenAIConfig @@ -16,8 +16,8 @@ from .streaming_utils import ChatGPTToolCallNormalizer class ChatGPTConfig(OpenAIConfig): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, custom_llm_provider: str = "openai", ) -> None: super().__init__() @@ -26,10 +26,10 @@ class ChatGPTConfig(OpenAIConfig): def _get_openai_compatible_provider_info( self, model: str, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, custom_llm_provider: str, - ) -> Tuple[Optional[str], Optional[str], str]: + ) -> tuple[str | None, str | None, str]: dynamic_api_base = self.authenticator.get_api_base() try: dynamic_api_key = self.authenticator.get_access_token() @@ -45,11 +45,11 @@ class ChatGPTConfig(OpenAIConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: validated_headers = super().validate_environment( headers, model, messages, optional_params, litellm_params, api_key, api_base diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index 8afef4b3828..31e180488d4 100644 --- a/litellm/llms/chatgpt/common_utils.py +++ b/litellm/llms/chatgpt/common_utils.py @@ -4,7 +4,7 @@ Constants and helpers for ChatGPT subscription OAuth. import os import platform -from typing import Any, Optional, Union +from typing import Any from uuid import uuid4 import httpx @@ -110,10 +110,10 @@ class ChatGPTAuthError(BaseLLMException): self, status_code, message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, - body: Optional[dict] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, + body: dict | None = None, ): super().__init__( status_code=status_code, @@ -227,8 +227,8 @@ def get_chatgpt_user_agent(originator: str) -> str: def get_chatgpt_default_headers( access_token: str, - account_id: Optional[str], - session_id: Optional[str] = None, + account_id: str | None, + session_id: str | None = None, ) -> dict: originator = get_chatgpt_originator() user_agent = get_chatgpt_user_agent(originator) @@ -250,7 +250,7 @@ def get_chatgpt_default_instructions() -> str: return os.getenv("CHATGPT_DEFAULT_INSTRUCTIONS") or CHATGPT_DEFAULT_INSTRUCTIONS -def _normalize_litellm_params(litellm_params: Optional[Any]) -> dict: +def _normalize_litellm_params(litellm_params: Any | None) -> dict: if litellm_params is None: return {} if isinstance(litellm_params, dict): @@ -268,7 +268,7 @@ def _normalize_litellm_params(litellm_params: Optional[Any]) -> dict: return {} -def get_chatgpt_session_id(litellm_params: Optional[Any]) -> Optional[str]: +def get_chatgpt_session_id(litellm_params: Any | None) -> str | None: params = _normalize_litellm_params(litellm_params) for key in ("litellm_session_id", "session_id"): value = params.get(key) @@ -286,5 +286,5 @@ def get_chatgpt_session_id(litellm_params: Optional[Any]) -> Optional[str]: return None -def ensure_chatgpt_session_id(litellm_params: Optional[Any]) -> str: +def ensure_chatgpt_session_id(litellm_params: Any | None) -> str: return get_chatgpt_session_id(litellm_params) or str(uuid4()) diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index 8b5fae4ef35..e08df3af508 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional +from typing import Any from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.core_helpers import process_response_headers @@ -42,7 +42,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: try: access_token = self.authenticator.get_access_token() @@ -144,13 +144,11 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): or "\ndata:" in body_text ) - def _extract_completed_response_from_sse( - self, body_text: str - ) -> tuple[Optional[ResponsesAPIResponse], Optional[str]]: + def _extract_completed_response_from_sse(self, body_text: str) -> tuple[ResponsesAPIResponse | None, str | None]: completed_response = None error_message = None - streamed_output_items: Dict[int, dict] = {} - text_only_output_items: Dict[int, dict] = {} + streamed_output_items: dict[int, dict] = {} + text_only_output_items: dict[int, dict] = {} for chunk in body_text.splitlines(): parsed_chunk = parse_sse_json_chunk(chunk) if parsed_chunk is None: @@ -177,7 +175,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): # output_index, but text-only items at indices without a # matching OUTPUT_ITEM_DONE must still be preserved (e.g. # providers that emit only OUTPUT_TEXT_DONE for some indices). - merged_items: Dict[int, dict] = {**text_only_output_items} + merged_items: dict[int, dict] = {**text_only_output_items} merged_items.update(streamed_output_items) completed_response = self._build_completed_response_from_chunk( parsed_chunk=parsed_chunk, @@ -196,8 +194,8 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): return completed_response, error_message def _build_completed_response_from_chunk( - self, parsed_chunk: Dict[str, Any], streamed_output_items: Dict[int, dict] - ) -> Optional[ResponsesAPIResponse]: + self, parsed_chunk: dict[str, Any], streamed_output_items: dict[int, dict] + ) -> ResponsesAPIResponse | None: response_payload = parsed_chunk.get("response") if not isinstance(response_payload, dict): return None @@ -211,7 +209,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): except Exception: return ResponsesAPIResponse.model_construct(**response_payload) - def _extract_error_message(self, parsed_chunk: Dict[str, Any]) -> Optional[str]: + def _extract_error_message(self, parsed_chunk: dict[str, Any]) -> str | None: error_obj = parsed_chunk.get("error") or (parsed_chunk.get("response") or {}).get("error") if error_obj is None: return None @@ -233,7 +231,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = api_base or self.authenticator.get_api_base() or CHATGPT_API_BASE diff --git a/litellm/llms/clarifai/chat/transformation.py b/litellm/llms/clarifai/chat/transformation.py index 95c0444924b..147c7986f2a 100644 --- a/litellm/llms/clarifai/chat/transformation.py +++ b/litellm/llms/clarifai/chat/transformation.py @@ -1,14 +1,14 @@ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.openai.common_utils import OpenAIError from litellm.secret_managers.main import get_secret_str -from litellm.types.utils import ModelResponse from litellm.types.llms.openai import ( AllMessageValues, ) -from litellm.llms.openai.common_utils import OpenAIError -from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.types.utils import ModelResponse from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -45,15 +45,15 @@ class ClarifaiConfig(OpenAIGPTConfig): ] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("CLARIFAI_API_KEY") @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or "https://api.clarifai.com/v2/ext/openai/v1" @staticmethod - def get_base_model(model: Optional[str] = None) -> Optional[str]: + def get_base_model(model: str | None = None) -> str | None: if model: user_id, app_id, model_id = model.split(".") return f"https://clarifai.com/{user_id}/{app_id}/models/{model_id}" @@ -61,9 +61,9 @@ class ClarifaiConfig(OpenAIGPTConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: """ Get API base and key for Clarifai provider. """ @@ -82,12 +82,12 @@ class ClarifaiConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the Clarifai response to a standard ModelResponse. @@ -106,7 +106,7 @@ class ClarifaiConfig(OpenAIGPTConfig): except Exception as e: raise OpenAIError( status_code=raw_response.status_code, - message=f"Failed to parse Clarifai response: {str(e)}", + message=f"Failed to parse Clarifai response: {e!s}", headers=raw_response.headers, ) from e @@ -117,9 +117,7 @@ class ClarifaiConfig(OpenAIGPTConfig): return response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get the appropriate error class for Clarifai errors. Since Clarifai is OpenAI-compatible, we use OpenAI error handling. diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py index df8ac884a32..c81499b3b2d 100644 --- a/litellm/llms/cloudflare/chat/transformation.py +++ b/litellm/llms/cloudflare/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - import httpx from litellm._logging import verbose_logger @@ -29,12 +27,12 @@ class CloudflareError(BaseLLMException): class CloudflareChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return super().get_complete_url( api_base=self._resolve_api_base(api_base), @@ -46,7 +44,7 @@ class CloudflareChatConfig(OpenAIGPTConfig): ) @staticmethod - def _resolve_api_base(api_base: Optional[str]) -> str: + def _resolve_api_base(api_base: str | None) -> str: if not api_base: account_id = normalize_nonempty_secret_str(get_secret_str("CLOUDFLARE_ACCOUNT_ID")) if account_id is None: @@ -66,11 +64,11 @@ class CloudflareChatConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: raise ValueError( @@ -86,9 +84,7 @@ class CloudflareChatConfig(OpenAIGPTConfig): api_base=api_base, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CloudflareError( status_code=status_code, message=error_message, diff --git a/litellm/llms/codestral/completion/handler.py b/litellm/llms/codestral/completion/handler.py index b369898f4ab..1261604e6a7 100644 --- a/litellm/llms/codestral/completion/handler.py +++ b/litellm/llms/codestral/completion/handler.py @@ -4,7 +4,6 @@ import json from collections.abc import Callable from functools import partial -from typing import List, Optional, Union import httpx # type: ignore @@ -28,8 +27,8 @@ class TextCompletionCodestralError(Exception): self, status_code, message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, ): self.status_code = status_code self.message = message @@ -79,14 +78,14 @@ class CodestralTextCompletion: def _validate_environment( self, - api_key: Optional[str], + api_key: str | None, user_headers: dict, ) -> dict: if api_key is None: raise ValueError("Missing CODESTRAL_API_Key - Please add CODESTRAL_API_Key to your environment variables") headers = { "content-type": "application/json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } if user_headers is not None and isinstance(user_headers, dict): headers = {**headers, **user_headers} @@ -121,7 +120,7 @@ class CodestralTextCompletion: logging_obj: LiteLLMLogging, optional_params: dict, api_key: str, - data: Union[dict, str], + data: dict | str, messages: list, print_verbose, encoding, @@ -146,7 +145,7 @@ class CodestralTextCompletion: raise TextCompletionCodestralError(message=response.text, status_code=422) _original_choices = completion_response.get("choices", []) - _choices: List[TextChoices] = [] + _choices: list[TextChoices] = [] for choice in _original_choices: # This is what 1 choice looks like from codestral API # { @@ -197,12 +196,12 @@ class CodestralTextCompletion: api_key: str, logging_obj, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, acompletion=None, litellm_params=None, logger_fn=None, headers: dict = {}, - ) -> Union[TextCompletionResponse, CustomStreamWrapper]: + ) -> TextCompletionResponse | CustomStreamWrapper: headers = self._validate_environment(api_key, headers) if optional_params.pop("custom_endpoint", None) is True: @@ -339,7 +338,7 @@ class CodestralTextCompletion: stream, data: dict, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params=None, logger_fn=None, headers={}, @@ -353,11 +352,11 @@ class CodestralTextCompletion: except httpx.HTTPStatusError as e: raise TextCompletionCodestralError( status_code=e.response.status_code, - message="HTTPStatusError - {}".format(e.response.text), + message=f"HTTPStatusError - {e.response.text}", ) except Exception as e: raise TextCompletionCodestralError( - status_code=500, message="{}".format(str(e)) + status_code=500, message=f"{e!s}" ) # don't use verbose_logger.exception, if exception is raised return self.process_text_completion_response( model=model, @@ -385,7 +384,7 @@ class CodestralTextCompletion: api_key, logging_obj, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, optional_params=None, litellm_params=None, logger_fn=None, diff --git a/litellm/llms/codestral/completion/transformation.py b/litellm/llms/codestral/completion/transformation.py index d4299ee2ebd..4b63f454faf 100644 --- a/litellm/llms/codestral/completion/transformation.py +++ b/litellm/llms/codestral/completion/transformation.py @@ -1,5 +1,4 @@ import json -from typing import Optional import litellm from litellm.llms.openai.completion.transformation import OpenAITextCompletionConfig @@ -11,23 +10,23 @@ class CodestralTextCompletionConfig(OpenAITextCompletionConfig): Reference: https://docs.mistral.ai/api/#operation/createFIMCompletion """ - suffix: Optional[str] = None - temperature: Optional[int] = None - max_tokens: Optional[int] = None - min_tokens: Optional[int] = None - stream: Optional[bool] = None - random_seed: Optional[int] = None + suffix: str | None = None + temperature: int | None = None + max_tokens: int | None = None + min_tokens: int | None = None + stream: bool | None = None + random_seed: int | None = None def __init__( self, - suffix: Optional[str] = None, - temperature: Optional[int] = None, - top_p: Optional[float] = None, - max_tokens: Optional[int] = None, - min_tokens: Optional[int] = None, - stream: Optional[bool] = None, - random_seed: Optional[int] = None, - stop: Optional[str] = None, + suffix: str | None = None, + temperature: int | None = None, + top_p: float | None = None, + max_tokens: int | None = None, + min_tokens: int | None = None, + stream: bool | None = None, + random_seed: int | None = None, + stop: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/cohere/chat/transformation.py b/litellm/llms/cohere/chat/transformation.py index b30817cb7c3..96c25668fdc 100644 --- a/litellm/llms/cohere/chat/transformation.py +++ b/litellm/llms/cohere/chat/transformation.py @@ -1,7 +1,7 @@ import json import time from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -27,7 +27,7 @@ class CohereError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[httpx.Headers] = None, + headers: httpx.Headers | None = None, ): self.status_code = status_code self.message = message @@ -66,47 +66,47 @@ class CohereChatConfig(BaseConfig): seed (int, optional): A seed to assist reproducibility of the model's response. """ - preamble: Optional[str] = None - chat_history: Optional[list] = None - generation_id: Optional[str] = None - response_id: Optional[str] = None - conversation_id: Optional[str] = None - prompt_truncation: Optional[str] = None - connectors: Optional[list] = None - search_queries_only: Optional[bool] = None - documents: Optional[list] = None - temperature: Optional[int] = None - max_tokens: Optional[int] = None - max_completion_tokens: Optional[int] = None - k: Optional[int] = None - p: Optional[int] = None - frequency_penalty: Optional[int] = None - presence_penalty: Optional[int] = None - tools: Optional[list] = None - tool_results: Optional[list] = None - seed: Optional[int] = None + preamble: str | None = None + chat_history: list | None = None + generation_id: str | None = None + response_id: str | None = None + conversation_id: str | None = None + prompt_truncation: str | None = None + connectors: list | None = None + search_queries_only: bool | None = None + documents: list | None = None + temperature: int | None = None + max_tokens: int | None = None + max_completion_tokens: int | None = None + k: int | None = None + p: int | None = None + frequency_penalty: int | None = None + presence_penalty: int | None = None + tools: list | None = None + tool_results: list | None = None + seed: int | None = None def __init__( self, - preamble: Optional[str] = None, - chat_history: Optional[list] = None, - generation_id: Optional[str] = None, - response_id: Optional[str] = None, - conversation_id: Optional[str] = None, - prompt_truncation: Optional[str] = None, - connectors: Optional[list] = None, - search_queries_only: Optional[bool] = None, - documents: Optional[list] = None, - temperature: Optional[int] = None, - max_tokens: Optional[int] = None, - max_completion_tokens: Optional[int] = None, - k: Optional[int] = None, - p: Optional[int] = None, - frequency_penalty: Optional[int] = None, - presence_penalty: Optional[int] = None, - tools: Optional[list] = None, - tool_results: Optional[list] = None, - seed: Optional[int] = None, + preamble: str | None = None, + chat_history: list | None = None, + generation_id: str | None = None, + response_id: str | None = None, + conversation_id: str | None = None, + prompt_truncation: str | None = None, + connectors: list | None = None, + search_queries_only: bool | None = None, + documents: list | None = None, + temperature: int | None = None, + max_tokens: int | None = None, + max_completion_tokens: int | None = None, + k: int | None = None, + p: int | None = None, + frequency_penalty: int | None = None, + presence_penalty: int | None = None, + tools: list | None = None, + tool_results: list | None = None, + seed: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -117,11 +117,11 @@ class CohereChatConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return cohere_validate_environment( headers=headers, @@ -131,7 +131,7 @@ class CohereChatConfig(BaseConfig): api_key=api_key, ) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "stream", "temperature", @@ -183,7 +183,7 @@ class CohereChatConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -222,12 +222,12 @@ class CohereChatConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: raw_response_json = raw_response.json() @@ -281,7 +281,7 @@ class CohereChatConfig(BaseConfig): def _construct_cohere_tool( self, - tools: Optional[list] = None, + tools: list | None = None, ): if tools is None: tools = [] @@ -350,9 +350,9 @@ class CohereChatConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return CohereModelResponseIterator( streaming_response=streaming_response, @@ -360,7 +360,5 @@ class CohereChatConfig(BaseConfig): json_mode=json_mode, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CohereError(status_code=status_code, message=error_message) diff --git a/litellm/llms/cohere/chat/v2_transformation.py b/litellm/llms/cohere/chat/v2_transformation.py index a08411745d9..5180c30a5a3 100644 --- a/litellm/llms/cohere/chat/v2_transformation.py +++ b/litellm/llms/cohere/chat/v2_transformation.py @@ -1,6 +1,6 @@ import time from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -52,45 +52,45 @@ class CohereV2ChatConfig(OpenAIGPTConfig): seed (int, optional): A seed to assist reproducibility of the model's response. """ - preamble: Optional[str] = None - chat_history: Optional[list] = None - generation_id: Optional[str] = None - response_id: Optional[str] = None - conversation_id: Optional[str] = None - prompt_truncation: Optional[str] = None - connectors: Optional[list] = None - search_queries_only: Optional[bool] = None - documents: Optional[list] = None - temperature: Optional[int] = None - max_tokens: Optional[int] = None - k: Optional[int] = None - p: Optional[int] = None - frequency_penalty: Optional[int] = None - presence_penalty: Optional[int] = None - tools: Optional[list] = None - tool_results: Optional[list] = None - seed: Optional[int] = None + preamble: str | None = None + chat_history: list | None = None + generation_id: str | None = None + response_id: str | None = None + conversation_id: str | None = None + prompt_truncation: str | None = None + connectors: list | None = None + search_queries_only: bool | None = None + documents: list | None = None + temperature: int | None = None + max_tokens: int | None = None + k: int | None = None + p: int | None = None + frequency_penalty: int | None = None + presence_penalty: int | None = None + tools: list | None = None + tool_results: list | None = None + seed: int | None = None def __init__( self, - preamble: Optional[str] = None, - chat_history: Optional[list] = None, - generation_id: Optional[str] = None, - response_id: Optional[str] = None, - conversation_id: Optional[str] = None, - prompt_truncation: Optional[str] = None, - connectors: Optional[list] = None, - search_queries_only: Optional[bool] = None, - documents: Optional[list] = None, - temperature: Optional[int] = None, - max_tokens: Optional[int] = None, - k: Optional[int] = None, - p: Optional[int] = None, - frequency_penalty: Optional[int] = None, - presence_penalty: Optional[int] = None, - tools: Optional[list] = None, - tool_results: Optional[list] = None, - seed: Optional[int] = None, + preamble: str | None = None, + chat_history: list | None = None, + generation_id: str | None = None, + response_id: str | None = None, + conversation_id: str | None = None, + prompt_truncation: str | None = None, + connectors: list | None = None, + search_queries_only: bool | None = None, + documents: list | None = None, + temperature: int | None = None, + max_tokens: int | None = None, + k: int | None = None, + p: int | None = None, + frequency_penalty: int | None = None, + presence_penalty: int | None = None, + tools: list | None = None, + tool_results: list | None = None, + seed: int | None = None, ) -> None: locals_ = locals() for key, value in locals_.items(): @@ -101,11 +101,11 @@ class CohereV2ChatConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return cohere_validate_environment( headers=headers, @@ -115,7 +115,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig): api_key=api_key, ) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "stream", "temperature", @@ -167,7 +167,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -186,12 +186,12 @@ class CohereV2ChatConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: raw_response_json = raw_response.json() @@ -210,7 +210,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig): ) ## ADD CITATIONS AS ANNOTATIONS - annotations: Optional[List[ChatCompletionAnnotation]] = None + annotations: list[ChatCompletionAnnotation] | None = None citations = None if "message" in cohere_v2_chat_response and "citations" in cohere_v2_chat_response["message"]: @@ -223,7 +223,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig): cohere_tools_response = cohere_v2_chat_response["message"].get("tool_calls", []) if cohere_tools_response is not None and cohere_tools_response != []: # convert cohere_tools_response to OpenAI response format - tool_calls: List[ChatCompletionToolCallChunk] = [] + tool_calls: list[ChatCompletionToolCallChunk] = [] for index, tool in enumerate(cohere_tools_response): tool_call: ChatCompletionToolCallChunk = { **tool, # type: ignore @@ -258,9 +258,9 @@ class CohereV2ChatConfig(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return CohereV2ModelResponseIterator( streaming_response=streaming_response, @@ -270,12 +270,12 @@ class CohereV2ChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Cohere v2 chat completion. @@ -285,12 +285,10 @@ class CohereV2ChatConfig(OpenAIGPTConfig): raise ValueError("api_base is required") return api_base - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CohereError(status_code=status_code, message=error_message) - def _translate_citations_to_openai_annotations(self, citations: List[dict]) -> List[ChatCompletionAnnotation]: + def _translate_citations_to_openai_annotations(self, citations: list[dict]) -> list[ChatCompletionAnnotation]: """ Transform Cohere citations to OpenAI annotations format. @@ -319,7 +317,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig): Returns: List of OpenAI ChatCompletionAnnotation objects (one per source) """ - annotations: List[ChatCompletionAnnotation] = [] + annotations: list[ChatCompletionAnnotation] = [] for citation in citations: start_index = citation.get("start", 0) diff --git a/litellm/llms/cohere/common_utils.py b/litellm/llms/cohere/common_utils.py index c03061ba18f..d4aa657977b 100644 --- a/litellm/llms/cohere/common_utils.py +++ b/litellm/llms/cohere/common_utils.py @@ -1,5 +1,5 @@ import json -from typing import List, Optional, Literal, Tuple +from typing import Literal from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -21,49 +21,48 @@ class CohereModelInfo(BaseLLMModelInfo): def get_provider_info( self, model: str, - ) -> Optional[ProviderSpecificModelInfo]: + ) -> ProviderSpecificModelInfo | None: """ Default values all models of this provider support. """ return None - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: """ Returns a list of models supported by this provider. """ return [] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key @staticmethod def get_api_base( - api_base: Optional[str] = None, - ) -> Optional[str]: + api_base: str | None = None, + ) -> str | None: return api_base def validate_environment( self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {} @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: """ Returns the base model name from the given model name. Some providers like bedrock - can receive model=`invoke/anthropic.claude-3-opus-20240229-v1:0` or `converse/anthropic.claude-3-opus-20240229-v1:0` This function will return `anthropic.claude-3-opus-20240229-v1:0` """ - pass @staticmethod def get_cohere_route(model: str) -> Literal["v1", "v2"]: @@ -87,9 +86,9 @@ class CohereModelInfo(BaseLLMModelInfo): def validate_environment( headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: """ Return headers to use for cohere chat completion request @@ -116,20 +115,20 @@ def validate_environment( class ModelResponseIterator: - def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): + def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): self.streaming_response = streaming_response self.response_iterator = self.streaming_response - self.content_blocks: List = [] + self.content_blocks: list = [] self.tool_index = -1 self.json_mode = json_mode def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: try: text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None provider_specific_fields = None index = int(chunk.get("index", 0)) @@ -217,10 +216,10 @@ class ModelResponseIterator: class CohereV2ModelResponseIterator: """V2-specific response iterator for Cohere streaming""" - def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): + def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): self.streaming_response = streaming_response self.response_iterator = self.streaming_response - self.content_blocks: List = [] + self.content_blocks: list = [] self.tool_index = -1 self.json_mode = json_mode @@ -235,7 +234,7 @@ class CohereV2ModelResponseIterator: return content return "" - def _parse_tool_call_delta(self, chunk: dict) -> Optional[ChatCompletionToolCallChunk]: + def _parse_tool_call_delta(self, chunk: dict) -> ChatCompletionToolCallChunk | None: """Parse tool-call-delta chunks to extract tool calls.""" delta = chunk.get("delta", {}) tool_calls = delta.get("tool_calls", []) @@ -250,7 +249,7 @@ class CohereV2ModelResponseIterator: } # type: ignore return None - def _parse_tool_plan_delta(self, chunk: dict) -> Optional[dict]: + def _parse_tool_plan_delta(self, chunk: dict) -> dict | None: """Parse tool-plan-delta events to extract tool plan.""" data = chunk.get("data", {}) delta = data.get("delta", {}) @@ -260,7 +259,7 @@ class CohereV2ModelResponseIterator: return {"tool_plan": tool_plan} return None - def _parse_citation_start(self, chunk: dict) -> Optional[dict]: + def _parse_citation_start(self, chunk: dict) -> dict | None: """Parse citation-start events to extract citations.""" data = chunk.get("data", {}) delta = data.get("delta", {}) @@ -277,7 +276,7 @@ class CohereV2ModelResponseIterator: return {"citations": [citation_data]} return None - def _parse_message_end(self, chunk: dict) -> Tuple[bool, str, Optional[ChatCompletionUsageBlock]]: + def _parse_message_end(self, chunk: dict) -> tuple[bool, str, ChatCompletionUsageBlock | None]: """Parse message-end events to extract finish info and usage.""" data = chunk.get("data", {}) delta = data.get("delta", {}) @@ -309,10 +308,10 @@ class CohereV2ModelResponseIterator: """ try: text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None provider_specific_fields = None index = int(chunk.get("index", 0)) diff --git a/litellm/llms/cohere/embed/handler.py b/litellm/llms/cohere/embed/handler.py index 4699d55c356..dea87711cb6 100644 --- a/litellm/llms/cohere/embed/handler.py +++ b/litellm/llms/cohere/embed/handler.py @@ -4,7 +4,7 @@ Legacy /v1/embedding handler for Bedrock Cohere. import json from collections.abc import Callable -from typing import Any, Optional, Union +from typing import Any import httpx @@ -24,7 +24,7 @@ from .v1_transformation import CohereEmbeddingConfig def validate_environment(api_key, headers: dict): # Create a lowercase key lookup to avoid duplicate headers with different cases # This is important when headers come from AWS signed requests (which use Title-Case) - existing_keys_lower = {k.lower(): k for k in headers.keys()} + existing_keys_lower = {k.lower(): k for k in headers} # Only add headers if they don't already exist (case-insensitive check) if "request-source" not in existing_keys_lower: @@ -49,17 +49,17 @@ class CohereError(Exception): async def async_embedding( model: str, - data: Union[dict, CohereEmbeddingRequest], + data: dict | CohereEmbeddingRequest, input: list, model_response: litellm.utils.EmbeddingResponse, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, logging_obj: LiteLLMLoggingObj, optional_params: dict, api_base: str, - api_key: Optional[str], + api_key: str | None, headers: dict, encoding: Callable, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, ): ## LOGGING logging_obj.pre_call( @@ -121,12 +121,12 @@ def embedding( optional_params: dict, headers: dict, encoding: Any, - data: Optional[Union[dict, CohereEmbeddingRequest]] = None, - complete_api_base: Optional[str] = None, - api_key: Optional[str] = None, - aembedding: Optional[bool] = None, - timeout: Optional[Union[float, httpx.Timeout]] = httpx.Timeout(None), - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + data: dict | CohereEmbeddingRequest | None = None, + complete_api_base: str | None = None, + api_key: str | None = None, + aembedding: bool | None = None, + timeout: float | httpx.Timeout | None = httpx.Timeout(None), + client: HTTPHandler | AsyncHTTPHandler | None = None, ): headers = validate_environment(api_key, headers=headers) embed_url = complete_api_base or "https://api.cohere.ai/v1/embed" diff --git a/litellm/llms/cohere/embed/transformation.py b/litellm/llms/cohere/embed/transformation.py index 3325e6be578..fc51a992b11 100644 --- a/litellm/llms/cohere/embed/transformation.py +++ b/litellm/llms/cohere/embed/transformation.py @@ -10,7 +10,7 @@ Convers Docs - https://docs.cohere.com/v2/reference/embed """ -from typing import Any, List, Optional, Union, cast +from typing import Any, cast import httpx @@ -38,7 +38,7 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): def __init__(self) -> None: pass - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["encoding_format", "dimensions"] def map_openai_params( @@ -62,11 +62,11 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: default_headers = { "Content-Type": "application/json", @@ -81,17 +81,17 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return api_base or "https://api.cohere.ai/v2/embed" def _transform_request( - self, model: str, input: List[str], inference_params: dict + self, model: str, input: list[str], inference_params: dict ) -> CohereEmbeddingRequestWithModel: is_encoded = False for input_str in input: @@ -128,19 +128,19 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): dict, self._transform_request( model=model, - input=cast(List[str], input) if isinstance(input, List) else [input], + input=cast(list[str], input) if isinstance(input, list) else [input], inference_params=optional_params, ), ) - def _calculate_usage(self, input: List[str], encoding: Any, meta: dict) -> Usage: + def _calculate_usage(self, input: list[str], encoding: Any, meta: dict) -> Usage: input_tokens = 0 - text_tokens: Optional[int] = meta.get("billed_units", {}).get("input_tokens") + text_tokens: int | None = meta.get("billed_units", {}).get("input_tokens") - image_tokens: Optional[int] = meta.get("billed_units", {}).get("images") + image_tokens: int | None = meta.get("billed_units", {}).get("images") - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if image_tokens is None and text_tokens is None: for text in input: input_tokens += len(encoding.encode(text)) @@ -164,9 +164,9 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): def _transform_response( self, response: httpx.Response, - api_key: Optional[str], + api_key: str | None, logging_obj: LiteLLMLoggingObj, - data: Union[dict, CohereEmbeddingRequest], + data: dict | CohereEmbeddingRequest, model_response: EmbeddingResponse, model: str, encoding: Any, @@ -217,7 +217,7 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -233,9 +233,7 @@ class CohereEmbeddingConfig(BaseEmbeddingConfig): input=logging_obj.model_call_details["input"], ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CohereError( status_code=status_code, message=error_message, diff --git a/litellm/llms/cohere/embed/v1_transformation.py b/litellm/llms/cohere/embed/v1_transformation.py index 3f0fcfc03ad..058e23e1a8a 100644 --- a/litellm/llms/cohere/embed/v1_transformation.py +++ b/litellm/llms/cohere/embed/v1_transformation.py @@ -2,7 +2,7 @@ Legacy /v1/embedding transformation logic for Bedrock Cohere. """ -from typing import Any, List, Optional, Union +from typing import Any import httpx @@ -24,7 +24,7 @@ class CohereEmbeddingConfig: def __init__(self) -> None: pass - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return ["encoding_format"] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: @@ -37,7 +37,7 @@ class CohereEmbeddingConfig: return "3" in model def _transform_request( - self, model: str, input: List[str], inference_params: dict + self, model: str, input: list[str], inference_params: dict ) -> CohereEmbeddingRequestWithModel: is_encoded = False for input_str in input: @@ -61,14 +61,14 @@ class CohereEmbeddingConfig: return transformed_request - def _calculate_usage(self, input: List[str], encoding: Any, meta: dict) -> Usage: + def _calculate_usage(self, input: list[str], encoding: Any, meta: dict) -> Usage: input_tokens = 0 - text_tokens: Optional[int] = meta.get("billed_units", {}).get("input_tokens") + text_tokens: int | None = meta.get("billed_units", {}).get("input_tokens") - image_tokens: Optional[int] = meta.get("billed_units", {}).get("images") + image_tokens: int | None = meta.get("billed_units", {}).get("images") - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if image_tokens is None and text_tokens is None: for text in input: input_tokens += len(encoding.encode(text)) @@ -92,9 +92,9 @@ class CohereEmbeddingConfig: def _transform_response( self, response: httpx.Response, - api_key: Optional[str], + api_key: str | None, logging_obj: LiteLLMLoggingObj, - data: Union[dict, CohereEmbeddingRequest], + data: dict | CohereEmbeddingRequest, model_response: EmbeddingResponse, model: str, encoding: Any, diff --git a/litellm/llms/cohere/rerank/guardrail_translation/__init__.py b/litellm/llms/cohere/rerank/guardrail_translation/__init__.py index 70b580facf5..066e646f4de 100644 --- a/litellm/llms/cohere/rerank/guardrail_translation/__init__.py +++ b/litellm/llms/cohere/rerank/guardrail_translation/__init__.py @@ -8,4 +8,4 @@ guardrail_translation_mappings = { CallTypes.arerank: CohereRerankHandler, } -__all__ = ["guardrail_translation_mappings", "CohereRerankHandler"] +__all__ = ["CohereRerankHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/cohere/rerank/guardrail_translation/handler.py b/litellm/llms/cohere/rerank/guardrail_translation/handler.py index 36ca3895d4a..632c415df99 100644 --- a/litellm/llms/cohere/rerank/guardrail_translation/handler.py +++ b/litellm/llms/cohere/rerank/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for the rerank endpoint. The handler processes only the 'query' parameter for guardrails. """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -42,7 +42,7 @@ class CohereRerankHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input text fields ('query' and 'instruction') by applying @@ -94,9 +94,9 @@ class CohereRerankHandler(BaseTranslation): self, response: "RerankResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response - not applicable for rerank. diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index e494e89fbf2..b12a019a569 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -50,15 +50,15 @@ class CohereRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map Cohere rerank params @@ -106,7 +106,7 @@ class CohereRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -148,7 +148,5 @@ class CohereRerankConfig(BaseRerankConfig): return RerankResponse(**raw_response_json) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + 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/cohere/rerank_v2/transformation.py b/litellm/llms/cohere/rerank_v2/transformation.py index 7c68a431a90..eea7b41c592 100644 --- a/litellm/llms/cohere/rerank_v2/transformation.py +++ b/litellm/llms/cohere/rerank_v2/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Union +from typing import Any from litellm.llms.cohere.rerank.transformation import CohereRerankConfig from litellm.types.rerank import OptionalRerankParams, RerankRequest @@ -42,15 +42,15 @@ class CohereRerankV2Config(CohereRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map Cohere rerank params @@ -70,7 +70,7 @@ class CohereRerankV2Config(CohereRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: diff --git a/litellm/llms/cometapi/chat/transformation.py b/litellm/llms/cometapi/chat/transformation.py index 56d230057de..73e4071ed10 100644 --- a/litellm/llms/cometapi/chat/transformation.py +++ b/litellm/llms/cometapi/chat/transformation.py @@ -6,7 +6,7 @@ Documentation: [CometAPI Documentation Link] """ from collections.abc import AsyncIterator, Iterator -from typing import Any, List, Optional, Tuple, Union +from typing import Any import httpx @@ -55,9 +55,9 @@ class CometAPIConfig(OpenAIGPTConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, - messages: List[AllMessageValues], - tools: Optional[List["ChatCompletionToolParam"]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + messages: list[AllMessageValues], + tools: list["ChatCompletionToolParam"] | None = None, + ) -> tuple[list[AllMessageValues], list["ChatCompletionToolParam"] | None]: """ Remove cache control flags from messages and tools if not supported """ @@ -67,7 +67,7 @@ class CometAPIConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -85,12 +85,12 @@ class CometAPIConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the CometAPI call. @@ -123,9 +123,7 @@ class CometAPIConfig(OpenAIGPTConfig): return f"{api_base}/v1/{endpoint}" return f"{api_base}/{endpoint}" - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Return CometAPI-specific error class """ @@ -137,9 +135,9 @@ class CometAPIConfig(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: """ Get model response iterator for streaming responses diff --git a/litellm/llms/cometapi/common_utils.py b/litellm/llms/cometapi/common_utils.py index 8cb0a304026..6991962fc88 100644 --- a/litellm/llms/cometapi/common_utils.py +++ b/litellm/llms/cometapi/common_utils.py @@ -3,5 +3,3 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException class CometAPIException(BaseLLMException): """CometAPI exception handling class""" - - pass diff --git a/litellm/llms/cometapi/embed/transformation.py b/litellm/llms/cometapi/embed/transformation.py index 2d481eb1bcb..703b9fa8205 100644 --- a/litellm/llms/cometapi/embed/transformation.py +++ b/litellm/llms/cometapi/embed/transformation.py @@ -2,8 +2,6 @@ CometAPI Embedding API support - OpenAI compatible """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -29,12 +27,12 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the CometAPI embedding endpoint. @@ -47,11 +45,11 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate and set up authentication headers for CometAPI. @@ -70,7 +68,7 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): return {**default_headers, **headers} - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Get the supported OpenAI parameters for embedding requests. CometAPI supports standard OpenAI embedding parameters. @@ -115,7 +113,7 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -144,9 +142,7 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get the appropriate error class for CometAPI exceptions. """ diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py index e78b50b2fab..bf22834837a 100644 --- a/litellm/llms/cometapi/image_generation/transformation.py +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -24,7 +24,7 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://api.cometapi.com" IMAGE_GENERATION_ENDPOINT: str = "v1/images/generations" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ https://api.cometapi.com/v1/images/generations """ @@ -45,8 +45,8 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # CometAPI uses OpenAI-compatible parameters, so we can pass them directly optional_params[k] = non_default_params[k] @@ -61,12 +61,12 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -86,13 +86,13 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("COMETAPI_KEY") or get_secret_str("COMETAPI_API_KEY") + final_api_key: str | None = api_key or get_secret_str("COMETAPI_KEY") or get_secret_str("COMETAPI_API_KEY") if not final_api_key: raise ValueError("COMETAPI_KEY or COMETAPI_API_KEY is not set") @@ -131,8 +131,8 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the image generation response to the litellm image response diff --git a/litellm/llms/compactifai/chat/transformation.py b/litellm/llms/compactifai/chat/transformation.py index 2dc1ade2f4e..30bee150147 100644 --- a/litellm/llms/compactifai/chat/transformation.py +++ b/litellm/llms/compactifai/chat/transformation.py @@ -2,14 +2,14 @@ CompactifAI chat completion transformation """ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.openai.common_utils import OpenAIError from litellm.secret_managers.main import get_secret_str from litellm.types.utils import ModelResponse -from litellm.llms.openai.common_utils import OpenAIError -from litellm.llms.base_llm.chat.transformation import BaseLLMException from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -29,9 +29,9 @@ class CompactifAIChatConfig(OpenAIGPTConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: """ Get API base and key for CompactifAI provider. """ @@ -46,12 +46,12 @@ class CompactifAIChatConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List, + messages: list, optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform CompactifAI response to LiteLLM format. @@ -86,9 +86,7 @@ class CompactifAIChatConfig(OpenAIGPTConfig): return returned_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get the appropriate error class for CompactifAI errors. Since CompactifAI is OpenAI-compatible, we use OpenAI error handling. diff --git a/litellm/llms/custom_httpx/aiohttp_handler.py b/litellm/llms/custom_httpx/aiohttp_handler.py index dd78f677795..4b611d5e8e6 100644 --- a/litellm/llms/custom_httpx/aiohttp_handler.py +++ b/litellm/llms/custom_httpx/aiohttp_handler.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import aiohttp import httpx # type: ignore @@ -36,9 +36,9 @@ DEFAULT_TIMEOUT = 600 class BaseLLMAIOHTTPHandler: def __init__( self, - client_session: Optional[aiohttp.ClientSession] = None, - transport: Optional[LiteLLMAiohttpTransport] = None, - connector: Optional[aiohttp.BaseConnector] = None, + client_session: aiohttp.ClientSession | None = None, + transport: LiteLLMAiohttpTransport | None = None, + connector: aiohttp.BaseConnector | None = None, ): self.client_session = client_session self._owns_session = client_session is None # Track if we own the session for cleanup @@ -49,7 +49,7 @@ class BaseLLMAIOHTTPHandler: self.connector = connector self._owns_connector = connector is None # Track if we own the connector for cleanup - def _get_or_create_transport(self) -> Optional[LiteLLMAiohttpTransport]: + def _get_or_create_transport(self) -> LiteLLMAiohttpTransport | None: """Get existing transport or create a new one if needed.""" if self.transport: return self.transport @@ -63,7 +63,7 @@ class BaseLLMAIOHTTPHandler: # If transport creation fails, return None (will use direct session) return None - def _get_connector(self) -> Optional[aiohttp.BaseConnector]: + def _get_connector(self) -> aiohttp.BaseConnector | None: """Get or create a connector for the client session.""" if self.connector: return self.connector @@ -94,7 +94,7 @@ class BaseLLMAIOHTTPHandler: session = aiohttp.ClientSession() return session - def _get_async_client_session(self, dynamic_client_session: Optional[ClientSession] = None) -> ClientSession: + def _get_async_client_session(self, dynamic_client_session: ClientSession | None = None) -> ClientSession: if dynamic_client_session: return dynamic_client_session elif self.client_session: @@ -152,20 +152,20 @@ class BaseLLMAIOHTTPHandler: async def _make_common_async_call( self, - async_client_session: Optional[ClientSession], + async_client_session: ClientSession | None, provider_config: BaseConfig, api_base: str, headers: dict, - data: Optional[dict], - timeout: Union[float, httpx.Timeout], + data: dict | None, + timeout: float | httpx.Timeout, litellm_params: dict, - form_data: Optional[FormData] = None, + form_data: FormData | None = None, stream: bool = False, ) -> aiohttp.ClientResponse: """Common implementation across stream + non-stream calls. Meant to ensure consistent error-handling.""" max_retry_on_unprocessable_entity_error = provider_config.max_retry_on_unprocessable_entity_error - response: Optional[aiohttp.ClientResponse] = None + response: aiohttp.ClientResponse | None = None async_client_session = self._get_async_client_session(dynamic_client_session=async_client_session) for i in range(max(max_retry_on_unprocessable_entity_error, 1)): @@ -201,16 +201,16 @@ class BaseLLMAIOHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, litellm_params: dict, stream: bool = False, - files: Optional[dict] = None, + files: dict | None = None, content: Any = None, - params: Optional[dict] = None, + params: dict | None = None, ) -> httpx.Response: max_retry_on_unprocessable_entity_error = provider_config.max_retry_on_unprocessable_entity_error - response: Optional[httpx.Response] = None + response: httpx.Response | None = None for i in range(max(max_retry_on_unprocessable_entity_error, 1)): try: @@ -254,7 +254,7 @@ class BaseLLMAIOHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, model: str, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, @@ -262,8 +262,8 @@ class BaseLLMAIOHTTPHandler: optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - client: Optional[ClientSession] = None, + api_key: str | None = None, + client: ClientSession | None = None, ): _response = await self._make_common_async_call( async_client_session=client, @@ -299,14 +299,14 @@ class BaseLLMAIOHTTPHandler: encoding, logging_obj: LiteLLMLoggingObj, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, acompletion: bool, - stream: Optional[bool] = False, + stream: bool | None = False, fake_stream: bool = False, - api_key: Optional[str] = None, - headers: Optional[dict] = {}, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler, ClientSession]] = None, + api_key: str | None = None, + headers: dict | None = {}, + client: HTTPHandler | AsyncHTTPHandler | ClientSession | None = None, ): provider_config = ProviderConfigManager.get_provider_chat_config( model=model, provider=litellm.LlmProviders(custom_llm_provider) @@ -431,10 +431,10 @@ class BaseLLMAIOHTTPHandler: messages: list, logging_obj, litellm_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, fake_stream: bool = False, - client: Optional[HTTPHandler] = None, - ) -> Tuple[Any, dict]: + client: HTTPHandler | None = None, + ) -> tuple[Any, dict]: if client is None or not isinstance(client, HTTPHandler): sync_httpx_client = _get_httpx_client() else: @@ -475,7 +475,7 @@ class BaseLLMAIOHTTPHandler: async def async_image_variations( self, - client: Optional[ClientSession], + client: ClientSession | None, provider_config: BaseImageVariationConfig, api_base: str, headers: dict, @@ -485,12 +485,12 @@ class BaseLLMAIOHTTPHandler: model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, api_key: str, - model: Optional[str], + model: str | None, image: FileTypes, optional_params: dict, ) -> ImageResponse: # create aiohttp form data if files in data - form_data: Optional[FormData] = None + form_data: FormData | None = None if "files" in data and "data" in data: form_data = FormData() for k, v in data["files"].items(): @@ -539,20 +539,20 @@ class BaseLLMAIOHTTPHandler: self, model_response: ImageResponse, api_key: str, - model: Optional[str], + model: str | None, image: FileTypes, timeout: float, custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, optional_params: dict, litellm_params: dict, - print_verbose: Optional[Callable] = None, - api_base: Optional[str] = None, + print_verbose: Callable | None = None, + api_base: str | None = None, aimage_variation: bool = False, logger_fn=None, client=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, + organization: str | None = None, + headers: dict | None = None, ) -> ImageResponse: if model is None: raise ValueError("model is required for non-openai image variations") diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index ac1785ed1ea..ac7a8908616 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -6,7 +6,7 @@ import ssl import typing import urllib.request from collections.abc import Callable -from typing import Any, ClassVar, Dict, Optional, Union +from typing import Any, ClassVar import aiohttp import aiohttp.client_exceptions @@ -18,7 +18,7 @@ import litellm from litellm._logging import verbose_logger from litellm.secret_managers.main import str_to_bool -AIOHTTP_EXC_MAP: Dict = { +AIOHTTP_EXC_MAP: dict = { # Order matters here, most specific exception first # Timeout related exceptions asyncio.TimeoutError: httpx.TimeoutException, @@ -116,7 +116,7 @@ class AiohttpResponseStream(httpx.AsyncByteStream): class AiohttpTransport(httpx.AsyncBaseTransport): def __init__( self, - client: Union[ClientSession, Callable[[], ClientSession]], + client: ClientSession | Callable[[], ClientSession], owns_session: bool = True, ) -> None: self.client = client @@ -125,7 +125,7 @@ class AiohttpTransport(httpx.AsyncBaseTransport): ######################################################### # Class variables for proxy settings ######################################################### - self.proxy_cache: Dict[str, Optional[str]] = {} + self.proxy_cache: dict[str, str | None] = {} async def aclose(self) -> None: if self._owns_session and isinstance(self.client, ClientSession): @@ -147,8 +147,8 @@ class LiteLLMAiohttpTransport(AiohttpTransport): def __init__( self, - client: Union[ClientSession, Callable[[], ClientSession]], - ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None, + client: ClientSession | Callable[[], ClientSession], + ssl_verify: bool | ssl.SSLContext | None = None, owns_session: bool = True, session_factory: Callable[[], ClientSession] | None = None, ): @@ -224,7 +224,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): session_loop = getattr(session, "_loop", None) try: - current_loop: Optional[asyncio.AbstractEventLoop] = asyncio.get_running_loop() + current_loop: asyncio.AbstractEventLoop | None = asyncio.get_running_loop() except RuntimeError: current_loop = None @@ -313,9 +313,9 @@ class LiteLLMAiohttpTransport(AiohttpTransport): client_session: ClientSession, request: httpx.Request, timeout: dict, - proxy: Optional[str], - sni_hostname: Optional[str], - ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None, + proxy: str | None, + sni_hostname: str | None, + ssl_verify: bool | ssl.SSLContext | None = None, ) -> ClientResponse: """ Helper function to make an aiohttp request with the given parameters. @@ -346,7 +346,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): # Only pass ssl kwarg when explicitly configured, to avoid # overriding the session/connector defaults with None (which is # not a valid value for aiohttp's ssl parameter). - request_kwargs: Dict[str, Any] = { + request_kwargs: dict[str, Any] = { "method": request.method, "url": YarlURL(str(request.url), encoded=True), "headers": request.headers, @@ -440,7 +440,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): return proxy - def _proxy_from_env(self, url: httpx.URL) -> typing.Optional[str]: + def _proxy_from_env(self, url: httpx.URL) -> str | None: """ Return proxy URL from env for the given request URL diff --git a/litellm/llms/custom_httpx/container_handler.py b/litellm/llms/custom_httpx/container_handler.py index 351b5e2216f..66774b10cd9 100644 --- a/litellm/llms/custom_httpx/container_handler.py +++ b/litellm/llms/custom_httpx/container_handler.py @@ -8,7 +8,7 @@ endpoint defined in endpoints.json, eliminating the need for individual handler import json from collections.abc import Coroutine from pathlib import Path -from typing import TYPE_CHECKING, Any, Dict, Optional, Type, Union +from typing import TYPE_CHECKING, Any import httpx @@ -33,21 +33,21 @@ if TYPE_CHECKING: # Response type mapping -RESPONSE_TYPES: Dict[str, Type] = { +RESPONSE_TYPES: dict[str, type] = { "ContainerFileListResponse": ContainerFileListResponse, "ContainerFileObject": ContainerFileObject, "DeleteContainerFileResponse": DeleteContainerFileResponse, } -def _load_endpoints_config() -> Dict: +def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" config_path = Path(__file__).parent.parent.parent / "containers" / "endpoints.json" with open(config_path) as f: return json.load(f) -def _get_endpoint_config(endpoint_name: str) -> Optional[Dict]: +def _get_endpoint_config(endpoint_name: str) -> dict | None: """Get config for a specific endpoint by name.""" config = _load_endpoints_config() for endpoint in config["endpoints"]: @@ -59,7 +59,7 @@ def _get_endpoint_config(endpoint_name: str) -> Optional[Dict]: def _build_url( api_base: str, path_template: str, - path_params: Dict[str, str], + path_params: dict[str, str], ) -> str: """Build the full URL by substituting path parameters. @@ -69,8 +69,7 @@ def _build_url( """ # api_base ends with /containers, path_template starts with /containers # So we need to strip /containers from the path - if path_template.startswith("/containers"): - path_template = path_template[len("/containers") :] + path_template = path_template.removeprefix("/containers") # Substitute path parameters for param, value in path_params.items(): @@ -91,8 +90,8 @@ def _build_url( def _build_query_params( query_param_names: list, - kwargs: Dict[str, Any], -) -> Dict[str, str]: + kwargs: dict[str, Any], +) -> dict[str, str]: """Build query parameters from kwargs.""" params = {} for param_name in query_param_names: @@ -104,7 +103,7 @@ def _build_query_params( def _prepare_multipart_file_upload( file: Any, - headers: Dict[str, Any], + headers: dict[str, Any], ) -> tuple: """ Prepare file and headers for multipart upload. @@ -144,13 +143,13 @@ class GenericContainerHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Union[float, httpx.Timeout] = 600, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, **kwargs, - ) -> Union[Any, Coroutine[Any, Any, Any]]: + ) -> Any | Coroutine[Any, Any, Any]: """ Generic handler for any container file endpoint. @@ -197,10 +196,10 @@ class GenericContainerHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, **kwargs, ) -> Any: """Synchronous request handler.""" @@ -302,10 +301,10 @@ class GenericContainerHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, **kwargs, ) -> Any: """Asynchronous request handler.""" diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 148da723399..046840e6fd0 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -10,11 +10,7 @@ from collections.abc import Callable, Mapping from typing import ( TYPE_CHECKING, Any, - Dict, - List, Optional, - Tuple, - Union, ) import certifi @@ -69,7 +65,7 @@ except Exception: _AIOHTTP_SUPPORTS_SOCKET_FACTORY = "socket_factory" in inspect.signature(TCPConnector.__init__).parameters -def _build_aiohttp_keepalive_socket_factory() -> Optional[Callable[[Tuple[Any, ...]], socket.socket]]: +def _build_aiohttp_keepalive_socket_factory() -> Callable[[tuple[Any, ...]], socket.socket] | None: """ Build a socket_factory that enables SO_KEEPALIVE on aiohttp TCP sockets. @@ -84,7 +80,7 @@ def _build_aiohttp_keepalive_socket_factory() -> Optional[Callable[[Tuple[Any, . if not AIOHTTP_SO_KEEPALIVE or not _AIOHTTP_SUPPORTS_SOCKET_FACTORY: return None - def factory(addr_info: Tuple[Any, ...]) -> socket.socket: + def factory(addr_info: tuple[Any, ...]) -> socket.socket: family, type_, proto = addr_info[0], addr_info[1], addr_info[2] sock = socket.socket(family=family, type=type_, proto=proto) sock.setblocking(False) @@ -144,9 +140,9 @@ _STREAMING_ERROR_BODY_READ_EXECUTOR = concurrent.futures.ThreadPoolExecutor( def _prepare_request_data_and_content( - data: Optional[Union[dict, str, bytes]] = None, + data: dict | str | bytes | None = None, content: Any = None, -) -> Tuple[Optional[Union[dict, Mapping]], Any]: +) -> tuple[dict | Mapping | None, Any]: """ Helper function to route data/content parameters correctly for httpx requests @@ -186,13 +182,13 @@ def _prepare_request_data_and_content( # Cache for SSL contexts to avoid creating duplicate contexts with the same configuration # Key: tuple of (cafile, ssl_security_level, ssl_ecdh_curve) # Value: ssl.SSLContext -_ssl_context_cache: Dict[Tuple[Optional[str], Optional[str], Optional[str]], ssl.SSLContext] = {} +_ssl_context_cache: dict[tuple[str | None, str | None, str | None], ssl.SSLContext] = {} def _create_ssl_context( - cafile: Optional[str], - ssl_security_level: Optional[str], - ssl_ecdh_curve: Optional[str], + cafile: str | None, + ssl_security_level: str | None, + ssl_ecdh_curve: str | None, ) -> ssl.SSLContext: """ Create an SSL context with the given configuration. @@ -238,8 +234,8 @@ def _create_ssl_context( def get_ssl_verify( - ssl_verify: Optional[Union[bool, str]] = None, -) -> Union[bool, str]: + ssl_verify: bool | str | None = None, +) -> bool | str: """ Common utility to resolve the SSL verification setting. Prioritizes: @@ -277,8 +273,8 @@ def get_ssl_verify( def get_ssl_configuration( - ssl_verify: Optional[VerifyTypes] = None, -) -> Union[bool, str, ssl.SSLContext]: + ssl_verify: VerifyTypes | None = None, +) -> bool | str | ssl.SSLContext: """ Unified SSL configuration function that handles ssl_context and ssl_verify logic. @@ -342,10 +338,10 @@ def get_ssl_configuration( return ssl_verify -_shared_realtime_ssl_context: Optional[Union[bool, str, ssl.SSLContext]] = None +_shared_realtime_ssl_context: bool | str | ssl.SSLContext | None = None -def get_shared_realtime_ssl_context() -> Union[bool, str, ssl.SSLContext]: +def get_shared_realtime_ssl_context() -> bool | str | ssl.SSLContext: """ Lazily create the SSL context reused by realtime websocket clients so we avoid import-order cycles during startup while keeping a single shared configuration. @@ -388,7 +384,7 @@ def _safe_get_response_text(response: httpx.Response) -> str: return "" -async def _safe_aread_response(response: httpx.Response, timeout: Optional[float] = None) -> bytes: +async def _safe_aread_response(response: httpx.Response, timeout: float | None = None) -> bytes: """Safely read async response body, falling back to empty bytes on errors.""" try: if timeout is not None: @@ -398,7 +394,7 @@ async def _safe_aread_response(response: httpx.Response, timeout: Optional[float return b"" -def _safe_read_response(response: httpx.Response, timeout: Optional[float] = None) -> bytes: +def _safe_read_response(response: httpx.Response, timeout: float | None = None) -> bytes: """Safely read sync response body, falling back to empty bytes on errors.""" try: if timeout is not None: @@ -454,7 +450,7 @@ async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> N class MaskedHTTPStatusError(httpx.HTTPStatusError): - def __init__(self, original_error, message: Optional[str] = None, text: Optional[str] = None): + def __init__(self, original_error, message: str | None = None, text: str | None = None): # Create a new error with the masked URL masked_url = mask_sensitive_info(str(original_error.request.url)) # Mask the original exception message too (it contains the full URL) @@ -508,11 +504,11 @@ class MaskedHTTPStatusError(httpx.HTTPStatusError): class AsyncHTTPHandler: def __init__( self, - timeout: Optional[Union[float, httpx.Timeout]] = None, - event_hooks: Optional[Mapping[str, List[Callable[..., Any]]]] = None, + timeout: float | httpx.Timeout | None = None, + event_hooks: Mapping[str, list[Callable[..., Any]]] | None = None, concurrent_limit=None, # Kept for backward compatibility, but ignored (no limits) - client_alias: Optional[str] = None, # name for client in logs - ssl_verify: Optional[VerifyTypes] = None, + client_alias: str | None = None, # name for client in logs + ssl_verify: VerifyTypes | None = None, shared_session: Optional["ClientSession"] = None, ): self.timeout = timeout @@ -527,9 +523,9 @@ class AsyncHTTPHandler: def create_client( self, - timeout: Optional[Union[float, httpx.Timeout]], - event_hooks: Optional[Mapping[str, List[Callable[..., Any]]]], - ssl_verify: Optional[VerifyTypes] = None, + timeout: float | httpx.Timeout | None, + event_hooks: Mapping[str, list[Callable[..., Any]]] | None, + ssl_verify: VerifyTypes | None = None, shared_session: Optional["ClientSession"] = None, ) -> httpx.AsyncClient: # Get unified SSL configuration @@ -576,10 +572,10 @@ class AsyncHTTPHandler: async def get( self, url: str, - params: Optional[dict] = None, - headers: Optional[dict] = None, - follow_redirects: Optional[bool] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + params: dict | None = None, + headers: dict | None = None, + follow_redirects: bool | None = None, + timeout: float | httpx.Timeout | None = None, ): # Set follow_redirects to UseClientDefault if None _follow_redirects = follow_redirects if follow_redirects is not None else USE_CLIENT_DEFAULT @@ -600,14 +596,14 @@ class AsyncHTTPHandler: async def post( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, stream: bool = False, - logging_obj: Optional[LiteLLMLoggingObject] = None, - files: Optional[RequestFiles] = None, + logging_obj: LiteLLMLoggingObject | None = None, + files: RequestFiles | None = None, content: Any = None, ): start_time = time.time() @@ -654,7 +650,7 @@ class AsyncHTTPHandler: error_response = getattr(e, "response", None) if error_response is not None: for key, value in error_response.headers.items(): - headers["response_headers-{}".format(key)] = value + headers[f"response_headers-{key}"] = value raise litellm.Timeout( message=f"Connection timed out. Timeout passed={timeout}, time taken={time_delta} seconds", @@ -670,11 +666,11 @@ class AsyncHTTPHandler: async def put( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, stream: bool = False, content: Any = None, ): @@ -718,7 +714,7 @@ class AsyncHTTPHandler: error_response = getattr(e, "response", None) if error_response is not None: for key, value in error_response.headers.items(): - headers["response_headers-{}".format(key)] = value + headers[f"response_headers-{key}"] = value raise litellm.Timeout( message=f"Connection timed out after {timeout} seconds.", @@ -734,11 +730,11 @@ class AsyncHTTPHandler: async def patch( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, stream: bool = False, content: Any = None, ): @@ -782,7 +778,7 @@ class AsyncHTTPHandler: error_response = getattr(e, "response", None) if error_response is not None: for key, value in error_response.headers.items(): - headers["response_headers-{}".format(key)] = value + headers[f"response_headers-{key}"] = value raise litellm.Timeout( message=f"Connection timed out after {timeout} seconds.", @@ -798,11 +794,11 @@ class AsyncHTTPHandler: async def delete( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, stream: bool = False, content: Any = None, ): @@ -850,10 +846,10 @@ class AsyncHTTPHandler: self, url: str, client: httpx.AsyncClient, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, stream: bool = False, content: Any = None, ): @@ -886,10 +882,10 @@ class AsyncHTTPHandler: @staticmethod def _create_async_transport( - ssl_context: Optional[ssl.SSLContext] = None, - ssl_verify: Optional[bool] = None, + ssl_context: ssl.SSLContext | None = None, + ssl_verify: bool | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Optional[Union[LiteLLMAiohttpTransport, AsyncHTTPTransport]]: + ) -> LiteLLMAiohttpTransport | AsyncHTTPTransport | None: """ - Creates a transport for httpx.AsyncClient - if litellm.force_ipv4 is True, it will return AsyncHTTPTransport with local_address="0.0.0.0" @@ -949,9 +945,9 @@ class AsyncHTTPHandler: @staticmethod def _get_ssl_connector_kwargs( - ssl_verify: Optional[bool] = None, - ssl_context: Optional[ssl.SSLContext] = None, - ) -> Dict[str, Any]: + ssl_verify: bool | None = None, + ssl_context: ssl.SSLContext | None = None, + ) -> dict[str, Any]: """ Helper method to get SSL connector initialization arguments for aiohttp TCPConnector. @@ -962,7 +958,7 @@ class AsyncHTTPHandler: Returns: Dict with appropriate SSL configuration for TCPConnector """ - connector_kwargs: Dict[str, Any] = { + connector_kwargs: dict[str, Any] = { "local_addr": ("0.0.0.0", 0) if litellm.force_ipv4 else None, } @@ -977,8 +973,8 @@ class AsyncHTTPHandler: @staticmethod def _create_aiohttp_transport( - ssl_verify: Optional[bool] = None, - ssl_context: Optional[ssl.SSLContext] = None, + ssl_verify: bool | None = None, + ssl_context: ssl.SSLContext | None = None, shared_session: Optional["ClientSession"] = None, ) -> LiteLLMAiohttpTransport: """ @@ -1004,7 +1000,7 @@ class AsyncHTTPHandler: # Determine SSL config to pass to transport for per-request override # This ensures ssl_verify works even with shared sessions ######################################################### - ssl_for_transport: Optional[Union[bool, ssl.SSLContext]] = None + ssl_for_transport: bool | ssl.SSLContext | None = None if ssl_context is not None: ssl_for_transport = ssl_context elif ssl_verify is False: @@ -1053,7 +1049,7 @@ class AsyncHTTPHandler: ) @staticmethod - def _create_httpx_transport() -> Optional[AsyncHTTPTransport]: + def _create_httpx_transport() -> AsyncHTTPTransport | None: """ Creates an AsyncHTTPTransport @@ -1069,13 +1065,12 @@ class AsyncHTTPHandler: class HTTPHandler: def __init__( self, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, concurrent_limit=None, # Kept for backward compatibility, but ignored (no limits) - client: Optional[httpx.Client] = None, - ssl_verify: Optional[Union[bool, str]] = None, - disable_default_headers: Optional[ - bool - ] = False, # arize phoenix returns different API responses when user agent header in request + client: httpx.Client | None = None, + ssl_verify: bool | str | None = None, + disable_default_headers: bool + | None = False, # arize phoenix returns different API responses when user agent header in request ): if timeout is None: timeout = _DEFAULT_TIMEOUT @@ -1112,10 +1107,10 @@ class HTTPHandler: def get( self, url: str, - params: Optional[dict] = None, - headers: Optional[dict] = None, - follow_redirects: Optional[bool] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + params: dict | None = None, + headers: dict | None = None, + follow_redirects: bool | None = None, + timeout: float | httpx.Timeout | None = None, ): # Set follow_redirects to UseClientDefault if None _follow_redirects = follow_redirects if follow_redirects is not None else USE_CLIENT_DEFAULT @@ -1133,7 +1128,7 @@ class HTTPHandler: return response @staticmethod - def extract_query_params(url: str) -> Dict[str, str]: + def extract_query_params(url: str) -> dict[str, str]: """ Parse a URL’s query-string into a dict. @@ -1148,15 +1143,15 @@ class HTTPHandler: def post( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, - json: Optional[Union[dict, str, List]] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, + data: dict | str | bytes | None = None, + json: dict | str | list | None = None, + params: dict | None = None, + headers: dict | None = None, stream: bool = False, - timeout: Optional[Union[float, httpx.Timeout]] = None, - files: Optional[Union[dict, RequestFiles]] = None, + timeout: float | httpx.Timeout | None = None, + files: dict | RequestFiles | None = None, content: Any = None, - logging_obj: Optional[LiteLLMLoggingObject] = None, + logging_obj: LiteLLMLoggingObject | None = None, ): try: # Prepare data/content parameters to prevent httpx DeprecationWarning (memory leak fix) @@ -1202,12 +1197,12 @@ class HTTPHandler: def patch( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, - json: Optional[Union[dict, str]] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, + data: dict | str | bytes | None = None, + json: dict | str | None = None, + params: dict | None = None, + headers: dict | None = None, stream: bool = False, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, content: Any = None, ): try: @@ -1252,12 +1247,12 @@ class HTTPHandler: def put( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, - json: Optional[Union[dict, str]] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, + data: dict | str | bytes | None = None, + json: dict | str | None = None, + params: dict | None = None, + headers: dict | None = None, stream: bool = False, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, content: Any = None, ): try: @@ -1301,11 +1296,11 @@ class HTTPHandler: def delete( self, url: str, - data: Optional[Union[dict, str, bytes]] = None, # type: ignore - json: Optional[dict] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + data: dict | str | bytes | None = None, # type: ignore + json: dict | None = None, + params: dict | None = None, + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, stream: bool = False, content: Any = None, ): @@ -1354,7 +1349,7 @@ class HTTPHandler: except Exception: pass - def _create_sync_transport(self) -> Optional[HTTPTransport]: + def _create_sync_transport(self) -> HTTPTransport | None: """ Create an HTTP transport with IPv4 only if litellm.force_ipv4 is True. Otherwise, return None. @@ -1368,8 +1363,8 @@ class HTTPHandler: def get_async_httpx_client( - llm_provider: Union[LlmProviders, httpxSpecialProvider], - params: Optional[dict] = None, + llm_provider: LlmProviders | httpxSpecialProvider, + params: dict | None = None, shared_session: Optional["ClientSession"] = None, ) -> AsyncHTTPHandler: """ @@ -1420,7 +1415,7 @@ def get_async_httpx_client( return _new_client -def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler: +def _get_httpx_client(params: dict | None = None) -> HTTPHandler: """ Retrieves the HTTP client from the cache If not present, creates a new client diff --git a/litellm/llms/custom_httpx/httpx_handler.py b/litellm/llms/custom_httpx/httpx_handler.py index a66d30c9007..c96a8889941 100644 --- a/litellm/llms/custom_httpx/httpx_handler.py +++ b/litellm/llms/custom_httpx/httpx_handler.py @@ -1,5 +1,4 @@ import os -from typing import Optional, Union import httpx @@ -39,16 +38,16 @@ class HTTPHandler: # Close the client when you're done with it await self.client.aclose() - async def get(self, url: str, params: Optional[dict] = None, headers: Optional[dict] = None): + async def get(self, url: str, params: dict | None = None, headers: dict | None = None): response = await self.client.get(url, params=params, headers=headers) return response async def post( self, url: str, - data: Optional[Union[dict, str]] = None, - params: Optional[dict] = None, - headers: Optional[dict] = None, + data: dict | str | None = None, + params: dict | None = None, + headers: dict | None = None, ): try: response = await self.client.post( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index f4edb8f0526..a203f0d6c8c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -8,11 +8,8 @@ from functools import lru_cache from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, TypeVar, Union, cast, @@ -194,7 +191,7 @@ def _google_genai_streaming_hidden_params( litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, response_headers: httpx.Headers, -) -> Dict[str, object]: +) -> dict[str, object]: """Pre-stream metadata for proxy response headers (mirrors CustomStreamWrapper._hidden_params).""" from litellm.litellm_core_utils.core_helpers import process_response_headers @@ -257,16 +254,16 @@ class BaseLLMHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, logging_obj: LiteLLMLoggingObj, stream: bool = False, - signed_json_body: Optional[bytes] = None, + signed_json_body: bytes | None = None, ) -> httpx.Response: """Common implementation across stream + non-stream calls. Meant to ensure consistent error-handling.""" max_retry_on_unprocessable_entity_error = provider_config.max_retry_on_unprocessable_entity_error - response: Optional[httpx.Response] = None + response: httpx.Response | None = None for i in range(max(max_retry_on_unprocessable_entity_error, 1)): try: response = await async_httpx_client.post( @@ -307,15 +304,15 @@ class BaseLLMHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, logging_obj: LiteLLMLoggingObj, stream: bool = False, - signed_json_body: Optional[bytes] = None, + signed_json_body: bytes | None = None, ) -> httpx.Response: max_retry_on_unprocessable_entity_error = provider_config.max_retry_on_unprocessable_entity_error - response: Optional[httpx.Response] = None + response: httpx.Response | None = None for i in range(max(max_retry_on_unprocessable_entity_error, 1)): try: @@ -357,7 +354,7 @@ class BaseLLMHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, model: str, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, @@ -365,10 +362,10 @@ class BaseLLMHTTPHandler: optional_params: dict, litellm_params: dict, encoding: object, - api_key: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, + api_key: str | None = None, + client: AsyncHTTPHandler | None = None, json_mode: bool = False, - signed_json_body: Optional[bytes] = None, + signed_json_body: bytes | None = None, shared_session: Optional["ClientSession"] = None, ): if client is None: @@ -427,25 +424,25 @@ class BaseLLMHTTPHandler: self, model: str, messages: list, - api_base: Optional[str], + api_base: str | None, custom_llm_provider: str, model_response: ModelResponse, encoding: object, logging_obj: LiteLLMLoggingObj, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, acompletion: bool, - stream: Optional[bool] = False, + stream: bool | None = False, fake_stream: bool = False, - api_key: Optional[str] = None, - headers: Optional[Dict[str, Any]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - provider_config: Optional[BaseConfig] = None, + api_key: str | None = None, + headers: dict[str, Any] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + provider_config: BaseConfig | None = None, shared_session: Optional["ClientSession"] = None, ): json_mode: bool = optional_params.pop("json_mode", False) - extra_body: Optional[dict] = optional_params.pop("extra_body", None) + extra_body: dict | None = optional_params.pop("extra_body", None) provider_config = provider_config or ProviderConfigManager.get_provider_chat_config( model=model, provider=litellm.LlmProviders(custom_llm_provider) @@ -479,7 +476,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, ) - data: Dict[str, object] = provider_config.transform_request( + data: dict[str, object] = provider_config.transform_request( model=model, messages=messages, optional_params=optional_params, @@ -645,18 +642,18 @@ class BaseLLMHTTPHandler: api_base: str, headers: dict, data: dict, - signed_json_body: Optional[bytes], + signed_json_body: bytes | None, original_data: dict, model: str, messages: list, logging_obj, optional_params: dict, litellm_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, fake_stream: bool = False, - client: Optional[HTTPHandler] = None, + client: HTTPHandler | None = None, json_mode: bool = False, - ) -> Tuple[object, dict]: + ) -> tuple[object, dict]: if client is None or not isinstance(client, HTTPHandler): sync_httpx_client = _get_httpx_client( { @@ -722,15 +719,15 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, headers: dict, provider_config: BaseConfig, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, data: dict, litellm_params: dict, optional_params: dict, fake_stream: bool = False, - client: Optional[AsyncHTTPHandler] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ): if provider_config.has_custom_stream_wrapper is True: return await provider_config.get_async_custom_stream_wrapper( @@ -781,14 +778,14 @@ class BaseLLMHTTPHandler: data: dict, messages: list, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params: dict, optional_params: dict, fake_stream: bool = False, - client: Optional[AsyncHTTPHandler] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, - ) -> Tuple[object, httpx.Headers]: + client: AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, + ) -> tuple[object, httpx.Headers]: """ Helper function for making an async call with stream. @@ -851,10 +848,10 @@ class BaseLLMHTTPHandler: def _add_stream_param_to_request_body( self, - data: Dict[str, object], + data: dict[str, object], provider_config: BaseConfig, fake_stream: bool, - ) -> Dict[str, object]: + ) -> dict[str, object]: """ Some providers like Bedrock invoke do not support the stream parameter in the request body, we only pass `stream` in the request body the provider supports it. """ @@ -875,14 +872,14 @@ class BaseLLMHTTPHandler: timeout: float, custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, - api_base: Optional[str], + api_base: str | None, optional_params: dict, litellm_params: dict, model_response: EmbeddingResponse, - api_key: Optional[str] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - aembedding: Optional[bool] = False, - headers: Optional[Dict[str, Any]] = None, + api_key: str | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + aembedding: bool | None = False, + headers: dict[str, Any] | None = None, ) -> EmbeddingResponse: provider_config = ProviderConfigManager.get_provider_embedding_config( model=model, provider=litellm.LlmProviders(custom_llm_provider) @@ -1004,10 +1001,10 @@ class BaseLLMHTTPHandler: logging_obj: LiteLLMLoggingObj, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - signed_body: Optional[bytes] = None, + api_key: str | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + signed_body: bytes | None = None, ) -> EmbeddingResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -1052,15 +1049,15 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, provider_config: BaseRerankConfig, - optional_rerank_params: Dict, - timeout: Optional[Union[float, httpx.Timeout]], + optional_rerank_params: dict, + timeout: float | httpx.Timeout | None, model_response: RerankResponse, _is_async: bool = False, - headers: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - litellm_params: Optional[Dict[str, Any]] = None, + headers: dict[str, object] | None = None, + api_key: str | None = None, + api_base: str | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + litellm_params: dict[str, Any] | None = None, ) -> RerankResponse: # get config from model, custom llm provider headers = provider_config.validate_environment( @@ -1146,9 +1143,9 @@ class BaseLLMHTTPHandler: model_response: RerankResponse, api_base: str, headers: dict, - api_key: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: str | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> RerankResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client(llm_provider=litellm.LlmProviders(custom_llm_provider)) @@ -1180,11 +1177,11 @@ class BaseLLMHTTPHandler: optional_params: dict, litellm_params: dict, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - headers: Optional[Dict[str, object]], + api_key: str | None, + api_base: str | None, + headers: dict[str, object] | None, provider_config: BaseAudioTranscriptionConfig, - ) -> Tuple[dict, str, Union[dict, bytes, None], Optional[dict]]: + ) -> tuple[dict, str, dict | bytes | None, dict | None]: """ Shared logic for preparing audio transcription requests. Returns: (headers, complete_url, data, files) @@ -1249,7 +1246,7 @@ class BaseLLMHTTPHandler: model_response: TranscriptionResponse, logging_obj: LiteLLMLoggingObj, optional_params: dict, - api_key: Optional[str], + api_key: str | None, ) -> TranscriptionResponse: """Shared logic for transforming audio transcription responses.""" return provider_config.transform_audio_transcription_response( @@ -1266,15 +1263,15 @@ class BaseLLMHTTPHandler: timeout: float, max_retries: int, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, atranscription: bool = False, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseAudioTranscriptionConfig] = None, + headers: dict[str, object] | None = None, + provider_config: BaseAudioTranscriptionConfig | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Union[TranscriptionResponse, Coroutine[object, object, TranscriptionResponse]]: + ) -> TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]: if provider_config is None: raise ValueError(f"No provider config found for model: {model} and provider: {custom_llm_provider}") @@ -1352,12 +1349,12 @@ class BaseLLMHTTPHandler: timeout: float, max_retries: int, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseAudioTranscriptionConfig] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + headers: dict[str, object] | None = None, + provider_config: BaseAudioTranscriptionConfig | None = None, shared_session: Optional["ClientSession"] = None, ) -> TranscriptionResponse: if provider_config is None: @@ -1417,15 +1414,15 @@ class BaseLLMHTTPHandler: def _prepare_ocr_request( self, model: str, - document: Dict[str, str], + document: dict[str, str], optional_params: dict, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - headers: Optional[Dict[str, object]], + api_key: str | None, + api_base: str | None, + headers: dict[str, object] | None, provider_config: BaseOCRConfig, litellm_params: dict, - ) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]: + ) -> tuple[dict[str, Any], str, dict[str, Any], None]: """ Shared logic for preparing OCR requests. Returns: (headers, complete_url, data, files) @@ -1483,15 +1480,15 @@ class BaseLLMHTTPHandler: async def _async_prepare_ocr_request( self, model: str, - document: Dict[str, str], + document: dict[str, str], optional_params: dict, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - headers: Optional[Dict[str, object]], + api_key: str | None, + api_base: str | None, + headers: dict[str, object] | None, provider_config: BaseOCRConfig, litellm_params: dict, - ) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]: + ) -> tuple[dict[str, Any], str, dict[str, Any], None]: """ Async version of _prepare_ocr_request for providers that need async transforms. Returns: (headers, complete_url, data, files) @@ -1563,19 +1560,19 @@ class BaseLLMHTTPHandler: def ocr( self, model: str, - document: Dict[str, str], + document: dict[str, str], optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, aocr: bool = False, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseOCRConfig] = None, - litellm_params: Optional[dict] = None, - ) -> Union[OCRResponse, Coroutine[object, object, OCRResponse]]: + headers: dict[str, object] | None = None, + provider_config: BaseOCRConfig | None = None, + litellm_params: dict | None = None, + ) -> OCRResponse | Coroutine[object, object, OCRResponse]: """ Sync OCR handler. """ @@ -1638,17 +1635,17 @@ class BaseLLMHTTPHandler: async def async_ocr( self, model: str, - document: Dict[str, str], + document: dict[str, str], optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseOCRConfig] = None, - litellm_params: Optional[dict] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + headers: dict[str, object] | None = None, + provider_config: BaseOCRConfig | None = None, + litellm_params: dict | None = None, ) -> OCRResponse: """ Async OCR handler. @@ -1699,18 +1696,18 @@ class BaseLLMHTTPHandler: def search( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, asearch: bool = False, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseSearchConfig] = None, - ) -> Union[SearchResponse, Coroutine[object, object, SearchResponse]]: + headers: dict[str, object] | None = None, + provider_config: BaseSearchConfig | None = None, + ) -> SearchResponse | Coroutine[object, object, SearchResponse]: """ Sync Search handler. """ @@ -1795,16 +1792,16 @@ class BaseLLMHTTPHandler: async def async_search( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, custom_llm_provider: str, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - headers: Optional[Dict[str, object]] = None, - provider_config: Optional[BaseSearchConfig] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + headers: dict[str, object] | None = None, + provider_config: BaseSearchConfig | None = None, ) -> SearchResponse: """ Async Search handler. @@ -1889,15 +1886,15 @@ class BaseLLMHTTPHandler: headers: dict, # str when the caller passes a pre-serialized (unsigned) body to avoid # re-dumping; bytes when a provider signed the request (e.g. Bedrock). - signed_json_body: Optional[Union[str, bytes]], + signed_json_body: str | bytes | None, request_body: dict, stream: bool, logging_obj: LiteLLMLoggingObj, provider_config: BaseAnthropicMessagesConfig, litellm_params: GenericLiteLLMParams, - api_key: Optional[str], + api_key: str | None, model: str, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ) -> httpx.Response: max_attempts = max(provider_config.max_retry_on_anthropic_messages_http_error, 1) litellm_params_dict = dict(litellm_params) @@ -1950,7 +1947,7 @@ class BaseLLMHTTPHandler: litellm_params: GenericLiteLLMParams, stream: bool, custom_llm_provider: str, - ) -> Optional[Union[float, httpx.Timeout]]: + ) -> float | httpx.Timeout | None: from litellm.litellm_core_utils.completion_timeout import CompletionTimeout from litellm.litellm_core_utils.request_timeout_resolver import ( get_configured_request_timeout, @@ -1974,19 +1971,19 @@ class BaseLLMHTTPHandler: async def async_anthropic_messages_handler( self, model: str, - messages: List[Dict], + messages: list[dict], anthropic_messages_provider_config: BaseAnthropicMessagesConfig, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - client: Optional[AsyncHTTPHandler] = None, - extra_headers: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - stream: Optional[bool] = False, - kwargs: Optional[Dict[str, Any]] = None, - ) -> Union[AnthropicMessagesResponse, AsyncIterator]: + client: AsyncHTTPHandler | None = None, + extra_headers: dict[str, object] | None = None, + api_key: str | None = None, + api_base: str | None = None, + stream: bool | None = False, + kwargs: dict[str, Any] | None = None, + ) -> AnthropicMessagesResponse | AsyncIterator: from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, ) @@ -1999,7 +1996,7 @@ class BaseLLMHTTPHandler: # Prepare headers kwargs = kwargs or {} provider_specific_header = cast( - Optional[litellm.types.utils.ProviderSpecificHeader], + litellm.types.utils.ProviderSpecificHeader | None, kwargs.get("provider_specific_header", None), ) provider_specific_headers = ProviderSpecificHeaderUtils.get_provider_specific_headers( @@ -2161,7 +2158,7 @@ class BaseLLMHTTPHandler: # used for logging + cost tracking logging_obj.model_call_details["httpx_response"] = response - initial_response: Union[AsyncIterator, AnthropicMessagesResponse] + initial_response: AsyncIterator | AnthropicMessagesResponse if stream: from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( AnthropicMessagesStreamingResponse, @@ -2333,22 +2330,19 @@ class BaseLLMHTTPHandler: def anthropic_messages_handler( self, model: str, - messages: List[Dict], + messages: list[dict], anthropic_messages_provider_config: BaseAnthropicMessagesConfig, - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, custom_llm_provider: str, _is_async: bool, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - stream: Optional[bool] = False, - kwargs: Optional[Dict[str, object]] = None, - ) -> Union[ - AnthropicMessagesResponse, - Coroutine[object, object, Union[AnthropicMessagesResponse, AsyncIterator]], - ]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_key: str | None = None, + api_base: str | None = None, + stream: bool | None = False, + kwargs: dict[str, object] | None = None, + ) -> AnthropicMessagesResponse | Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator]: """ LLM HTTP Handler for Anthropic Messages """ @@ -2374,14 +2368,14 @@ class BaseLLMHTTPHandler: self, *, model: str, - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, custom_llm_provider: str, response_api_optional_request_params: dict[str, Any], litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, ) -> tuple[ str, - Union[str, ResponseInputParam], + str | ResponseInputParam, str, dict[str, Any], GenericLiteLLMParams, @@ -2433,7 +2427,7 @@ class BaseLLMHTTPHandler: return ( str(modified_kwargs["model"]) if "model" in modified_kwargs else model, cast( - Union[str, ResponseInputParam], + str | ResponseInputParam, modified_kwargs["input"] if "input" in modified_kwargs else input, ), ( @@ -2448,25 +2442,25 @@ class BaseLLMHTTPHandler: def response_api_handler( self, model: str, - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, responses_api_provider_config: BaseResponsesAPIConfig, response_api_optional_request_params: dict[str, Any], custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Mapping[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, + litellm_metadata: dict[str, object] | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Union[ - ResponsesAPIResponse, - BaseResponsesAPIStreamingIterator, - Coroutine[object, object, Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]], - ]: + ) -> ( + ResponsesAPIResponse + | BaseResponsesAPIStreamingIterator + | Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator] + ): """ Handles responses API requests. When _is_async=True, returns a coroutine instead of making the call directly. @@ -2548,7 +2542,7 @@ class BaseLLMHTTPHandler: # Preserve the OpenAI-style request context (not sent to the provider) for streaming # hooks/metadata; the streaming iterator now consumes this to run deployment hooks # with the same info as chat, including litellm_params. - request_context: Dict[str, object] = {"input": input} + request_context: dict[str, object] = {"input": input} try: request_context.update(response_api_optional_request_params) except Exception: @@ -2578,7 +2572,7 @@ class BaseLLMHTTPHandler: stream=stream, fake_stream=fake_stream, ) - body_kwargs: Dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} + body_kwargs: dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} ## LOGGING logging_obj.pre_call( @@ -2662,20 +2656,20 @@ class BaseLLMHTTPHandler: async def async_response_api_handler( self, model: str, - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, + response_api_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Mapping[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, + litellm_metadata: dict[str, object] | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: + ) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: """ Async version of the responses API handler. Uses async HTTP client to make requests. @@ -2725,7 +2719,7 @@ class BaseLLMHTTPHandler: # Preserve the OpenAI-style request context (not sent to the provider) for streaming # hooks/metadata; the streaming iterator now consumes this to run deployment hooks # with the same info as chat, including litellm_params. - request_context: Dict[str, object] = {"input": input} + request_context: dict[str, object] = {"input": input} try: request_context.update(response_api_optional_request_params) except Exception: @@ -2752,7 +2746,7 @@ class BaseLLMHTTPHandler: stream=stream, fake_stream=fake_stream, ) - body_kwargs: Dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} + body_kwargs: dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} ## LOGGING logging_obj.pre_call( @@ -2851,11 +2845,11 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> DeleteResponseResult: @@ -2907,7 +2901,7 @@ class BaseLLMHTTPHandler: }, ) - delete_kwargs: Dict[str, Any] = { + delete_kwargs: dict[str, Any] = { "url": url, "headers": headers, "timeout": timeout, @@ -2935,14 +2929,14 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, - ) -> Union[DeleteResponseResult, Coroutine[object, object, DeleteResponseResult]]: + ) -> DeleteResponseResult | Coroutine[object, object, DeleteResponseResult]: """ Async version of the responses API handler. Uses async HTTP client to make requests. @@ -2997,7 +2991,7 @@ class BaseLLMHTTPHandler: }, ) - delete_kwargs: Dict[str, Any] = { + delete_kwargs: dict[str, Any] = { "url": url, "headers": headers, "timeout": timeout, @@ -3025,14 +3019,14 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, - ) -> Union[ResponsesAPIResponse, Coroutine[object, object, ResponsesAPIResponse]]: + ) -> ResponsesAPIResponse | Coroutine[object, object, ResponsesAPIResponse]: """ Get a response by ID Uses GET /v1/responses/{response_id} endpoint in the responses API @@ -3106,11 +3100,11 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> ResponsesAPIResponse: """ @@ -3182,18 +3176,18 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + custom_llm_provider: str | None = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, - ) -> Union[Dict, Coroutine[object, object, Dict]]: + ) -> dict | Coroutine[object, object, dict]: if _is_async: return self.async_list_responses_input_items( response_id=response_id, @@ -3268,17 +3262,17 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + custom_llm_provider: str | None = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Dict: + ) -> dict: if client is None or not isinstance(client, AsyncHTTPHandler): verbose_logger.debug( f"Creating HTTP client for list_input_items with shared_session: {id(shared_session) if shared_session else None}" @@ -3341,7 +3335,7 @@ class BaseLLMHTTPHandler: response: httpx.Response, upload_url_location: str, upload_url_key: str = "upload_url", - ) -> tuple[Optional[str], Optional[dict]]: + ) -> tuple[str | None, dict | None]: """ Extract upload URL from initial file creation response. @@ -3374,13 +3368,13 @@ class BaseLLMHTTPHandler: litellm_params: dict, provider_config: BaseFilesConfig, headers: dict, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, logging_obj: LiteLLMLoggingObj, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Union[OpenAIFileObject, Coroutine[object, object, OpenAIFileObject]]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> OpenAIFileObject | Coroutine[object, object, OpenAIFileObject]: """ Creates a file using Gemini's two-step upload process """ @@ -3482,7 +3476,7 @@ class BaseLLMHTTPHandler: ): # Handle pre-signed requests (e.g., from Bedrock S3 uploads) # Type narrowing: this is a plain dict, not TwoStepFileUploadConfig - presigned_request = cast(Dict[str, Any], transformed_request) + presigned_request = cast(dict[str, Any], transformed_request) upload_response = getattr(sync_httpx_client, presigned_request["method"].lower())( url=presigned_request["url"], headers=presigned_request["headers"], @@ -3528,7 +3522,7 @@ class BaseLLMHTTPHandler: elif isinstance(transformed_request, dict) and "file" in transformed_request: # Handle multipart form-data uploads (e.g., Anthropic Files API) # The dict contains tuples suitable for httpx's `files` parameter - file_request = cast(Dict[str, Any], transformed_request) + file_request = cast(dict[str, Any], transformed_request) upload_response = sync_httpx_client.post( url=api_base, headers=headers, @@ -3560,8 +3554,8 @@ class BaseLLMHTTPHandler: headers: dict, api_base: str, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ): """ Creates a file using Gemini's two-step upload process @@ -3644,7 +3638,7 @@ class BaseLLMHTTPHandler: ): # Handle pre-signed requests (e.g., from Bedrock S3 uploads) # Type narrowing: this is a plain dict, not TwoStepFileUploadConfig - presigned_request = cast(Dict[str, Any], transformed_request) + presigned_request = cast(dict[str, Any], transformed_request) upload_response = await getattr(async_httpx_client, presigned_request["method"].lower())( url=presigned_request["url"], headers=presigned_request["headers"], @@ -3730,13 +3724,13 @@ class BaseLLMHTTPHandler: *, client: HTTPHandler, url: str, - base_headers: Dict[str, str], + base_headers: dict[str, str], body_stream: BaseFileUploadStream, content_type: str, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, ) -> httpx.Response: headers = {**base_headers, "Content-Type": content_type} - kwargs: Dict[str, Any] = { + kwargs: dict[str, Any] = { "headers": headers, "content": self._iter_in_blocks(body_stream.iter_bytes(), self._MEDIA_UPLOAD_BLOCK_SIZE), } @@ -3751,10 +3745,10 @@ class BaseLLMHTTPHandler: *, client: AsyncHTTPHandler, url: str, - base_headers: Dict[str, str], + base_headers: dict[str, str], body_stream: BaseFileUploadStream, content_type: str, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, ) -> httpx.Response: """Stream the transformed body straight to a single media upload. Each block is produced on a worker thread (the transform never runs on the @@ -3773,7 +3767,7 @@ class BaseLLMHTTPHandler: break yield cast(bytes, block) - kwargs: Dict[str, Any] = {"headers": headers, "content": _abody()} + kwargs: dict[str, Any] = {"headers": headers, "content": _abody()} if timeout is not None: kwargs["timeout"] = timeout resp = await client.client.post(url, **kwargs) @@ -3787,13 +3781,13 @@ class BaseLLMHTTPHandler: litellm_params: dict, provider_config: "BaseBatchesConfig", headers: dict, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, logging_obj: "LiteLLMLoggingObj", _is_async: bool = False, - client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - model: Optional[str] = None, + client: Union["HTTPHandler", "AsyncHTTPHandler"] | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, ) -> Union["LiteLLMBatch", Coroutine[object, object, "LiteLLMBatch"]]: """ Creates a batch using provider-specific batch creation process @@ -3899,13 +3893,13 @@ class BaseLLMHTTPHandler: litellm_params: dict, provider_config: "BaseBatchesConfig", headers: dict, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, logging_obj: "LiteLLMLoggingObj", _is_async: bool = False, - client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - model: Optional[str] = None, + client: Union["HTTPHandler", "AsyncHTTPHandler"] | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, ) -> Union["LiteLLMBatch", Coroutine[object, object, "LiteLLMBatch"]]: """ Retrieve a batch using provider-specific configuration. @@ -3981,16 +3975,16 @@ class BaseLLMHTTPHandler: async def async_create_batch( self, - transformed_request: Union[bytes, str, dict], + transformed_request: bytes | str | dict, litellm_params: dict, provider_config: "BaseBatchesConfig", headers: dict, api_base: str, logging_obj: "LiteLLMLoggingObj", - client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Union["HTTPHandler", "AsyncHTTPHandler"] | None = None, + timeout: float | httpx.Timeout | None = None, create_batch_data: Optional["CreateBatchRequest"] = None, - model: Optional[str] = None, + model: str | None = None, ): """ Async version of create_batch @@ -4060,16 +4054,16 @@ class BaseLLMHTTPHandler: async def async_retrieve_batch( self, - transformed_request: Union[bytes, str, dict], + transformed_request: bytes | str | dict, litellm_params: dict, provider_config: "BaseBatchesConfig", headers: dict, - api_base: Optional[str], + api_base: str | None, logging_obj: "LiteLLMLoggingObj", - client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - batch_id: Optional[str] = None, - model: Optional[str] = None, + client: Union["HTTPHandler", "AsyncHTTPHandler"] | None = None, + timeout: float | httpx.Timeout | None = None, + batch_id: str | None = None, + model: str | None = None, ): """ Async version of retrieve_batch @@ -4142,14 +4136,14 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, - ) -> Union[ResponsesAPIResponse, Coroutine[object, object, ResponsesAPIResponse]]: + ) -> ResponsesAPIResponse | Coroutine[object, object, ResponsesAPIResponse]: """ Async version of the responses API handler. Uses async HTTP client to make requests. @@ -4222,11 +4216,11 @@ class BaseLLMHTTPHandler: responses_api_provider_config: BaseResponsesAPIConfig, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> ResponsesAPIResponse: @@ -4295,17 +4289,17 @@ class BaseLLMHTTPHandler: model: str, input: Union[str, "ResponseInputParam"], responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, - ) -> Union[ResponsesAPIResponse, Coroutine[object, object, ResponsesAPIResponse]]: + ) -> ResponsesAPIResponse | Coroutine[object, object, ResponsesAPIResponse]: """ Handler for the compact responses API. """ @@ -4362,7 +4356,7 @@ class BaseLLMHTTPHandler: api_key=litellm_params.api_key, model=model, ) - body_kwargs: Dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} + body_kwargs: dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} ## LOGGING logging_obj.pre_call( @@ -4394,14 +4388,14 @@ class BaseLLMHTTPHandler: model: str, input: Union[str, "ResponseInputParam"], responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> ResponsesAPIResponse: @@ -4453,7 +4447,7 @@ class BaseLLMHTTPHandler: api_key=litellm_params.api_key, model=model, ) - body_kwargs: Dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} + body_kwargs: dict[str, Any] = {"data": signed_body} if signed_body is not None else {"json": data} ## LOGGING logging_obj.pre_call( @@ -4488,9 +4482,9 @@ class BaseLLMHTTPHandler: headers: dict, logging_obj: LiteLLMLoggingObj, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Union[OpenAIFileObject, Coroutine[object, object, OpenAIFileObject]]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> OpenAIFileObject | Coroutine[object, object, OpenAIFileObject]: """ Retrieve file metadata by ID """ @@ -4555,8 +4549,8 @@ class BaseLLMHTTPHandler: litellm_params: dict, headers: dict, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ) -> OpenAIFileObject: """ Async retrieve file metadata by ID @@ -4612,8 +4606,8 @@ class BaseLLMHTTPHandler: headers: dict, logging_obj: LiteLLMLoggingObj, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ) -> Union["FileDeleted", Coroutine[object, object, "FileDeleted"]]: """ Delete a file by ID @@ -4679,8 +4673,8 @@ class BaseLLMHTTPHandler: litellm_params: dict, headers: dict, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ) -> "FileDeleted": """ Async delete a file by ID @@ -4730,15 +4724,15 @@ class BaseLLMHTTPHandler: def list_files( self, - purpose: Optional[str], + purpose: str | None, provider_config: BaseFilesConfig, litellm_params: dict, headers: dict, logging_obj: LiteLLMLoggingObj, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Union[List[OpenAIFileObject], Coroutine[object, object, List[OpenAIFileObject]]]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> list[OpenAIFileObject] | Coroutine[object, object, list[OpenAIFileObject]]: """ List all files """ @@ -4798,14 +4792,14 @@ class BaseLLMHTTPHandler: async def async_list_files( self, - purpose: Optional[str], + purpose: str | None, provider_config: BaseFilesConfig, litellm_params: dict, headers: dict, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> List[OpenAIFileObject]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> list[OpenAIFileObject]: """ Async list all files """ @@ -4860,8 +4854,8 @@ class BaseLLMHTTPHandler: headers: dict, logging_obj: LiteLLMLoggingObj, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ) -> Union["HttpxBinaryResponseContent", Coroutine[object, object, "HttpxBinaryResponseContent"]]: """ Retrieve file content by ID @@ -4934,8 +4928,8 @@ class BaseLLMHTTPHandler: litellm_params: dict, headers: dict, logging_obj: LiteLLMLoggingObj, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, ) -> "HttpxBinaryResponseContent": """ Async retrieve file content by ID @@ -4995,7 +4989,7 @@ class BaseLLMHTTPHandler: stream: bool, data: dict, fake_stream: bool, - ) -> Tuple[bool, dict]: + ) -> tuple[bool, dict]: """ Handles preparing a request when `fake_stream` is True. """ @@ -5006,7 +5000,7 @@ class BaseLLMHTTPHandler: return stream, data @staticmethod - def _get_agentic_loop_settings(kwargs: Dict) -> Tuple[int, int, List[str]]: + def _get_agentic_loop_settings(kwargs: dict) -> tuple[int, int, list[str]]: depth = int(kwargs.get("_agentic_loop_depth", 0) or 0) max_loops = int(kwargs.get("max_agentic_loops", 3) or 3) fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or []) @@ -5045,7 +5039,7 @@ class BaseLLMHTTPHandler: @staticmethod def _check_agentic_loop_safety( tool_calls: object, - fingerprints: List[str], + fingerprints: list[str], depth: int, max_loops: int, model: str, @@ -5077,17 +5071,17 @@ class BaseLLMHTTPHandler: self, plan: AgenticLoopPlan, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj", - kwargs: Dict, + kwargs: dict, depth: int, max_loops: int, - fingerprints: List[str], + fingerprints: list[str], fingerprint: str, stream: bool = False, callback: Optional["CustomLogger"] = None, - ) -> Union[AnthropicMessagesResponse, AsyncIterator[object]]: + ) -> AnthropicMessagesResponse | AsyncIterator[object]: from litellm.anthropic_interface import messages as anthropic_messages patch = plan.request_patch or AgenticLoopRequestPatch() @@ -5106,7 +5100,7 @@ class BaseLLMHTTPHandler: max_tokens = patch.max_tokens if max_tokens is None: - max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None)) + max_tokens = cast(int | None, optional_params.pop("max_tokens", None)) else: optional_params.pop("max_tokens", None) if max_tokens is None: @@ -5126,7 +5120,7 @@ class BaseLLMHTTPHandler: kwargs_for_followup["max_agentic_loops"] = max_loops kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint] - response: Union[AnthropicMessagesResponse, AsyncIterator[object]] = await anthropic_messages.acreate( + response: AnthropicMessagesResponse | AsyncIterator[object] = await anthropic_messages.acreate( **{ "max_tokens": max_tokens, "messages": patch.messages, @@ -5166,7 +5160,7 @@ class BaseLLMHTTPHandler: fingerprints: list[str], fingerprint: str, callback: Optional["CustomLogger"] = None, - ) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: + ) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: patch = plan.request_patch or AgenticLoopRequestPatch() if patch.messages is None: raise ValueError("Agentic loop plan missing patched responses input") @@ -5197,7 +5191,7 @@ class BaseLLMHTTPHandler: kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint] try: - response: Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator] = await litellm.aresponses( + response: ResponsesAPIResponse | BaseResponsesAPIStreamingIterator = await litellm.aresponses( model=patch.model or model, input=patch.messages, **optional_params, @@ -5283,13 +5277,13 @@ class BaseLLMHTTPHandler: self, plan: AgenticLoopPlan, model: str, - messages: List[Dict], - optional_params: Dict, - kwargs: Dict, + messages: list[dict], + optional_params: dict, + kwargs: dict, custom_llm_provider: str, depth: int, max_loops: int, - fingerprints: List[str], + fingerprints: list[str], fingerprint: str, ) -> Any: patch = plan.request_patch or AgenticLoopRequestPatch() @@ -5376,15 +5370,15 @@ class BaseLLMHTTPHandler: self, response: Any, model: str, - messages: List[Dict], + messages: list[dict], anthropic_messages_provider_config: "BaseAnthropicMessagesConfig", - anthropic_messages_optional_request_params: Dict, + anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, custom_llm_provider: str, - kwargs: Dict, + kwargs: dict, api_surface: str = "anthropic_messages", - ) -> Optional[Any]: + ) -> Any | None: """ Call agentic completion hooks for all custom loggers (Anthropic Messages API). @@ -5547,13 +5541,13 @@ class BaseLLMHTTPHandler: self, response: Any, model: str, - messages: List[Dict], - optional_params: Dict, + messages: list[dict], + optional_params: dict, logging_obj: "LiteLLMLoggingObj", stream: bool, custom_llm_provider: str, - kwargs: Dict, - ) -> Optional[Any]: + kwargs: dict, + ) -> Any | None: """ Call agentic chat completion hooks for all custom loggers (Chat Completions API). @@ -5665,9 +5659,7 @@ class BaseLLMHTTPHandler: fingerprint=fingerprint, ) except Exception as e: - verbose_logger.exception( - f"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: {str(e)}" - ) + verbose_logger.exception(f"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: {e!s}") # Check if we need to convert response to fake stream for chat completions # This happens when: @@ -5755,7 +5747,7 @@ class BaseLLMHTTPHandler: ) @staticmethod - def _append_query_params(url: str, query_params: Optional[RealtimeQueryParams]) -> str: + def _append_query_params(url: str, query_params: RealtimeQueryParams | None) -> str: """Append query_params to url, skipping keys already present in the URL.""" if not query_params: return url @@ -5800,7 +5792,7 @@ class BaseLLMHTTPHandler: ) if exc is not None ) - last_exc: Optional[BaseException] = None + last_exc: BaseException | None = None for _ in range(max_attempts): try: return await websockets_module.connect( @@ -5828,13 +5820,13 @@ class BaseLLMHTTPHandler: logging_obj: LiteLLMLoggingObj, provider_config: BaseRealtimeConfig, headers: dict, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - client: Optional[Any] = None, - timeout: Optional[float] = None, - user_api_key_dict: Optional[Any] = None, - litellm_metadata: Optional[Dict[str, object]] = None, - query_params: Optional[RealtimeQueryParams] = None, + api_base: str | None = None, + api_key: str | None = None, + client: Any | None = None, + timeout: float | None = None, + user_api_key_dict: Any | None = None, + litellm_metadata: dict[str, object] | None = None, + query_params: RealtimeQueryParams | None = None, ): import websockets from websockets.asyncio.client import ClientConnection @@ -5855,7 +5847,7 @@ class BaseLLMHTTPHandler: ssl_context.verify_mode = ssl.CERT_NONE backend_ws = await self._open_realtime_backend_ws(websockets, url, headers, ssl_context) async with backend_ws: - _request_data: Dict[str, Any] = {} + _request_data: dict[str, Any] = {} if litellm_metadata: _request_data["litellm_metadata"] = litellm_metadata realtime_streaming = RealTimeStreaming( @@ -5877,7 +5869,7 @@ class BaseLLMHTTPHandler: # auto-response disable can be folded into this one setup: Gemini # rejects a second setup, so a follow-up disable would be dropped # and the guardrail bypassed. - _session_config: Optional[str] = None + _session_config: str | None = None if provider_config.requires_session_configuration(): _session_config = provider_config.session_configuration_request(model) if _session_config: @@ -5914,7 +5906,7 @@ class BaseLLMHTTPHandler: except Exception as e: verbose_logger.exception(f"Error connecting to backend: {e}") try: - await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {str(e)}")) + await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e!s}")) except RuntimeError as close_error: if "already completed" in str(close_error) or "websocket.close" in str(close_error): # The WebSocket is already closed or the response is completed, so we can ignore this error @@ -5927,14 +5919,14 @@ class BaseLLMHTTPHandler: self, api_base: str, api_key: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - provider_config: Optional[Any] = None, - model: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_version: Optional[str] = None, + timeout: float | httpx.Timeout, + provider_config: Any | None = None, + model: str | None = None, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_version: str | None = None, ) -> httpx.Response: """ Forward POST /v1/realtime/client_secrets to upstream provider. @@ -5960,14 +5952,14 @@ class BaseLLMHTTPHandler: self, api_base: str, api_key: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - provider_config: Optional[Any] = None, - model: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_version: Optional[str] = None, + timeout: float | httpx.Timeout, + provider_config: Any | None = None, + model: str | None = None, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_version: str | None = None, ) -> httpx.Response: """Forward POST /v1/realtime/transcription_sessions to upstream provider.""" return await self._async_realtime_session_post( @@ -5989,14 +5981,14 @@ class BaseLLMHTTPHandler: endpoint: Literal["client_secrets", "transcription_sessions"], api_base: str, api_key: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - provider_config: Optional[Any] = None, - model: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_version: Optional[str] = None, + timeout: float | httpx.Timeout, + provider_config: Any | None = None, + model: str | None = None, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_version: str | None = None, ) -> httpx.Response: """ Shared POST flow for the realtime HTTP session endpoints @@ -6019,7 +6011,7 @@ class BaseLLMHTTPHandler: ) else: url = provider_config.get_complete_url(api_base=api_base, model=model or "", api_version=api_version) - headers: Dict[str, Any] = provider_config.validate_environment( + headers: dict[str, Any] = provider_config.validate_environment( headers={}, model=model or "", api_key=api_key ) else: @@ -6063,13 +6055,13 @@ class BaseLLMHTTPHandler: openai_ephemeral_key: str, sdp_body: bytes, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - provider_config: Optional[Any] = None, - model: Optional[str] = None, - session_config: Optional[Dict[str, object]] = None, - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - api_version: Optional[str] = None, + timeout: float | httpx.Timeout, + provider_config: Any | None = None, + model: str | None = None, + session_config: dict[str, object] | None = None, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + api_version: str | None = None, ) -> httpx.Response: """ Forward POST /v1/realtime/calls (SDP exchange) to upstream provider. @@ -6090,7 +6082,7 @@ class BaseLLMHTTPHandler: if provider_config is not None: url = provider_config.get_realtime_calls_url(api_base=api_base, model=model or "", api_version=api_version) - headers: Dict[str, Any] = provider_config.get_realtime_calls_headers(ephemeral_key=openai_ephemeral_key) + headers: dict[str, Any] = provider_config.get_realtime_calls_headers(ephemeral_key=openai_ephemeral_key) else: url = f"{api_base.rstrip('/')}/v1/realtime/calls" headers = { @@ -6144,14 +6136,14 @@ class BaseLLMHTTPHandler: model: str, websocket: Any, logging_obj: LiteLLMLoggingObj, - responses_api_provider_config: Optional[BaseResponsesAPIConfig], - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Optional[float] = None, - user_api_key_dict: Optional[Any] = None, - litellm_metadata: Optional[Dict[str, object]] = None, - custom_llm_provider: Optional[str] = None, - first_message: Optional[str] = None, + responses_api_provider_config: BaseResponsesAPIConfig | None, + api_base: str | None = None, + api_key: str | None = None, + timeout: float | None = None, + user_api_key_dict: Any | None = None, + litellm_metadata: dict[str, object] | None = None, + custom_llm_provider: str | None = None, + first_message: str | None = None, **kwargs: Any, ): """ @@ -6258,7 +6250,7 @@ class BaseLLMHTTPHandler: yield backend async with _backend_connection() as backend_ws: - _request_data: Dict[str, Any] = {} + _request_data: dict[str, Any] = {} if litellm_metadata: _request_data["litellm_metadata"] = litellm_metadata @@ -6311,7 +6303,7 @@ class BaseLLMHTTPHandler: except Exception as e: verbose_logger.exception(f"Error in responses WS: {e}") try: - await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {str(e)}")) + await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e!s}")) except RuntimeError as close_error: if "already completed" in str(close_error) or "websocket.close" in str(close_error): pass @@ -6322,23 +6314,20 @@ class BaseLLMHTTPHandler: self, model: str, image: Any, - prompt: Optional[str], + prompt: str | None, image_edit_provider_config: BaseImageEditConfig, - image_edit_optional_request_params: Dict, + image_edit_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - ) -> Union[ - ImageResponse, - Coroutine[object, object, ImageResponse], - ]: + litellm_metadata: dict[str, object] | None = None, + ) -> ImageResponse | Coroutine[object, object, ImageResponse]: """ Handles image edit requests. @@ -6442,18 +6431,18 @@ class BaseLLMHTTPHandler: self, model: str, image: FileTypes, - prompt: Optional[str], + prompt: str | None, image_edit_provider_config: BaseImageEditConfig, - image_edit_optional_request_params: Dict, + image_edit_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, + litellm_metadata: dict[str, object] | None = None, ) -> ImageResponse: """ Async version of the image edit handler. @@ -6542,22 +6531,19 @@ class BaseLLMHTTPHandler: model: str, prompt: str, image_generation_provider_config: BaseImageGenerationConfig, - image_generation_optional_request_params: Dict, + image_generation_optional_request_params: dict, custom_llm_provider: str, - litellm_params: Dict, + litellm_params: dict, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - ) -> Union[ - ImageResponse, - Coroutine[object, object, ImageResponse], - ]: + litellm_metadata: dict[str, object] | None = None, + api_key: str | None = None, + ) -> ImageResponse | Coroutine[object, object, ImageResponse]: """ Handles image generation requests. When _is_async=True, returns a coroutine instead of making the call directly. @@ -6669,17 +6655,17 @@ class BaseLLMHTTPHandler: model: str, prompt: str, image_generation_provider_config: BaseImageGenerationConfig, - image_generation_optional_request_params: Dict, + image_generation_optional_request_params: dict, custom_llm_provider: str, - litellm_params: Dict, + litellm_params: dict, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, + litellm_metadata: dict[str, object] | None = None, + api_key: str | None = None, ) -> ImageResponse: """ Async version of the image generation handler. @@ -6777,22 +6763,19 @@ class BaseLLMHTTPHandler: model: str, prompt: str, video_generation_provider_config: BaseVideoConfig, - video_generation_optional_request_params: Dict, + video_generation_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - ) -> Union[ - VideoObject, - Coroutine[object, object, VideoObject], - ]: + litellm_metadata: dict[str, object] | None = None, + api_key: str | None = None, + ) -> VideoObject | Coroutine[object, object, VideoObject]: """ Handles video generation requests. When _is_async=True, returns a coroutine instead of making the call directly. @@ -6901,17 +6884,17 @@ class BaseLLMHTTPHandler: model: str, prompt: str, video_generation_provider_config: "BaseVideoConfig", - video_generation_optional_request_params: Dict, + video_generation_optional_request_params: dict, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, fake_stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, + litellm_metadata: dict[str, object] | None = None, + api_key: str | None = None, ) -> VideoObject: """ Async version of the video generation handler. @@ -7005,13 +6988,13 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + api_key: str | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - variant: Optional[str] = None, - ) -> Union[bytes, Coroutine[object, object, bytes]]: + variant: str | None = None, + ) -> bytes | Coroutine[object, object, bytes]: """ Handle video content download requests. """ @@ -7095,11 +7078,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - api_key: Optional[str] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - variant: Optional[str] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + api_key: str | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + variant: str | None = None, ) -> bytes: """ Async version of the video content download handler. @@ -7174,12 +7157,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Handler for video remix requests. @@ -7273,11 +7256,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Async version of the video remix handler. @@ -7356,11 +7339,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if _is_async: return self.async_video_create_character_handler( @@ -7440,10 +7423,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -7511,11 +7494,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if _is_async: return self.async_video_get_character_handler( @@ -7580,10 +7563,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -7639,12 +7622,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if _is_async: return self.async_video_edit_handler( @@ -7748,11 +7731,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -7845,12 +7828,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if _is_async: return self.async_video_extension_handler( @@ -7934,11 +7917,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -8002,19 +7985,19 @@ class BaseLLMHTTPHandler: def video_list_handler( self, - after: Optional[str], - limit: Optional[int], - order: Optional[str], + after: str | None, + limit: int | None, + order: str | None, video_list_provider_config, custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Handler for video list requests. @@ -8056,18 +8039,18 @@ class BaseLLMHTTPHandler: async def async_video_list_handler( self, - after: Optional[str], - limit: Optional[int], - order: Optional[str], + after: str | None, + limit: int | None, + order: str | None, video_list_provider_config: BaseVideoConfig, custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Async version of the video list handler. @@ -8144,10 +8127,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Async version of the video delete handler. @@ -8220,12 +8203,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, _is_async: bool = False, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Handler for video status requests. @@ -8325,11 +8308,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params, logging_obj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[float] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | None = None, client=None, - api_key: Optional[str] = None, + api_key: str | None = None, ): """ Async version of the video status handler. @@ -8411,14 +8394,14 @@ class BaseLLMHTTPHandler: def container_create_handler( self, name: str, - container_create_request_params: Dict, + container_create_request_params: dict, container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> Union["ContainerObject", Coroutine[object, object, "ContainerObject"]]: if _is_async: # Return the async coroutine if called with _is_async=True @@ -8498,13 +8481,13 @@ class BaseLLMHTTPHandler: async def async_container_create_handler( self, name: str, - container_create_request_params: Dict, + container_create_request_params: dict, container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "ContainerObject": # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -8576,14 +8559,14 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> Union["ContainerListResponse", Coroutine[object, object, "ContainerListResponse"]]: if _is_async: # Return the async coroutine if called with _is_async=True @@ -8666,13 +8649,13 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "ContainerListResponse": # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -8744,11 +8727,11 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> Union["ContainerObject", Coroutine[object, object, "ContainerObject"]]: if _is_async: # Return the async coroutine if called with _is_async=True @@ -8832,10 +8815,10 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "ContainerObject": # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -8909,11 +8892,11 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> Union["DeleteContainerResult", Coroutine[object, object, "DeleteContainerResult"]]: if _is_async: # Return the async coroutine if called with _is_async=True @@ -8997,10 +8980,10 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "DeleteContainerResult": # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -9074,14 +9057,14 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> Union["ContainerFileListResponse", Coroutine[object, object, "ContainerFileListResponse"]]: if _is_async: return self.async_container_file_list_handler( @@ -9166,13 +9149,13 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "ContainerFileListResponse": # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -9246,11 +9229,11 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - ) -> Union[bytes, Coroutine[object, object, bytes]]: + client: HTTPHandler | AsyncHTTPHandler | None = None, + ) -> bytes | Coroutine[object, object, bytes]: if _is_async: return self.async_container_file_content_handler( container_id=container_id, @@ -9332,9 +9315,9 @@ class BaseLLMHTTPHandler: container_provider_config: "BaseContainerConfig", litellm_params: GenericLiteLLMParams, logging_obj: "LiteLLMLoggingObj", - extra_headers: Optional[Dict[str, object]] = None, - timeout: Union[float, httpx.Timeout] = 600, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout = 600, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> bytes: # For async calls, use async HTTP client if client is None or not isinstance(client, AsyncHTTPHandler): @@ -9405,16 +9388,16 @@ class BaseLLMHTTPHandler: async def async_vector_store_search_handler( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, vector_store_provider_config: BaseVectorStoreConfig, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, ) -> VectorStoreSearchResponse: if client is None or not isinstance(client, AsyncHTTPHandler): @@ -9464,7 +9447,7 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), extra_body=extra_body, ) - all_optional_params: Dict[str, Any] = dict(litellm_params) + all_optional_params: dict[str, Any] = dict(litellm_params) all_optional_params.update(vector_store_search_optional_params or {}) headers, signed_json_body = vector_store_provider_config.sign_request( headers=headers, @@ -9503,18 +9486,18 @@ class BaseLLMHTTPHandler: def vector_store_search_handler( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, vector_store_provider_config: BaseVectorStoreConfig, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreSearchResponse, Coroutine[object, object, VectorStoreSearchResponse]]: + ) -> VectorStoreSearchResponse | Coroutine[object, object, VectorStoreSearchResponse]: if _is_async: return self.async_vector_store_search_handler( vector_store_id=vector_store_id, @@ -9560,7 +9543,7 @@ class BaseLLMHTTPHandler: extra_body=extra_body, ) - all_optional_params: Dict[str, Any] = dict(litellm_params) + all_optional_params: dict[str, Any] = dict(litellm_params) all_optional_params.update(vector_store_search_optional_params or {}) headers, signed_json_body = vector_store_provider_config.sign_request( @@ -9603,10 +9586,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, ) -> VectorStoreCreateResponse: if client is None or not isinstance(client, AsyncHTTPHandler): @@ -9663,12 +9646,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreCreateResponse, Coroutine[object, object, VectorStoreCreateResponse]]: + ) -> VectorStoreCreateResponse | Coroutine[object, object, VectorStoreCreateResponse]: if _is_async: return self.async_vector_store_create_handler( vector_store_create_optional_params=vector_store_create_optional_params, @@ -9733,10 +9716,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreCreateResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -9786,12 +9769,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreCreateResponse, Coroutine[object, object, VectorStoreCreateResponse]]: + ) -> VectorStoreCreateResponse | Coroutine[object, object, VectorStoreCreateResponse]: if _is_async: return self.async_vector_store_retrieve_handler( vector_store_id=vector_store_id, @@ -9845,18 +9828,18 @@ class BaseLLMHTTPHandler: async def async_vector_store_list_handler( self, - after: Optional[str], - before: Optional[str], - limit: Optional[int], - order: Optional[str], + after: str | None, + before: str | None, + limit: int | None, + order: str | None, vector_store_provider_config: BaseVectorStoreConfig, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -9880,7 +9863,7 @@ class BaseLLMHTTPHandler: url = api_base - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if after is not None: params["after"] = after if before is not None: @@ -9909,18 +9892,18 @@ class BaseLLMHTTPHandler: def vector_store_list_handler( self, - after: Optional[str], - before: Optional[str], - limit: Optional[int], - order: Optional[str], + after: str | None, + before: str | None, + limit: int | None, + order: str | None, vector_store_provider_config: BaseVectorStoreConfig, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, ): if _is_async: @@ -9958,7 +9941,7 @@ class BaseLLMHTTPHandler: url = api_base - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if after is not None: params["after"] = after if before is not None: @@ -9993,10 +9976,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreCreateResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10021,7 +10004,7 @@ class BaseLLMHTTPHandler: encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url = f"{api_base}/{encoded_vector_store_id}" - request_body: Dict[str, Any] = dict(vector_store_update_optional_params) + request_body: dict[str, Any] = dict(vector_store_update_optional_params) # Clean metadata to only include string values (OpenAI requirement) if "metadata" in request_body and request_body["metadata"] is not None: @@ -10059,12 +10042,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreCreateResponse, Coroutine[object, object, VectorStoreCreateResponse]]: + ) -> VectorStoreCreateResponse | Coroutine[object, object, VectorStoreCreateResponse]: if _is_async: return self.async_vector_store_update_handler( vector_store_id=vector_store_id, @@ -10099,7 +10082,7 @@ class BaseLLMHTTPHandler: encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url = f"{api_base}/{encoded_vector_store_id}" - request_body: Dict[str, Any] = dict(vector_store_update_optional_params) + request_body: dict[str, Any] = dict(vector_store_update_optional_params) # Clean metadata to only include string values (OpenAI requirement) if "metadata" in request_body and request_body["metadata"] is not None: @@ -10136,10 +10119,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ): if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10187,10 +10170,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, ): if _is_async: @@ -10254,10 +10237,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileObject: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10319,12 +10302,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreFileObject, Coroutine[object, object, VectorStoreFileObject]]: + ) -> VectorStoreFileObject | Coroutine[object, object, VectorStoreFileObject]: if _is_async: return self.async_vector_store_file_create_handler( vector_store_id=vector_store_id, @@ -10396,10 +10379,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileListResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10460,12 +10443,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_query: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_query: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreFileListResponse, Coroutine[object, object, VectorStoreFileListResponse]]: + ) -> VectorStoreFileListResponse | Coroutine[object, object, VectorStoreFileListResponse]: if _is_async: return self.async_vector_store_file_list_handler( vector_store_id=vector_store_id, @@ -10536,9 +10519,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileObject: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10595,11 +10578,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreFileObject, Coroutine[object, object, VectorStoreFileObject]]: + ) -> VectorStoreFileObject | Coroutine[object, object, VectorStoreFileObject]: if _is_async: return self.async_vector_store_file_retrieve_handler( vector_store_id=vector_store_id, @@ -10665,9 +10648,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileContentResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10726,14 +10709,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[ - VectorStoreFileContentResponse, - Coroutine[object, object, VectorStoreFileContentResponse], - ]: + ) -> VectorStoreFileContentResponse | Coroutine[object, object, VectorStoreFileContentResponse]: if _is_async: return self.async_vector_store_file_content_handler( vector_store_id=vector_store_id, @@ -10802,10 +10782,10 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileObject: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -10868,12 +10848,12 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[VectorStoreFileObject, Coroutine[object, object, VectorStoreFileObject]]: + ) -> VectorStoreFileObject | Coroutine[object, object, VectorStoreFileObject]: if _is_async: return self.async_vector_store_file_update_handler( vector_store_id=vector_store_id, @@ -10946,9 +10926,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> VectorStoreFileDeleteResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( @@ -11005,14 +10985,11 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, str]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - ) -> Union[ - VectorStoreFileDeleteResponse, - Coroutine[object, object, VectorStoreFileDeleteResponse], - ]: + ) -> VectorStoreFileDeleteResponse | Coroutine[object, object, VectorStoreFileDeleteResponse]: if _is_async: return self.async_vector_store_file_delete_handler( vector_store_id=vector_store_id, @@ -11077,19 +11054,19 @@ class BaseLLMHTTPHandler: model: str, contents: Any, generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig, - generate_content_config_dict: Dict, + generate_content_config_dict: dict, tools: Any, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, _is_async: bool = False, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - system_instruction: Optional[Any] = None, + litellm_metadata: dict[str, object] | None = None, + system_instruction: Any | None = None, ) -> Any: """ Handles Google GenAI generate content requests. @@ -11209,18 +11186,18 @@ class BaseLLMHTTPHandler: model: str, contents: Any, generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig, - generate_content_config_dict: Dict, + generate_content_config_dict: dict, tools: Any, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + extra_headers: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, stream: bool = False, - litellm_metadata: Optional[Dict[str, object]] = None, - system_instruction: Optional[Any] = None, + litellm_metadata: dict[str, object] | None = None, + system_instruction: Any | None = None, ) -> Any: """ Async version of the generate content handler. @@ -11326,15 +11303,15 @@ class BaseLLMHTTPHandler: self, model: str, input: str, - voice: Optional[str], + voice: str | None, text_to_speech_provider_config: BaseTextToSpeechConfig, - text_to_speech_optional_params: Dict, + text_to_speech_optional_params: dict, custom_llm_provider: str, - litellm_params: Dict, + litellm_params: dict, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, ) -> Union[ "HttpxBinaryResponseContent", @@ -11441,15 +11418,15 @@ class BaseLLMHTTPHandler: self, model: str, input: str, - voice: Optional[str], + voice: str | None, text_to_speech_provider_config: BaseTextToSpeechConfig, - text_to_speech_optional_params: Dict, + text_to_speech_optional_params: dict, custom_llm_provider: str, - litellm_params: Dict, + litellm_params: dict, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, object]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + timeout: float | httpx.Timeout, + extra_headers: dict[str, object] | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ) -> "HttpxBinaryResponseContent": """ Async version of the text-to-speech handler. @@ -11542,9 +11519,9 @@ class BaseLLMHTTPHandler: def _prepare_skill_multipart_request( self, - request_body: Dict, + request_body: dict, headers: dict, - ) -> tuple[Optional[Dict], Optional[list]]: + ) -> tuple[dict | None, list | None]: """ Helper to prepare multipart/form-data request for skills API. @@ -11559,8 +11536,7 @@ class BaseLLMHTTPHandler: return None, None # Remove content-type header if present - httpx will set it automatically for multipart - if "content-type" in headers: - del headers["content-type"] + headers.pop("content-type", None) # Prepare files for multipart upload files = [] @@ -11575,14 +11551,14 @@ class BaseLLMHTTPHandler: def create_skill_handler( self, url: str, - request_body: Dict, + request_body: dict, skills_api_provider_config: "BaseSkillsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Skill", Coroutine[object, object, "Skill"]]: @@ -11641,14 +11617,14 @@ class BaseLLMHTTPHandler: async def async_create_skill_handler( self, url: str, - request_body: Dict, + request_body: dict, skills_api_provider_config: "BaseSkillsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Skill": """Async create a skill""" @@ -11697,14 +11673,14 @@ class BaseLLMHTTPHandler: def list_skills_handler( self, url: str, - query_params: Dict, + query_params: dict, skills_api_provider_config: "BaseSkillsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["ListSkillsResponse", Coroutine[object, object, "ListSkillsResponse"]]: @@ -11756,14 +11732,14 @@ class BaseLLMHTTPHandler: async def async_list_skills_handler( self, url: str, - query_params: Dict, + query_params: dict, skills_api_provider_config: "BaseSkillsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "ListSkillsResponse": """Async list skills""" @@ -11807,9 +11783,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Skill", Coroutine[object, object, "Skill"]]: @@ -11863,9 +11839,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Skill": """Async get a skill""" @@ -11908,9 +11884,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["DeleteSkillResponse", Coroutine[object, object, "DeleteSkillResponse"]]: @@ -11964,9 +11940,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "DeleteSkillResponse": """Async delete a skill""" @@ -12009,14 +11985,14 @@ class BaseLLMHTTPHandler: def create_eval_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Eval", Coroutine[object, object, "Eval"]]: @@ -12068,14 +12044,14 @@ class BaseLLMHTTPHandler: async def async_create_eval_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Eval": """Async create an eval""" @@ -12115,14 +12091,14 @@ class BaseLLMHTTPHandler: def list_evals_handler( self, url: str, - query_params: Dict, + query_params: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["ListEvalsResponse", Coroutine[object, object, "ListEvalsResponse"]]: @@ -12174,14 +12150,14 @@ class BaseLLMHTTPHandler: async def async_list_evals_handler( self, url: str, - query_params: Dict, + query_params: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "ListEvalsResponse": """Async list evals""" @@ -12225,9 +12201,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Eval", Coroutine[object, object, "Eval"]]: @@ -12281,9 +12257,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Eval": """Async get an eval""" @@ -12322,14 +12298,14 @@ class BaseLLMHTTPHandler: def update_eval_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Eval", Coroutine[object, object, "Eval"]]: @@ -12381,14 +12357,14 @@ class BaseLLMHTTPHandler: async def async_update_eval_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Eval": """Async update an eval""" @@ -12432,9 +12408,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["DeleteEvalResponse", Coroutine[object, object, "DeleteEvalResponse"]]: @@ -12488,9 +12464,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "DeleteEvalResponse": """Async delete an eval""" @@ -12533,9 +12509,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["CancelEvalResponse", Coroutine[object, object, "CancelEvalResponse"]]: @@ -12589,9 +12565,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "CancelEvalResponse": """Async cancel an eval""" @@ -12634,14 +12610,14 @@ class BaseLLMHTTPHandler: def create_run_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Run", Coroutine[object, object, "Run"]]: @@ -12693,14 +12669,14 @@ class BaseLLMHTTPHandler: async def async_create_run_handler( self, url: str, - request_body: Dict, + request_body: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Run": """Async create a run""" @@ -12740,14 +12716,14 @@ class BaseLLMHTTPHandler: def list_runs_handler( self, url: str, - query_params: Dict, + query_params: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["ListRunsResponse", Coroutine[object, object, "ListRunsResponse"]]: @@ -12799,14 +12775,14 @@ class BaseLLMHTTPHandler: async def async_list_runs_handler( self, url: str, - query_params: Dict, + query_params: dict, evals_api_provider_config: "BaseEvalsAPIConfig", custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "ListRunsResponse": """Async list runs""" @@ -12850,9 +12826,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["Run", Coroutine[object, object, "Run"]]: @@ -12906,9 +12882,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "Run": """Async get a run""" @@ -12951,9 +12927,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["CancelRunResponse", Coroutine[object, object, "CancelRunResponse"]]: @@ -13007,9 +12983,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "CancelRunResponse": """Async cancel a run""" @@ -13052,9 +13028,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, shared_session: Optional["ClientSession"] = None, ) -> Union["RunDeleteResponse", Coroutine[object, object, "RunDeleteResponse"]]: @@ -13108,9 +13084,9 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, - extra_headers: Optional[Dict[str, object]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + extra_headers: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, shared_session: Optional["ClientSession"] = None, ) -> "RunDeleteResponse": """Async delete a run""" diff --git a/litellm/llms/custom_httpx/mock_transport.py b/litellm/llms/custom_httpx/mock_transport.py index ad93cc134ee..e9248b92209 100644 --- a/litellm/llms/custom_httpx/mock_transport.py +++ b/litellm/llms/custom_httpx/mock_transport.py @@ -9,7 +9,6 @@ so the full proxy -> router -> OpenAI SDK -> httpx path is exercised. import json import time import uuid -from typing import Tuple import httpx @@ -64,7 +63,7 @@ class MockOpenAITransport(httpx.AsyncBaseTransport, httpx.BaseTransport): """ @staticmethod - def _parse_request(request: httpx.Request) -> Tuple[str, bool]: + def _parse_request(request: httpx.Request) -> tuple[str, bool]: """Extract model from the request body.""" try: body = json.loads(request.content) diff --git a/litellm/llms/custom_llm.py b/litellm/llms/custom_llm.py index a5475da077b..fcd41d11499 100644 --- a/litellm/llms/custom_llm.py +++ b/litellm/llms/custom_llm.py @@ -12,7 +12,6 @@ from collections.abc import AsyncIterator, Callable, Coroutine, Iterator from typing import ( TYPE_CHECKING, Any, - Optional, Union, ) @@ -59,8 +58,8 @@ class CustomLLM(BaseLLM): litellm_params=None, logger_fn=None, headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, ) -> Union[ModelResponse, "CustomStreamWrapper"]: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -80,8 +79,8 @@ class CustomLLM(BaseLLM): litellm_params=None, logger_fn=None, headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, ) -> Iterator[GenericStreamingChunk]: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -101,12 +100,9 @@ class CustomLLM(BaseLLM): litellm_params=None, logger_fn=None, headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, - ) -> Union[ - Coroutine[Any, Any, Union[ModelResponse, "CustomStreamWrapper"]], - Union[ModelResponse, "CustomStreamWrapper"], - ]: + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, + ) -> Coroutine[Any, Any, Union[ModelResponse, "CustomStreamWrapper"]] | Union[ModelResponse, "CustomStreamWrapper"]: raise CustomLLMError(status_code=500, message="Not implemented yet!") async def astreaming( @@ -125,8 +121,8 @@ class CustomLLM(BaseLLM): litellm_params=None, logger_fn=None, headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> AsyncIterator[GenericStreamingChunk]: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -134,13 +130,13 @@ class CustomLLM(BaseLLM): self, model: str, prompt: str, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, model_response: ImageResponse, optional_params: dict, logging_obj: Any, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, ) -> ImageResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -149,12 +145,12 @@ class CustomLLM(BaseLLM): model: str, prompt: str, model_response: ImageResponse, - api_key: Optional[str], # dynamically set api_key - https://docs.litellm.ai/docs/set_keys#api_key - api_base: Optional[str], # dynamically set api_base - https://docs.litellm.ai/docs/set_keys#api_base + api_key: str | None, # dynamically set api_key - https://docs.litellm.ai/docs/set_keys#api_key + api_base: str | None, # dynamically set api_base - https://docs.litellm.ai/docs/set_keys#api_base optional_params: dict, logging_obj: Any, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -166,9 +162,9 @@ class CustomLLM(BaseLLM): print_verbose: Callable, logging_obj: Any, optional_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, litellm_params=None, ) -> EmbeddingResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -181,9 +177,9 @@ class CustomLLM(BaseLLM): print_verbose: Callable, logging_obj: Any, optional_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, litellm_params=None, ) -> EmbeddingResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -192,14 +188,14 @@ class CustomLLM(BaseLLM): self, model: str, image: Any, - prompt: Optional[str], + prompt: str | None, model_response: ImageResponse, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, optional_params: dict, logging_obj: Any, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[HTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | None = None, ) -> ImageResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") @@ -207,19 +203,19 @@ class CustomLLM(BaseLLM): self, model: str, image: Any, - prompt: Optional[str], + prompt: str | None, model_response: ImageResponse, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, optional_params: dict, logging_obj: Any, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, + timeout: float | httpx.Timeout | None = None, + client: AsyncHTTPHandler | None = None, ) -> ImageResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") -def custom_chat_llm_router(async_fn: bool, stream: Optional[bool], custom_llm: CustomLLM): +def custom_chat_llm_router(async_fn: bool, stream: bool | None, custom_llm: CustomLLM): """ Routes call to CustomLLM completion/acompletion/streaming/astreaming functions, based on call type diff --git a/litellm/llms/dashscope/chat/transformation.py b/litellm/llms/dashscope/chat/transformation.py index def3a367384..a639b302a84 100644 --- a/litellm/llms/dashscope/chat/transformation.py +++ b/litellm/llms/dashscope/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to DashScope's `/v1/chat/complet """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, overload +from typing import Any, Literal, overload from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam @@ -15,9 +15,9 @@ class DashScopeChatConfig(OpenAIGPTConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, - messages: List[AllMessageValues], - tools: Optional[List[ChatCompletionToolParam]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List[ChatCompletionToolParam]]]: + messages: list[AllMessageValues], + tools: list[ChatCompletionToolParam] | None = None, + ) -> tuple[list[AllMessageValues], list[ChatCompletionToolParam] | None]: """ Override to preserve cache_control for DashScope. DashScope supports cache_control - don't strip it. @@ -26,28 +26,28 @@ class DashScopeChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: if is_async: return super()._transform_messages(messages=messages, model=model, is_async=True) else: return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = ( api_base or get_secret_str("DASHSCOPE_API_BASE") or "https://dashscope.aliyuncs.com/compatible-mode/v1" ) # type: ignore @@ -56,12 +56,12 @@ class DashScopeChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ If api_base is not provided, use the default DashScope /chat/completions endpoint. diff --git a/litellm/llms/dashscope/common_utils.py b/litellm/llms/dashscope/common_utils.py index b3b89cbbebf..9a7dd4da8d3 100644 --- a/litellm/llms/dashscope/common_utils.py +++ b/litellm/llms/dashscope/common_utils.py @@ -2,8 +2,6 @@ Common utilities for the DashScope LLM provider. """ -from typing import Optional - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -16,7 +14,7 @@ class DashScopeError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[httpx.Headers] = None, + headers: httpx.Headers | None = None, ): self.status_code = status_code self.message = message diff --git a/litellm/llms/dashscope/cost_calculator.py b/litellm/llms/dashscope/cost_calculator.py index 2732b97cd35..8106a97f2ea 100644 --- a/litellm/llms/dashscope/cost_calculator.py +++ b/litellm/llms/dashscope/cost_calculator.py @@ -5,7 +5,6 @@ Handles tiered pricing and prompt caching scenarios. """ from dataclasses import dataclass -from typing import List, Optional, Tuple from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import calculate_tiered_cost from litellm.types.utils import ModelInfo, Usage @@ -46,7 +45,7 @@ def _extract_token_breakdown(usage: Usage) -> TokenBreakdown: def _calculate_prompt_cost( breakdown: TokenBreakdown, model_info: ModelInfo, - tiered_pricing: Optional[List[dict]], + tiered_pricing: list[dict] | None, ) -> float: """Calculate total prompt cost including cached tokens.""" if tiered_pricing: @@ -78,7 +77,7 @@ def _calculate_prompt_cost( def _calculate_completion_cost( breakdown: TokenBreakdown, model_info: ModelInfo, - tiered_pricing: Optional[List[dict]], + tiered_pricing: list[dict] | None, ) -> float: """Calculate total completion cost including reasoning tokens.""" if tiered_pricing: @@ -107,7 +106,7 @@ def _calculate_completion_cost( return (breakdown.completion_tokens * output_cost) + (breakdown.reasoning_tokens * reasoning_cost) -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculate cost per token for Dashscope models. diff --git a/litellm/llms/dashscope/embed/transformation.py b/litellm/llms/dashscope/embed/transformation.py index 070e2f57667..55722ce35d1 100644 --- a/litellm/llms/dashscope/embed/transformation.py +++ b/litellm/llms/dashscope/embed/transformation.py @@ -11,8 +11,6 @@ Endpoint Docs - https://help.aliyun.com/zh/model-studio/text-embedding-synchronous-api """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -38,7 +36,7 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): def __init__(self) -> None: pass - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: # DashScope's compatible-mode embeddings API accepts the same params as OpenAI. # `dimensions` / `encoding_format` are only honored by text-embedding-v3 / v4; # earlier versions silently ignore them server-side. @@ -66,11 +64,11 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("DASHSCOPE_API_KEY") @@ -86,12 +84,12 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base = api_base or get_secret_str("DASHSCOPE_API_BASE") or DEFAULT_API_BASE base = base.rstrip("/") @@ -122,7 +120,7 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -132,7 +130,7 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): except Exception as e: raise DashScopeError( status_code=raw_response.status_code, - message=f"Failed to parse DashScope response as JSON: {str(e)}", + message=f"Failed to parse DashScope response as JSON: {e!s}", ) logging_obj.post_call( @@ -176,7 +174,7 @@ class DashScopeEmbeddingConfig(BaseEmbeddingConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: if isinstance(headers, dict): headers = httpx.Headers(headers) diff --git a/litellm/llms/dashscope/image_generation/transformation.py b/litellm/llms/dashscope/image_generation/transformation.py index 094e06d1269..5cf7037f01a 100644 --- a/litellm/llms/dashscope/image_generation/transformation.py +++ b/litellm/llms/dashscope/image_generation/transformation.py @@ -23,7 +23,7 @@ Response format: } """ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -62,7 +62,7 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro). """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "size"] def map_openai_params( @@ -88,12 +88,12 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE @@ -101,11 +101,11 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: final_api_key = api_key or get_secret_str("DASHSCOPE_API_KEY") if not final_api_key: @@ -152,8 +152,8 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform DashScope response to litellm ImageResponse. diff --git a/litellm/llms/dashscope/rerank/transformation.py b/litellm/llms/dashscope/rerank/transformation.py index 365e15fdd7a..a5306083ffb 100644 --- a/litellm/llms/dashscope/rerank/transformation.py +++ b/litellm/llms/dashscope/rerank/transformation.py @@ -22,7 +22,7 @@ as supported only for gte-rerank-v2 / qwen3-vl-rerank. Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api """ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -109,15 +109,15 @@ class DashScopeRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: # qwen3-rerank accepts query/documents/top_n/return_documents. The # rest (rank_fields, max_*_per_doc) are silently dropped. params: OptionalRerankParams = OptionalRerankParams( @@ -133,7 +133,7 @@ class DashScopeRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -142,7 +142,7 @@ class DashScopeRerankConfig(BaseRerankConfig): if "documents" not in optional_rerank_params: raise ValueError("documents is required for DashScope rerank") - request: Dict[str, Any] = { + request: dict[str, Any] = { "model": model, "query": optional_rerank_params["query"], "documents": optional_rerank_params["documents"], @@ -201,9 +201,9 @@ class DashScopeRerankConfig(BaseRerankConfig): # plus, when return_documents=true was sent: # "document": {"text": "..."} # which already matches LiteLLM's RerankResponseDocument shape. - transformed_results: List[dict] = [] + transformed_results: list[dict] = [] for r in results: - item: Dict[str, Any] = { + item: dict[str, Any] = { "index": r["index"], "relevance_score": r["relevance_score"], } @@ -231,7 +231,7 @@ class DashScopeRerankConfig(BaseRerankConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: if isinstance(headers, dict): headers = httpx.Headers(headers) diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 85a2eccde2f..9f6f669a264 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -7,11 +7,7 @@ from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( TYPE_CHECKING, Any, - List, Literal, - Optional, - Tuple, - Union, cast, overload, ) @@ -109,7 +105,7 @@ def _split_parallel_tool_calls(messages: list[AllMessageValues]) -> list[AllMess def _expand( assistant: ChatCompletionAssistantMessage, - calls_by_id: dict[Optional[str], ChatCompletionAssistantToolCall], + calls_by_id: dict[str | None, ChatCompletionAssistantToolCall], tool_messages: list[ChatCompletionToolMessage], ) -> Iterator[AllMessageValues]: for position, tool_message in enumerate(tool_messages): @@ -158,21 +154,21 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): Reference: https://docs.databricks.com/en/machine-learning/foundation-models/api-reference.html#chat-request """ - max_tokens: Optional[int] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - stop: Optional[Union[List[str], str]] = None - n: Optional[int] = None + max_tokens: int | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + stop: list[str] | str | None = None + n: int | None = None def __init__( self, - max_tokens: Optional[int] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - stop: Optional[Union[List[str], str]] = None, - n: Optional[int] = None, + max_tokens: int | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + stop: list[str] | str | None = None, + n: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -180,14 +176,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "databricks" @classmethod def get_config(cls): return super().get_config() - def get_required_params(self) -> List[ProviderField]: + def get_required_params(self) -> list[ProviderField]: """For a given provider, return it's required fields with a description""" return [ ProviderField( @@ -208,11 +204,11 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: # Check for custom user agent in optional_params or environment # This allows partners building on LiteLLM to set their own telemetry @@ -239,18 +235,18 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = self._get_api_base(api_base) complete_url = f"{api_base}/chat/completions" return complete_url - def get_supported_openai_params(self, model: Optional[str] = None) -> list: + def get_supported_openai_params(self, model: str | None = None) -> list: return [ "stream", "stop", @@ -266,9 +262,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): "thinking", ] - def convert_anthropic_tool_to_databricks_tool( - self, tool: Optional[AllAnthropicToolsValues] - ) -> Optional[DatabricksTool]: + def convert_anthropic_tool_to_databricks_tool(self, tool: AllAnthropicToolsValues | None) -> DatabricksTool | None: if tool is None: return None @@ -281,14 +275,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): # Only add description if it exists description = tool.get("description") if description is not None: - function_params["description"] = cast(Union[dict, str], description) + function_params["description"] = cast(dict | str, description) return DatabricksTool( type="function", function=function_params, ) - def _map_openai_to_dbrx_tool(self, model: str, tools: List) -> List[DatabricksTool]: + def _map_openai_to_dbrx_tool(self, model: str, tools: list) -> list[DatabricksTool]: # if not claude, send as is if "claude" not in model: return tools @@ -303,10 +297,10 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): def map_response_format_to_databricks_tool( self, model: str, - value: Optional[dict], + value: dict | None, optional_params: dict, is_thinking_enabled: bool, - ) -> Optional[DatabricksTool]: + ) -> DatabricksTool | None: if value is None: return None @@ -318,9 +312,9 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, # allows overrides to selectively run this - messages: List[AllMessageValues], - tools: Optional[List["ChatCompletionToolParam"]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + messages: list[AllMessageValues], + tools: list["ChatCompletionToolParam"] | None = None, + ) -> tuple[list[AllMessageValues], list["ChatCompletionToolParam"] | None]: """ Override the parent class method to preserve cache_control for models on Databricks. Databricks supports Anthropic-style cache control for Claude models. @@ -383,7 +377,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): else: optional_params["thinking"] = mapped_thinking if AnthropicConfig._is_adaptive_thinking_model(model, "databricks"): - mapped_effort: Optional[str] = None + mapped_effort: str | None = None if isinstance(reasoning_effort_value, str): mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort_value) if mapped_effort is None: @@ -412,20 +406,20 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Databricks does not support: - 'name' in user message. @@ -478,8 +472,8 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @staticmethod def extract_content_str( - content: Optional[AllDatabricksContentValues], - ) -> Optional[str]: + content: AllDatabricksContentValues | None, + ) -> str | None: if content is None: return None if isinstance(content, str): @@ -496,18 +490,18 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @staticmethod def extract_reasoning_content( - content: Optional[AllDatabricksContentValues], - ) -> Tuple[ - Optional[str], - Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], + content: AllDatabricksContentValues | None, + ) -> tuple[ + str | None, + list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None, ]: """ Extract and return the reasoning content and thinking blocks """ if content is None: return None, None - thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None - reasoning_content: Optional[str] = None + thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None + reasoning_content: str | None = None if isinstance(content, list): for item in content: if item.get("type") == "reasoning": @@ -529,8 +523,8 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @staticmethod def extract_citations( - content: Optional[AllDatabricksContentValues], - ) -> Optional[List[Any]]: + content: AllDatabricksContentValues | None, + ) -> list[Any] | None: if content is None: return None citations = [] @@ -541,9 +535,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): citations.append([{**citation, "supported_text": text} for citation in citations_item]) return citations or None - def _transform_dbrx_choices( - self, choices: List[DatabricksChoice], json_mode: Optional[bool] = None - ) -> List[Choices]: + def _transform_dbrx_choices(self, choices: list[DatabricksChoice], json_mode: bool | None = None) -> list[Choices]: transformed_choices = [] for choice in choices: @@ -559,14 +551,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): if fixed_tool_calls is not None: tool_calls = fixed_tool_calls - translated_message: Optional[Message] = None - finish_reason: Optional[str] = None + translated_message: Message | None = None + finish_reason: str | None = None if tool_calls and _should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=json_mode, ): # to support response_format on claude models - json_mode_content_str: Optional[str] = str(tool_calls[0]["function"].get("arguments", "")) or None + json_mode_content_str: str | None = str(tool_calls[0]["function"].get("arguments", "")) or None if json_mode_content_str is not None: translated_message = Message(content=json_mode_content_str) finish_reason = "stop" @@ -614,12 +606,12 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: # Redact sensitive data before logging to prevent credential leakage redacted_request_data = self.redact_sensitive_data(request_data) @@ -638,7 +630,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): except Exception as e: response_headers = getattr(raw_response, "headers", None) raise DatabricksException( - message="Unable to get json response - {}, Original Response: {}".format(str(e), raw_response.text), + message=f"Unable to get json response - {e!s}, Original Response: {raw_response.text}", status_code=raw_response.status_code, headers=response_headers, ) @@ -659,9 +651,9 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return DatabricksChatResponseIterator( streaming_response=streaming_response, @@ -673,9 +665,9 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): class DatabricksChatResponseIterator(BaseModelResponseIterator): def __init__( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): super().__init__(streaming_response, sync_stream) diff --git a/litellm/llms/databricks/common_utils.py b/litellm/llms/databricks/common_utils.py index 908aa56a4d6..2fb7cacb9bf 100644 --- a/litellm/llms/databricks/common_utils.py +++ b/litellm/llms/databricks/common_utils.py @@ -12,7 +12,7 @@ Authentication priority: import os import re -from typing import Any, Dict, Literal, Optional, Tuple +from typing import Any, Literal from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -97,7 +97,7 @@ class DatabricksBase: return data @classmethod - def redact_headers_for_logging(cls, headers: Dict[str, str]) -> Dict[str, str]: + def redact_headers_for_logging(cls, headers: dict[str, str]) -> dict[str, str]: """ Create a copy of headers with sensitive values redacted for safe logging. @@ -133,7 +133,7 @@ class DatabricksBase: return redacted @staticmethod - def _build_user_agent(custom_user_agent: Optional[str] = None) -> str: + def _build_user_agent(custom_user_agent: str | None = None) -> str: """ Build the User-Agent string for Databricks API calls. @@ -176,7 +176,7 @@ class DatabricksBase: # Default: just litellm return f"litellm/{version}" - def _get_api_base(self, api_base: Optional[str]) -> str: + def _get_api_base(self, api_base: str | None) -> str: """ Get the Databricks API base URL. @@ -245,7 +245,7 @@ class DatabricksBase: except requests.RequestException as e: raise DatabricksException( status_code=500, - message=f"OAuth M2M token request failed: {str(e)}", + message=f"OAuth M2M token request failed: {e!s}", ) if response.status_code != 200: @@ -258,8 +258,8 @@ class DatabricksBase: return token_data["access_token"] def _get_databricks_credentials( - self, api_key: Optional[str], api_base: Optional[str], headers: Optional[dict] - ) -> Tuple[str, dict]: + self, api_key: str | None, api_base: str | None, headers: dict | None + ) -> tuple[str, dict]: """ Get Databricks credentials using the Databricks SDK. @@ -303,13 +303,13 @@ class DatabricksBase: def databricks_validate_environment( self, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, endpoint_type: Literal["chat_completions", "embeddings"], - custom_endpoint: Optional[bool], - headers: Optional[dict], - custom_user_agent: Optional[str] = None, - ) -> Tuple[str, dict]: + custom_endpoint: bool | None, + headers: dict | None, + custom_user_agent: str | None = None, + ) -> tuple[str, dict]: """ Validate and configure the Databricks environment. @@ -372,12 +372,12 @@ class DatabricksBase: if headers is None: headers = { - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", } else: if api_key is not None: - headers.update({"Authorization": "Bearer {}".format(api_key)}) + headers.update({"Authorization": f"Bearer {api_key}"}) if api_key is not None: headers["Authorization"] = f"Bearer {api_key}" @@ -389,8 +389,8 @@ class DatabricksBase: verbose_logger.debug(f"Databricks request headers: {self.redact_headers_for_logging(headers)}") if endpoint_type == "chat_completions" and custom_endpoint is not True: - api_base = "{}/chat/completions".format(api_base) + api_base = f"{api_base}/chat/completions" elif endpoint_type == "embeddings" and custom_endpoint is not True: - api_base = "{}/embeddings".format(api_base) + api_base = f"{api_base}/embeddings" return api_base, headers diff --git a/litellm/llms/databricks/cost_calculator.py b/litellm/llms/databricks/cost_calculator.py index 9db151538b5..7413a04c731 100644 --- a/litellm/llms/databricks/cost_calculator.py +++ b/litellm/llms/databricks/cost_calculator.py @@ -3,13 +3,11 @@ Helper util for handling databricks-specific cost calculation - e.g.: handling 'dbrx-instruct-*' """ -from typing import Tuple - from litellm.types.utils import Usage from litellm.utils import get_model_info -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -29,9 +27,12 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: "meta-llama-3.1-405b-instruct" ): base_model = "databricks-meta-llama-3-1-405b-instruct" - elif model.startswith("databricks/mixtral-8x7b-instruct-v0.1") or model.startswith("mixtral-8x7b-instruct-v0.1"): - base_model = "databricks-mixtral-8x7b-instruct" - elif model.startswith("databricks/mixtral-8x7b-instruct-v0.1") or model.startswith("mixtral-8x7b-instruct-v0.1"): + elif ( + model.startswith("databricks/mixtral-8x7b-instruct-v0.1") + or model.startswith("mixtral-8x7b-instruct-v0.1") + or model.startswith("databricks/mixtral-8x7b-instruct-v0.1") + or model.startswith("mixtral-8x7b-instruct-v0.1") + ): base_model = "databricks-mixtral-8x7b-instruct" elif model.startswith("databricks/bge-large-en") or model.startswith("bge-large-en"): base_model = "databricks-bge-large-en" diff --git a/litellm/llms/databricks/embed/handler.py b/litellm/llms/databricks/embed/handler.py index 227824f72d0..fbd1dbc6b98 100644 --- a/litellm/llms/databricks/embed/handler.py +++ b/litellm/llms/databricks/embed/handler.py @@ -3,7 +3,6 @@ Calling logic for Databricks embeddings """ import os -from typing import Optional from litellm.utils import EmbeddingResponse @@ -18,14 +17,14 @@ class DatabricksEmbeddingHandler(OpenAILikeEmbeddingHandler, DatabricksBase): input: list, timeout: float, logging_obj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, optional_params: dict, - model_response: Optional[EmbeddingResponse] = None, + model_response: EmbeddingResponse | None = None, client=None, aembedding=None, - custom_endpoint: Optional[bool] = None, - headers: Optional[dict] = None, + custom_endpoint: bool | None = None, + headers: dict | None = None, ) -> EmbeddingResponse: # Check for custom user agent in optional_params or environment # This allows partners building on LiteLLM to set their own telemetry diff --git a/litellm/llms/databricks/embed/transformation.py b/litellm/llms/databricks/embed/transformation.py index 53e3b30dd21..8c0e9ae01a4 100644 --- a/litellm/llms/databricks/embed/transformation.py +++ b/litellm/llms/databricks/embed/transformation.py @@ -3,7 +3,6 @@ Translates from OpenAI's `/v1/embeddings` to Databricks' `/embeddings` """ import types -from typing import Optional class DatabricksEmbeddingConfig: @@ -11,11 +10,11 @@ class DatabricksEmbeddingConfig: Reference: https://learn.microsoft.com/en-us/azure/databricks/machine-learning/foundation-models/api-reference#--embedding-task """ - instruction: Optional[str] = ( + instruction: str | None = ( None # An optional instruction to pass to the embedding model. BGE Authors recommend 'Represent this sentence for searching relevant passages:' for retrieval queries ) - def __init__(self, instruction: Optional[str] = None) -> None: + def __init__(self, instruction: str | None = None) -> None: locals_ = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: diff --git a/litellm/llms/databricks/responses/transformation.py b/litellm/llms/databricks/responses/transformation.py index 090fef5ac82..d6d6b600fa0 100644 --- a/litellm/llms/databricks/responses/transformation.py +++ b/litellm/llms/databricks/responses/transformation.py @@ -8,7 +8,7 @@ Reference: https://docs.databricks.com/aws/en/machine-learning/foundation-model- """ import os -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any from litellm.llms.databricks.common_utils import DatabricksBase from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -42,7 +42,7 @@ class DatabricksResponsesAPIConfig(DatabricksBase, OpenAIResponsesAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or os.getenv("DATABRICKS_API_KEY") @@ -65,7 +65,7 @@ class DatabricksResponsesAPIConfig(DatabricksBase, OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = api_base or os.getenv("DATABRICKS_API_BASE") @@ -76,11 +76,11 @@ class DatabricksResponsesAPIConfig(DatabricksBase, OpenAIResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Transform request for Databricks Responses API. @@ -88,8 +88,7 @@ class DatabricksResponsesAPIConfig(DatabricksBase, OpenAIResponsesAPIConfig): then delegates to OpenAI's transformation. """ # Strip provider prefix if present (e.g., "databricks/databricks-gpt-5-nano" -> "databricks-gpt-5-nano") - if model.startswith("databricks/"): - model = model[len("databricks/") :] + model = model.removeprefix("databricks/") return super().transform_responses_api_request( model=model, diff --git a/litellm/llms/databricks/streaming_utils.py b/litellm/llms/databricks/streaming_utils.py index a6a45719fe6..74216888111 100644 --- a/litellm/llms/databricks/streaming_utils.py +++ b/litellm/llms/databricks/streaming_utils.py @@ -1,5 +1,4 @@ import json -from typing import Optional import litellm from litellm import verbose_logger @@ -20,10 +19,10 @@ class ModelResponseIterator: processed_chunk = litellm.ModelResponseStream(**chunk) text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None # Usage-only final chunk (OpenAI ``stream_options.include_usage``) # arrives with an empty ``choices`` list — return usage without @@ -75,7 +74,7 @@ class ModelResponseIterator: is_finished = True finish_reason = processed_chunk.choices[0].finish_reason - usage_chunk: Optional[Usage] = getattr(processed_chunk, "usage", None) + usage_chunk: Usage | None = getattr(processed_chunk, "usage", None) if usage_chunk is not None: usage = ChatCompletionUsageBlock( prompt_tokens=usage_chunk.prompt_tokens, diff --git a/litellm/llms/dataforseo/search/transformation.py b/litellm/llms/dataforseo/search/transformation.py index 97a2539b3df..2eb92ede878 100644 --- a/litellm/llms/dataforseo/search/transformation.py +++ b/litellm/llms/dataforseo/search/transformation.py @@ -4,7 +4,7 @@ Calls DataForSEO SERP API to search the web. DataForSEO API Reference: https://docs.dataforseo.com/v3/serp/google/organic/live/advanced/?bash """ -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal import httpx @@ -40,11 +40,11 @@ class DataForSEOSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate DataForSEO environment and set up authentication. @@ -91,9 +91,9 @@ class DataForSEOSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -105,11 +105,11 @@ class DataForSEOSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, **kwargs, - ) -> Union[Dict, List[Dict]]: + ) -> dict | list[dict]: """ Transform Search request to DataForSEO SERP API format. @@ -126,7 +126,7 @@ class DataForSEOSearchConfig(BaseSearchConfig): List[Dict]: Request body for DataForSEO API (array of task objects as required by API) """ # DataForSEO expects an array of task objects - task: Dict[str, Any] = {} + task: dict[str, Any] = {} # Convert query to string if it's a list if isinstance(query, list): diff --git a/litellm/llms/datarobot/chat/transformation.py b/litellm/llms/datarobot/chat/transformation.py index 75bbfc19b69..0857ddc321c 100644 --- a/litellm/llms/datarobot/chat/transformation.py +++ b/litellm/llms/datarobot/chat/transformation.py @@ -4,9 +4,10 @@ Support for OpenAI's `/v1/chat/completions` endpoint. Calls done in OpenAI/openai.py as DataRobot is openai-compatible. """ -from typing import Optional, Tuple -from litellm.secret_managers.main import get_secret_str from urllib.parse import urlparse, urlunparse + +from litellm.secret_managers.main import get_secret_str + from ...openai_like.chat.transformation import OpenAILikeChatConfig LLMGW_PATH = "/genai/llmgw/chat/completions" @@ -14,7 +15,7 @@ LLMGW_PATH = "/genai/llmgw/chat/completions" class DataRobotConfig(OpenAILikeChatConfig): @staticmethod - def _resolve_api_key(api_key: Optional[str] = None) -> str: + def _resolve_api_key(api_key: str | None = None) -> str: """Attempt to ensure that the API key is set, preferring the user-provided key over the secret manager key (``DATAROBOT_API_TOKEN``). @@ -23,7 +24,7 @@ class DataRobotConfig(OpenAILikeChatConfig): return api_key or get_secret_str("DATAROBOT_API_TOKEN") or "fake-api-key" @staticmethod - def _resolve_api_base(api_base: Optional[str] = None) -> Optional[str]: + def _resolve_api_base(api_base: str | None = None) -> str | None: """Attempt to ensure that the API base is set, preferring the user-provided key over the secret manager key (``DATAROBOT_ENDPOINT``). @@ -54,8 +55,8 @@ class DataRobotConfig(OpenAILikeChatConfig): return urlunparse(updated_parsed) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: """Attempts to ensure that the API base and key are set, preferring user-provided values, before falling back to secret manager values (``DATAROBOT_ENDPOINT`` and ``DATAROBOT_API_TOKEN`` respectively). @@ -69,12 +70,12 @@ class DataRobotConfig(OpenAILikeChatConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the API call. Datarobot's API base is set to diff --git a/litellm/llms/deepgram/audio_transcription/transformation.py b/litellm/llms/deepgram/audio_transcription/transformation.py index b05fba3b5ca..034c41c79fb 100644 --- a/litellm/llms/deepgram/audio_transcription/transformation.py +++ b/litellm/llms/deepgram/audio_transcription/transformation.py @@ -2,7 +2,6 @@ Translates from OpenAI's `/v1/audio/transcriptions` to Deepgram's `/v1/listen` """ -from typing import List, Optional, Union from urllib.parse import urlencode from httpx import Headers, Response @@ -24,7 +23,7 @@ from ..common_utils import DeepgramException class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: return ["language"] def map_openai_params( @@ -40,7 +39,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): optional_params[k] = v return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return DeepgramException(message=error_message, status_code=status_code, headers=headers) def transform_audio_transcription_request( @@ -123,7 +122,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return response except Exception as e: - raise ValueError(f"Error transforming Deepgram response: {str(e)}\nResponse: {raw_response.text}") + raise ValueError(f"Error transforming Deepgram response: {e!s}\nResponse: {raw_response.text}") def _reconstruct_diarized_transcript(self, words: list) -> str: """ @@ -165,12 +164,12 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: api_base = get_secret_str("DEEPGRAM_API_BASE") or "https://api.deepgram.com/v1" @@ -233,11 +232,11 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("DEEPGRAM_API_KEY") return { diff --git a/litellm/llms/deepinfra/chat/transformation.py b/litellm/llms/deepinfra/chat/transformation.py index 0dd1b73f9fa..1ed9436df54 100644 --- a/litellm/llms/deepinfra/chat/transformation.py +++ b/litellm/llms/deepinfra/chat/transformation.py @@ -1,6 +1,6 @@ import json from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, cast, overload +from typing import Any, Literal, cast, overload import litellm from litellm.constants import MIN_NON_ZERO_TEMPERATURE @@ -17,38 +17,38 @@ class DeepInfraConfig(OpenAIGPTConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "deepinfra" - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None - tools: Optional[list] = None - tool_choice: Optional[Union[str, dict]] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None + tools: list | None = None + tool_choice: str | dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, - tools: Optional[list] = None, - tool_choice: Optional[Union[str, dict]] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -105,9 +105,7 @@ class DeepInfraConfig(OpenAIGPTConfig): value = None else: raise litellm.utils.UnsupportedParamsError( - message="Deepinfra doesn't support tool_choice={}. To drop unsupported openai params from the call, set `litellm.drop_params = True`".format( - value - ), + message=f"Deepinfra doesn't support tool_choice={value}. To drop unsupported openai params from the call, set `litellm.drop_params = True`", status_code=400, ) elif param == "max_completion_tokens": @@ -117,7 +115,7 @@ class DeepInfraConfig(OpenAIGPTConfig): optional_params[param] = value return optional_params - def _transform_tool_message_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def _transform_tool_message_content(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ Transform tool message content from array to string format for DeepInfra compatibility. @@ -155,20 +153,20 @@ class DeepInfraConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Transform messages for DeepInfra compatibility. Handles both sync and async transformations. @@ -193,8 +191,8 @@ class DeepInfraConfig(OpenAIGPTConfig): return self._transform_tool_message_content(parent_result) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # deepinfra is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1 api_base = api_base or get_secret_str("DEEPINFRA_API_BASE") or "https://api.deepinfra.com/v1/openai" dynamic_api_key = api_key or get_secret_str("DEEPINFRA_API_KEY") diff --git a/litellm/llms/deepinfra/rerank/transformation.py b/litellm/llms/deepinfra/rerank/transformation.py index 82069e4e195..f43aa74659b 100644 --- a/litellm/llms/deepinfra/rerank/transformation.py +++ b/litellm/llms/deepinfra/rerank/transformation.py @@ -2,7 +2,7 @@ Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format. """ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -93,15 +93,15 @@ class DeepinfraRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: # Start with the basic parameters optional_rerank_params = {} if query: @@ -127,7 +127,7 @@ class DeepinfraRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -210,9 +210,7 @@ class DeepinfraRerankConfig(BaseRerankConfig): def get_supported_cohere_rerank_params(self, model: str) -> list: return ["query", "documents"] - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: # Deepinfra errors may come as JSON: {"detail": {"error": "..."}} import json diff --git a/litellm/llms/deepseek/chat/transformation.py b/litellm/llms/deepseek/chat/transformation.py index d04105ad614..60145aa9c9f 100644 --- a/litellm/llms/deepseek/chat/transformation.py +++ b/litellm/llms/deepseek/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to DeepSeek's `/v1/chat/completi """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, cast, overload +from typing import Any, Literal, cast, overload import litellm from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -60,7 +60,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): return optional_params - def _fill_reasoning_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def _fill_reasoning_content(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ DeepSeek thinking mode requires `reasoning_content` to be passed back on every assistant message in multi-turn conversations. If it is missing, @@ -72,7 +72,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): (LiteLLM stores provider-specific response fields there). 2. Otherwise inject a single space — the minimum value the API accepts. """ - result: List[AllMessageValues] = [] + result: list[AllMessageValues] = [] for msg in messages: if msg.get("role") == "assistant" and not msg.get("reasoning_content"): patched = dict(cast(dict, msg)) @@ -102,20 +102,20 @@ class DeepSeekChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ DeepSeek does not support content in list format. """ @@ -208,7 +208,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -236,7 +236,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -257,20 +257,20 @@ class DeepSeekChatConfig(OpenAIGPTConfig): ) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta" # type: ignore dynamic_api_key = api_key or get_secret_str("DEEPSEEK_API_KEY") return api_base, dynamic_api_key def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ If api_base is not provided, use the default DeepSeek /chat/completions endpoint. diff --git a/litellm/llms/deepseek/cost_calculator.py b/litellm/llms/deepseek/cost_calculator.py index 312bd5bdeab..5a0c065f0ee 100644 --- a/litellm/llms/deepseek/cost_calculator.py +++ b/litellm/llms/deepseek/cost_calculator.py @@ -4,13 +4,11 @@ Cost calculator for DeepSeek Chat models. Handles prompt caching scenario. """ -from typing import Tuple - from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.types.utils import Usage -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/deepseek/messages/transformation.py b/litellm/llms/deepseek/messages/transformation.py index ddbbe7c2107..9ce33dcd135 100644 --- a/litellm/llms/deepseek/messages/transformation.py +++ b/litellm/llms/deepseek/messages/transformation.py @@ -2,7 +2,7 @@ DeepSeek Anthropic-compatible messages transformation config. """ -from typing import Any, Dict, List, Optional, Tuple +from typing import Any import litellm from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -23,18 +23,18 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "deepseek" def should_strip_billing_metadata(self) -> bool: return True @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("DEEPSEEK_API_KEY") or litellm.api_key @staticmethod - def get_api_base(api_base: Optional[str] = None) -> str: + def get_api_base(api_base: str | None = None) -> str: return ( api_base or get_secret_str("DEEPSEEK_ANTHROPIC_API_BASE") @@ -46,12 +46,12 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: dynamic_api_key = self.get_api_key(api_key=api_key) if "x-api-key" not in headers and "authorization" not in headers and dynamic_api_key is not None: @@ -72,23 +72,20 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base_url = self.get_api_base(api_base=api_base).rstrip("/") if base_url.endswith("/v1/messages") and "/anthropic/" in base_url: return base_url - if base_url.endswith("/v1/messages"): - base_url = base_url[: -len("/v1/messages")] - if base_url.endswith("/v1"): - base_url = base_url[: -len("/v1")] - if base_url.endswith("/beta"): - base_url = base_url[: -len("/beta")] + base_url = base_url.removesuffix("/v1/messages") + base_url = base_url.removesuffix("/v1") + base_url = base_url.removesuffix("/beta") if not base_url.endswith("/anthropic") and "/anthropic/" not in base_url: base_url = f"{base_url}/anthropic" @@ -113,11 +110,11 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: anthropic_messages_request = super().transform_anthropic_messages_request( model=model, messages=messages, diff --git a/litellm/llms/deprecated_providers/aleph_alpha.py b/litellm/llms/deprecated_providers/aleph_alpha.py index a9061eac453..343f7475288 100644 --- a/litellm/llms/deprecated_providers/aleph_alpha.py +++ b/litellm/llms/deprecated_providers/aleph_alpha.py @@ -2,7 +2,6 @@ import json import time import types from collections.abc import Callable -from typing import Optional import httpx # type: ignore @@ -74,71 +73,71 @@ class AlephAlphaConfig: - `control_log_additive` (boolean; default value: true): Method of applying control to attention scores. """ - maximum_tokens: Optional[int] = litellm.max_tokens # aleph alpha requires max tokens - minimum_tokens: Optional[int] = None - echo: Optional[bool] = None - temperature: Optional[int] = None - top_k: Optional[int] = None - top_p: Optional[int] = None - presence_penalty: Optional[int] = None - frequency_penalty: Optional[int] = None - sequence_penalty: Optional[int] = None - sequence_penalty_min_length: Optional[int] = None - repetition_penalties_include_prompt: Optional[bool] = None - repetition_penalties_include_completion: Optional[bool] = None - use_multiplicative_presence_penalty: Optional[bool] = None - use_multiplicative_frequency_penalty: Optional[bool] = None - use_multiplicative_sequence_penalty: Optional[bool] = None - penalty_bias: Optional[str] = None - penalty_exceptions_include_stop_sequences: Optional[bool] = None - best_of: Optional[int] = None - n: Optional[int] = None - logit_bias: Optional[dict] = None - log_probs: Optional[int] = None - stop_sequences: Optional[list] = None - tokens: Optional[bool] = None - raw_completion: Optional[bool] = None - disable_optimizations: Optional[bool] = None - completion_bias_inclusion: Optional[list] = None - completion_bias_exclusion: Optional[list] = None - completion_bias_inclusion_first_token_only: Optional[bool] = None - completion_bias_exclusion_first_token_only: Optional[bool] = None - contextual_control_threshold: Optional[int] = None - control_log_additive: Optional[bool] = None + maximum_tokens: int | None = litellm.max_tokens # aleph alpha requires max tokens + minimum_tokens: int | None = None + echo: bool | None = None + temperature: int | None = None + top_k: int | None = None + top_p: int | None = None + presence_penalty: int | None = None + frequency_penalty: int | None = None + sequence_penalty: int | None = None + sequence_penalty_min_length: int | None = None + repetition_penalties_include_prompt: bool | None = None + repetition_penalties_include_completion: bool | None = None + use_multiplicative_presence_penalty: bool | None = None + use_multiplicative_frequency_penalty: bool | None = None + use_multiplicative_sequence_penalty: bool | None = None + penalty_bias: str | None = None + penalty_exceptions_include_stop_sequences: bool | None = None + best_of: int | None = None + n: int | None = None + logit_bias: dict | None = None + log_probs: int | None = None + stop_sequences: list | None = None + tokens: bool | None = None + raw_completion: bool | None = None + disable_optimizations: bool | None = None + completion_bias_inclusion: list | None = None + completion_bias_exclusion: list | None = None + completion_bias_inclusion_first_token_only: bool | None = None + completion_bias_exclusion_first_token_only: bool | None = None + contextual_control_threshold: int | None = None + control_log_additive: bool | None = None def __init__( self, - maximum_tokens: Optional[int] = None, - minimum_tokens: Optional[int] = None, - echo: Optional[bool] = None, - temperature: Optional[int] = None, - top_k: Optional[int] = None, - top_p: Optional[int] = None, - presence_penalty: Optional[int] = None, - frequency_penalty: Optional[int] = None, - sequence_penalty: Optional[int] = None, - sequence_penalty_min_length: Optional[int] = None, - repetition_penalties_include_prompt: Optional[bool] = None, - repetition_penalties_include_completion: Optional[bool] = None, - use_multiplicative_presence_penalty: Optional[bool] = None, - use_multiplicative_frequency_penalty: Optional[bool] = None, - use_multiplicative_sequence_penalty: Optional[bool] = None, - penalty_bias: Optional[str] = None, - penalty_exceptions_include_stop_sequences: Optional[bool] = None, - best_of: Optional[int] = None, - n: Optional[int] = None, - logit_bias: Optional[dict] = None, - log_probs: Optional[int] = None, - stop_sequences: Optional[list] = None, - tokens: Optional[bool] = None, - raw_completion: Optional[bool] = None, - disable_optimizations: Optional[bool] = None, - completion_bias_inclusion: Optional[list] = None, - completion_bias_exclusion: Optional[list] = None, - completion_bias_inclusion_first_token_only: Optional[bool] = None, - completion_bias_exclusion_first_token_only: Optional[bool] = None, - contextual_control_threshold: Optional[int] = None, - control_log_additive: Optional[bool] = None, + maximum_tokens: int | None = None, + minimum_tokens: int | None = None, + echo: bool | None = None, + temperature: int | None = None, + top_k: int | None = None, + top_p: int | None = None, + presence_penalty: int | None = None, + frequency_penalty: int | None = None, + sequence_penalty: int | None = None, + sequence_penalty_min_length: int | None = None, + repetition_penalties_include_prompt: bool | None = None, + repetition_penalties_include_completion: bool | None = None, + use_multiplicative_presence_penalty: bool | None = None, + use_multiplicative_frequency_penalty: bool | None = None, + use_multiplicative_sequence_penalty: bool | None = None, + penalty_bias: str | None = None, + penalty_exceptions_include_stop_sequences: bool | None = None, + best_of: int | None = None, + n: int | None = None, + logit_bias: dict | None = None, + log_probs: int | None = None, + stop_sequences: list | None = None, + tokens: bool | None = None, + raw_completion: bool | None = None, + disable_optimizations: bool | None = None, + completion_bias_inclusion: list | None = None, + completion_bias_exclusion: list | None = None, + completion_bias_inclusion_first_token_only: bool | None = None, + completion_bias_exclusion_first_token_only: bool | None = None, + contextual_control_threshold: int | None = None, + control_log_additive: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/deprecated_providers/palm.py b/litellm/llms/deprecated_providers/palm.py index 45ff8870075..146d007e0bd 100644 --- a/litellm/llms/deprecated_providers/palm.py +++ b/litellm/llms/deprecated_providers/palm.py @@ -3,7 +3,6 @@ import time import traceback import types from collections.abc import Callable -from typing import Optional import httpx @@ -44,23 +43,23 @@ class PalmConfig: - `max_output_tokens` (int): Sets the maximum number of tokens to be returned in the output """ - context: Optional[str] = None - examples: Optional[list] = None - temperature: Optional[float] = None - candidate_count: Optional[int] = None - top_k: Optional[int] = None - top_p: Optional[float] = None - max_output_tokens: Optional[int] = None + context: str | None = None + examples: list | None = None + temperature: float | None = None + candidate_count: int | None = None + top_k: int | None = None + top_p: float | None = None + max_output_tokens: int | None = None def __init__( self, - context: Optional[str] = None, - examples: Optional[list] = None, - temperature: Optional[float] = None, - candidate_count: Optional[int] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - max_output_tokens: Optional[int] = None, + context: str | None = None, + examples: list | None = None, + temperature: float | None = None, + candidate_count: int | None = None, + top_k: int | None = None, + top_p: float | None = None, + max_output_tokens: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/docker_model_runner/chat/transformation.py b/litellm/llms/docker_model_runner/chat/transformation.py index 2db070699dd..09afb85c91f 100644 --- a/litellm/llms/docker_model_runner/chat/transformation.py +++ b/litellm/llms/docker_model_runner/chat/transformation.py @@ -5,7 +5,7 @@ Docker Model Runner API Reference: https://docs.docker.com/ai/model-runner/api-r """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, overload +from typing import Any, Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -26,20 +26,20 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Docker Model Runner is OpenAI-compatible, so we use standard message transformation. """ @@ -50,8 +50,8 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: """ Get API base and key for Docker Model Runner. @@ -67,12 +67,12 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Build the complete URL for Docker Model Runner API. diff --git a/litellm/llms/duckduckgo/search/transformation.py b/litellm/llms/duckduckgo/search/transformation.py index 0ef21222a29..0d28dfa59fe 100644 --- a/litellm/llms/duckduckgo/search/transformation.py +++ b/litellm/llms/duckduckgo/search/transformation.py @@ -4,7 +4,7 @@ Calls DuckDuckGo's Instant Answer API to search the web. DuckDuckGo API Reference: https://duckduckgo.com/api """ -from typing import Dict, List, Literal, Optional, TypedDict, Union +from typing import Literal, TypedDict from urllib.parse import urlencode import httpx @@ -56,11 +56,11 @@ class DuckDuckGoSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. DuckDuckGo Instant Answer API does not require authentication. @@ -71,9 +71,9 @@ class DuckDuckGoSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -92,10 +92,10 @@ class DuckDuckGoSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to DuckDuckGo API format. diff --git a/litellm/llms/e2b/sandbox/transformation.py b/litellm/llms/e2b/sandbox/transformation.py index a78f1d8541e..c74682309c1 100644 --- a/litellm/llms/e2b/sandbox/transformation.py +++ b/litellm/llms/e2b/sandbox/transformation.py @@ -8,15 +8,15 @@ Talks to e2b's REST API directly over httpx (no e2b SDK dependency): """ import json -from typing import Union, cast +from typing import cast import httpx from litellm.llms.base_llm.sandbox.transformation import ( + SANDBOX_MAX_OUTPUT_BYTES, BaseSandboxConfig, CodeExecutionResult, ContainerHandle, - SANDBOX_MAX_OUTPUT_BYTES, ) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -94,7 +94,7 @@ class E2BSandboxConfig(BaseSandboxConfig): async def arun_code( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, code: str, api_key: str | None = None, env_vars: dict | None = None, @@ -132,7 +132,7 @@ class E2BSandboxConfig(BaseSandboxConfig): async def adelete_sandbox( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, api_key: str | None = None, api_base: str | None = None, client: AsyncHTTPHandler | None = None, @@ -156,7 +156,7 @@ class E2BSandboxConfig(BaseSandboxConfig): return 200 <= response.status_code < 300 @staticmethod - def _as_handle(container: Union[ContainerHandle, str]) -> ContainerHandle: + def _as_handle(container: ContainerHandle | str) -> ContainerHandle: if isinstance(container, ContainerHandle): return container handle = ContainerHandle(id=str(container), provider="e2b", domain=E2B_DEFAULT_DOMAIN) diff --git a/litellm/llms/elevenlabs/audio_transcription/transformation.py b/litellm/llms/elevenlabs/audio_transcription/transformation.py index 68d1b5e16dd..a33e221dafd 100644 --- a/litellm/llms/elevenlabs/audio_transcription/transformation.py +++ b/litellm/llms/elevenlabs/audio_transcription/transformation.py @@ -2,8 +2,6 @@ Translates from OpenAI's `/v1/audio/transcriptions` to ElevenLabs's `/v1/speech-to-text` """ -from typing import List, Optional, Union - from httpx import Headers, Response import litellm @@ -28,7 +26,7 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def custom_llm_provider(self) -> str: return litellm.LlmProviders.ELEVENLABS.value - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: return ["language", "temperature"] def map_openai_params( @@ -48,7 +46,7 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): optional_params[k] = v return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return ElevenLabsException(message=error_message, status_code=status_code, headers=headers) def transform_audio_transcription_request( @@ -146,16 +144,16 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return response except Exception as e: - raise ValueError(f"Error transforming ElevenLabs response: {str(e)}\nResponse: {raw_response.text}") + raise ValueError(f"Error transforming ElevenLabs response: {e!s}\nResponse: {raw_response.text}") def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: api_base = get_secret_str("ELEVENLABS_API_BASE") or "https://api.elevenlabs.io" @@ -170,11 +168,11 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("ELEVENLABS_API_KEY") if api_key is None: diff --git a/litellm/llms/elevenlabs/text_to_speech/transformation.py b/litellm/llms/elevenlabs/text_to_speech/transformation.py index b5b7799a3e9..46ecdef7b6e 100644 --- a/litellm/llms/elevenlabs/text_to_speech/transformation.py +++ b/litellm/llms/elevenlabs/text_to_speech/transformation.py @@ -4,7 +4,7 @@ Elevenlabs Text-to-Speech transformation Maps OpenAI TTS spec to Elevenlabs TTS API """ -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from urllib.parse import urlencode import httpx @@ -80,13 +80,13 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): def _resolve_voice_id( self, - voice: Optional[Union[str, Dict[str, Any]]], - params: Dict[str, Any], + voice: str | dict[str, Any] | None, + params: dict[str, Any], ) -> str: """ Determine the ElevenLabs voice_id based on provided voice input or parameters. """ - mapped_voice: Optional[str] = None + mapped_voice: str | None = None if isinstance(voice, str) and voice.strip(): mapped_voice = self._extract_voice_id(voice) @@ -112,20 +112,20 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Optional[Dict[str, Any]] = None, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict[str, Any] | None = None, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to ElevenLabs TTS parameters """ - mapped_params: Dict[str, Any] = {} - query_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} + query_params: dict[str, Any] = {} # Work on a copy so we don't mutate the caller's dictionary params = dict(optional_params) if optional_params else {} - passthrough_kwargs: Dict[str, Any] = kwargs if kwargs is not None else {} + passthrough_kwargs: dict[str, Any] = kwargs if kwargs is not None else {} # Extract voice identifier mapped_voice = self._resolve_voice_id(voice, params) @@ -140,7 +140,7 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): # Drop it to avoid sending unsupported keys unless caller already provided voice_settings. speed = params.pop("speed", None) if speed is not None: - speed_value: Optional[float] + speed_value: float | None try: speed_value = float(speed) except (TypeError, ValueError): @@ -167,8 +167,8 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate Azure environment and set up authentication headers @@ -187,16 +187,16 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): return headers - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return ElevenLabsException(message=error_message, status_code=status_code, headers=headers) def transform_text_to_speech_request( self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ @@ -205,7 +205,7 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): params = dict(optional_params) if optional_params else {} extra_body = params.pop("extra_body", None) - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "text": input, "model_id": model, } @@ -229,10 +229,10 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): def _add_elevenlabs_specific_params( self, mapped_voice: str, - query_params: Dict[str, Any], - mapped_params: Dict[str, Any], - kwargs: Optional[Dict[str, Any]], - remaining_params: Dict[str, Any], + query_params: dict[str, Any], + mapped_params: dict[str, Any], + kwargs: dict[str, Any] | None, + remaining_params: dict[str, Any], ) -> None: if kwargs is None: kwargs = {} @@ -291,7 +291,7 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/exa_ai/search/transformation.py b/litellm/llms/exa_ai/search/transformation.py index 93fbdeff990..87efe82453f 100644 --- a/litellm/llms/exa_ai/search/transformation.py +++ b/litellm/llms/exa_ai/search/transformation.py @@ -4,7 +4,7 @@ Calls Exa AI's /search endpoint to search the web. Exa AI API Reference: https://docs.exa.ai/reference/search """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -33,15 +33,15 @@ class ExaAISearchRequest(_ExaAISearchRequestRequired, total=False): category: str # Optional - data category ('company', 'research paper', 'news', 'pdf', 'github', 'tweet', 'personal site', 'linkedin profile', 'financial report') userLocation: str # Optional - two-letter ISO country code numResults: int # Optional - number of results (max 100), default 10 - includeDomains: List[str] # Optional - list of domains to include - excludeDomains: List[str] # Optional - list of domains to exclude + includeDomains: list[str] # Optional - list of domains to include + excludeDomains: list[str] # Optional - list of domains to exclude startCrawlDate: str # Optional - crawl date filter (ISO 8601 format) endCrawlDate: str # Optional - crawl date filter (ISO 8601 format) startPublishedDate: str # Optional - published date filter (ISO 8601 format) endPublishedDate: str # Optional - published date filter (ISO 8601 format) - includeText: List[str] # Optional - strings that must be present in webpage text - excludeText: List[str] # Optional - strings that must not be present in webpage text - context: Union[bool, dict] # Optional - format results for LLMs + includeText: list[str] # Optional - strings that must be present in webpage text + excludeText: list[str] # Optional - strings that must not be present in webpage text + context: bool | dict # Optional - format results for LLMs moderation: bool # Optional - enable content moderation, default false contents: dict # Optional - content retrieval options @@ -55,11 +55,11 @@ class ExaAISearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -78,9 +78,9 @@ class ExaAISearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -96,10 +96,10 @@ class ExaAISearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Exa AI API format. diff --git a/litellm/llms/fal_ai/__init__.py b/litellm/llms/fal_ai/__init__.py index 0de526a8eb7..10e9c43d6da 100644 --- a/litellm/llms/fal_ai/__init__.py +++ b/litellm/llms/fal_ai/__init__.py @@ -13,15 +13,15 @@ from .image_generation import ( ) __all__ = [ - "cost_calculator", "FalAIBaseConfig", - "FalAIImageGenerationConfig", - "FalAIImagen4Config", - "FalAIRecraftV3Config", "FalAIBriaConfig", "FalAIFluxProV11Config", "FalAIFluxProV11UltraConfig", "FalAIFluxSchnellConfig", + "FalAIImageGenerationConfig", + "FalAIImagen4Config", + "FalAIRecraftV3Config", "FalAIStableDiffusionConfig", + "cost_calculator", "get_fal_ai_image_generation_config", ] diff --git a/litellm/llms/fal_ai/image_generation/__init__.py b/litellm/llms/fal_ai/image_generation/__init__.py index d31524510b8..2b1e579f010 100644 --- a/litellm/llms/fal_ai/image_generation/__init__.py +++ b/litellm/llms/fal_ai/image_generation/__init__.py @@ -3,34 +3,34 @@ from litellm.llms.base_llm.image_generation.transformation import ( ) from .bria_transformation import FalAIBriaConfig +from .bytedance_transformation import ( + FalAIBytedanceDreaminaV31Config, + FalAIBytedanceSeedreamV3Config, +) from .flux_pro_v11_transformation import FalAIFluxProV11Config from .flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig from .flux_schnell_transformation import FalAIFluxSchnellConfig +from .ideogram_v3_transformation import FalAIIdeogramV3Config from .imagen4_transformation import FalAIImagen4Config from .nano_banana_transformation import FalAINanoBananaConfig from .recraft_v3_transformation import FalAIRecraftV3Config -from .ideogram_v3_transformation import FalAIIdeogramV3Config from .stable_diffusion_transformation import FalAIStableDiffusionConfig from .transformation import FalAIBaseConfig, FalAIImageGenerationConfig -from .bytedance_transformation import ( - FalAIBytedanceSeedreamV3Config, - FalAIBytedanceDreaminaV31Config, -) __all__ = [ "FalAIBaseConfig", + "FalAIBriaConfig", + "FalAIBytedanceDreaminaV31Config", + "FalAIBytedanceSeedreamV3Config", + "FalAIFluxProV11Config", + "FalAIFluxProV11UltraConfig", + "FalAIFluxSchnellConfig", + "FalAIIdeogramV3Config", "FalAIImageGenerationConfig", "FalAIImagen4Config", "FalAINanoBananaConfig", "FalAIRecraftV3Config", - "FalAIBriaConfig", - "FalAIFluxProV11Config", - "FalAIFluxProV11UltraConfig", - "FalAIFluxSchnellConfig", "FalAIStableDiffusionConfig", - "FalAIBytedanceSeedreamV3Config", - "FalAIBytedanceDreaminaV31Config", - "FalAIIdeogramV3Config", ] diff --git a/litellm/llms/fal_ai/image_generation/bria_transformation.py b/litellm/llms/fal_ai/image_generation/bria_transformation.py index 7bdfa860c5d..1d601a601fc 100644 --- a/litellm/llms/fal_ai/image_generation/bria_transformation.py +++ b/litellm/llms/fal_ai/image_generation/bria_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -28,7 +28,7 @@ class FalAIBriaConfig(FalAIBaseConfig): IMAGE_GENERATION_ENDPOINT: str = "bria/text-to-image/3.2" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for Bria 3.2. """ @@ -60,8 +60,8 @@ class FalAIBriaConfig(FalAIBaseConfig): "size": "aspect_ratio", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # Use mapped parameter name if exists mapped_key = param_mapping.get(k, k) @@ -186,8 +186,8 @@ class FalAIBriaConfig(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the Bria 3.2 response to litellm ImageResponse format. diff --git a/litellm/llms/fal_ai/image_generation/bytedance_transformation.py b/litellm/llms/fal_ai/image_generation/bytedance_transformation.py index b52d08dd9e4..db70e8fc078 100644 --- a/litellm/llms/fal_ai/image_generation/bytedance_transformation.py +++ b/litellm/llms/fal_ai/image_generation/bytedance_transformation.py @@ -36,8 +36,8 @@ class FalAIBytedanceBaseConfig(FalAIFluxProV11UltraConfig): "size": "image_size", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: mapped_key = param_mapping.get(k, k) mapped_value = non_default_params[k] diff --git a/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py b/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py index 5226419a29e..14d89d7c8b7 100644 --- a/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_pro_v11_transformation.py @@ -44,8 +44,8 @@ class FalAIFluxProV11Config(FalAIFluxProV11UltraConfig): "size": "image_size", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: mapped_key = param_mapping.get(k, k) mapped_value = non_default_params[k] 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 fb980905a28..465e658d453 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,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -28,7 +28,7 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig): IMAGE_GENERATION_ENDPOINT: str = "fal-ai/flux-pro/v1.1-ultra" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for Flux Pro v1.1-ultra. """ @@ -62,8 +62,8 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig): "size": "aspect_ratio", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # Use mapped parameter name if exists mapped_key = param_mapping.get(k, k) @@ -193,8 +193,8 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the Flux Pro v1.1-ultra response to litellm ImageResponse format. diff --git a/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py b/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py index 7a59fae6c1a..e3aa620405a 100644 --- a/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_schnell_transformation.py @@ -41,8 +41,8 @@ class FalAIFluxSchnellConfig(FalAIFluxProV11UltraConfig): "size": "image_size", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: mapped_key = param_mapping.get(k, k) mapped_value = non_default_params[k] 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 500a4b20ef2..55b47aa4365 100644 --- a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -38,7 +38,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): "1024x1536": "portrait_16_9", } - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Ideogram v3 accepts the core OpenAI image parameters. """ @@ -62,7 +62,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): + for k in non_default_params: if k in optional_params: continue @@ -149,8 +149,8 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Parse Ideogram v3 responses which contain a list of File objects. diff --git a/litellm/llms/fal_ai/image_generation/imagen4_transformation.py b/litellm/llms/fal_ai/image_generation/imagen4_transformation.py index 1b111c98987..2b7d6ed9b50 100644 --- a/litellm/llms/fal_ai/image_generation/imagen4_transformation.py +++ b/litellm/llms/fal_ai/image_generation/imagen4_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -31,7 +31,7 @@ class FalAIImagen4Config(FalAIBaseConfig): IMAGE_GENERATION_ENDPOINT: str = "fal-ai/imagen4/preview" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for Imagen4. """ @@ -64,8 +64,8 @@ class FalAIImagen4Config(FalAIBaseConfig): "size": "aspect_ratio", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # Use mapped parameter name if exists mapped_key = param_mapping.get(k, k) @@ -181,8 +181,8 @@ class FalAIImagen4Config(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the Imagen4 response to litellm ImageResponse format. diff --git a/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py b/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py index 0a8ba3699bb..658af1cd20b 100644 --- a/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py +++ b/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional - from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams @@ -18,7 +16,7 @@ class FalAINanoBananaConfig(FalAIBaseConfig): Documentation: https://fal.ai/models/fal-ai/nano-banana """ - SUPPORTED_ASPECT_RATIOS: List[str] = [ + SUPPORTED_ASPECT_RATIOS: list[str] = [ "21:9", "16:9", "3:2", @@ -33,18 +31,18 @@ class FalAINanoBananaConfig(FalAIBaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base_url: str = (api_base or get_secret_str("FAL_AI_API_BASE") or self.DEFAULT_BASE_URL).rstrip("/") endpoint = model if model.startswith("fal-ai/") else f"fal-ai/{model}" return f"{base_url}/{endpoint}" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "response_format", "size"] def map_openai_params( 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 2ce36d9c1ea..5233db7673f 100644 --- a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -28,7 +28,7 @@ class FalAIRecraftV3Config(FalAIBaseConfig): IMAGE_GENERATION_ENDPOINT: str = "fal-ai/recraft/v3/text-to-image" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for Recraft v3. """ @@ -60,8 +60,8 @@ class FalAIRecraftV3Config(FalAIBaseConfig): "size": "image_size", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # Use mapped parameter name if exists mapped_key = param_mapping.get(k, k) @@ -171,8 +171,8 @@ class FalAIRecraftV3Config(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the Recraft v3 response to litellm ImageResponse format. 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 bc7a3839bd3..1786a8d4065 100644 --- a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py +++ b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -32,12 +32,12 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request. @@ -63,7 +63,7 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): complete_url = f"{complete_url}/{endpoint}" return complete_url - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for Stable Diffusion models. """ @@ -97,8 +97,8 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): "size": "image_size", } - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: # Use mapped parameter name if exists mapped_key = param_mapping.get(k, k) @@ -207,8 +207,8 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the Stable Diffusion response to litellm ImageResponse format. diff --git a/litellm/llms/fal_ai/image_generation/transformation.py b/litellm/llms/fal_ai/image_generation/transformation.py index 07eb2cc4cc4..985b577e6a0 100644 --- a/litellm/llms/fal_ai/image_generation/transformation.py +++ b/litellm/llms/fal_ai/image_generation/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -31,12 +31,12 @@ class FalAIBaseConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -54,13 +54,13 @@ class FalAIBaseConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("FAL_AI_API_KEY") + final_api_key: str | None = api_key or get_secret_str("FAL_AI_API_KEY") if not final_api_key: raise ValueError("FAL_AI_API_KEY is not set") @@ -77,8 +77,8 @@ class FalAIBaseConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the image generation response to the litellm image response @@ -122,7 +122,7 @@ class FalAIImageGenerationConfig(FalAIBaseConfig): Default Fal AI image generation configuration for generic models. """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for fal.ai image generation """ @@ -140,8 +140,8 @@ class FalAIImageGenerationConfig(FalAIBaseConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: diff --git a/litellm/llms/fastcrw/search/transformation.py b/litellm/llms/fastcrw/search/transformation.py index 6de9ef642fb..b5d60f232b7 100644 --- a/litellm/llms/fastcrw/search/transformation.py +++ b/litellm/llms/fastcrw/search/transformation.py @@ -8,7 +8,7 @@ or cloud). The search response uses the Firecrawl-compatible envelope fastCRW API Reference: https://fastcrw.com/docs/rest-api """ -from typing import Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -48,8 +48,8 @@ class FastCRWSearchConfig(BaseSearchConfig): def validate_environment( self, headers: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, **kwargs, ) -> dict: """ @@ -70,9 +70,9 @@ class FastCRWSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[dict, list[dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -88,7 +88,7 @@ class FastCRWSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, list[str]], + query: str | list[str], optional_params: dict, **kwargs, ) -> dict: diff --git a/litellm/llms/featherless_ai/chat/transformation.py b/litellm/llms/featherless_ai/chat/transformation.py index cf11c72c326..297bf42c0f3 100644 --- a/litellm/llms/featherless_ai/chat/transformation.py +++ b/litellm/llms/featherless_ai/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Tuple, Union - import litellm from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.secret_managers.main import get_secret_str @@ -12,35 +10,35 @@ class FeatherlessAIConfig(OpenAIGPTConfig): The class `FeatherlessAI` provides configuration for the FeatherlessAI's Chat Completions API interface. Below are the parameters: """ - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None - tool_choice: Optional[str] = None - tools: Optional[list] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None + tool_choice: str | None = None + tools: list | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, - tool_choice: Optional[str] = None, - tools: Optional[list] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, + tool_choice: str | None = None, + tools: list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -98,8 +96,8 @@ class FeatherlessAIConfig(OpenAIGPTConfig): return optional_params def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # FeatherlessAI is openai compatible, set to custom_openai and use FeatherlessAI's endpoint api_base = ( api_base @@ -117,8 +115,8 @@ class FeatherlessAIConfig(OpenAIGPTConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if not api_key: raise ValueError("Missing Featherless AI API Key") diff --git a/litellm/llms/firecrawl/search/transformation.py b/litellm/llms/firecrawl/search/transformation.py index 7aac6d7e7dd..1785cf8d3f2 100644 --- a/litellm/llms/firecrawl/search/transformation.py +++ b/litellm/llms/firecrawl/search/transformation.py @@ -4,7 +4,7 @@ Calls Firecrawl's /search endpoint to search the web. Firecrawl API Reference: https://docs.firecrawl.dev/api-reference/endpoint/search """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -30,14 +30,14 @@ class FirecrawlSearchRequest(_FirecrawlSearchRequestRequired, total=False): """ limit: int # Optional - maximum number of results to return (default 5, max 100) - sources: List[str] # Optional - sources to search ('web', 'images', 'news'), default ['web'] - categories: List[Dict[str, str]] # Optional - categories to filter by (github, research, pdf) + sources: list[str] # Optional - sources to search ('web', 'images', 'news'), default ['web'] + categories: list[dict[str, str]] # Optional - categories to filter by (github, research, pdf) tbs: str # Optional - time-based search parameter location: str # Optional - location parameter for geo-targeting country: str # Optional - ISO country code (default 'US') timeout: int # Optional - timeout in milliseconds (default 60000) ignoreInvalidURLs: bool # Optional - exclude invalid URLs (default false) - scrapeOptions: Dict # Optional - options for scraping search results + scrapeOptions: dict # Optional - options for scraping search results class FirecrawlSearchConfig(BaseSearchConfig): @@ -49,11 +49,11 @@ class FirecrawlSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -72,9 +72,9 @@ class FirecrawlSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -90,10 +90,10 @@ class FirecrawlSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Firecrawl API format. diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index e1228ef05d7..9fcb81e00e3 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -2,11 +2,7 @@ import json from collections.abc import AsyncIterator, Iterator from typing import ( Any, - List, Literal, - Optional, - Tuple, - Union, cast, ) @@ -76,42 +72,42 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): The class `FireworksAIConfig` provides configuration for the Fireworks's Chat Completions API interface. Below are the parameters: """ - tools: Optional[list] = None - tool_choice: Optional[Union[str, dict]] = None - max_tokens: Optional[int] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - frequency_penalty: Optional[int] = None - presence_penalty: Optional[int] = None - n: Optional[int] = None - stop: Optional[Union[str, list]] = None - response_format: Optional[dict] = None - user: Optional[str] = None - logprobs: Optional[int] = None - reasoning_effort: Optional[str] = None + tools: list | None = None + tool_choice: str | dict | None = None + max_tokens: int | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + frequency_penalty: int | None = None + presence_penalty: int | None = None + n: int | None = None + stop: str | list | None = None + response_format: dict | None = None + user: str | None = None + logprobs: int | None = None + reasoning_effort: str | None = None - prompt_truncate_len: Optional[int] = None - context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None + prompt_truncate_len: int | None = None + context_length_exceeded_behavior: Literal["error", "truncate"] | None = None def __init__( self, - tools: Optional[list] = None, - tool_choice: Optional[Union[str, dict]] = None, - max_tokens: Optional[int] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - frequency_penalty: Optional[int] = None, - presence_penalty: Optional[int] = None, - n: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - response_format: Optional[dict] = None, - user: Optional[str] = None, - logprobs: Optional[int] = None, - reasoning_effort: Optional[str] = None, - prompt_truncate_len: Optional[int] = None, - context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None, + tools: list | None = None, + tool_choice: str | dict | None = None, + max_tokens: int | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + frequency_penalty: int | None = None, + presence_penalty: int | None = None, + n: int | None = None, + stop: str | list | None = None, + response_format: dict | None = None, + user: str | None = None, + logprobs: int | None = None, + reasoning_effort: str | None = None, + prompt_truncate_len: int | None = None, + context_length_exceeded_behavior: Literal["error", "truncate"] | None = None, ) -> None: OpenAIGPTConfig.__init__( self, @@ -136,7 +132,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, api_key: str | None = None, @@ -281,7 +277,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): return optional_params - def _transform_tools(self, tools: List[OpenAIChatCompletionToolParam]) -> List[OpenAIChatCompletionToolParam]: + def _transform_tools(self, tools: list[OpenAIChatCompletionToolParam]) -> list[OpenAIChatCompletionToolParam]: for tool in tools: if tool.get("type") != "function": continue @@ -293,8 +289,8 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): return tools def _transform_messages_helper( - self, messages: List[AllMessageValues], model: str, litellm_params: dict - ) -> List[AllMessageValues]: + self, messages: list[AllMessageValues], model: str, litellm_params: dict + ) -> list[AllMessageValues]: """ Strip fields not permitted by FireworksAI from messages. """ @@ -349,25 +345,24 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): # mutation_generation): the generation counter is bumped on every # register_model / reload path, so add+remove or in-place value # replacement (which can leave id and len unchanged) still invalidates. - _fireworks_index_cache: Optional[Tuple[int, int, List[Tuple[str, dict]]]] = None + _fireworks_index_cache: tuple[int, int, list[tuple[str, dict]]] | None = None @classmethod - def _get_fireworks_index(cls) -> List[Tuple[str, dict]]: + def _get_fireworks_index(cls) -> list[tuple[str, dict]]: model_cost = litellm.model_cost signature = (id(model_cost), get_model_cost_mutation_generation()) cached = cls._fireworks_index_cache if cached is not None and cached[0] == signature[0] and cached[1] == signature[1]: return cached[2] - index: List[Tuple[str, dict]] = [] + index: list[tuple[str, dict]] = [] for key, model_info in model_cost.items(): if not key.startswith("fireworks_ai/"): continue if not isinstance(model_info, dict): continue key_short = key[len("fireworks_ai/") :] - if key_short.startswith("accounts/fireworks/models/"): - key_short = key_short[len("accounts/fireworks/models/") :] + key_short = key_short.removeprefix("accounts/fireworks/models/") if not key_short: continue index.append((key_short, model_info)) @@ -392,13 +387,11 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): @staticmethod def _short_model_name(model: str) -> str: short_name = model - if short_name.startswith("fireworks_ai/"): - short_name = short_name[len("fireworks_ai/") :] - if short_name.startswith("accounts/fireworks/models/"): - short_name = short_name[len("accounts/fireworks/models/") :] + short_name = short_name.removeprefix("fireworks_ai/") + short_name = short_name.removeprefix("accounts/fireworks/models/") return short_name - def _get_model_cost_capability_exact(self, model: str, capability: str) -> Optional[bool]: + def _get_model_cost_capability_exact(self, model: str, capability: str) -> bool | None: short_name = self._short_model_name(model) candidate_keys = ( model, @@ -408,10 +401,10 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): for candidate_key in candidate_keys: model_info = litellm.model_cost.get(candidate_key) if model_info is not None and model_info.get(capability) is not None: - return cast(Optional[bool], model_info.get(capability)) + return cast(bool | None, model_info.get(capability)) return None - def _get_model_cost_capability(self, model: str, capability: str) -> Optional[bool]: + def _get_model_cost_capability(self, model: str, capability: str) -> bool | None: exact = self._get_model_cost_capability_exact(model=model, capability=capability) if exact is not None: return exact @@ -426,7 +419,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): # custom deployment. short_name = self._short_model_name(model) matches = [ - (key_short, cast(Optional[bool], model_info.get(capability))) + (key_short, cast(bool | None, model_info.get(capability))) for key_short, model_info in self._get_fireworks_index() if model_info.get(capability) is not None and self._matches_on_hyphen_boundary(short_name, key_short) ] @@ -465,7 +458,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -499,7 +492,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): def _handle_message_content_with_tool_calls( self, message: Message, - tool_calls: Optional[List[ChatCompletionToolParam]], + tool_calls: list[ChatCompletionToolParam] | None, ) -> Message: """ Fireworks AI sends tool calls in the content field instead of tool_calls @@ -528,12 +521,12 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -549,7 +542,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): except Exception as e: response_headers = getattr(raw_response, "headers", None) raise FireworksAIException( - message="Unable to get json response - {}, Original Response: {}".format(str(e), raw_response.text), + message=f"Unable to get json response - {e!s}, Original Response: {raw_response.text}", status_code=raw_response.status_code, headers=response_headers, ) @@ -579,9 +572,9 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return FireworksAIChatCompletionStreamingHandler( streaming_response=streaming_response, @@ -590,8 +583,8 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): ) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("FIREWORKS_API_BASE") or "https://api.fireworks.ai/inference/v1" # type: ignore dynamic_api_key = api_key or ( get_secret_str("FIREWORKS_API_KEY") @@ -601,7 +594,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): ) return api_base, dynamic_api_key - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None): + def get_models(self, api_key: str | None = None, api_base: str | None = None): api_base, api_key = self._get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) if api_base is None or api_key is None: raise ValueError( @@ -615,8 +608,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): ) base = api_base.rstrip("/") - if base.endswith("/v1"): - base = base[: -len("/v1")] + base = base.removesuffix("/v1") response = litellm.module_level_client.get( url=f"{base}/v1/accounts/{account_id}/models", headers={"Authorization": f"Bearer {api_key}"}, @@ -632,7 +624,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): return ["fireworks_ai/" + model["name"] for model in models] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or ( get_secret_str("FIREWORKS_API_KEY") or get_secret_str("FIREWORKS_AI_API_KEY") diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index 51ed8afbbd2..f3f933322c0 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - from httpx import Headers from litellm.secret_managers.main import get_secret_str @@ -34,14 +32,14 @@ class FireworksAIMixin: Common Base Config functions across Fireworks AI Endpoints """ - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return FireworksAIException( status_code=status_code, message=error_message, headers=headers, ) - def _get_api_key(self, api_key: Optional[str]) -> Optional[str]: + def _get_api_key(self, api_key: str | None) -> str | None: dynamic_api_key = api_key or ( get_secret_str("FIREWORKS_API_KEY") or get_secret_str("FIREWORKS_AI_API_KEY") @@ -54,17 +52,17 @@ class FireworksAIMixin: self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = self._get_api_key(api_key) if api_key is None: raise ValueError("FIREWORKS_API_KEY is not set") - auth_headers = {"Authorization": "Bearer {}".format(api_key), **headers} + auth_headers = {"Authorization": f"Bearer {api_key}", **headers} content_type_header = ( {} if any(key.lower() == "content-type" for key in auth_headers) else {"Content-Type": "application/json"} ) diff --git a/litellm/llms/fireworks_ai/completion/transformation.py b/litellm/llms/fireworks_ai/completion/transformation.py index 3ac77288c70..a0c22483cd9 100644 --- a/litellm/llms/fireworks_ai/completion/transformation.py +++ b/litellm/llms/fireworks_ai/completion/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Union - from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUserMessage from ...base_llm.completion.transformation import BaseTextCompletionConfig @@ -44,7 +42,7 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig def transform_text_completion_request( self, model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], optional_params: dict, headers: dict, ) -> dict: diff --git a/litellm/llms/fireworks_ai/cost_calculator.py b/litellm/llms/fireworks_ai/cost_calculator.py index 682adf5a8ff..db2f314f885 100644 --- a/litellm/llms/fireworks_ai/cost_calculator.py +++ b/litellm/llms/fireworks_ai/cost_calculator.py @@ -2,8 +2,6 @@ For calculating cost of fireworks ai serverless inference models. """ -from typing import Tuple - from litellm.constants import ( FIREWORKS_AI_4_B, FIREWORKS_AI_16_B, @@ -54,7 +52,7 @@ def get_base_model_for_pricing(model_name: str) -> str: return "fireworks-ai-default" -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index 393a6c5a8e5..7979eeeba42 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -4,7 +4,7 @@ Fireworks AI Rerank API transformation Reference: https://docs.fireworks.ai/inference-api-reference/rerank """ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -37,9 +37,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): # Remove trailing slashes and ensure clean base URL api_base = api_base.rstrip("/") if not api_base.endswith("/rerank"): - if api_base.endswith("/v1"): - api_base = f"{api_base}/rerank" - elif api_base.endswith("/inference/v1"): + if api_base.endswith("/v1") or api_base.endswith("/inference/v1"): api_base = f"{api_base}/rerank" else: api_base = f"{api_base}/inference/v1/rerank" @@ -60,19 +58,19 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map Cohere rerank params to Fireworks AI rerank params """ - params: Dict[str, Any] = { + params: dict[str, Any] = { "query": query, "documents": documents, } @@ -126,7 +124,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -180,7 +178,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): raw_response_json = raw_response.json() except Exception as e: raise self.get_error_class( - error_message=f"Failed to parse response: {str(e)}", + error_message=f"Failed to parse response: {e!s}", status_code=raw_response.status_code, headers=raw_response.headers, ) @@ -213,12 +211,12 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) # Extract results - Fireworks AI uses "data" instead of "results" - _results: List[dict] | None = raw_response_json.get("data") or raw_response_json.get("results") + _results: list[dict] | None = raw_response_json.get("data") or raw_response_json.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") - rerank_results: List[RerankResponseResult] = [] + rerank_results: list[RerankResponseResult] = [] for result in _results: # Validate required fields exist diff --git a/litellm/llms/gdc/chat/transformation.py b/litellm/llms/gdc/chat/transformation.py index 61631920a64..0416d246ea1 100644 --- a/litellm/llms/gdc/chat/transformation.py +++ b/litellm/llms/gdc/chat/transformation.py @@ -220,7 +220,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig): AttributeError, ) as e: raise litellm.utils.AuthenticationError( - message=f"Failed to load service account credentials from api_key: {str(e)}", + message=f"Failed to load service account credentials from api_key: {e!s}", llm_provider="gdc", model=model, ) from e diff --git a/litellm/llms/gemini/agents/transformation.py b/litellm/llms/gemini/agents/transformation.py index 9e1f6935da4..68acde77efd 100644 --- a/litellm/llms/gemini/agents/transformation.py +++ b/litellm/llms/gemini/agents/transformation.py @@ -9,7 +9,7 @@ Proxies the Gemini v1beta Agents API: GET /v1beta/agents/{name}/versions list versions """ -from typing import Any, Dict, Optional, Tuple, Union +from typing import Any import httpx @@ -65,7 +65,7 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def api_version(self) -> str: return "v1beta" - def _base_url(self, api_base: Optional[str]) -> str: + def _base_url(self, api_base: str | None) -> str: return f"{GeminiModelInfo.get_api_base(api_base)}/{self.api_version}" # ------------------------------------------------------------------ # @@ -76,7 +76,7 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> Exception: return GeminiError( message=error_message, @@ -86,16 +86,16 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def get_complete_url( self, - api_base: Optional[str], - litellm_params: Dict[str, Any], + api_base: str | None, + litellm_params: dict[str, Any], ) -> str: return f"{self._base_url(api_base)}/agents" def validate_environment( self, - headers: Dict[str, str], - litellm_params: Dict[str, Any], - ) -> Dict[str, str]: + headers: dict[str, str], + litellm_params: dict[str, Any], + ) -> dict[str, str]: headers = dict(headers) headers["Content-Type"] = "application/json" explicit_api_key = litellm_params.get("api_key") @@ -132,9 +132,9 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def transform_create_request( self, name: str, - litellm_params: Dict[str, Any], - ) -> Dict[str, Any]: - body: Dict[str, Any] = {"name": name} + litellm_params: dict[str, Any], + ) -> dict[str, Any]: + body: dict[str, Any] = {"name": name} for key in _GEMINI_AGENT_BODY_KEYS: value = litellm_params.get(key) if value is not None: @@ -154,7 +154,7 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): """ self._raise_for_status(raw_response) try: - data: Dict[str, Any] = raw_response.json() + data: dict[str, Any] = raw_response.json() except Exception: verbose_logger.warning( "GeminiAgentsConfig: non-JSON create response (status=%d).", @@ -173,11 +173,11 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def transform_list_request( self, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: url = f"{self._base_url(api_base)}/agents" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if litellm_params.get("page_size"): params["pageSize"] = litellm_params["page_size"] if litellm_params.get("page_token"): @@ -206,9 +206,9 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def transform_get_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: url = f"{self._base_url(api_base)}/agents/{name}" return url, {} @@ -235,8 +235,8 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def transform_delete_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], + api_base: str | None, + litellm_params: dict[str, Any], ) -> str: return f"{self._base_url(api_base)}/agents/{name}" @@ -261,11 +261,11 @@ class GeminiAgentsConfig(BaseAgentsAPIConfig): def transform_list_versions_request( self, name: str, - api_base: Optional[str], - litellm_params: Dict[str, Any], - ) -> Tuple[str, Dict[str, Any]]: + api_base: str | None, + litellm_params: dict[str, Any], + ) -> tuple[str, dict[str, Any]]: url = f"{self._base_url(api_base)}/agents/{name}/versions" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if litellm_params.get("page_size"): params["pageSize"] = litellm_params["page_size"] if litellm_params.get("page_token"): diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index 94130ac4a6e..d9c500fdbf4 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -1,7 +1,6 @@ -from typing import List, Optional, cast +from typing import cast import litellm - from litellm.litellm_core_utils.prompt_templates.factory import ( convert_generic_image_chunk_to_openai_image_obj, convert_to_anthropic_image_obj, @@ -42,25 +41,25 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): Note: Please make sure to modify the default parameters as required for your use case. """ - temperature: Optional[float] = None - max_output_tokens: Optional[int] = None - top_p: Optional[float] = None - top_k: Optional[int] = None - response_mime_type: Optional[str] = None - response_schema: Optional[dict] = None - candidate_count: Optional[int] = None - stop_sequences: Optional[list] = None + temperature: float | None = None + max_output_tokens: int | None = None + top_p: float | None = None + top_k: int | None = None + response_mime_type: str | None = None + response_schema: dict | None = None + candidate_count: int | None = None + stop_sequences: list | None = None def __init__( self, - temperature: Optional[float] = None, - max_output_tokens: Optional[int] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - response_mime_type: Optional[str] = None, - response_schema: Optional[dict] = None, - candidate_count: Optional[int] = None, - stop_sequences: Optional[list] = None, + temperature: float | None = None, + max_output_tokens: int | None = None, + top_p: float | None = None, + top_k: int | None = None, + response_mime_type: str | None = None, + response_schema: dict | None = None, + candidate_count: int | None = None, + stop_sequences: list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -74,7 +73,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): def is_model_gemini_audio_model(self, model: str) -> bool: return "tts" in model - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: supported_params = [ "temperature", "top_p", @@ -105,10 +104,10 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): def _transform_messages( self, - messages: List[AllMessageValues], - model: Optional[str] = None, - litellm_params: Optional[dict] = None, - ) -> List[ContentType]: + messages: list[AllMessageValues], + model: str | None = None, + litellm_params: dict | None = None, + ) -> list[ContentType]: """ Google AI Studio Gemini does not support HTTP/HTTPS URLs for files. Convert them to base64 data instead. @@ -116,13 +115,13 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): for message in messages: _message_content = message.get("content") if _message_content is not None and isinstance(_message_content, list): - _parts: List[PartType] = [] + _parts: list[PartType] = [] for element in _message_content: if element.get("type") == "image_url": img_element = element - _image_url: Optional[str] = None - format: Optional[str] = None - detail: Optional[str] = None + _image_url: str | None = None + format: str | None = None + detail: str | None = None if isinstance(img_element.get("image_url"), dict): _image_url = img_element["image_url"].get("url") # type: ignore format = img_element["image_url"].get("format") # type: ignore diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index 3367c65f8d5..15c9915f85e 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -3,7 +3,7 @@ import datetime import json import math from collections.abc import Sequence -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -15,7 +15,7 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import TokenCountResponse -GEMINI_IMAGE_ASPECT_RATIOS: Dict[str, float] = { +GEMINI_IMAGE_ASPECT_RATIOS: dict[str, float] = { "1:1": 1 / 1, "1:4": 1 / 4, "1:8": 1 / 8, @@ -34,7 +34,7 @@ GEMINI_IMAGE_ASPECT_RATIOS: Dict[str, float] = { # Supported aspect ratio dimensions from Google Gemini image generation docs: # https://ai.google.dev/gemini-api/docs/image-generation#aspect_ratios_and_image_size -GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: Dict[tuple[int, int], str] = { +GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: dict[tuple[int, int], str] = { (512, 512): "1:1", (1024, 1024): "1:1", (2048, 2048): "1:1", @@ -96,7 +96,7 @@ GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: Dict[tuple[int, int], str] = { } -def map_openai_size_to_gemini_image_config(size: str, model: str) -> Optional[Dict[str, str]]: +def map_openai_size_to_gemini_image_config(size: str, model: str) -> dict[str, str] | None: dimensions = _parse_openai_image_size(size) if dimensions is None: return None @@ -129,16 +129,16 @@ def is_gemini_image_model(model: str) -> bool: def map_openai_image_params_to_gemini( - params: Dict[str, Any], + params: dict[str, Any], model: str, supported_params: Sequence[str], - optional_params: Optional[Dict[str, Any]] = None, + optional_params: dict[str, Any] | None = None, parse_image_config_string: bool = False, -) -> Dict[str, Any]: +) -> dict[str, Any]: optional_params = optional_params or {} filtered_params = {key: value for key, value in params.items() if key in supported_params} - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} if "n" in filtered_params and "n" not in optional_params: mapped_params["sampleCount"] = filtered_params["n"] @@ -175,14 +175,14 @@ def map_openai_image_params_to_gemini( return mapped_params -def _dedupe_gemini_search_tools(tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]: +def _dedupe_gemini_search_tools(tools: list[dict[str, Any]]) -> list[dict[str, Any]]: from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) search_tool_keys = VertexGeminiConfig._search_tool_keys() seen_search_keys: set[str] = set() - deduped_tools: List[Dict[str, Any]] = [] + deduped_tools: list[dict[str, Any]] = [] for tool in tools: if not isinstance(tool, dict): @@ -203,7 +203,7 @@ def _dedupe_gemini_search_tools(tools: List[Dict[str, Any]]) -> List[Dict[str, A return deduped_tools -def _has_gemini_search_tool(tools: List[Any]) -> bool: +def _has_gemini_search_tool(tools: list[Any]) -> bool: from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) @@ -213,9 +213,9 @@ def _has_gemini_search_tool(tools: List[Any]) -> bool: def map_gemini_image_tools_params( - non_default_params: Dict[str, Any], - mapped_params: Dict[str, Any], -) -> Dict[str, Any]: + non_default_params: dict[str, Any], + mapped_params: dict[str, Any], +) -> dict[str, Any]: from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) @@ -246,13 +246,13 @@ def map_gemini_image_tools_params( def get_gemini_image_web_search_requests( - response_data: Dict[str, Any], -) -> Optional[int]: + response_data: dict[str, Any], +) -> int | None: from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) - grounding_metadata: List[Dict[str, Any]] = [] + grounding_metadata: list[dict[str, Any]] = [] for candidate in response_data.get("candidates", []): if not isinstance(candidate, dict): continue @@ -267,11 +267,11 @@ def get_gemini_image_web_search_requests( def get_gemini_image_generation_config( model: str, - optional_params: Dict[str, Any], -) -> Dict[str, Any]: - generation_config: Dict[str, Any] = {"response_modalities": ["IMAGE", "TEXT"]} + optional_params: dict[str, Any], +) -> dict[str, Any]: + generation_config: dict[str, Any] = {"response_modalities": ["IMAGE", "TEXT"]} - image_config: Dict[str, Any] = {} + image_config: dict[str, Any] = {} if isinstance(optional_params.get("imageConfig"), dict): image_config.update(optional_params["imageConfig"]) @@ -295,7 +295,7 @@ def get_gemini_image_generation_config( return generation_config -def _parse_openai_image_size(size: str) -> Optional[tuple[int, int]]: +def _parse_openai_image_size(size: str) -> tuple[int, int] | None: if size == "auto": return None @@ -346,11 +346,11 @@ class GeminiModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """Google AI Studio sends api key via x-goog-api-key header""" return headers @@ -360,18 +360,18 @@ class GeminiModelInfo(BaseLLMModelInfo): return "v1beta" @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or get_secret_str("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com" @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or (get_secret_str("GOOGLE_API_KEY")) or (get_secret_str("GEMINI_API_KEY")) @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: return model.replace("gemini/", "") - def process_model_name(self, models: List[Dict[str, str]]) -> List[str]: + def process_model_name(self, models: list[dict[str, str]]) -> list[str]: litellm_model_names = [] for model in models: stripped_model_name = model["name"].replace("models/", "") @@ -379,7 +379,7 @@ class GeminiModelInfo(BaseLLMModelInfo): litellm_model_names.append(litellm_model_name) return litellm_model_names - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base = GeminiModelInfo.get_api_base(api_base) api_key = GeminiModelInfo.get_api_key(api_key) endpoint = f"/{self.api_version}/models" @@ -403,12 +403,10 @@ class GeminiModelInfo(BaseLLMModelInfo): litellm_model_names = self.process_model_name(models) return litellm_model_names - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return GeminiError(status_code=status_code, message=error_message, headers=headers) - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create a token counter for this provider. @@ -419,7 +417,7 @@ class GeminiModelInfo(BaseLLMModelInfo): return GoogleAIStudioTokenCounter() -def encode_unserializable_types(data: Dict[str, object], depth: int = 0) -> Dict[str, object]: +def encode_unserializable_types(data: dict[str, object], depth: int = 0) -> dict[str, object]: """Converts unserializable types in dict to json.dumps() compatible types. This function is called in models.py after calling convert_to_dict(). The @@ -457,7 +455,7 @@ def encode_unserializable_types(data: Dict[str, object], depth: int = 0) -> Dict return processed_data -def get_api_key_from_env() -> Optional[str]: +def get_api_key_from_env() -> str | None: return get_secret_str("GOOGLE_API_KEY") or get_secret_str("GEMINI_API_KEY") @@ -466,7 +464,7 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: from litellm.types.utils import LlmProviders @@ -475,13 +473,13 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter): async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: import copy from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter diff --git a/litellm/llms/gemini/cost_calculator.py b/litellm/llms/gemini/cost_calculator.py index f69cfe03270..3dc8976e6b9 100644 --- a/litellm/llms/gemini/cost_calculator.py +++ b/litellm/llms/gemini/cost_calculator.py @@ -4,13 +4,13 @@ This file is used to calculate the cost of the Gemini API. Handles the context caching for Gemini API. """ -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING if TYPE_CHECKING: from litellm.types.utils import ModelInfo, Usage -def cost_per_token(model: str, usage: "Usage", service_tier: Optional[str] = None) -> Tuple[float, float]: +def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index 27df584d476..ed82a37e47b 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -47,7 +47,7 @@ class GoogleAIStudioTokenCounter: return cleaned_contents - def _construct_url(self, model: str, api_base: Optional[str] = None) -> str: + def _construct_url(self, model: str, api_base: str | None = None) -> str: """ Construct the URL for the Google Gen AI Studio countTokens endpoint. """ @@ -56,12 +56,12 @@ class GoogleAIStudioTokenCounter: async def validate_environment( self, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - headers: Optional[Dict[str, Any]] = None, + api_base: str | None = None, + api_key: str | None = None, + headers: dict[str, Any] | None = None, model: str = "", - litellm_params: Optional[Dict[str, Any]] = None, - ) -> Tuple[Dict[str, Any], str]: + litellm_params: dict[str, Any] | None = None, + ) -> tuple[dict[str, Any], str]: """ Returns a Tuple of headers and url for the Google Gen AI Studio countTokens endpoint. """ @@ -81,11 +81,11 @@ class GoogleAIStudioTokenCounter: self, contents: Any, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Count tokens using Google Gen AI Studio countTokens endpoint. @@ -155,8 +155,8 @@ class GoogleAIStudioTokenCounter: status_code=e.response.status_code, ) from e except httpx.RequestError as e: - error_msg = f"Request to Google Gen AI Studio failed: {str(e)}" + error_msg = f"Request to Google Gen AI Studio failed: {e!s}" raise litellm.APIConnectionError(message=error_msg, llm_provider="gemini", model=model) from e except Exception as e: - error_msg = f"Unexpected error during token counting: {str(e)}" + error_msg = f"Unexpected error during token counting: {e!s}" raise Exception(error_msg) from e diff --git a/litellm/llms/gemini/files/transformation.py b/litellm/llms/gemini/files/transformation.py index a18dc152cb6..f91737ae613 100644 --- a/litellm/llms/gemini/files/transformation.py +++ b/litellm/llms/gemini/files/transformation.py @@ -5,15 +5,15 @@ For vertex ai, check out the vertex_ai/files/handler.py file. """ import time -from typing import Any, List, Literal, Optional +from typing import Any, Literal from urllib.parse import urlparse import httpx from openai.types.file_deleted import FileDeleted from litellm._logging import verbose_logger -from litellm.litellm_core_utils.url_utils import encode_url_path_segment 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 from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, LiteLLMLoggingObj, @@ -43,11 +43,11 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): self, headers: dict[Any, Any], model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict[Any, Any], litellm_params: dict[Any, Any], - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict[Any, Any]: """ Validate environment and add Gemini API key to headers. @@ -62,12 +62,12 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -86,10 +86,10 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): if not final_api_key: raise ValueError("api_key is required") - url = "{}/{}".format(api_base, endpoint) + url = f"{api_base}/{endpoint}" return url - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: return [] def map_openai_params( @@ -155,7 +155,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -190,8 +190,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): status_details=None, ) except Exception as e: - verbose_logger.exception(f"Error parsing file upload response: {str(e)}") - raise ValueError(f"Error parsing file upload response: {str(e)}") + verbose_logger.exception(f"Error parsing file upload response: {e!s}") + raise ValueError(f"Error parsing file upload response: {e!s}") def transform_retrieve_file_request( self, @@ -294,8 +294,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): status_details=(str(response_json.get("error", "")) if gemini_state == "FAILED" else None), ) except Exception as e: - verbose_logger.exception(f"Error parsing file retrieve response: {str(e)}") - raise ValueError(f"Error parsing file retrieve response: {str(e)}") + verbose_logger.exception(f"Error parsing file retrieve response: {e!s}") + raise ValueError(f"Error parsing file retrieve response: {e!s}") def transform_delete_file_request( self, @@ -362,12 +362,12 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): else: raise ValueError(f"Failed to delete file: {raw_response.text}") except Exception as e: - verbose_logger.exception(f"Error parsing file delete response: {str(e)}") - raise ValueError(f"Error parsing file delete response: {str(e)}") + verbose_logger.exception(f"Error parsing file delete response: {e!s}") + raise ValueError(f"Error parsing file delete response: {e!s}") def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: @@ -378,7 +378,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: raise NotImplementedError("GoogleAIStudioFilesHandler does not support file listing") def transform_file_content_request( diff --git a/litellm/llms/gemini/google_genai/transformation.py b/litellm/llms/gemini/google_genai/transformation.py index 68f30308621..528b69f01a5 100644 --- a/litellm/llms/gemini/google_genai/transformation.py +++ b/litellm/llms/gemini/google_genai/transformation.py @@ -3,7 +3,7 @@ Transformation for Calling Google models in their native format. """ from copy import deepcopy -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, cast import httpx @@ -54,7 +54,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): super().__init__() VertexLLM.__init__(self) - def get_supported_generate_content_optional_params(self, model: str) -> List[str]: + def get_supported_generate_content_optional_params(self, model: str) -> list[str]: """ Get the list of supported Google GenAI parameters for the model. @@ -101,7 +101,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): self, generate_content_config_dict: GenerateContentConfigDict, model: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map Google GenAI parameters to provider-specific format. @@ -117,7 +117,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): _snake_to_camel, ) - _generate_content_config_dict: Dict[str, Any] = {} + _generate_content_config_dict: dict[str, Any] = {} supported_google_genai_params = self.get_supported_generate_content_optional_params(model) # Create a set with both camelCase and snake_case versions for faster lookup supported_params_set = set(supported_google_genai_params) @@ -145,10 +145,10 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): def validate_environment( self, - api_key: Optional[str], - headers: Optional[dict], + api_key: str | None, + headers: dict | None, model: str, - litellm_params: Optional[Union[GenericLiteLLMParams, dict]], + litellm_params: GenericLiteLLMParams | dict | None, ) -> dict: default_headers = { "Content-Type": "application/json", @@ -164,7 +164,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): return default_headers - def _get_google_ai_studio_api_key(self, litellm_params: dict) -> Optional[str]: + def _get_google_ai_studio_api_key(self, litellm_params: dict) -> str | None: return ( litellm_params.pop("api_key", None) or litellm_params.pop("gemini_api_key", None) @@ -175,7 +175,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): def _get_common_auth_components( self, litellm_params: dict, - ) -> Tuple[Any, Optional[str], Optional[str]]: + ) -> tuple[Any, str | None, str | None]: """ Get common authentication components used by both sync and async methods. @@ -190,14 +190,14 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): def _build_final_headers_and_url( self, model: str, - auth_header: Optional[str], - vertex_project: Optional[str], - vertex_location: Optional[str], + auth_header: str | None, + vertex_project: str | None, + vertex_location: str | None, vertex_credentials: Any, stream: bool, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Build final headers and API URL from auth components. """ @@ -227,11 +227,11 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): def sync_get_auth_token_and_url( self, - api_base: Optional[str], + api_base: str | None, model: str, litellm_params: dict, stream: bool, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Sync version of get_auth_token_and_url. """ @@ -260,11 +260,11 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): async def get_auth_token_and_url( self, - api_base: Optional[str], + api_base: str | None, model: str, litellm_params: dict, stream: bool, - ) -> Tuple[dict, str]: + ) -> tuple[dict, str]: """ Get the complete URL for the request. @@ -300,7 +300,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): ) @staticmethod - def _normalize_response_schema(generate_content_config_dict: Dict, model: str) -> None: + def _normalize_response_schema(generate_content_config_dict: dict, model: str) -> None: schema_key = next( (k for k in ("responseSchema", "response_schema") if k in generate_content_config_dict), None, @@ -335,9 +335,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): self, model: str, contents: GenerateContentContentListUnionDict, - tools: Optional[ToolConfigDict], - generate_content_config_dict: Dict, - system_instruction: Optional[Any] = None, + tools: ToolConfigDict | None, + generate_content_config_dict: dict, + system_instruction: Any | None = None, ) -> dict: from litellm.types.google_genai.main import ( GenerateContentConfigDict, @@ -391,7 +391,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): return GenerateContentResponse(**response) - def convert_citation_sources_to_citations(self, response: Dict) -> Dict: + def convert_citation_sources_to_citations(self, response: dict) -> dict: """ Convert citation sources to citations. API's camelCase citationSources becomes the SDK's snake_case citations diff --git a/litellm/llms/gemini/image_edit/__init__.py b/litellm/llms/gemini/image_edit/__init__.py index cb097d3eee6..7db93dd8914 100644 --- a/litellm/llms/gemini/image_edit/__init__.py +++ b/litellm/llms/gemini/image_edit/__init__.py @@ -1,9 +1,9 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig -from .transformation import GeminiImageEditConfig from .cost_calculator import cost_calculator +from .transformation import GeminiImageEditConfig -__all__ = ["GeminiImageEditConfig", "get_gemini_image_edit_config", "cost_calculator"] +__all__ = ["GeminiImageEditConfig", "cost_calculator", "get_gemini_image_edit_config"] def get_gemini_image_edit_config(model: str) -> BaseImageEditConfig: diff --git a/litellm/llms/gemini/image_edit/transformation.py b/litellm/llms/gemini/image_edit/transformation.py index 78d682395bb..118842c23cd 100644 --- a/litellm/llms/gemini/image_edit/transformation.py +++ b/litellm/llms/gemini/image_edit/transformation.py @@ -1,6 +1,6 @@ import base64 from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -34,9 +34,9 @@ else: class GeminiImageEditConfig(BaseImageEditConfig): DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta" - SUPPORTED_PARAMS: List[str] = ["n", "size", "imageConfig"] + SUPPORTED_PARAMS: list[str] = ["n", "size", "imageConfig"] - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return list(self.SUPPORTED_PARAMS) def map_openai_params( @@ -44,7 +44,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: return map_openai_image_params_to_gemini( params=image_edit_optional_params, # type: ignore[arg-type] model=model, @@ -56,11 +56,11 @@ class GeminiImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("GEMINI_API_KEY") + final_api_key: str | None = api_key or get_secret_str("GEMINI_API_KEY") if not final_api_key: raise ValueError("GEMINI_API_KEY is not set") @@ -75,7 +75,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: base_url = api_base or get_secret_str("GEMINI_API_BASE") or self.DEFAULT_BASE_URL @@ -85,12 +85,12 @@ class GeminiImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( # type: ignore[override] self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict[str, Any], + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict[str, Any], Optional[RequestFiles]]: + ) -> tuple[dict[str, Any], RequestFiles | None]: inline_parts = self._prepare_inline_image_parts(image) if image else [] if not inline_parts: raise ValueError("Gemini image edit requires at least one image.") @@ -106,7 +106,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): } ] - request_body: Dict[str, Any] = {"contents": contents} + request_body: dict[str, Any] = {"contents": contents} request_body["generationConfig"] = get_gemini_image_generation_config( model=model, @@ -133,7 +133,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): ) candidates = response_json.get("candidates", []) - data_list: List[ImageObject] = [] + data_list: list[ImageObject] = [] for candidate in candidates: content = candidate.get("content", {}) @@ -148,19 +148,19 @@ class GeminiImageEditConfig(BaseImageEditConfig): ) ) - model_response.data = cast(List[OpenAIImage], data_list) + model_response.data = cast(list[OpenAIImage], data_list) if "usageMetadata" in response_json: model_response.usage = transform_gemini_image_usage(response_json["usageMetadata"]) return model_response - def _prepare_inline_image_parts(self, image: Union[FileTypes, List[FileTypes]]) -> List[Dict[str, Any]]: - images: List[FileTypes] + def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, Any]]: + images: list[FileTypes] if isinstance(image, list): images = image else: images = [image] - inline_parts: List[Dict[str, Any]] = [] + inline_parts: list[dict[str, Any]] = [] for img in images: if img is None: continue diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py index dcdec46edca..d2d3a84b492 100644 --- a/litellm/llms/gemini/image_generation/transformation.py +++ b/litellm/llms/gemini/image_generation/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -34,7 +34,7 @@ else: class GoogleImageGenConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Google AI Imagen API supported parameters https://ai.google.dev/gemini-api/docs/imagen @@ -63,12 +63,12 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -93,13 +93,13 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("GEMINI_API_KEY") + final_api_key: str | None = api_key or get_secret_str("GEMINI_API_KEY") if not final_api_key: raise ValueError("GEMINI_API_KEY is not set") @@ -172,8 +172,8 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Google AI Imagen response to litellm ImageResponse format diff --git a/litellm/llms/gemini/interactions/transformation.py b/litellm/llms/gemini/interactions/transformation.py index 7443720f496..c8d14564e22 100644 --- a/litellm/llms/gemini/interactions/transformation.py +++ b/litellm/llms/gemini/interactions/transformation.py @@ -12,12 +12,11 @@ Schema versioning: litellm.use_legacy_interactions_schema = True. Remove flag after June 8, 2026. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx import litellm - from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.url_utils import encode_url_path_segment @@ -57,7 +56,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): def api_version(self) -> str: return "v1beta" - def get_supported_params(self, model: str) -> List[str]: + def get_supported_params(self, model: str) -> list[str]: """Per OpenAPI spec CreateModelInteractionParams.""" return [ "model", @@ -80,7 +79,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: """Google AI Studio uses x-goog-api-key header for authentication.""" headers = headers or {} @@ -102,11 +101,11 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): def get_complete_url( self, - api_base: Optional[str], - model: Optional[str], - agent: Optional[str] = None, - litellm_params: Optional[dict] = None, - stream: Optional[bool] = None, + api_base: str | None, + model: str | None, + agent: str | None = None, + litellm_params: dict | None = None, + stream: bool | None = None, ) -> str: """POST /{api_version}/interactions""" litellm_params = litellm_params or {} @@ -123,13 +122,13 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): def transform_request( self, - model: Optional[str], - agent: Optional[str], - input: Optional[InteractionInput], + model: str | None, + agent: str | None, + input: InteractionInput | None, optional_params: InteractionsAPIOptionalRequestParams, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Build request body per OpenAPI spec. @@ -144,7 +143,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): """ use_legacy: bool = litellm.use_legacy_interactions_schema - request_body: Dict[str, Any] = {} + request_body: dict[str, Any] = {} # Model or Agent (one required) if model: @@ -190,7 +189,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): and (not isinstance(response_format, dict) or "mime_type" not in response_format) ): # Wrap the legacy schema into the new polymorphic format. - new_rf: Dict[str, Any] = { + new_rf: dict[str, Any] = { "type": "text", "mime_type": response_mime_type, } @@ -202,7 +201,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): request_body["response_format"] = response_format # image_config moves out of generation_config into response_format. - generation_config: Optional[Dict[str, Any]] = optional_params.get("generation_config") + generation_config: dict[str, Any] | None = optional_params.get("generation_config") if generation_config is not None: image_config = None if isinstance(generation_config, dict): @@ -216,7 +215,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): if image_config is not None: # Move image_config to response_format with type=image. - image_rf: Dict[str, Any] = {"type": "image", **image_config} + image_rf: dict[str, Any] = {"type": "image", **image_config} existing_rf = request_body.get("response_format") if existing_rf is None: request_body["response_format"] = image_rf @@ -230,7 +229,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): def transform_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, ) -> InteractionsAPIResponse: @@ -258,7 +257,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): def transform_streaming_response( self, - model: Optional[str], + model: str | None, parsed_chunk: dict, logging_obj: LiteLLMLoggingObj, ) -> InteractionsAPIStreamingResponse: @@ -274,7 +273,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """GET /{api_version}/interactions/{interaction_id}""" resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): @@ -308,7 +307,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """DELETE /{api_version}/interactions/{interaction_id}""" resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): @@ -339,7 +338,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """POST /{api_version}/interactions/{interaction_id}:cancel (if supported)""" resolved_api_base = GeminiModelInfo.get_api_base(api_base) if not GeminiModelInfo.get_api_key(litellm_params.api_key): diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index bc2145fd832..6631c9d9ec7 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -4,7 +4,7 @@ This file contains the transformation logic for the Gemini realtime API. import json from collections import OrderedDict -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast import litellm from litellm import verbose_logger @@ -60,7 +60,7 @@ from litellm.utils import get_empty_usage from ..common_utils import encode_unserializable_types, get_api_key_from_env -MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]] = { +MAP_GEMINI_FIELD_TO_OPENAI_EVENT: dict[str, OpenAIRealtimeEventTypes | ResponsesAPIStreamEvents] = { "setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED, "serverContent.generationComplete": OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE, "serverContent.turnComplete": OpenAIRealtimeEventTypes.RESPONSE_DONE, @@ -77,10 +77,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def __init__(self): super().__init__() - self._tool_call_id_to_name: "OrderedDict[str, str]" = OrderedDict() + self._tool_call_id_to_name: OrderedDict[str, str] = OrderedDict() # Gemini Live sometimes emits usageMetadata in a standalone frame between # turns; buffer it here so the next response.done carries the token counts. - self._pending_usage_metadata: Optional[dict] = None + self._pending_usage_metadata: dict | None = None def is_setup_message(self, msg_obj: dict) -> bool: return "setup" in msg_obj @@ -93,7 +93,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return True @staticmethod - def _usage_detail_alias(details: Any, defaults: Dict[str, int]) -> Dict[str, Any]: + def _usage_detail_alias(details: Any, defaults: dict[str, int]) -> dict[str, Any]: if not isinstance(details, dict): return dict(defaults) return { @@ -102,7 +102,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): } @staticmethod - def _add_pipecat_usage_detail_aliases(usage_dict: Dict[str, Any]) -> Dict[str, Any]: + def _add_pipecat_usage_detail_aliases(usage_dict: dict[str, Any]) -> dict[str, Any]: usage_dict.setdefault( "input_token_details", GeminiRealtimeConfig._usage_detail_alias( @@ -119,10 +119,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) return usage_dict - def validate_environment(self, headers: dict, model: str, api_key: Optional[str] = None) -> dict: + def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict: return headers - def get_complete_url(self, api_base: Optional[str], model: str, api_key: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_key: str | None = None) -> str: """ Example output: "BACKEND_WS_URL = "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent""; @@ -164,7 +164,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): raise ValueError(f"Unexpected part type: {part}") raise ValueError(f"Unexpected model turn event, no 'parts' key: {model_turn}") - def map_generation_complete_event(self, delta_type: Optional[ALL_DELTA_TYPES]) -> OpenAIRealtimeEventTypes: + def map_generation_complete_event(self, delta_type: ALL_DELTA_TYPES | None) -> OpenAIRealtimeEventTypes: if delta_type == "text": return OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE elif delta_type == "audio": @@ -181,7 +181,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return mime_types.get(input_audio_format, "application/octet-stream") - def _manual_turn_detection_enabled(self, session_configuration_request: Optional[str]) -> bool: + def _manual_turn_detection_enabled(self, session_configuration_request: str | None) -> bool: if not session_configuration_request: return False try: @@ -191,7 +191,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): except (json.JSONDecodeError, TypeError, AttributeError): return False - def _handle_input_audio_buffer_commit_or_end(self, session_configuration_request: Optional[str]) -> List[str]: + def _handle_input_audio_buffer_commit_or_end(self, session_configuration_request: str | None) -> list[str]: """Map OpenAI buffer commit/end to Gemini Live turn-boundary signals.""" if self._manual_turn_detection_enabled(session_configuration_request): realtime_input_dict: BidiGenerateContentRealtimeInput = { @@ -227,7 +227,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): automatic_activity_dection["silenceDurationMs"] = value["silence_duration_ms"] return automatic_activity_dection - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "instructions", "temperature", @@ -251,7 +251,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): optional_params["generationConfig"]["maxOutputTokens"] = value elif key == "modalities": optional_params["generationConfig"]["responseModalities"] = [ - modality.upper() for modality in cast(List[str], value) + modality.upper() for modality in cast(list[str], value) ] elif key == "tools": from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( @@ -295,7 +295,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return optional_params @staticmethod - def _extract_turn_detection(session: dict) -> Optional[dict]: + def _extract_turn_detection(session: dict) -> dict | None: """Extract turn_detection from a session.update payload. Handles both the flat beta shape (``session.turn_detection``) and the @@ -383,7 +383,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return without_text if without_text else ["AUDIO"] @staticmethod - def _finalize_gemini_live_setup(model: str, setup: Dict[str, Any]) -> Dict[str, Any]: + def _finalize_gemini_live_setup(model: str, setup: dict[str, Any]) -> dict[str, Any]: """Drop fields Gemini Live native-audio rejects on ``setup``.""" generation_config = setup.get("generationConfig") if isinstance(generation_config, dict): @@ -400,8 +400,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): self, json_message: dict, model: str, - session_configuration_request: Optional[str], - ) -> List[str]: + session_configuration_request: str | None, + ) -> list[str]: """ Handle session.update by sending setup to Gemini. @@ -453,7 +453,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): verbose_logger.debug("Gemini Realtime: Ignoring session.update (setup already sent)") return [] - def _handle_conversation_item(self, json_message: dict) -> List[str]: + def _handle_conversation_item(self, json_message: dict) -> list[str]: """ Handle conversation.item.create for user text or function call output. @@ -467,7 +467,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return self._handle_function_call_output(item) return self._handle_user_text_content(item) - def _handle_function_call_output(self, item: dict) -> List[str]: + def _handle_function_call_output(self, item: dict) -> list[str]: """Transform function_call_output to Gemini toolResponse format.""" call_id = item.get("call_id", "") output = item.get("output", "{}") @@ -501,7 +501,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return [json.dumps(tool_response_message)] - def _handle_user_text_content(self, item: dict) -> List[str]: + def _handle_user_text_content(self, item: dict) -> list[str]: """Transform user text content to Gemini clientContent format.""" content_list = item.get("content", []) text_parts = [c.get("text", "") for c in content_list if isinstance(c, dict) and c.get("type") == "input_text"] @@ -522,8 +522,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): self, message: str, model: str, - session_configuration_request: Optional[str] = None, - ) -> List[str]: + session_configuration_request: str | None = None, + ) -> list[str]: realtime_input_dict: BidiGenerateContentRealtimeInput = {} try: json_message = json.loads(message) @@ -534,7 +534,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): message_str = str(message) raise ValueError(f"Invalid JSON message: {message_str}") - messages: List[str] = [] + messages: list[str] = [] msg_type = json_message.get("type") if msg_type == "session.update": @@ -553,7 +553,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): realtime_input_dict = cast( BidiGenerateContentRealtimeInput, - encode_unserializable_types(cast(Dict[str, object], realtime_input_dict)), + encode_unserializable_types(cast(dict[str, object], realtime_input_dict)), ) gemini_msg = json.dumps({"realtimeInput": realtime_input_dict}) @@ -573,7 +573,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): self, model: str, logging_session_id: str, - session_configuration_request: Optional[str] = None, + session_configuration_request: str | None = None, ) -> OpenAIRealtimeStreamSessionEvents: if session_configuration_request: session_configuration_request_dict: BidiGenerateContentSetup = json.loads( @@ -585,7 +585,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): _model = session_configuration_request_dict.get("model") or model generation_config = session_configuration_request_dict.get("generationConfig", {}) or {} gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) - _modalities = [modality.lower() for modality in cast(List[str], gemini_modalities)] + _modalities = [modality.lower() for modality in cast(list[str], gemini_modalities)] _system_instruction = session_configuration_request_dict.get("systemInstruction") session = OpenAIRealtimeStreamSession( id=logging_session_id, @@ -610,7 +610,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def _is_new_content_delta( self, - previous_messages: Optional[List[OpenAIRealtimeEvents]] = None, + previous_messages: list[OpenAIRealtimeEvents] | None = None, ) -> bool: if previous_messages is None or len(previous_messages) == 0: return True @@ -624,8 +624,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): output_item_id: str, conversation_id: str, delta_type: ALL_DELTA_TYPES, - session_configuration_request: Optional[str] = None, - ) -> List[OpenAIRealtimeEvents]: + session_configuration_request: str | None = None, + ) -> list[OpenAIRealtimeEvents]: session_configuration_request_dict: BidiGenerateContentSetup = {} if session_configuration_request is not None: try: @@ -634,15 +634,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): session_configuration_request_dict = {} generation_config = session_configuration_request_dict.get("generationConfig", {}) gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) - _modalities = [modality.lower() for modality in cast(List[str], gemini_modalities)] + _modalities = [modality.lower() for modality in cast(list[str], gemini_modalities)] _temperature = generation_config.get("temperature") _max_output_tokens = generation_config.get("maxOutputTokens") - response_items: List[OpenAIRealtimeEvents] = [] + response_items: list[OpenAIRealtimeEvents] = [] response_created = OpenAIRealtimeStreamResponseBaseObject( type="response.created", - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", response={ "object": "realtime.response", "id": response_id, @@ -660,7 +660,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ## - return response.output_item.added response_output_item_added = OpenAIRealtimeStreamResponseOutputItemAdded( type="response.output_item.added", - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", response_id=response_id, output_index=0, item={ @@ -682,7 +682,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): OpenAIRealtimeEvents, { "type": "conversation.item.added", - "event_id": "event_{}".format(uuid.uuid4()), + "event_id": f"event_{uuid.uuid4()}", "previous_item_id": None, "item": { "id": output_item_id, @@ -700,7 +700,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): type="response.content_part.added", content_index=0, output_index=0, - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", item_id=output_item_id, part=( { @@ -739,7 +739,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return OpenAIRealtimeResponseDelta( type=("response.output_text.delta" if delta_type == "text" else "response.output_audio.delta"), content_index=0, - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", item_id=output_item_id, output_index=0, response_id=response_id, @@ -748,24 +748,24 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def transform_content_done_event( self, - delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], - current_output_item_id: Optional[str], - current_response_id: Optional[str], + delta_chunks: list[OpenAIRealtimeResponseDelta] | None, + current_output_item_id: str | None, + current_response_id: str | None, delta_type: ALL_DELTA_TYPES, - ) -> Union[OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone]: + ) -> OpenAIRealtimeResponseTextDone | OpenAIRealtimeResponseAudioDone: if delta_chunks: delta = "".join([delta_chunk["delta"] for delta_chunk in delta_chunks]) else: delta = "" if current_output_item_id is None: - current_output_item_id = "item_{}".format(uuid.uuid4()) + current_output_item_id = f"item_{uuid.uuid4()}" if current_response_id is None: - current_response_id = "resp_{}".format(uuid.uuid4()) + current_response_id = f"resp_{uuid.uuid4()}" if delta_type == "text": return OpenAIRealtimeResponseTextDone( type="response.output_text.done", content_index=0, - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", item_id=current_output_item_id, output_index=0, response_id=current_response_id, @@ -775,7 +775,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return OpenAIRealtimeResponseAudioDone( type="response.output_audio.done", content_index=0, - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", item_id=current_output_item_id, output_index=0, response_id=current_response_id, @@ -783,27 +783,27 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def return_additional_content_done_events( self, - current_output_item_id: Optional[str], - current_response_id: Optional[str], - delta_done_event: Union[OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone], + current_output_item_id: str | None, + current_response_id: str | None, + delta_done_event: OpenAIRealtimeResponseTextDone | OpenAIRealtimeResponseAudioDone, delta_type: ALL_DELTA_TYPES, - ) -> List[OpenAIRealtimeEvents]: + ) -> list[OpenAIRealtimeEvents]: """ - return response.content_part.done - return response.output_item.done """ if current_output_item_id is None: - current_output_item_id = "item_{}".format(uuid.uuid4()) + current_output_item_id = f"item_{uuid.uuid4()}" if current_response_id is None: - current_response_id = "resp_{}".format(uuid.uuid4()) - returned_items: List[OpenAIRealtimeEvents] = [] + current_response_id = f"resp_{uuid.uuid4()}" + returned_items: list[OpenAIRealtimeEvents] = [] - delta_done_event_text = cast(Optional[str], delta_done_event.get("text")) + delta_done_event_text = cast(str | None, delta_done_event.get("text")) # response.content_part.done response_content_part_done = OpenAIRealtimeContentPartDone( type="response.content_part.done", content_index=0, - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", item_id=current_output_item_id, output_index=0, part=( @@ -820,7 +820,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): # response.output_item.done response_output_item_done = OpenAIRealtimeOutputItemDone( type="response.output_item.done", - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", output_index=0, response_id=current_response_id, item={ @@ -844,7 +844,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): returned_items.append(response_output_item_done) return returned_items - def _consume_usage_metadata_for_response_done(self, frame: dict) -> Optional[dict]: + def _consume_usage_metadata_for_response_done(self, frame: dict) -> dict | None: """Pop usageMetadata from the frame (authoritative) or drain the pending buffer. Uses pop so a frame with both ``toolCall`` and ``turnComplete`` can't @@ -861,16 +861,16 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def transform_tool_call_events( self, tool_call_message: dict, - response_id: Optional[str] = None, - output_item_id: Optional[str] = None, - ) -> List[OpenAIRealtimeFunctionCallArgumentsDone]: + response_id: str | None = None, + output_item_id: str | None = None, + ) -> list[OpenAIRealtimeFunctionCallArgumentsDone]: function_calls = tool_call_message.get("functionCalls", []) resolved_response_id = response_id or f"resp_{uuid.uuid4()}" resolved_output_item_id = output_item_id or f"item_{uuid.uuid4()}" verbose_logger.debug(f"Gemini Realtime: Transforming {len(function_calls)} tool call(s) to OpenAI format") - events: List[OpenAIRealtimeFunctionCallArgumentsDone] = [] + events: list[OpenAIRealtimeFunctionCallArgumentsDone] = [] for idx, fc in enumerate(function_calls): call_id = fc.get("id", "") or f"call_{uuid.uuid4().hex[:16]}" name = fc.get("name", "") @@ -909,9 +909,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def update_current_delta_chunks( self, - transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]], - current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], - ) -> Optional[List[OpenAIRealtimeResponseDelta]]: + transformed_message: OpenAIRealtimeEvents | list[OpenAIRealtimeEvents], + current_delta_chunks: list[OpenAIRealtimeResponseDelta] | None, + ) -> list[OpenAIRealtimeResponseDelta] | None: try: if isinstance(transformed_message, list): current_delta_chunks = [] @@ -939,9 +939,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def update_current_item_chunks( self, - transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]], - current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]], - ) -> Optional[List[OpenAIRealtimeOutputItemDone]]: + transformed_message: OpenAIRealtimeEvents | list[OpenAIRealtimeEvents], + current_item_chunks: list[OpenAIRealtimeOutputItemDone] | None, + ) -> list[OpenAIRealtimeOutputItemDone] | None: try: if isinstance(transformed_message, list): current_item_chunks = [] @@ -966,15 +966,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def transform_response_done_event( self, message: BidiGenerateContentServerMessage, - current_response_id: Optional[str], - current_conversation_id: Optional[str], - output_items: Optional[List[OpenAIRealtimeOutputItemDone]], - session_configuration_request: Optional[str] = None, + current_response_id: str | None, + current_conversation_id: str | None, + output_items: list[OpenAIRealtimeOutputItemDone] | None, + session_configuration_request: str | None = None, ) -> OpenAIRealtimeDoneEvent: if current_conversation_id is None: - current_conversation_id = "conv_{}".format(uuid.uuid4()) + current_conversation_id = f"conv_{uuid.uuid4()}" if current_response_id is None: - current_response_id = "resp_{}".format(uuid.uuid4()) + current_response_id = f"resp_{uuid.uuid4()}" if session_configuration_request: session_configuration_request_dict: BidiGenerateContentSetup = json.loads( @@ -987,7 +987,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): temperature = generation_config.get("temperature") max_output_tokens = generation_config.get("maxOutputTokens") gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) - _modalities = [modality.lower() for modality in cast(List[str], gemini_modalities)] + _modalities = [modality.lower() for modality in cast(list[str], gemini_modalities)] resolved_usage_metadata = self._consume_usage_metadata_for_response_done(cast(dict, message)) if resolved_usage_metadata is not None: _chat_completion_usage = VertexGeminiConfig._calculate_usage( @@ -1006,7 +1006,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): self._add_pipecat_usage_detail_aliases(_usage_dict) response_done_event = OpenAIRealtimeDoneEvent( type="response.done", - event_id="event_{}".format(uuid.uuid4()), + event_id=f"event_{uuid.uuid4()}", response=OpenAIRealtimeResponseDoneObject( object="realtime.response", id=current_response_id, @@ -1038,16 +1038,16 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): current_delta_chunks = realtime_response_transform_input["current_delta_chunks"] session_configuration_request = realtime_response_transform_input["session_configuration_request"] - returned_message: List[OpenAIRealtimeEvents] = [] + returned_message: list[OpenAIRealtimeEvents] = [] if ( openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA ): - current_response_id = current_response_id or "resp_{}".format(uuid.uuid4()) + current_response_id = current_response_id or f"resp_{uuid.uuid4()}" if not current_output_item_id: # send the list of standard 'new' content.delta events - current_output_item_id = "item_{}".format(uuid.uuid4()) - current_conversation_id = current_conversation_id or "conv_{}".format(uuid.uuid4()) + current_output_item_id = f"item_{uuid.uuid4()}" + current_conversation_id = current_conversation_id or f"conv_{uuid.uuid4()}" returned_message = self.return_new_content_delta_events( session_configuration_request=session_configuration_request, response_id=current_response_id, @@ -1102,15 +1102,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): self, key: str, value: Any, - current_delta_type: Optional[ALL_DELTA_TYPES], - ) -> Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]: + current_delta_type: ALL_DELTA_TYPES | None, + ) -> OpenAIRealtimeEventTypes | ResponsesAPIStreamEvents: if isinstance(value, dict): model_turn_event = value.get("modelTurn") generation_complete_event = value.get("generationComplete") else: model_turn_event = None generation_complete_event = None - openai_event: Optional[Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]] = None + openai_event: OpenAIRealtimeEventTypes | ResponsesAPIStreamEvents | None = None if model_turn_event: # check if model turn event openai_event = self.map_model_turn_event(model_turn_event) elif generation_complete_event: @@ -1142,7 +1142,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def transform_realtime_response( self, - message: Union[str, bytes], + message: str | bytes, model: str, logging_obj: LiteLLMLoggingObj, realtime_response_transform_input: RealtimeResponseTransformInput, @@ -1172,8 +1172,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): current_delta_chunks = realtime_response_transform_input["current_delta_chunks"] session_configuration_request = realtime_response_transform_input["session_configuration_request"] current_item_chunks = realtime_response_transform_input["current_item_chunks"] - current_delta_type: Optional[ALL_DELTA_TYPES] = realtime_response_transform_input["current_delta_type"] - returned_message: List[OpenAIRealtimeEvents] = [] + current_delta_type: ALL_DELTA_TYPES | None = realtime_response_transform_input["current_delta_type"] + returned_message: list[OpenAIRealtimeEvents] = [] server_content = json_message.get("serverContent") if isinstance(server_content, dict): @@ -1184,9 +1184,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): OpenAIRealtimeEvents, { "type": "conversation.item.input_audio_transcription.completed", - "event_id": "event_{}".format(uuid.uuid4()), + "event_id": f"event_{uuid.uuid4()}", "transcript": input_tx["text"], - "item_id": "item_{}".format(uuid.uuid4()), + "item_id": f"item_{uuid.uuid4()}", "content_index": 0, }, ) @@ -1195,10 +1195,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): output_tx = server_content.get("outputTranscription") if isinstance(output_tx, dict) and output_tx.get("text"): if current_response_id is None: - current_response_id = "resp_{}".format(uuid.uuid4()) + current_response_id = f"resp_{uuid.uuid4()}" if current_output_item_id is None: - current_output_item_id = "item_{}".format(uuid.uuid4()) - current_conversation_id = current_conversation_id or "conv_{}".format(uuid.uuid4()) + current_output_item_id = f"item_{uuid.uuid4()}" + current_conversation_id = current_conversation_id or f"conv_{uuid.uuid4()}" returned_message.extend( self.return_new_content_delta_events( session_configuration_request=session_configuration_request, @@ -1213,7 +1213,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): OpenAIRealtimeEvents, { "type": "response.output_audio_transcript.delta", - "event_id": "event_{}".format(uuid.uuid4()), + "event_id": f"event_{uuid.uuid4()}", "transcript": output_tx["text"], "item_id": current_output_item_id, "content_index": 0, @@ -1282,7 +1282,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): tool_call_modalities = [ modality.lower() for modality in cast( - List[str], + list[str], tool_call_generation_config.get("responseModalities", ["AUDIO"]), ) ] @@ -1572,7 +1572,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ``` """ - response_modalities: List[GeminiResponseModalities] = ["AUDIO"] + response_modalities: list[GeminiResponseModalities] = ["AUDIO"] output_audio_transcription = False # if "audio" in model: ## UNCOMMENT THIS WHEN AUDIO IS SUPPORTED # output_audio_transcription = True diff --git a/litellm/llms/gemini/vector_stores/transformation.py b/litellm/llms/gemini/vector_stores/transformation.py index f98cb0e5b0c..051c0c544f5 100644 --- a/litellm/llms/gemini/vector_stores/transformation.py +++ b/litellm/llms/gemini/vector_stores/transformation.py @@ -5,7 +5,7 @@ Implements the transformation between LiteLLM's unified vector store API and Google Gemini's File Search API. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -43,7 +43,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): def __init__(self) -> None: super().__init__() self.model_info = GeminiModelInfo() - self._cached_api_key: Optional[str] = None + self._cached_api_key: str | None = None def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials: """Gemini uses x-goog-api-key header for authentication.""" @@ -61,11 +61,11 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): "write": [("POST", "/fileSearchStores")], } - def get_supported_openai_params(self, model: str) -> List[VECTOR_STORE_OPENAI_PARAMS]: + def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]: """Supported parameters for Gemini File Search.""" return ["max_num_results", "filters"] - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """Validate and set up headers for Gemini API.""" headers = headers or {} headers.setdefault("Content-Type", "application/json") @@ -77,7 +77,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): return headers - def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str: + def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str: """ Get the complete base URL for Gemini API. @@ -94,7 +94,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): api_version = "v1beta" return f"{api_base}/{api_version}" - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]) -> GeminiError: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> GeminiError: """Return Gemini-specific error class.""" return GeminiError( status_code=status_code, @@ -105,13 +105,13 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform search request to Gemini's generateContent format. @@ -133,7 +133,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): url = f"{api_base}/models/{model}:generateContent" # Build file_search tool configuration (using snake_case as per Gemini docs) - file_search_config: Dict[str, Any] = {"file_search_store_names": [vector_store_id]} + file_search_config: dict[str, Any] = {"file_search_store_names": [vector_store_id]} # Add metadata filter if provided metadata_filter = vector_store_search_optional_params.get("filters") @@ -152,7 +152,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): file_search_config["metadata_filter"] = metadata_filter # Build request body - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "contents": [{"parts": [{"text": query}]}], "tools": [{"file_search": file_search_config}], } @@ -179,7 +179,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): """ try: response_data = response.json() - results: List[VectorStoreSearchResult] = [] + results: list[VectorStoreSearchResult] = [] # Extract candidates and grounding metadata candidates = response_data.get("candidates", []) @@ -256,7 +256,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): except Exception as e: raise self.get_error_class( - error_message=f"Failed to parse Gemini response: {str(e)}", + error_message=f"Failed to parse Gemini response: {e!s}", status_code=response.status_code, headers=response.headers, ) @@ -265,7 +265,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform create request to Gemini's fileSearchStores format. """ @@ -273,7 +273,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): # API key is passed via x-goog-api-key header (set in validate_environment) - request_body: Dict[str, Any] = {} + request_body: dict[str, Any] = {} # Add display name if provided name = vector_store_create_optional_params.get("name") @@ -327,7 +327,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): except Exception as e: raise self.get_error_class( - error_message=f"Failed to parse Gemini create response: {str(e)}", + error_message=f"Failed to parse Gemini create response: {e!s}", status_code=response.status_code, headers=response.headers, ) diff --git a/litellm/llms/gemini/videos/transformation.py b/litellm/llms/gemini/videos/transformation.py index 4a9b3830ec5..fa30d663798 100644 --- a/litellm/llms/gemini/videos/transformation.py +++ b/litellm/llms/gemini/videos/transformation.py @@ -1,5 +1,5 @@ import base64 -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -34,7 +34,7 @@ else: BaseLLMException = Any -def _convert_image_to_gemini_format(image_file) -> Dict[str, str]: +def _convert_image_to_gemini_format(image_file) -> dict[str, str]: """ Convert image file to Gemini format with base64 encoding and MIME type. @@ -55,8 +55,8 @@ def _convert_image_to_gemini_format(image_file) -> Dict[str, str]: def _usage_video_resolution_from_parameters( - parameters: Dict[str, Any], -) -> Optional[str]: + parameters: dict[str, Any], +) -> str | None: """Normalize Veo ``parameters.resolution`` for usage and cost tracking.""" res = parameters.get("resolution") if res is None or res == "": @@ -75,7 +75,7 @@ class GeminiVideoConfig(BaseVideoConfig): 4. Download video using file API """ - _OPENAI_VIDEO_SIZE_TO_ASPECT_RATIO: Dict[str, str] = { + _OPENAI_VIDEO_SIZE_TO_ASPECT_RATIO: dict[str, str] = { "1280x720": "16:9", "1920x1080": "16:9", "720x1280": "9:16", @@ -97,7 +97,7 @@ class GeminiVideoConfig(BaseVideoConfig): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map OpenAI-style parameters to Veo format. @@ -111,7 +111,7 @@ class GeminiVideoConfig(BaseVideoConfig): All other params are passed through as-is to support Gemini-specific parameters. """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} # Get supported OpenAI params (exclude "model" and "prompt" which are handled separately) supported_openai_params = self.get_supported_openai_params(model) @@ -151,7 +151,7 @@ class GeminiVideoConfig(BaseVideoConfig): return mapped_params - def _convert_size_to_aspect_ratio(self, size: str) -> Optional[str]: + def _convert_size_to_aspect_ratio(self, size: str) -> str | None: """ Convert OpenAI size format to Veo aspectRatio format. @@ -164,7 +164,7 @@ class GeminiVideoConfig(BaseVideoConfig): return self._OPENAI_VIDEO_SIZE_TO_ASPECT_RATIO.get(size, "16:9") - def _convert_size_to_resolution(self, size: str) -> Optional[str]: + def _convert_size_to_resolution(self, size: str) -> str | None: """ Map OpenAI ``size`` (WxH) to Veo ``resolution`` for presets in ``_OPENAI_VIDEO_SIZE_TO_ASPECT_RATIO`` (720p / 1080p from the smaller edge). @@ -188,8 +188,8 @@ class GeminiVideoConfig(BaseVideoConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: """ Validate environment and add Gemini API key to headers. @@ -218,7 +218,7 @@ class GeminiVideoConfig(BaseVideoConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -242,10 +242,10 @@ class GeminiVideoConfig(BaseVideoConfig): model: str, prompt: str, api_base: str, - video_create_optional_request_params: Dict, + video_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles, str]: + ) -> tuple[dict, RequestFiles, str]: """ Transform the video creation request for Veo API. @@ -293,8 +293,8 @@ class GeminiVideoConfig(BaseVideoConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: """ Transform the Veo video creation response. @@ -336,7 +336,7 @@ class GeminiVideoConfig(BaseVideoConfig): model=model, ) - usage_data: Dict[str, Any] = {} + usage_data: dict[str, Any] = {} if request_data: parameters = request_data.get("parameters", {}) duration = parameters.get("durationSeconds") or DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS @@ -358,7 +358,7 @@ class GeminiVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video status retrieve request for Veo API. @@ -367,7 +367,7 @@ class GeminiVideoConfig(BaseVideoConfig): """ operation_name = extract_original_video_id(video_id) url = f"{api_base.rstrip('/')}/v1beta/{operation_name}" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} return url, params @@ -375,7 +375,7 @@ class GeminiVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """ Transform the Veo operation status response. @@ -428,8 +428,8 @@ class GeminiVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - variant: Optional[str] = None, - ) -> Tuple[str, Dict]: + variant: str | None = None, + ) -> tuple[str, dict]: """ Transform the video content request for Veo API. @@ -458,7 +458,7 @@ class GeminiVideoConfig(BaseVideoConfig): generated_samples = operation_response.response.generateVideoResponse.generatedSamples download_url = generated_samples[0].video.uri - params: Dict[str, Any] = {} + params: dict[str, Any] = {} return download_url, params @@ -480,8 +480,8 @@ class GeminiVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Video remix is not supported by Veo API. """ @@ -493,7 +493,7 @@ class GeminiVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """Video remix is not supported.""" raise NotImplementedError("Video remix is not supported by Google Veo.") @@ -503,11 +503,11 @@ class GeminiVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Video list is not supported by Veo API. """ @@ -520,8 +520,8 @@ class GeminiVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - ) -> Dict[str, str]: + custom_llm_provider: str | None = None, + ) -> dict[str, str]: """Video list is not supported.""" raise NotImplementedError("Video list is not supported by Google Veo.") @@ -531,7 +531,7 @@ class GeminiVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Video delete is not supported by Veo API. """ @@ -595,9 +595,7 @@ class GeminiVideoConfig(BaseVideoConfig): def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None): raise NotImplementedError("video extension is not supported for Gemini") - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ..common_utils import GeminiError return GeminiError( diff --git a/litellm/llms/gigachat/authenticator.py b/litellm/llms/gigachat/authenticator.py index e61015a4a21..f5bced63869 100644 --- a/litellm/llms/gigachat/authenticator.py +++ b/litellm/llms/gigachat/authenticator.py @@ -7,7 +7,6 @@ Based on official GigaChat SDK authentication flow. import time import uuid -from typing import Optional, Tuple import httpx @@ -38,10 +37,8 @@ _token_cache = InMemoryCache() class GigaChatAuthError(BaseLLMException): """GigaChat authentication error.""" - pass - -def _get_credentials() -> Optional[str]: +def _get_credentials() -> str | None: """Get GigaChat credentials from environment.""" return get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY") @@ -62,9 +59,9 @@ def _get_http_client() -> HTTPHandler: def get_access_token( - credentials: Optional[str] = None, - scope: Optional[str] = None, - auth_url: Optional[str] = None, + credentials: str | None = None, + scope: str | None = None, + auth_url: str | None = None, ) -> str: """ Get valid access token, using cache if available. @@ -112,9 +109,9 @@ def get_access_token( async def get_access_token_async( - credentials: Optional[str] = None, - scope: Optional[str] = None, - auth_url: Optional[str] = None, + credentials: str | None = None, + scope: str | None = None, + auth_url: str | None = None, ) -> str: """Async version of get_access_token.""" credentials = credentials or _get_credentials() @@ -151,7 +148,7 @@ def _request_token_sync( credentials: str, scope: str, auth_url: str, -) -> Tuple[str, int]: +) -> tuple[str, int]: """ Request new access token from GigaChat OAuth endpoint (sync). @@ -180,7 +177,7 @@ def _request_token_sync( except httpx.RequestError as e: raise GigaChatAuthError( status_code=500, - message=f"GigaChat authentication request failed: {str(e)}", + message=f"GigaChat authentication request failed: {e!s}", ) @@ -188,7 +185,7 @@ async def _request_token_async( credentials: str, scope: str, auth_url: str, -) -> Tuple[str, int]: +) -> tuple[str, int]: """Async version of _request_token_sync.""" headers = { "Authorization": f"Basic {credentials}", @@ -215,11 +212,11 @@ async def _request_token_async( except httpx.RequestError as e: raise GigaChatAuthError( status_code=500, - message=f"GigaChat authentication request failed: {str(e)}", + message=f"GigaChat authentication request failed: {e!s}", ) -def _parse_token_response(response: httpx.Response) -> Tuple[str, int]: +def _parse_token_response(response: httpx.Response) -> tuple[str, int]: """Parse OAuth token response.""" data = response.json() diff --git a/litellm/llms/gigachat/chat/__init__.py b/litellm/llms/gigachat/chat/__init__.py index 3e030497a1a..eb9492b90b3 100644 --- a/litellm/llms/gigachat/chat/__init__.py +++ b/litellm/llms/gigachat/chat/__init__.py @@ -2,8 +2,8 @@ GigaChat Chat Module """ -from .transformation import GigaChatConfig, GigaChatError from .streaming import GigaChatModelResponseIterator +from .transformation import GigaChatConfig, GigaChatError __all__ = [ "GigaChatConfig", diff --git a/litellm/llms/gigachat/chat/streaming.py b/litellm/llms/gigachat/chat/streaming.py index 4f10f8bb658..4092ec1880b 100644 --- a/litellm/llms/gigachat/chat/streaming.py +++ b/litellm/llms/gigachat/chat/streaming.py @@ -4,7 +4,7 @@ GigaChat Streaming Response Handler import json import uuid -from typing import Any, Optional +from typing import Any from litellm.types.llms.openai import ( ChatCompletionToolCallChunk, @@ -20,7 +20,7 @@ class GigaChatModelResponseIterator: self, streaming_response: Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): self.streaming_response = streaming_response self.response_iterator = self.streaming_response @@ -29,9 +29,9 @@ class GigaChatModelResponseIterator: def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: """Parse a single streaming chunk from GigaChat.""" text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False - finish_reason: Optional[str] = None + finish_reason: str | None = None choices = chunk.get("choices", []) if not choices: @@ -90,8 +90,7 @@ class GigaChatModelResponseIterator: chunk = self.response_iterator.__next__() if isinstance(chunk, str): # Parse SSE format: data: {...} - if chunk.startswith("data: "): - chunk = chunk[6:] + chunk = chunk.removeprefix("data: ") if chunk.strip() == "[DONE]": raise StopIteration try: @@ -117,8 +116,7 @@ class GigaChatModelResponseIterator: chunk = await self.response_iterator.__anext__() if isinstance(chunk, str): # Parse SSE format - if chunk.startswith("data: "): - chunk = chunk[6:] + chunk = chunk.removeprefix("data: ") if chunk.strip() == "[DONE]": raise StopAsyncIteration try: diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py index 6c9bbea2970..4007588cfc5 100644 --- a/litellm/llms/gigachat/chat/transformation.py +++ b/litellm/llms/gigachat/chat/transformation.py @@ -8,7 +8,7 @@ import json import time import uuid from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -45,8 +45,6 @@ def is_valid_json(value: str) -> bool: class GigaChatError(BaseLLMException): """GigaChat API error.""" - pass - class GigaChatConfig(BaseConfig): """ @@ -63,36 +61,36 @@ class GigaChatConfig(BaseConfig): stream: Enable streaming """ - temperature: Optional[float] = None - top_p: Optional[float] = None - max_tokens: Optional[int] = None - repetition_penalty: Optional[float] = None - profanity_check: Optional[bool] = None + temperature: float | None = None + top_p: float | None = None + max_tokens: int | None = None + repetition_penalty: float | None = None + profanity_check: bool | None = None def __init__( self, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - max_tokens: Optional[int] = None, - repetition_penalty: Optional[float] = None, - profanity_check: Optional[bool] = None, + temperature: float | None = None, + top_p: float | None = None, + max_tokens: int | None = None, + repetition_penalty: float | None = None, + profanity_check: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) # Instance variables for current request context - self._current_credentials: Optional[str] = None - self._current_api_base: Optional[str] = None + self._current_credentials: str | None = None + self._current_api_base: str | None = None def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """Get complete API URL for chat completions.""" base = api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL @@ -102,11 +100,11 @@ class GigaChatConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Set up headers with OAuth token. @@ -125,7 +123,7 @@ class GigaChatConfig(BaseConfig): return headers - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """Return list of supported OpenAI parameters.""" return [ "stream", @@ -198,7 +196,7 @@ class GigaChatConfig(BaseConfig): return optional_params - def _convert_tools_to_functions(self, tools: List[dict]) -> List[dict]: + def _convert_tools_to_functions(self, tools: list[dict]) -> list[dict]: """Convert OpenAI tools format to GigaChat functions format.""" functions = [] for tool in tools: @@ -213,7 +211,7 @@ class GigaChatConfig(BaseConfig): ) return functions - def _map_tool_choice(self, tool_choice: Union[str, dict]) -> Optional[Union[str, dict]]: + def _map_tool_choice(self, tool_choice: str | dict) -> str | dict | None: """ Map OpenAI tool_choice to GigaChat function_call format. @@ -253,7 +251,7 @@ class GigaChatConfig(BaseConfig): # Default to None (don't set function_call) return None - def _upload_image(self, image_url: str) -> Optional[str]: + def _upload_image(self, image_url: str) -> str | None: """ Upload image to GigaChat and return file_id. @@ -276,7 +274,7 @@ class GigaChatConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -311,7 +309,7 @@ class GigaChatConfig(BaseConfig): return request_data - def _transform_messages(self, messages: List[AllMessageValues]) -> List[dict]: + def _transform_messages(self, messages: list[AllMessageValues]) -> list[dict]: """Transform OpenAI messages to GigaChat format.""" transformed = [] @@ -390,12 +388,12 @@ class GigaChatConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """Transform GigaChat response to OpenAI format.""" try: @@ -480,7 +478,7 @@ class GigaChatConfig(BaseConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: """Return GigaChat error class.""" return GigaChatError( @@ -491,9 +489,9 @@ class GigaChatConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): """Return streaming response iterator.""" from .streaming import GigaChatModelResponseIterator diff --git a/litellm/llms/gigachat/embedding/transformation.py b/litellm/llms/gigachat/embedding/transformation.py index 0da6565050e..84bb867dd3b 100644 --- a/litellm/llms/gigachat/embedding/transformation.py +++ b/litellm/llms/gigachat/embedding/transformation.py @@ -6,14 +6,13 @@ API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/res """ import types -from typing import List, Optional, Tuple, Union import httpx from litellm import LlmProviders +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse @@ -26,8 +25,6 @@ GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" class GigaChatEmbeddingError(BaseLLMException): """GigaChat Embedding API error.""" - pass - class GigaChatEmbeddingConfig(BaseEmbeddingConfig): """ @@ -57,7 +54,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """GigaChat embeddings don't support additional parameters.""" return [] @@ -73,9 +70,9 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[str, Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str | None, str | None]: """ Returns provider info for GigaChat. @@ -87,12 +84,12 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """Get the complete URL for embeddings endpoint.""" base = api_base or GIGACHAT_BASE_URL @@ -123,8 +120,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): input_list = [input] # Remove gigachat/ prefix from model if present - if model.startswith("gigachat/"): - model = model[9:] + model = model.removeprefix("gigachat/") return { "model": model, @@ -137,7 +133,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -184,11 +180,11 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Set up headers with OAuth token for GigaChat. @@ -202,9 +198,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): } return {**default_headers, **headers} - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """Return GigaChat-specific error class.""" return GigaChatEmbeddingError( status_code=status_code, diff --git a/litellm/llms/gigachat/file_handler.py b/litellm/llms/gigachat/file_handler.py index 200428a747a..ee16a6c5870 100644 --- a/litellm/llms/gigachat/file_handler.py +++ b/litellm/llms/gigachat/file_handler.py @@ -9,7 +9,6 @@ import base64 import hashlib import re import uuid -from typing import Dict, Optional, Tuple from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( @@ -24,7 +23,7 @@ from .authenticator import get_access_token, get_access_token_async GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" # Simple in-memory cache for file IDs -_file_cache: Dict[str, str] = {} +_file_cache: dict[str, str] = {} def _get_url_hash(url: str) -> str: @@ -32,7 +31,7 @@ def _get_url_hash(url: str) -> str: return hashlib.sha256(url.encode()).hexdigest() -def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]: +def _parse_data_url(data_url: str) -> tuple[bytes, str, str] | None: """ Parse data URL (base64 image). @@ -51,7 +50,7 @@ def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]: return content_bytes, content_type, ext -def _download_image_sync(url: str) -> Tuple[bytes, str, str]: +def _download_image_sync(url: str) -> tuple[bytes, str, str]: """Download image from URL synchronously.""" client = _get_httpx_client(params={"ssl_verify": False}) response = client.get(url) @@ -63,7 +62,7 @@ def _download_image_sync(url: str) -> Tuple[bytes, str, str]: return response.content, content_type, ext -async def _download_image_async(url: str) -> Tuple[bytes, str, str]: +async def _download_image_async(url: str) -> tuple[bytes, str, str]: """Download image from URL asynchronously.""" client = get_async_httpx_client( llm_provider=LlmProviders.GIGACHAT, @@ -80,9 +79,9 @@ async def _download_image_async(url: str) -> Tuple[bytes, str, str]: def upload_file_sync( image_url: str, - credentials: Optional[str] = None, - api_base: Optional[str] = None, -) -> Optional[str]: + credentials: str | None = None, + api_base: str | None = None, +) -> str | None: """ Upload file to GigaChat and return file_id (sync). @@ -145,9 +144,9 @@ def upload_file_sync( async def upload_file_async( image_url: str, - credentials: Optional[str] = None, - api_base: Optional[str] = None, -) -> Optional[str]: + credentials: str | None = None, + api_base: str | None = None, +) -> str | None: """ Upload file to GigaChat and return file_id (async). diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 9fefc5df0c5..2cb099edfb4 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -2,7 +2,7 @@ import json import os import time from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any import httpx @@ -54,7 +54,7 @@ class Authenticator: access_token = f.read().strip() if access_token: return access_token - except IOError: + except OSError: verbose_logger.warning("No existing access token found or error reading file") for attempt in range(3): @@ -64,11 +64,11 @@ class Authenticator: try: with open(self.access_token_file, "w") as f: f.write(access_token) - except IOError: + except OSError: verbose_logger.error("Error saving access token to file") return access_token except (GetDeviceCodeError, GetAccessTokenError, RefreshAPIKeyError) as e: - verbose_logger.warning(f"Failed attempt {attempt + 1}: {str(e)}") + verbose_logger.warning(f"Failed attempt {attempt + 1}: {e!s}") continue raise GetAccessTokenError( @@ -97,10 +97,10 @@ class Authenticator: message="API key expired", status_code=401, ) - except IOError: + except OSError: verbose_logger.warning("No API key file found or error opening file") except (json.JSONDecodeError, KeyError) as e: - verbose_logger.warning(f"Error reading API key from file: {str(e)}") + verbose_logger.warning(f"Error reading API key from file: {e!s}") except APIKeyExpiredError: pass # Already logged in the try block @@ -116,19 +116,19 @@ class Authenticator: message="API key response missing token", status_code=401, ) - except IOError as e: - verbose_logger.error(f"Error saving API key to file: {str(e)}") + except OSError as e: + verbose_logger.error(f"Error saving API key to file: {e!s}") raise GetAPIKeyError( - message=f"Failed to save API key: {str(e)}", + message=f"Failed to save API key: {e!s}", status_code=500, ) except RefreshAPIKeyError as e: raise GetAPIKeyError( - message=f"Failed to refresh API key: {str(e)}", + message=f"Failed to refresh API key: {e!s}", status_code=401, ) - def get_api_base(self) -> Optional[str]: + def get_api_base(self) -> str | None: """ Get the API endpoint from the api-key.json file. @@ -141,11 +141,11 @@ class Authenticator: endpoints = api_key_info.get("endpoints", {}) api_endpoint = endpoints.get("api") return api_endpoint - except (IOError, json.JSONDecodeError, KeyError) as e: - verbose_logger.warning(f"Error reading API endpoint from file: {str(e)}") + except (OSError, json.JSONDecodeError, KeyError) as e: + verbose_logger.warning(f"Error reading API endpoint from file: {e!s}") return None - def _refresh_api_key(self) -> Dict[str, Any]: + def _refresh_api_key(self) -> dict[str, Any]: """ Refresh the API key using the access token. @@ -173,9 +173,9 @@ class Authenticator: else: verbose_logger.warning(f"API key response missing token: {response_json}") except httpx.HTTPStatusError as e: - verbose_logger.error(f"HTTP error refreshing API key (attempt {attempt + 1}/{max_retries}): {str(e)}") + verbose_logger.error(f"HTTP error refreshing API key (attempt {attempt + 1}/{max_retries}): {e!s}") except Exception as e: - verbose_logger.error(f"Unexpected error refreshing API key: {str(e)}") + verbose_logger.error(f"Unexpected error refreshing API key: {e!s}") raise RefreshAPIKeyError( message="Failed to refresh API key after maximum retries", @@ -187,7 +187,7 @@ class Authenticator: if not os.path.exists(self.token_dir): os.makedirs(self.token_dir, exist_ok=True) - def _get_github_headers(self, access_token: Optional[str] = None) -> Dict[str, str]: + def _get_github_headers(self, access_token: str | None = None) -> dict[str, str]: """ Generate standard GitHub headers for API requests. @@ -213,7 +213,7 @@ class Authenticator: return headers - def _get_device_code(self) -> Dict[str, str]: + def _get_device_code(self) -> dict[str, str]: """ Get a device code for GitHub authentication. @@ -245,21 +245,21 @@ class Authenticator: return resp_json except httpx.HTTPStatusError as e: - verbose_logger.error(f"HTTP error getting device code: {str(e)}") + verbose_logger.error(f"HTTP error getting device code: {e!s}") raise GetDeviceCodeError( - message=f"Failed to get device code: {str(e)}", + message=f"Failed to get device code: {e!s}", status_code=400, ) except json.JSONDecodeError as e: - verbose_logger.error(f"Error decoding JSON response: {str(e)}") + verbose_logger.error(f"Error decoding JSON response: {e!s}") raise GetDeviceCodeError( - message=f"Failed to decode device code response: {str(e)}", + message=f"Failed to decode device code response: {e!s}", status_code=400, ) except Exception as e: - verbose_logger.error(f"Unexpected error getting device code: {str(e)}") + verbose_logger.error(f"Unexpected error getting device code: {e!s}") raise GetDeviceCodeError( - message=f"Failed to get device code: {str(e)}", + message=f"Failed to get device code: {e!s}", status_code=400, ) @@ -304,21 +304,21 @@ class Authenticator: else: verbose_logger.warning(f"Unexpected response: {resp_json}") except httpx.HTTPStatusError as e: - verbose_logger.error(f"HTTP error polling for access token: {str(e)}") + verbose_logger.error(f"HTTP error polling for access token: {e!s}") raise GetAccessTokenError( - message=f"Failed to get access token: {str(e)}", + message=f"Failed to get access token: {e!s}", status_code=400, ) except json.JSONDecodeError as e: - verbose_logger.error(f"Error decoding JSON response: {str(e)}") + verbose_logger.error(f"Error decoding JSON response: {e!s}") raise GetAccessTokenError( - message=f"Failed to decode access token response: {str(e)}", + message=f"Failed to decode access token response: {e!s}", status_code=400, ) except Exception as e: - verbose_logger.error(f"Unexpected error polling for access token: {str(e)}") + verbose_logger.error(f"Unexpected error polling for access token: {e!s}") raise GetAccessTokenError( - message=f"Failed to get access token: {str(e)}", + message=f"Failed to get access token: {e!s}", status_code=400, ) diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 2cc05227948..35e1bf4cdd7 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -1,7 +1,6 @@ import json -from typing import Any, List, Tuple - import os +from typing import Any import httpx @@ -35,7 +34,7 @@ class GithubCopilotConfig(OpenAIConfig): api_base: str | None, api_key: str | None, custom_llm_provider: str, - ) -> Tuple[str | None, str | None, str]: + ) -> tuple[str | None, str | None, str]: dynamic_api_base = ( api_base or self.authenticator.get_api_base() @@ -82,7 +81,7 @@ class GithubCopilotConfig(OpenAIConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, api_key: str | None = None, @@ -135,7 +134,7 @@ class GithubCopilotConfig(OpenAIConfig): return base_params - def _determine_initiator(self, messages: List[AllMessageValues]) -> str: + def _determine_initiator(self, messages: list[AllMessageValues]) -> str: """ Determine if request is user or agent initiated based on message roles. Returns 'agent' if any message has role 'tool' or 'assistant', otherwise 'user'. @@ -146,7 +145,7 @@ class GithubCopilotConfig(OpenAIConfig): return "agent" return "user" - def _has_vision_content(self, messages: List[AllMessageValues]) -> bool: + def _has_vision_content(self, messages: list[AllMessageValues]) -> bool: """ Check if any message contains vision content (images). Returns True if any message has content with vision-related types, otherwise False. @@ -172,8 +171,8 @@ class GithubCopilotConfig(OpenAIConfig): @staticmethod def _parse_anthropic_native_content( - content_blocks: List[Any], - ) -> Tuple[str, List[ChatCompletionToolCallChunk], List[Any] | None]: + content_blocks: list[Any], + ) -> tuple[str, list[ChatCompletionToolCallChunk], list[Any] | None]: """ Parse Anthropic-native content blocks into OpenAI-compatible fields. @@ -219,8 +218,8 @@ class GithubCopilotConfig(OpenAIConfig): return response_json content = "" - tool_calls: List[ChatCompletionToolCallChunk] = [] - thinking_blocks: List[Any] | None = None + tool_calls: list[ChatCompletionToolCallChunk] = [] + thinking_blocks: list[Any] | None = None raw_content = response_json.get("content") if isinstance(raw_content, list): content, tool_calls, thinking_blocks = cls._parse_anthropic_native_content(raw_content) @@ -275,7 +274,7 @@ class GithubCopilotConfig(OpenAIConfig): model_response: "ModelResponse", logging_obj: Any, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, diff --git a/litellm/llms/github_copilot/common_utils.py b/litellm/llms/github_copilot/common_utils.py index 2413cdd63d7..7f45f60983c 100644 --- a/litellm/llms/github_copilot/common_utils.py +++ b/litellm/llms/github_copilot/common_utils.py @@ -2,7 +2,6 @@ Constants for Copilot integration """ -from typing import Optional, Union from uuid import uuid4 import httpx @@ -22,10 +21,10 @@ class GithubCopilotError(BaseLLMException): self, status_code, message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, - body: Optional[dict] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, + body: dict | None = None, ): super().__init__( status_code=status_code, diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index d4014ec6242..75c2d0e8c40 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -7,9 +7,8 @@ Implementation based on analysis of the copilot-api project by caozhiyuan: https://github.com/caozhiyuan/copilot-api """ -from typing import TYPE_CHECKING, Any, Optional - import os +from typing import TYPE_CHECKING, Any import httpx @@ -22,8 +21,8 @@ from litellm.utils import convert_to_model_response_object from ..authenticator import Authenticator from ..common_utils import ( - GetAPIKeyError, DEFAULT_GITHUB_COPILOT_API_BASE, + GetAPIKeyError, get_copilot_default_headers, ) @@ -53,8 +52,8 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for GitHub Copilot API. @@ -89,12 +88,12 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for GitHub Copilot Embedding API endpoint. @@ -144,7 +143,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, diff --git a/litellm/llms/github_copilot/messages/transformation.py b/litellm/llms/github_copilot/messages/transformation.py index 4d7b003c48f..f2182a3664a 100644 --- a/litellm/llms/github_copilot/messages/transformation.py +++ b/litellm/llms/github_copilot/messages/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Optional +from typing import Any from litellm.exceptions import AuthenticationError from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -26,7 +26,7 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): self.authenticator = Authenticator() @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "github_copilot" def handles_web_search_natively(self) -> bool: @@ -54,9 +54,9 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: """ Validate environment for GitHub Copilot and add Copilot-specific headers. @@ -99,12 +99,12 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Return the complete URL for GitHub Copilot /v1/messages endpoint. diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 0393d6a9d64..170ad938efb 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -8,9 +8,8 @@ Implementation based on analysis of the copilot-api project by caozhiyuan: https://github.com/caozhiyuan/copilot-api """ -from typing import TYPE_CHECKING, Any, Dict, Optional, Union - import os +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -96,7 +95,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): def __init__(self) -> None: super().__init__() self.authenticator = Authenticator() - self._stream_item_ids_by_output_index: Dict[int, str] = {} + self._stream_item_ids_by_output_index: dict[int, str] = {} @property def custom_llm_provider(self) -> LlmProviders: @@ -116,7 +115,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map parameters for GitHub Copilot Responses API. @@ -184,7 +183,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: """ Validate environment and set up headers for GitHub Copilot API. @@ -243,7 +242,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -263,7 +262,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Return the responses endpoint return f"{effective_api_base}/responses" - def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]: + def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]: """ Handle reasoning items for GitHub Copilot, preserving encrypted_content. @@ -281,7 +280,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Filter out None values for known problematic fields, # but preserve encrypted_content even if it exists - filtered_item: Dict[str, Any] = {} + filtered_item: dict[str, Any] = {} for k, v in item.items(): # Always include encrypted_content if present (even if None) if k == "encrypted_content": @@ -303,9 +302,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # ==================== Helper Methods ==================== - def _get_input_from_params( - self, litellm_params: Optional[GenericLiteLLMParams] - ) -> Optional[Union[str, ResponseInputParam]]: + def _get_input_from_params(self, litellm_params: GenericLiteLLMParams | None) -> str | ResponseInputParam | None: """ Extract input parameter from litellm_params. @@ -323,7 +320,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # If not found, return None and let the API handle it return None - def _get_initiator(self, input_param: Union[str, ResponseInputParam]) -> str: + def _get_initiator(self, input_param: str | ResponseInputParam) -> str: """ Determine X-Initiator header value based on input analysis. @@ -359,7 +356,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Default to user-initiated return "user" - def _has_vision_input(self, input_param: Union[str, ResponseInputParam]) -> bool: + def _has_vision_input(self, input_param: str | ResponseInputParam) -> bool: """ Check if input contains vision content (images). diff --git a/litellm/llms/google_pse/search/transformation.py b/litellm/llms/google_pse/search/transformation.py index 52d4baba955..c2798c824a9 100644 --- a/litellm/llms/google_pse/search/transformation.py +++ b/litellm/llms/google_pse/search/transformation.py @@ -4,7 +4,7 @@ Calls Google Programmable Search Engine (PSE) API to search the web. Google PSE API Reference: https://developers.google.com/custom-search/v1/reference/rest/v1/cse/list """ -from typing import Dict, List, Literal, Optional, TypedDict, Union +from typing import Literal, TypedDict import httpx @@ -70,11 +70,11 @@ class GooglePSESearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. @@ -103,9 +103,9 @@ class GooglePSESearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -128,13 +128,13 @@ class GooglePSESearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, api_base: str | None = None, - search_engine_id: Optional[str] = None, + search_engine_id: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Google PSE API format. diff --git a/litellm/llms/gradient_ai/chat/transformation.py b/litellm/llms/gradient_ai/chat/transformation.py index e81c09d5cf3..1cf556988ac 100644 --- a/litellm/llms/gradient_ai/chat/transformation.py +++ b/litellm/llms/gradient_ai/chat/transformation.py @@ -1,4 +1,4 @@ -from typing import List, Optional, Tuple, Union, Dict, Literal +from typing import Literal from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( @@ -12,35 +12,35 @@ GRADIENT_AI_SERVERLESS_ENDPOINT = "https://inference.do-ai.run" class GradientAIConfig(OpenAILikeChatConfig): - k: Optional[int] = None - kb_filters: Optional[List[Dict]] = None - filter_kb_content_by_query_metadata: Optional[bool] = None - instruction_override: Optional[str] = None - include_functions_info: Optional[bool] = None - include_retrieval_info: Optional[bool] = None - include_guardrails_info: Optional[bool] = None - provide_citations: Optional[bool] = None - retrieval_method: Optional[Literal["rewrite", "step_back", "sub_queries", "none"]] = None + k: int | None = None + kb_filters: list[dict] | None = None + filter_kb_content_by_query_metadata: bool | None = None + instruction_override: str | None = None + include_functions_info: bool | None = None + include_retrieval_info: bool | None = None + include_guardrails_info: bool | None = None + provide_citations: bool | None = None + retrieval_method: Literal["rewrite", "step_back", "sub_queries", "none"] | None = None def __init__( self, - frequency_penalty: Optional[float] = None, - max_tokens: Optional[int] = None, - max_completion_tokens: Optional[int] = None, - presence_penalty: Optional[float] = None, - retrieval_method: Optional[str] = None, - stop: Optional[Union[str, List[str]]] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - k: Optional[int] = None, - kb_filters: Optional[List[Dict]] = None, - filter_kb_content_by_query_metadata: Optional[bool] = None, - instruction_override: Optional[str] = None, - include_functions_info: Optional[bool] = None, - include_retrieval_info: Optional[bool] = None, - include_guardrails_info: Optional[bool] = None, - provide_citations: Optional[bool] = None, + frequency_penalty: float | None = None, + max_tokens: int | None = None, + max_completion_tokens: int | None = None, + presence_penalty: float | None = None, + retrieval_method: str | None = None, + stop: str | list[str] | None = None, + stream: bool | None = None, + temperature: float | None = None, + top_p: float | None = None, + k: int | None = None, + kb_filters: list[dict] | None = None, + filter_kb_content_by_query_metadata: bool | None = None, + instruction_override: str | None = None, + include_functions_info: bool | None = None, + include_retrieval_info: bool | None = None, + include_guardrails_info: bool | None = None, + provide_citations: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -79,11 +79,11 @@ class GradientAIConfig(OpenAILikeChatConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ): api_key = api_key or get_secret_str("GRADIENT_AI_API_KEY") if api_key is None: @@ -96,12 +96,12 @@ class GradientAIConfig(OpenAILikeChatConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: gradient_ai_endpoint = get_secret_str("GRADIENT_AI_AGENT_ENDPOINT") complete_url = f"{GRADIENT_AI_SERVERLESS_ENDPOINT}/v1/chat/completions" @@ -114,8 +114,8 @@ class GradientAIConfig(OpenAILikeChatConfig): return complete_url def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: gradient_ai_endpoint = get_secret_str("GRADIENT_AI_AGENT_ENDPOINT") if not api_base and not gradient_ai_endpoint: diff --git a/litellm/llms/groq/chat/handler.py b/litellm/llms/groq/chat/handler.py index 0cddc13243f..bd38e4c14e9 100644 --- a/litellm/llms/groq/chat/handler.py +++ b/litellm/llms/groq/chat/handler.py @@ -3,7 +3,7 @@ Handles the chat completion request for groq """ from collections.abc import Callable -from typing import List, Optional, Union, cast +from typing import cast from httpx._config import Timeout @@ -31,20 +31,20 @@ class GroqChatCompletion(OpenAILikeChatHandler): model_response: ModelResponse, print_verbose: Callable, encoding, - api_key: Optional[str], + api_key: str | None, logging_obj, optional_params: dict, acompletion=None, litellm_params=None, logger_fn=None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - custom_endpoint: Optional[bool] = None, - streaming_decoder: Optional[CustomStreamingDecoder] = None, + headers: dict | None = None, + timeout: float | Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + custom_endpoint: bool | None = None, + streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ): - messages = GroqChatConfig()._transform_messages(messages=cast(List[AllMessageValues], messages), model=model) + messages = GroqChatConfig()._transform_messages(messages=cast(list[AllMessageValues], messages), model=model) if optional_params.get("stream") is True: fake_stream = GroqChatConfig()._should_fake_stream(optional_params) diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index 2af71b77fcd..64537e33d0e 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -5,11 +5,7 @@ Translate from OpenAI's `/v1/chat/completions` to Groq's `/v1/chat/completions` from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( Any, - List, Literal, - Optional, - Tuple, - Union, cast, overload, ) @@ -37,35 +33,35 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig class GroqChatConfig(OpenAILikeChatConfig): - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None - tools: Optional[list] = None - tool_choice: Optional[Union[str, dict]] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None + tools: list | None = None + tool_choice: str | dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, - tools: Optional[list] = None, - tool_choice: Optional[Union[str, dict]] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -73,7 +69,7 @@ class GroqChatConfig(OpenAILikeChatConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "groq" @classmethod @@ -82,9 +78,9 @@ class GroqChatConfig(OpenAILikeChatConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return GroqChatCompletionStreamingHandler( streaming_response=streaming_response, @@ -109,20 +105,20 @@ class GroqChatConfig(OpenAILikeChatConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: for idx, message in enumerate(messages): """ 1. Don't pass 'null' function_call assistant message to groq - https://github.com/BerriAI/litellm/issues/5839 @@ -145,8 +141,8 @@ class GroqChatConfig(OpenAILikeChatConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # groq is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.groq.com/openai/v1 api_base = api_base or get_secret_str("GROQ_API_BASE") or "https://api.groq.com/openai/v1" # type: ignore dynamic_api_key = api_key or get_secret_str("GROQ_API_KEY") @@ -194,7 +190,7 @@ class GroqChatConfig(OpenAILikeChatConfig): if self._should_fake_stream(non_default_params): optional_params["fake_stream"] = True if _response_format is not None and isinstance(_response_format, dict): - json_schema: Optional[dict] = None + json_schema: dict | None = None if "response_schema" in _response_format: json_schema = _response_format["response_schema"] elif "json_schema" in _response_format: @@ -253,12 +249,12 @@ class GroqChatConfig(OpenAILikeChatConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: model_response = super().transform_response( model=model, @@ -280,7 +276,7 @@ class GroqChatConfig(OpenAILikeChatConfig): setattr(model_response, "service_tier", mapped_service_tier) return model_response - def _map_groq_service_tier(self, original_service_tier: Optional[str]) -> Literal["auto", "default", "flex"]: + def _map_groq_service_tier(self, original_service_tier: str | None) -> Literal["auto", "default", "flex"]: """ Ensure groq service tier is OpenAI compatible. """ diff --git a/litellm/llms/groq/stt/transformation.py b/litellm/llms/groq/stt/transformation.py index b467fab14f6..0a473b3d09c 100644 --- a/litellm/llms/groq/stt/transformation.py +++ b/litellm/llms/groq/stt/transformation.py @@ -3,41 +3,40 @@ Translate from OpenAI's `/v1/audio/transcriptions` to Groq's `/v1/audio/transcri """ import types -from typing import List, Optional, Union import litellm class GroqSTTConfig: - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None - tools: Optional[list] = None - tool_choice: Optional[Union[str, dict]] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None + tools: list | None = None + tool_choice: str | dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, - tools: Optional[list] = None, - tool_choice: Optional[Union[str, dict]] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -70,7 +69,7 @@ class GroqSTTConfig: "language", ] - def get_supported_openai_response_formats_stt(self) -> List[str]: + def get_supported_openai_response_formats_stt(self) -> list[str]: return ["json", "verbose_json", "text"] def map_openai_params_stt( @@ -90,9 +89,7 @@ class GroqSTTConfig: pass else: raise litellm.utils.UnsupportedParamsError( - message="Groq doesn't support response_format={}. To drop unsupported openai params from the call, set `litellm.drop_params = True`".format( - value - ), + message=f"Groq doesn't support response_format={value}. To drop unsupported openai params from the call, set `litellm.drop_params = True`", status_code=400, ) else: diff --git a/litellm/llms/heroku/chat/transformation.py b/litellm/llms/heroku/chat/transformation.py index 2935f73c762..fd0c29b080b 100644 --- a/litellm/llms/heroku/chat/transformation.py +++ b/litellm/llms/heroku/chat/transformation.py @@ -6,7 +6,7 @@ this is OpenAI compatible - no translation needed / occurs import os from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, overload +from typing import Any, Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -23,20 +23,20 @@ class HerokuError(Exception): class HerokuChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Heroku does not support content in list format. See: https://devcenter.heroku.com/articles/heroku-inference-api-v1-chat-completions#content-object @@ -48,8 +48,8 @@ class HerokuChatConfig(OpenAIGPTConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or os.getenv("HEROKU_API_BASE") api_key = api_key or os.getenv("HEROKU_API_KEY") @@ -57,12 +57,12 @@ class HerokuChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base, _ = self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 8b5dff1fed8..9f0d6985ec7 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -6,12 +6,7 @@ import json from collections.abc import Coroutine from typing import ( Any, - Dict, - List, Literal, - Optional, - Tuple, - Union, cast, overload, ) @@ -38,12 +33,12 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class HostedVLLMChatConfig(OpenAIGPTConfig): - def _convert_custom_tools_to_function_tools(self, tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, Any]]: """ vLLM chat completions currently accepts only OpenAI function tools. Convert custom tools into function tools so request validation does not fail. """ - converted_tools: List[Dict[str, Any]] = [] + converted_tools: list[dict[str, Any]] = [] for idx, tool in enumerate(tools): if not isinstance(tool, dict): converted_tools.append(tool) @@ -73,7 +68,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): "required": ["input"], } - function_tool: Dict[str, Any] = { + function_tool: dict[str, Any] = { "type": "function", "function": { "name": str(tool_name), @@ -87,7 +82,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): return converted_tools - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: params = super().get_supported_openai_params(model) params.extend(["reasoning_effort", "thinking"]) return params @@ -119,8 +114,8 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): return super().map_openai_params(non_default_params, optional_params, model, drop_params) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") dynamic_api_key = api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key" return api_base, dynamic_api_key @@ -157,20 +152,20 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Support translating: - video files from file_id or file_data to video_url @@ -235,7 +230,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): elif message["role"] == "user": message_content = message.get("content") if message_content and isinstance(message_content, list): - replaced_content_items: List[Tuple[int, ChatCompletionFileObject]] = [] + replaced_content_items: list[tuple[int, ChatCompletionFileObject]] = [] for idx, content_item in enumerate(message_content): if content_item.get("type") == "file": content_item = cast(ChatCompletionFileObject, content_item) diff --git a/litellm/llms/hosted_vllm/embedding/transformation.py b/litellm/llms/hosted_vllm/embedding/transformation.py index 9c3e8c6c7cc..ce42bd9de19 100644 --- a/litellm/llms/hosted_vllm/embedding/transformation.py +++ b/litellm/llms/hosted_vllm/embedding/transformation.py @@ -7,7 +7,7 @@ VLLM is OpenAI-compatible and supports embeddings via the /v1/embeddings endpoin Docs: https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html """ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -29,8 +29,6 @@ else: class HostedVLLMEmbeddingError(BaseLLMException): """Exception class for Hosted VLLM Embedding errors.""" - pass - class HostedVLLMEmbeddingConfig(BaseEmbeddingConfig): """ @@ -43,11 +41,11 @@ class HostedVLLMEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Hosted VLLM API. @@ -68,12 +66,12 @@ class HostedVLLMEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Hosted VLLM Embedding API endpoint. @@ -122,7 +120,7 @@ class HostedVLLMEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -167,9 +165,7 @@ class HostedVLLMEmbeddingConfig(BaseEmbeddingConfig): optional_params[param] = value return optional_params - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """ Get the error class for Hosted VLLM errors. """ diff --git a/litellm/llms/hosted_vllm/rerank/transformation.py b/litellm/llms/hosted_vllm/rerank/transformation.py index 77504eba04a..3b108d87c8a 100644 --- a/litellm/llms/hosted_vllm/rerank/transformation.py +++ b/litellm/llms/hosted_vllm/rerank/transformation.py @@ -2,7 +2,7 @@ Transformation logic for Hosted VLLM rerank """ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -28,7 +28,7 @@ class HostedVLLMRerankError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] | None = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) @@ -70,15 +70,15 @@ class HostedVLLMRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map parameters for Hosted VLLM rerank """ @@ -127,7 +127,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -168,9 +168,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig): return self._transform_response(raw_response_json) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return HostedVLLMRerankError(message=error_message, status_code=status_code, headers=headers) def _transform_response(self, response: dict) -> RerankResponse: @@ -181,12 +179,12 @@ class HostedVLLMRerankConfig(BaseRerankConfig): rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) # Extract results - _results: List[dict] | None = response.get("results") + _results: list[dict] | None = response.get("results") if _results is None: raise ValueError(f"No results found in the response={response}") - rerank_results: List[RerankResponseResult] = [] + rerank_results: list[RerankResponseResult] = [] for result in _results: # Validate required fields exist diff --git a/litellm/llms/hosted_vllm/responses/transformation.py b/litellm/llms/hosted_vllm/responses/transformation.py index d79690292aa..916293b0e35 100644 --- a/litellm/llms/hosted_vllm/responses/transformation.py +++ b/litellm/llms/hosted_vllm/responses/transformation.py @@ -6,8 +6,6 @@ so this config enables direct routing instead of falling back to the chat completions → responses conversion pipeline. """ -from typing import Optional - from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -32,7 +30,7 @@ class HostedVLLMResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = ( @@ -47,7 +45,7 @@ class HostedVLLMResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") diff --git a/litellm/llms/hosted_vllm/transcriptions/transformation.py b/litellm/llms/hosted_vllm/transcriptions/transformation.py index e726ee33abf..79f9c3efdc9 100644 --- a/litellm/llms/hosted_vllm/transcriptions/transformation.py +++ b/litellm/llms/hosted_vllm/transcriptions/transformation.py @@ -2,8 +2,6 @@ Transformation logic for Hosted VLLM rerank """ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.audio_transcription.transformation import ( @@ -21,7 +19,7 @@ class HostedVLLMAudioTranscriptionError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) @@ -32,12 +30,12 @@ class HostedVLLMAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: # Remove trailing slashes and ensure clean base URL diff --git a/litellm/llms/huggingface/chat/transformation.py b/litellm/llms/huggingface/chat/transformation.py index 353d3abac6b..da1ebd7c23a 100644 --- a/litellm/llms/huggingface/chat/transformation.py +++ b/litellm/llms/huggingface/chat/transformation.py @@ -1,6 +1,6 @@ import logging import os -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -47,11 +47,11 @@ class HuggingFaceChatConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], - optional_params: Dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: default_headers = { "content-type": "application/json", @@ -63,12 +63,10 @@ class HuggingFaceChatConfig(OpenAIGPTConfig): return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return HuggingFaceError(status_code=status_code, message=error_message, headers=headers) - def get_base_url(self, model: str, base_url: Optional[str]) -> Optional[str]: + def get_base_url(self, model: str, base_url: str | None) -> str | None: """ Get the API base for the Huggingface API. @@ -82,12 +80,12 @@ class HuggingFaceChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the API call. @@ -125,7 +123,7 @@ class HuggingFaceChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/huggingface/common_utils.py b/litellm/llms/huggingface/common_utils.py index 9ab4367c9b3..9dbdf05d0ec 100644 --- a/litellm/llms/huggingface/common_utils.py +++ b/litellm/llms/huggingface/common_utils.py @@ -1,6 +1,6 @@ import os from functools import lru_cache -from typing import Literal, Optional, Union +from typing import Literal import httpx @@ -14,9 +14,9 @@ class HuggingFaceError(BaseLLMException): self, status_code, message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, ): super().__init__( status_code=status_code, @@ -96,7 +96,7 @@ def _fetch_inference_provider_mapping(model: str) -> dict: status_code = 500 headers = {} raise HuggingFaceError( - message=f"Failed to fetch provider mapping: {str(e)}", + message=f"Failed to fetch provider mapping: {e!s}", status_code=status_code, headers=headers, ) diff --git a/litellm/llms/huggingface/embedding/handler.py b/litellm/llms/huggingface/embedding/handler.py index 6aa5706c40b..e7ce9bcf1ae 100644 --- a/litellm/llms/huggingface/embedding/handler.py +++ b/litellm/llms/huggingface/embedding/handler.py @@ -1,7 +1,7 @@ import json import os from collections.abc import Callable -from typing import Any, Dict, List, Literal, Optional, Union, get_args +from typing import Any, Literal, get_args import httpx @@ -29,29 +29,29 @@ hf_tasks_embeddings = ( ) -def get_hf_task_embedding_for_model(model: str, task_type: Optional[str], api_base: str) -> Optional[str]: +def get_hf_task_embedding_for_model(model: str, task_type: str | None, api_base: str) -> str | None: if task_type is not None: if task_type in get_args(hf_tasks_embeddings): return task_type else: - raise Exception("Invalid task_type={}. Expected one of={}".format(task_type, hf_tasks_embeddings)) + raise Exception(f"Invalid task_type={task_type}. Expected one of={hf_tasks_embeddings}") http_client = HTTPHandler(concurrent_limit=1) model_info = http_client.get(url=f"{api_base}/api/models/{model}") model_info_dict = model_info.json() - pipeline_tag: Optional[str] = model_info_dict.get("pipeline_tag", None) + pipeline_tag: str | None = model_info_dict.get("pipeline_tag", None) return pipeline_tag -async def async_get_hf_task_embedding_for_model(model: str, task_type: Optional[str], api_base: str) -> Optional[str]: +async def async_get_hf_task_embedding_for_model(model: str, task_type: str | None, api_base: str) -> str | None: if task_type is not None: if task_type in get_args(hf_tasks_embeddings): return task_type else: - raise Exception("Invalid task_type={}. Expected one of={}".format(task_type, hf_tasks_embeddings)) + raise Exception(f"Invalid task_type={task_type}. Expected one of={hf_tasks_embeddings}") http_client = get_async_httpx_client( llm_provider=litellm.LlmProviders.HUGGINGFACE, ) @@ -60,19 +60,19 @@ async def async_get_hf_task_embedding_for_model(model: str, task_type: Optional[ model_info_dict = model_info.json() - pipeline_tag: Optional[str] = model_info_dict.get("pipeline_tag", None) + pipeline_tag: str | None = model_info_dict.get("pipeline_tag", None) return pipeline_tag class HuggingFaceEmbedding(BaseLLM): - _client_session: Optional[httpx.Client] = None - _aclient_session: Optional[httpx.AsyncClient] = None + _client_session: httpx.Client | None = None + _aclient_session: httpx.AsyncClient | None = None def __init__(self) -> None: super().__init__() - def _transform_input_on_pipeline_tag(self, input: List, pipeline_tag: Optional[str]) -> dict: + def _transform_input_on_pipeline_tag(self, input: list, pipeline_tag: str | None) -> dict: if pipeline_tag is None: return {"inputs": input} if pipeline_tag == "sentence-similarity" or pipeline_tag == "similarity": @@ -94,9 +94,9 @@ class HuggingFaceEmbedding(BaseLLM): async def _async_transform_input( self, model: str, - task_type: Optional[str], + task_type: str | None, embed_url: str, - input: List, + input: list, optional_params: dict, ) -> dict: hf_task = await async_get_hf_task_embedding_for_model(model=model, task_type=task_type, api_base=HF_HUB_URL) @@ -134,13 +134,13 @@ class HuggingFaceEmbedding(BaseLLM): def _transform_input( self, - input: List, + input: list, model: str, call_type: Literal["sync", "async"], optional_params: dict, embed_url: str, ) -> dict: - data: Dict = {} + data: dict = {} ## TRANSFORMATION ## if "sentence-transformers" in model: @@ -172,7 +172,7 @@ class HuggingFaceEmbedding(BaseLLM): embeddings: dict, model_response: EmbeddingResponse, model: str, - input: List, + input: list, encoding: Any, ) -> EmbeddingResponse: output_data = [] @@ -187,15 +187,7 @@ class HuggingFaceEmbedding(BaseLLM): ) else: for idx, embedding in enumerate(embeddings): - if isinstance(embedding, float): - output_data.append( - { - "object": "embedding", - "index": idx, - "embedding": embedding, # flatten list returned from hf - } - ) - elif isinstance(embedding, list) and isinstance(embedding[0], float): + if isinstance(embedding, float) or isinstance(embedding, list) and isinstance(embedding[0], float): output_data.append( { "object": "embedding", @@ -236,14 +228,14 @@ class HuggingFaceEmbedding(BaseLLM): model: str, input: list, model_response: litellm.utils.EmbeddingResponse, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, optional_params: dict, api_base: str, - api_key: Optional[str], + api_key: str | None, headers: dict, encoding: Callable, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, ): ## TRANSFORMATION ## data = self._transform_input( @@ -303,11 +295,11 @@ class HuggingFaceEmbedding(BaseLLM): litellm_params: dict, logging_obj: LiteLLMLoggingObj, encoding: Callable, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Union[float, httpx.Timeout] = httpx.Timeout(None), - aembedding: Optional[bool] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout = httpx.Timeout(None), + aembedding: bool | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, headers={}, ) -> EmbeddingResponse: super().embedding() diff --git a/litellm/llms/huggingface/embedding/transformation.py b/litellm/llms/huggingface/embedding/transformation.py index 6f27e3115eb..78e80e679bf 100644 --- a/litellm/llms/huggingface/embedding/transformation.py +++ b/litellm/llms/huggingface/embedding/transformation.py @@ -2,7 +2,7 @@ import json import os import time from copy import deepcopy -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -40,40 +40,40 @@ class HuggingFaceEmbeddingConfig(BaseConfig): Reference: https://huggingface.github.io/text-generation-inference/#/Text%20Generation%20Inference/compat_generate """ - hf_task: Optional[hf_tasks] = ( + hf_task: hf_tasks | None = ( None # litellm-specific param, used to know the api spec to use when calling huggingface api ) - best_of: Optional[int] = None - decoder_input_details: Optional[bool] = None - details: Optional[bool] = True # enables returning logprobs + best of - max_new_tokens: Optional[int] = None - repetition_penalty: Optional[float] = None - return_full_text: Optional[bool] = False # by default don't return the input as part of the output - seed: Optional[int] = None - temperature: Optional[float] = None - top_k: Optional[int] = None - top_n_tokens: Optional[int] = None - top_p: Optional[int] = None - truncate: Optional[int] = None - typical_p: Optional[float] = None - watermark: Optional[bool] = None + best_of: int | None = None + decoder_input_details: bool | None = None + details: bool | None = True # enables returning logprobs + best of + max_new_tokens: int | None = None + repetition_penalty: float | None = None + return_full_text: bool | None = False # by default don't return the input as part of the output + seed: int | None = None + temperature: float | None = None + top_k: int | None = None + top_n_tokens: int | None = None + top_p: int | None = None + truncate: int | None = None + typical_p: float | None = None + watermark: bool | None = None def __init__( self, - best_of: Optional[int] = None, - decoder_input_details: Optional[bool] = None, - details: Optional[bool] = None, - max_new_tokens: Optional[int] = None, - repetition_penalty: Optional[float] = None, - return_full_text: Optional[bool] = None, - seed: Optional[int] = None, - temperature: Optional[float] = None, - top_k: Optional[int] = None, - top_n_tokens: Optional[int] = None, - top_p: Optional[int] = None, - truncate: Optional[int] = None, - typical_p: Optional[float] = None, - watermark: Optional[bool] = None, + best_of: int | None = None, + decoder_input_details: bool | None = None, + details: bool | None = None, + max_new_tokens: int | None = None, + repetition_penalty: float | None = None, + return_full_text: bool | None = None, + seed: int | None = None, + temperature: float | None = None, + top_k: int | None = None, + top_n_tokens: int | None = None, + top_p: int | None = None, + truncate: int | None = None, + typical_p: float | None = None, + watermark: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -101,11 +101,11 @@ class HuggingFaceEmbeddingConfig(BaseConfig): def map_openai_params( self, - non_default_params: Dict, - optional_params: Dict, + non_default_params: dict, + optional_params: dict, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: for param, value in non_default_params.items(): # temperature, top_p, n, stream, stop, max_tokens, n, presence_penalty default to None if param == "temperature": @@ -136,7 +136,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): return optional_params - def get_hf_api_key(self) -> Optional[str]: + def get_hf_api_key(self) -> str | None: return get_secret_str("HUGGINGFACE_API_KEY") def read_tgi_conv_models(self): @@ -180,7 +180,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): except Exception: return set(), set() - def get_hf_task_for_model(self, model: str) -> Tuple[hf_tasks, str]: + def get_hf_task_for_model(self, model: str) -> tuple[hf_tasks, str]: # read text file, cast it to set # read the file called "huggingface_llms_metadata/hf_text_generation_models.txt" if model.split("/")[0] in hf_task_list: @@ -200,7 +200,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -208,7 +208,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): task = litellm_params.get("task", None) ## VALIDATE API FORMAT if task is None or not isinstance(task, str) or task not in hf_task_list: - raise Exception("Invalid hf task - {}. Valid formats - {}.".format(task, hf_tasks)) + raise Exception(f"Invalid hf task - {task}. Valid formats - {hf_tasks}.") ## Load Config config = litellm.HuggingFaceEmbeddingConfig.get_config() @@ -318,14 +318,14 @@ class HuggingFaceEmbeddingConfig(BaseConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, - messages: List[AllMessageValues], - optional_params: Dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: default_headers = { "content-type": "application/json", } @@ -337,9 +337,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): headers = {**headers, **default_headers} return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return HuggingFaceError(status_code=status_code, message=error_message, headers=headers) def _convert_streamed_response_to_complete_response( @@ -348,8 +346,8 @@ class HuggingFaceEmbeddingConfig(BaseConfig): logging_obj: LoggingClass, model: str, data: dict, - api_key: Optional[str] = None, - ) -> List[Dict[str, Any]]: + api_key: str | None = None, + ) -> list[dict[str, Any]]: streamed_response = CustomStreamWrapper( completion_stream=response.iter_lines(), model=model, @@ -359,7 +357,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): content = "" for chunk in streamed_response: content += chunk["choices"][0]["delta"]["content"] - completion_response: List[Dict[str, Any]] = [{"generated_text": content}] + completion_response: list[dict[str, Any]] = [{"generated_text": content}] ## LOGGING logging_obj.post_call( input=data, @@ -371,12 +369,12 @@ class HuggingFaceEmbeddingConfig(BaseConfig): def convert_to_model_response_object( self, - completion_response: Union[List[Dict[str, Any]], Dict[str, Any]], + completion_response: list[dict[str, Any]] | dict[str, Any], model_response: ModelResponse, - task: Optional[hf_tasks], + task: hf_tasks | None, optional_params: dict, encoding: Any, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, ): if task is None: @@ -479,13 +477,13 @@ class HuggingFaceEmbeddingConfig(BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## Some servers might return streaming responses even though stream was not set to true. (e.g. Baseten) task = litellm_params.get("task", None) diff --git a/litellm/llms/huggingface/rerank/transformation.py b/litellm/llms/huggingface/rerank/transformation.py index cdad77a9815..c29dc5b3fb4 100644 --- a/litellm/llms/huggingface/rerank/transformation.py +++ b/litellm/llms/huggingface/rerank/transformation.py @@ -1,5 +1,5 @@ import os -from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from typing_extensions import TypedDict @@ -42,11 +42,10 @@ class HuggingFaceRerankResponse(TypedDict): """Type definition for HuggingFace rerank API complete response.""" # The response is a list of HuggingFaceRerankResponseItem - pass # Type alias for the actual response structure -HuggingFaceRerankResponseList = List[HuggingFaceRerankResponseItem] +HuggingFaceRerankResponseList = list[HuggingFaceRerankResponseItem] class HuggingFaceRerankConfig(BaseRerankConfig): @@ -93,15 +92,15 @@ class HuggingFaceRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: optional_rerank_params = {} if non_default_params is not None: for k, v in non_default_params.items(): @@ -145,7 +144,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Union[OptionalRerankParams, dict], + optional_rerank_params: OptionalRerankParams | dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -255,16 +254,14 @@ class HuggingFaceRerankConfig(BaseRerankConfig): meta=rerank_meta, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return HuggingFaceError(message=error_message, status_code=status_code) def get_api_credentials( self, api_key: str | None = None, api_base: str | None = None, - ) -> Tuple[str | None, str | None]: + ) -> tuple[str | None, str | None]: """ Get API key and base URL from multiple sources. Returns tuple of (api_key, api_base). diff --git a/litellm/llms/hyperbolic/chat/transformation.py b/litellm/llms/hyperbolic/chat/transformation.py index 48af9fa68a0..60b84770dce 100644 --- a/litellm/llms/hyperbolic/chat/transformation.py +++ b/litellm/llms/hyperbolic/chat/transformation.py @@ -2,8 +2,6 @@ Translate from OpenAI's `/v1/chat/completions` to Hyperbolic's `/v1/chat/completions` """ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai_like.chat.transformation import OpenAILikeChatConfig @@ -15,12 +13,12 @@ class HyperbolicChatConfig(OpenAILikeChatConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "hyperbolic" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # Hyperbolic is openai compatible, we just need to set the api_base api_base = ( api_base diff --git a/litellm/llms/inception/chat/transformation.py b/litellm/llms/inception/chat/transformation.py index 4c8af768047..f5bc8975ef9 100644 --- a/litellm/llms/inception/chat/transformation.py +++ b/litellm/llms/inception/chat/transformation.py @@ -6,8 +6,6 @@ diffusion LLMs through an OpenAI-compatible API, so we only need to point the OpenAI-like handler at the Inception API base and pick up the Inception API key. """ -from typing import List, Optional, Tuple - import litellm from litellm.secret_managers.main import get_secret_str @@ -20,10 +18,10 @@ class InceptionChatConfig(OpenAILikeChatConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "inception" - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "max_tokens", "max_completion_tokens", @@ -42,8 +40,8 @@ class InceptionChatConfig(OpenAILikeChatConfig): ] def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: passed_api_base = api_base api_base = api_base or get_secret_str("INCEPTION_API_BASE") or "https://api.inceptionlabs.ai/v1" # type: ignore dynamic_api_key = api_key diff --git a/litellm/llms/inception/completion/transformation.py b/litellm/llms/inception/completion/transformation.py index 1035042f6bf..3244709ab4e 100644 --- a/litellm/llms/inception/completion/transformation.py +++ b/litellm/llms/inception/completion/transformation.py @@ -8,13 +8,11 @@ Inception's FIM endpoint is OpenAI text-completion compatible: it takes a (see the `text-completion-inception` branch in `main.py`). """ -from typing import List - from litellm.llms.openai.completion.transformation import OpenAITextCompletionConfig class InceptionTextCompletionConfig(OpenAITextCompletionConfig): - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "suffix", "max_tokens", diff --git a/litellm/llms/infinity/common_utils.py b/litellm/llms/infinity/common_utils.py index cf52309ad84..e23fe4a0d37 100644 --- a/litellm/llms/infinity/common_utils.py +++ b/litellm/llms/infinity/common_utils.py @@ -1,11 +1,10 @@ -from typing import Union import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException class InfinityError(BaseLLMException): - def __init__(self, status_code: int, message: str, headers: Union[dict, httpx.Headers] = {}): + def __init__(self, status_code: int, message: str, headers: dict | httpx.Headers = {}): self.status_code = status_code self.message = message self.request = httpx.Request(method="POST", url="https://github.com/michaelfeil/infinity") diff --git a/litellm/llms/infinity/embedding/transformation.py b/litellm/llms/infinity/embedding/transformation.py index fd75887baa3..e2524101cf1 100644 --- a/litellm/llms/infinity/embedding/transformation.py +++ b/litellm/llms/infinity/embedding/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -22,12 +20,12 @@ class InfinityEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: raise ValueError("api_base is required for Infinity embeddings") @@ -41,11 +39,11 @@ class InfinityEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("INFINITY_API_KEY") @@ -109,7 +107,7 @@ class InfinityEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -131,7 +129,5 @@ class InfinityEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return InfinityError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/infinity/rerank/transformation.py b/litellm/llms/infinity/rerank/transformation.py index 94746da4609..de224e9636f 100644 --- a/litellm/llms/infinity/rerank/transformation.py +++ b/litellm/llms/infinity/rerank/transformation.py @@ -4,12 +4,10 @@ Transformation logic from Cohere's /v1/rerank format to Infinity's `/v1/rerank` Why separate file? Make it easy to see how transformation works """ -from litellm._uuid import uuid -from typing import List, Optional - import httpx import litellm +from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.cohere.rerank.transformation import CohereRerankConfig from litellm.secret_managers.main import get_secret_str @@ -28,9 +26,9 @@ from ..common_utils import InfinityError class InfinityRerankConfig(CohereRerankConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, model: str, - optional_params: Optional[dict] = None, + optional_params: dict | None = None, ) -> str: if api_base is None: raise ValueError("api_base is required for Infinity rerank") @@ -44,8 +42,8 @@ class InfinityRerankConfig(CohereRerankConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - optional_params: Optional[dict] = None, + api_key: str | None = None, + optional_params: dict | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("INFINITY_API_KEY") or get_secret_str("INFINITY_API_KEY") or litellm.infinity_key @@ -69,7 +67,7 @@ class InfinityRerankConfig(CohereRerankConfig): raw_response: httpx.Response, model_response: RerankResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -94,7 +92,7 @@ class InfinityRerankConfig(CohereRerankConfig): ) rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) - cohere_results: List[RerankResponseResult] = [] + cohere_results: list[RerankResponseResult] = [] if raw_response_json.get("results"): for result in raw_response_json.get("results"): _rerank_response = RerankResponseResult( diff --git a/litellm/llms/jina_ai/embedding/transformation.py b/litellm/llms/jina_ai/embedding/transformation.py index 80927a59a64..fc4909fd392 100644 --- a/litellm/llms/jina_ai/embedding/transformation.py +++ b/litellm/llms/jina_ai/embedding/transformation.py @@ -7,15 +7,15 @@ Docs - https://jina.ai/embeddings/ """ import types -from typing import List, Optional, Tuple, Union, cast +from typing import cast import httpx from litellm import LlmProviders -from litellm.secret_managers.main import get_secret_str -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.llms.base_llm import BaseEmbeddingConfig from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm import BaseEmbeddingConfig +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse from litellm.utils import is_base64_encoded @@ -54,7 +54,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): and v is not None } - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["dimensions"] def map_openai_params( @@ -70,9 +70,9 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[str, Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str | None, str | None]: """ Returns: Tuple[str, Optional[str], Optional[str]]: @@ -91,12 +91,12 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return f"{api_base}/embeddings" if api_base else "https://api.jina.ai/v1/embeddings" @@ -108,8 +108,8 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): headers: dict, ) -> dict: data = {"model": model, **optional_params} - input = cast(List[str], input) if isinstance(input, List) else [input] - if any((is_base64_encoded(x) for x in input)): + input = cast(list[str], input) if isinstance(input, list) else [input] + if any(is_base64_encoded(x) for x in input): transformed_input = [] for value in input: if isinstance(value, str): @@ -129,7 +129,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -148,11 +148,11 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: default_headers = { "Content-Type": "application/json", @@ -162,9 +162,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): headers = {**default_headers, **headers} return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return JinaAIError( status_code=status_code, message=error_message, diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 7f4c0709bdd..710aced0a88 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -6,7 +6,7 @@ Why separate file? Make it easy to see how transformation works Docs - https://jina.ai/reranker """ -from typing import Any, Dict, List, Tuple, Union +from typing import Any from httpx import URL, Response @@ -38,15 +38,15 @@ class JinaAIRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: optional_params = {} supported_params = self.get_supported_cohere_rerank_params(model) for k, v in non_default_params.items(): @@ -77,10 +77,10 @@ class JinaAIRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, - headers: Dict, + optional_rerank_params: dict, + headers: dict, litellm_params: dict | None = None, - ) -> Dict: + ) -> dict: return {"model": model, **optional_rerank_params} def transform_rerank_response( @@ -90,9 +90,9 @@ class JinaAIRerankConfig(BaseRerankConfig): model_response: RerankResponse, logging_obj: LiteLLMLoggingObj, api_key: str | None = None, - request_data: Dict = {}, - optional_params: Dict = {}, - litellm_params: Dict = {}, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, ) -> RerankResponse: if raw_response.status_code != 200: raise Exception(raw_response.text) @@ -105,7 +105,7 @@ class JinaAIRerankConfig(BaseRerankConfig): _tokens = RerankTokens(**_json_response.get("usage", {})) rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) - _results: List[dict] | None = _json_response.get("results") + _results: list[dict] | None = _json_response.get("results") if _results is None: raise ValueError(f"No results found in the response={_json_response}") @@ -135,11 +135,11 @@ class JinaAIRerankConfig(BaseRerankConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, optional_params: dict | None = None, - ) -> Dict: + ) -> dict: if api_key is None: raise ValueError("api_key is required. Set via `api_key` parameter or `JINA_API_KEY` environment variable.") return { @@ -154,7 +154,7 @@ class JinaAIRerankConfig(BaseRerankConfig): custom_llm_provider: str | None = None, billed_units: RerankBilledUnits | None = None, model_info: ModelInfo | None = None, - ) -> Tuple[float, float]: + ) -> tuple[float, float]: """ Jina AI reranker is priced at $0.000000018 per token. """ diff --git a/litellm/llms/lambda_ai/chat/transformation.py b/litellm/llms/lambda_ai/chat/transformation.py index 96d1dad1416..9bfa1fae840 100644 --- a/litellm/llms/lambda_ai/chat/transformation.py +++ b/litellm/llms/lambda_ai/chat/transformation.py @@ -2,8 +2,6 @@ Translate from OpenAI's `/v1/chat/completions` to Lambda's `/v1/chat/completions` """ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai_like.chat.transformation import OpenAILikeChatConfig @@ -15,12 +13,12 @@ class LambdaAIChatConfig(OpenAILikeChatConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "lambda_ai" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # Lambda AI is openai compatible, we just need to set the api_base api_base = ( api_base or get_secret_str("LAMBDA_API_BASE") or "https://api.lambda.ai/v1" # Default Lambda API base URL diff --git a/litellm/llms/langflow/a2a.py b/litellm/llms/langflow/a2a.py index dbe3e02401d..c2eed4a24f6 100644 --- a/litellm/llms/langflow/a2a.py +++ b/litellm/llms/langflow/a2a.py @@ -1,15 +1,15 @@ import hashlib -from typing import Any, Dict, Optional +from typing import Any -def get_session_id_from_a2a_params(params: Dict[str, Any]) -> Optional[str]: +def get_session_id_from_a2a_params(params: dict[str, Any]) -> str | None: message = params.get("message", {}) if isinstance(message, dict): return message.get("contextId") return getattr(message, "contextId", None) -def scope_session_to_principal(session_id: str, principal: Optional[str]) -> str: +def scope_session_to_principal(session_id: str, principal: str | None) -> str: """ Bind a client-supplied A2A contextId to the authenticated principal. @@ -26,10 +26,10 @@ def scope_session_to_principal(session_id: str, principal: Optional[str]) -> str def merge_a2a_session_into_litellm_params( - litellm_params: Dict[str, Any], - params: Dict[str, Any], - principal: Optional[str] = None, -) -> Dict[str, Any]: + litellm_params: dict[str, Any], + params: dict[str, Any], + principal: str | None = None, +) -> dict[str, Any]: merged = dict(litellm_params) session_id = get_session_id_from_a2a_params(params) if session_id and "session_id" not in merged: diff --git a/litellm/llms/langflow/chat/transformation.py b/litellm/llms/langflow/chat/transformation.py index 73fa49f492b..69af32fc840 100644 --- a/litellm/llms/langflow/chat/transformation.py +++ b/litellm/llms/langflow/chat/transformation.py @@ -1,6 +1,6 @@ """LangFlow run API: POST {api_base}/api/v1/run/{flow_id}""" -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from urllib.parse import quote import httpx @@ -29,8 +29,6 @@ else: class LangFlowError(BaseLLMException): """Exception class for LangFlow API errors.""" - pass - class LangFlowConfig(BaseConfig): """ @@ -45,16 +43,16 @@ class LangFlowConfig(BaseConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: from litellm.secret_managers.main import get_secret_str api_base = api_base or get_secret_str("LANGFLOW_API_BASE") or "http://localhost:7860" api_key = api_key or get_secret_str("LANGFLOW_API_KEY") return api_base, api_key - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["stream"] def map_openai_params( @@ -89,12 +87,12 @@ class LangFlowConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: raise ValueError( @@ -105,7 +103,7 @@ class LangFlowConfig(BaseConfig): flow_id = quote(self._get_flow_id(model, optional_params), safe="") return f"{api_base}/api/v1/run/{flow_id}" - def _get_last_user_message(self, messages: List[AllMessageValues]) -> str: + def _get_last_user_message(self, messages: list[AllMessageValues]) -> str: """Extract the text of the last user message to use as input_value.""" for msg in reversed(messages): if msg.get("role") == "user": @@ -140,7 +138,7 @@ class LangFlowConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -160,7 +158,7 @@ class LangFlowConfig(BaseConfig): input_value = self._get_last_user_message(messages) - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "input_value": input_value, "input_type": optional_params.get("input_type", "chat"), "output_type": optional_params.get("output_type", "chat"), @@ -173,7 +171,7 @@ class LangFlowConfig(BaseConfig): verbose_logger.debug(f"LangFlow request payload: {payload}") return payload - def _extract_content_from_response(self, response_json: dict) -> Optional[str]: + def _extract_content_from_response(self, response_json: dict) -> str | None: """ Extract the assistant text from a LangFlow run response. @@ -222,12 +220,12 @@ class LangFlowConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: response_json = raw_response.json() @@ -277,11 +275,11 @@ class LangFlowConfig(BaseConfig): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: self._reject_caller_tweaks(request_data) return headers, None @@ -289,11 +287,11 @@ class LangFlowConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers["Content-Type"] = "application/json" @@ -302,9 +300,7 @@ class LangFlowConfig(BaseConfig): return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return LangFlowError(status_code=status_code, message=error_message) @property @@ -313,8 +309,8 @@ class LangFlowConfig(BaseConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: return stream is True diff --git a/litellm/llms/langgraph/chat/sse_iterator.py b/litellm/llms/langgraph/chat/sse_iterator.py index 2eb17b4d4b4..895bbdca656 100644 --- a/litellm/llms/langgraph/chat/sse_iterator.py +++ b/litellm/llms/langgraph/chat/sse_iterator.py @@ -6,16 +6,12 @@ Handles Server-Sent Events (SSE) streaming responses from LangGraph. import json import uuid -from typing import TYPE_CHECKING, Optional import httpx from litellm._logging import verbose_logger from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices -if TYPE_CHECKING: - pass - class LangGraphSSEStreamIterator: """ @@ -44,7 +40,7 @@ class LangGraphSSEStreamIterator: self.async_line_iterator = self.response.aiter_lines() return self - def _parse_sse_line(self, line: str) -> Optional[ModelResponseStream]: + def _parse_sse_line(self, line: str) -> ModelResponseStream | None: """ Parse a single SSE line and return a ModelResponse chunk if applicable. @@ -71,7 +67,7 @@ class LangGraphSSEStreamIterator: return None - def _process_data(self, data) -> Optional[ModelResponseStream]: + def _process_data(self, data) -> ModelResponseStream | None: """ Process parsed data from SSE stream. @@ -101,7 +97,7 @@ class LangGraphSSEStreamIterator: return None - def _process_messages_event(self, payload) -> Optional[ModelResponseStream]: + def _process_messages_event(self, payload) -> ModelResponseStream | None: """ Process a messages event from the stream. @@ -116,9 +112,7 @@ class LangGraphSSEStreamIterator: content = msg.get("content", "") # Only return AI messages with content - if msg_type == "ai" and content: - return self._create_content_chunk(content) - elif msg_type == "AIMessageChunk" and content: + if msg_type == "ai" and content or msg_type == "AIMessageChunk" and content: return self._create_content_chunk(content) elif isinstance(item, dict): msg_type = item.get("type", "") @@ -128,7 +122,7 @@ class LangGraphSSEStreamIterator: return None - def _process_metadata_event(self, payload) -> Optional[ModelResponseStream]: + def _process_metadata_event(self, payload) -> ModelResponseStream | None: """ Process a metadata event, which may signal the end of the stream. """ @@ -202,7 +196,7 @@ class LangGraphSSEStreamIterator: except httpx.StreamClosed: raise StopIteration except Exception as e: - verbose_logger.error(f"Error in LangGraph SSE stream: {str(e)}") + verbose_logger.error(f"Error in LangGraph SSE stream: {e!s}") raise StopIteration async def __anext__(self) -> ModelResponseStream: @@ -230,5 +224,5 @@ class LangGraphSSEStreamIterator: except httpx.StreamClosed: raise StopAsyncIteration except Exception as e: - verbose_logger.error(f"Error in LangGraph SSE stream: {str(e)}") + verbose_logger.error(f"Error in LangGraph SSE stream: {e!s}") raise StopAsyncIteration diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index 77b5cfbc3fa..a40c08738f9 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -9,7 +9,7 @@ Non-streaming endpoint: POST /runs/wait """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Optional, Union, cast import httpx @@ -38,8 +38,6 @@ else: class LangGraphError(BaseLLMException): """Exception class for LangGraph API errors.""" - pass - class LangGraphConfig(BaseConfig): """ @@ -54,9 +52,9 @@ class LangGraphConfig(BaseConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: """ Get LangGraph API base and key from params or environment. @@ -71,7 +69,7 @@ class LangGraphConfig(BaseConfig): return api_base, api_key - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ LangGraph supports minimal OpenAI params since it's an agent runtime. """ @@ -91,12 +89,12 @@ class LangGraphConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the LangGraph request. @@ -135,7 +133,7 @@ class LangGraphConfig(BaseConfig): return parts[1] return model - def _convert_messages_to_langgraph_format(self, messages: List[AllMessageValues]) -> List[Dict[str, Any]]: + def _convert_messages_to_langgraph_format(self, messages: list[AllMessageValues]) -> list[dict[str, Any]]: """ Convert OpenAI-format messages to LangGraph format. @@ -144,7 +142,7 @@ class LangGraphConfig(BaseConfig): Preserves per-message ``metadata`` when present (e.g. A2A ``skillId``). """ - langgraph_messages: List[Dict[str, Any]] = [] + langgraph_messages: list[dict[str, Any]] = [] for msg in messages: role = msg.get("role", "user") content = msg.get("content", "") @@ -167,7 +165,7 @@ class LangGraphConfig(BaseConfig): if not isinstance(content, str): content = str(content) - langgraph_message: Dict[str, Any] = { + langgraph_message: dict[str, Any] = { "role": langgraph_role, "content": content, } @@ -182,7 +180,7 @@ class LangGraphConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -202,7 +200,7 @@ class LangGraphConfig(BaseConfig): assistant_id = self._get_assistant_id(model, optional_params) langgraph_messages = self._convert_messages_to_langgraph_format(messages) - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "assistant_id": assistant_id, "input": {"messages": langgraph_messages}, } @@ -283,9 +281,9 @@ class LangGraphConfig(BaseConfig): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, "AsyncHTTPHandler"]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for synchronous streaming. @@ -343,8 +341,8 @@ class LangGraphConfig(BaseConfig): data: dict, messages: list, client: Optional["AsyncHTTPHandler"] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for asynchronous streaming. @@ -412,12 +410,12 @@ class LangGraphConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the LangGraph response to LiteLLM ModelResponse format. @@ -453,14 +451,14 @@ class LangGraphConfig(BaseConfig): ) setattr(model_response, "usage", usage) except Exception as e: - verbose_logger.warning(f"Failed to calculate token usage: {str(e)}") + verbose_logger.warning(f"Failed to calculate token usage: {e!s}") return model_response except Exception as e: - verbose_logger.error(f"Error processing LangGraph response: {str(e)}") + verbose_logger.error(f"Error processing LangGraph response: {e!s}") raise LangGraphError( - message=f"Error processing response: {str(e)}", + message=f"Error processing response: {e!s}", status_code=raw_response.status_code, ) @@ -468,11 +466,11 @@ class LangGraphConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate and set up environment for LangGraph requests. @@ -485,16 +483,14 @@ class LangGraphConfig(BaseConfig): return headers - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return LangGraphError(status_code=status_code, message=error_message) def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ LangGraph has native streaming support, so we don't need to fake stream. diff --git a/litellm/llms/lemonade/chat/transformation.py b/litellm/llms/lemonade/chat/transformation.py index f10dbf49f66..5c27c5152c2 100644 --- a/litellm/llms/lemonade/chat/transformation.py +++ b/litellm/llms/lemonade/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to Lemonade's `/v1/chat/completions` """ -from typing import Any, List, Optional, Tuple, Union +from typing import Any from urllib.parse import quote import httpx @@ -22,35 +22,35 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig class LemonadeChatConfig(OpenAILikeChatConfig): _DEFAULT_API_KEY = "lemonade" - repeat_penalty: Optional[float] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - max_completion_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - response_format: Optional[dict] = None - tools: Optional[list] = None + repeat_penalty: float | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + max_completion_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + response_format: dict | None = None + tools: list | None = None def __init__( self, - repeat_penalty: Optional[float] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_completion_tokens: Optional[int] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - response_format: Optional[dict] = None, - tools: Optional[list] = None, + repeat_penalty: float | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_completion_tokens: int | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + response_format: dict | None = None, + tools: list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -58,14 +58,14 @@ class LemonadeChatConfig(OpenAILikeChatConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "lemonade" @classmethod def get_config(cls): return super().get_config() - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None): + def get_models(self, api_key: str | None = None, api_base: str | None = None): """ Get available models from Lemonade API. @@ -105,7 +105,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig): return ["lemonade/" + model["id"] for model in model_list] @staticmethod - def _get_positive_int(value: Any) -> Optional[int]: + def _get_positive_int(value: Any) -> int | None: if isinstance(value, bool): return None if isinstance(value, int) and value > 0: @@ -133,7 +133,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig): return provider_specific_entry - def _get_context_window(self, model_info: dict) -> Optional[int]: + def _get_context_window(self, model_info: dict) -> int | None: provider_specific_entry = self._get_provider_specific_entry(model_info) recipe_options = provider_specific_entry.get("recipe_options") if not isinstance(recipe_options, dict): @@ -165,8 +165,8 @@ class LemonadeChatConfig(OpenAILikeChatConfig): def get_model_info( self, model: str, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, ) -> Any: if model.startswith("lemonade/"): model = model.split("/", 1)[1] @@ -203,8 +203,8 @@ class LemonadeChatConfig(OpenAILikeChatConfig): return model_info_response def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # lemonade is openai compatible, we just need to set this to custom_openai and have the api_base be lemonade's endpoint passed_api_base = api_base api_base = api_base or get_secret_str("LEMONADE_API_BASE") or "http://localhost:8000/api/v1" # type: ignore @@ -213,7 +213,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig): key = api_key or litellm.lemonade_key or get_secret_str("LEMONADE_API_KEY") or self._DEFAULT_API_KEY return api_base, key - def _get_auth_headers(self, api_key: Optional[str]) -> dict: + def _get_auth_headers(self, api_key: str | None) -> dict: if api_key is None or api_key == self._DEFAULT_API_KEY: return {} return {"Authorization": f"Bearer {api_key}"} @@ -225,12 +225,12 @@ class LemonadeChatConfig(OpenAILikeChatConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: model_response = super().transform_response( model=model, diff --git a/litellm/llms/lemonade/cost_calculator.py b/litellm/llms/lemonade/cost_calculator.py index 74d62da8759..f2a3025c99d 100644 --- a/litellm/llms/lemonade/cost_calculator.py +++ b/litellm/llms/lemonade/cost_calculator.py @@ -5,15 +5,13 @@ Since Lemonade is a local/self-hosted service, all costs default to 0. This prevents cost calculation errors when using models not in model_prices_and_context_window.json """ -from typing import Tuple - from litellm.types.utils import Usage def cost_per_token( model: str, usage: Usage, -) -> Tuple[float, float]: +) -> tuple[float, float]: """ Calculate cost per token for Lemonade models. diff --git a/litellm/llms/linkup/search/transformation.py b/litellm/llms/linkup/search/transformation.py index a68231fa867..942f648daaa 100644 --- a/litellm/llms/linkup/search/transformation.py +++ b/litellm/llms/linkup/search/transformation.py @@ -4,7 +4,7 @@ Calls Linkup's /search endpoint to search the web. Linkup API Reference: https://docs.linkup.so/pages/documentation/api-reference/endpoint/post-search """ -from typing import Dict, List, Literal, Optional, TypedDict, Union +from typing import Literal, TypedDict import httpx @@ -36,8 +36,8 @@ class LinkupSearchRequest(_LinkupSearchRequestRequired, total=False): includeImages: bool # Optional - Include images in results (default false) fromDate: str # Optional - Start date for results (YYYY-MM-DD) toDate: str # Optional - End date for results (YYYY-MM-DD) - includeDomains: List[str] # Optional - Domains to search on (max 100) - excludeDomains: List[str] # Optional - Domains to exclude + includeDomains: list[str] # Optional - Domains to search on (max 100) + excludeDomains: list[str] # Optional - Domains to exclude includeInlineCitations: bool # Optional - Include inline citations (default false) maxResults: int # Optional - Maximum number of results to return @@ -51,11 +51,11 @@ class LinkupSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -74,9 +74,9 @@ class LinkupSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -92,10 +92,10 @@ class LinkupSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Linkup API format. diff --git a/litellm/llms/litellm_proxy/chat/transformation.py b/litellm/llms/litellm_proxy/chat/transformation.py index eee0ec6fa08..afb4be85d85 100644 --- a/litellm/llms/litellm_proxy/chat/transformation.py +++ b/litellm/llms/litellm_proxy/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions` """ -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS from litellm.secret_managers.main import get_secret_bool, get_secret_str @@ -15,7 +15,7 @@ if TYPE_CHECKING: class LiteLLMProxyChatConfig(OpenAIGPTConfig): - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: params_list = super().get_supported_openai_params(model) params_list.extend(OPENAI_CHAT_COMPLETION_PARAMS) return params_list @@ -36,13 +36,13 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): return optional_params def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("LITELLM_PROXY_API_BASE") # type: ignore dynamic_api_key = api_key or get_secret_str("LITELLM_PROXY_API_KEY") return api_base, dynamic_api_key - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base, api_key = self._get_openai_compatible_provider_info(api_base, api_key) if api_base is None: raise ValueError("api_base not set for LiteLLM Proxy route. Set in env via `LITELLM_PROXY_API_BASE`") @@ -50,12 +50,12 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): return [f"litellm_proxy/{model}" for model in models] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("LITELLM_PROXY_API_KEY") @staticmethod def _should_use_litellm_proxy_by_default( - litellm_params: Optional[LiteLLM_Params] = None, + litellm_params: LiteLLM_Params | None = None, ): """ Returns True if litellm proxy should be used by default for a given request @@ -79,8 +79,8 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): @staticmethod def litellm_proxy_get_custom_llm_provider_info( - model: str, api_base: Optional[str] = None, api_key: Optional[str] = None - ) -> Tuple[str, str, Optional[str], Optional[str]]: + model: str, api_base: str | None = None, api_key: str | None = None + ) -> tuple[str, str, str | None, str | None]: """ Force use litellm proxy for all models @@ -114,7 +114,7 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, headers: dict, @@ -129,7 +129,7 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): async def async_transform_request( self, model: str, - messages: List["AllMessageValues"], + messages: list["AllMessageValues"], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/litellm_proxy/image_edit/transformation.py b/litellm/llms/litellm_proxy/image_edit/transformation.py index 94825cffeae..9e5a2a2c343 100644 --- a/litellm/llms/litellm_proxy/image_edit/transformation.py +++ b/litellm/llms/litellm_proxy/image_edit/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig from litellm.secret_managers.main import get_secret_str @@ -11,15 +9,15 @@ class LiteLLMProxyImageEditConfig(OpenAIImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("LITELLM_PROXY_API_KEY") headers.update({"Authorization": f"Bearer {api_key}"}) return headers - def get_complete_url(self, model: str, api_base: Optional[str], litellm_params: dict) -> str: + def get_complete_url(self, model: str, api_base: str | None, litellm_params: dict) -> str: api_base = api_base or get_secret_str("LITELLM_PROXY_API_BASE") if api_base is None: raise ValueError("api_base not set for LiteLLM Proxy route. Set in env via `LITELLM_PROXY_API_BASE`") diff --git a/litellm/llms/litellm_proxy/image_generation/transformation.py b/litellm/llms/litellm_proxy/image_generation/transformation.py index 5fad663d126..fcddeabfc85 100644 --- a/litellm/llms/litellm_proxy/image_generation/transformation.py +++ b/litellm/llms/litellm_proxy/image_generation/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.llms.openai.image_generation.gpt_transformation import ( GPTImageGenerationConfig, ) @@ -16,8 +14,8 @@ class LiteLLMProxyImageGenerationConfig(GPTImageGenerationConfig): messages, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("LITELLM_PROXY_API_KEY") headers.update({"Authorization": f"Bearer {api_key}"}) @@ -25,12 +23,12 @@ class LiteLLMProxyImageGenerationConfig(GPTImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = api_base or get_secret_str("LITELLM_PROXY_API_BASE") if api_base is None: diff --git a/litellm/llms/litellm_proxy/responses/transformation.py b/litellm/llms/litellm_proxy/responses/transformation.py index e5bbaa78d1d..2928b07665a 100644 --- a/litellm/llms/litellm_proxy/responses/transformation.py +++ b/litellm/llms/litellm_proxy/responses/transformation.py @@ -5,8 +5,6 @@ LiteLLM Proxy supports the OpenAI Responses API natively when the underlying mod This config enables pass-through behavior to the proxy's /v1/responses endpoint. """ -from typing import Optional - from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.utils import LlmProviders @@ -26,7 +24,7 @@ class LiteLLMProxyResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/litellm_proxy/skills/__init__.py b/litellm/llms/litellm_proxy/skills/__init__.py index 5fb29e96bb9..6267dfebbdc 100644 --- a/litellm/llms/litellm_proxy/skills/__init__.py +++ b/litellm/llms/litellm_proxy/skills/__init__.py @@ -38,17 +38,17 @@ from litellm.llms.litellm_proxy.skills.transformation import ( ) __all__ = [ + "DEFAULT_MAX_ITERATIONS", + "DEFAULT_SANDBOX_TIMEOUT", + "LITELLM_CODE_EXECUTION_TOOL", + "CodeExecutionHandler", + "LiteLLMInternalTools", "LiteLLMSkillsHandler", "LiteLLMSkillsTransformationHandler", "SkillPromptInjectionHandler", "SkillsSandboxExecutor", - "CodeExecutionHandler", - "LiteLLMInternalTools", - "LITELLM_CODE_EXECUTION_TOOL", - "get_litellm_code_execution_tool", - "code_execution_handler", - "has_code_execution_tool", "add_code_execution_tool", - "DEFAULT_MAX_ITERATIONS", - "DEFAULT_SANDBOX_TIMEOUT", + "code_execution_handler", + "get_litellm_code_execution_tool", + "has_code_execution_tool", ] diff --git a/litellm/llms/litellm_proxy/skills/code_execution.py b/litellm/llms/litellm_proxy/skills/code_execution.py index 4ac3311921d..c99698a5c8e 100644 --- a/litellm/llms/litellm_proxy/skills/code_execution.py +++ b/litellm/llms/litellm_proxy/skills/code_execution.py @@ -14,7 +14,7 @@ Generated files are returned directly in the response - no separate storage need import base64 import json from enum import Enum -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger @@ -30,7 +30,7 @@ class LiteLLMInternalTools(str, Enum): CODE_EXECUTION = "litellm_code_execution" -def get_litellm_code_execution_tool() -> Dict[str, Any]: +def get_litellm_code_execution_tool() -> dict[str, Any]: """ Returns the litellm_code_execution tool definition in OpenAI format. @@ -51,7 +51,7 @@ def get_litellm_code_execution_tool() -> Dict[str, Any]: } -def get_litellm_code_execution_tool_anthropic() -> Dict[str, Any]: +def get_litellm_code_execution_tool_anthropic() -> dict[str, Any]: """ Returns the litellm_code_execution tool definition in Anthropic/messages API format. @@ -84,8 +84,8 @@ class CodeExecutionHandler: def __init__( self, - max_iterations: Optional[int] = None, - sandbox_timeout: Optional[int] = None, + max_iterations: int | None = None, + sandbox_timeout: int | None = None, ): from litellm.llms.litellm_proxy.skills.constants import ( DEFAULT_MAX_ITERATIONS, @@ -98,12 +98,12 @@ class CodeExecutionHandler: async def execute_with_code_execution( self, model: str, - messages: List[Dict], - tools: List[Dict], - skill_files: Dict[str, bytes], - skill_id: Optional[str] = None, + messages: list[dict], + tools: list[dict], + skill_files: dict[str, bytes], + skill_id: str | None = None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Execute an LLM call with automatic code execution handling. @@ -134,8 +134,8 @@ class CodeExecutionHandler: ) current_messages = list(messages) - generated_files: List[Dict[str, Any]] = [] # Files returned directly - execution_results: List[Dict] = [] + generated_files: list[dict[str, Any]] = [] # Files returned directly + execution_results: list[dict] = [] executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) response: Any = None # Initialize to avoid possibly unbound error @@ -155,7 +155,7 @@ class CodeExecutionHandler: stop_reason = response.choices[0].finish_reason # type: ignore # Build assistant message for conversation history - assistant_msg_dict: Dict[str, Any] = { + assistant_msg_dict: dict[str, Any] = { "role": "assistant", "content": assistant_message.content, } @@ -239,7 +239,7 @@ class CodeExecutionHandler: tool_result += f"\n\nError:\n{exec_result['error']}" except Exception as e: - tool_result = f"Code execution failed: {str(e)}" + tool_result = f"Code execution failed: {e!s}" execution_results.append( { "iteration": iteration, @@ -278,7 +278,7 @@ class CodeExecutionHandler: } -def has_code_execution_tool(tools: Optional[List[Dict]]) -> bool: +def has_code_execution_tool(tools: list[dict] | None) -> bool: """Check if litellm_code_execution tool is in the tools list.""" if not tools: return False @@ -289,7 +289,7 @@ def has_code_execution_tool(tools: Optional[List[Dict]]) -> bool: return False -def add_code_execution_tool(tools: Optional[List[Dict]]) -> List[Dict]: +def add_code_execution_tool(tools: list[dict] | None) -> list[dict]: """Add litellm_code_execution tool if not already present.""" tools = tools or [] if not has_code_execution_tool(tools): diff --git a/litellm/llms/litellm_proxy/skills/handler.py b/litellm/llms/litellm_proxy/skills/handler.py index 6f5ae261d2e..cc307917af4 100644 --- a/litellm/llms/litellm_proxy/skills/handler.py +++ b/litellm/llms/litellm_proxy/skills/handler.py @@ -6,7 +6,7 @@ Used by the transformation layer and skills injection hook. """ import uuid -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache @@ -61,8 +61,8 @@ class LiteLLMSkillsHandler: @staticmethod async def create_skill( data: NewSkillRequest, - user_id: Optional[str] = None, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_id: str | None = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ) -> LiteLLM_SkillsTable: prisma_client = await LiteLLMSkillsHandler._get_prisma_client() @@ -76,7 +76,7 @@ class LiteLLMSkillsHandler: # this module FastAPI-free per the project layering rule. raise ValueError("Unable to record skill ownership: caller has no identity scope.") - skill_data: Dict[str, Any] = { + skill_data: dict[str, Any] = { "skill_id": skill_id, "display_title": data.display_title, "description": data.description, @@ -109,13 +109,13 @@ class LiteLLMSkillsHandler: async def list_skills( limit: int = 20, offset: int = 0, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - ) -> List[LiteLLM_SkillsTable]: + user_api_key_dict: UserAPIKeyAuth | None = None, + ) -> list[LiteLLM_SkillsTable]: prisma_client = await LiteLLMSkillsHandler._get_prisma_client() verbose_logger.debug(f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}") - find_many_kwargs: Dict[str, Any] = { + find_many_kwargs: dict[str, Any] = { "take": limit, "skip": offset, "order": {"created_at": "desc"}, @@ -130,7 +130,7 @@ class LiteLLMSkillsHandler: return [_prisma_skill_to_litellm(s) for s in skills] @staticmethod - async def _load_skill(skill_id: str) -> Optional[Any]: + async def _load_skill(skill_id: str) -> Any | None: """Cache-first read of the Prisma skill row. Owner-scope filtering happens on the cached row, so the cache is per-skill not per-caller. """ @@ -148,7 +148,7 @@ class LiteLLMSkillsHandler: @staticmethod async def get_skill( skill_id: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ) -> LiteLLM_SkillsTable: verbose_logger.debug(f"LiteLLMSkillsHandler: Getting skill {skill_id}") @@ -163,8 +163,8 @@ class LiteLLMSkillsHandler: @staticmethod async def delete_skill( skill_id: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - ) -> Dict[str, str]: + user_api_key_dict: UserAPIKeyAuth | None = None, + ) -> dict[str, str]: prisma_client = await LiteLLMSkillsHandler._get_prisma_client() verbose_logger.debug(f"LiteLLMSkillsHandler: Deleting skill {skill_id}") @@ -180,8 +180,8 @@ class LiteLLMSkillsHandler: @staticmethod async def fetch_skill_from_db( skill_id: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - ) -> Optional[LiteLLM_SkillsTable]: + user_api_key_dict: UserAPIKeyAuth | None = None, + ) -> LiteLLM_SkillsTable | None: """Skills-injection-hook helper: returns None instead of raising on not-found / not-authorized so the hook can silently skip.""" try: diff --git a/litellm/llms/litellm_proxy/skills/prompt_injection.py b/litellm/llms/litellm_proxy/skills/prompt_injection.py index 8be6f105845..244d4196404 100644 --- a/litellm/llms/litellm_proxy/skills/prompt_injection.py +++ b/litellm/llms/litellm_proxy/skills/prompt_injection.py @@ -8,7 +8,7 @@ and injection into the system prompt for non-Anthropic models. import posixpath import zipfile from io import BytesIO -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.proxy._types import LiteLLM_SkillsTable @@ -25,7 +25,7 @@ class SkillPromptInjectionHandler: - Create execute_code tool definition """ - def extract_skill_content(self, skill: LiteLLM_SkillsTable) -> Optional[str]: + def extract_skill_content(self, skill: LiteLLM_SkillsTable) -> str | None: """ Extract skill content from the stored zip file. @@ -71,7 +71,7 @@ class SkillPromptInjectionHandler: return skill.instructions - def extract_all_files(self, skill: LiteLLM_SkillsTable) -> Dict[str, bytes]: + def extract_all_files(self, skill: LiteLLM_SkillsTable) -> dict[str, bytes]: """ Extract ALL files from skill ZIP for code execution. @@ -84,7 +84,7 @@ class SkillPromptInjectionHandler: Returns: Dict mapping file paths to binary content """ - files: Dict[str, bytes] = {} + files: dict[str, bytes] = {} if not skill.file_content: return files @@ -124,7 +124,7 @@ class SkillPromptInjectionHandler: return files def inject_skill_content_to_messages( - self, data: dict, skill_contents: List[str], use_anthropic_format: bool = False + self, data: dict, skill_contents: list[str], use_anthropic_format: bool = False ) -> dict: """ Inject skill content into the system prompt. @@ -181,7 +181,7 @@ class SkillPromptInjectionHandler: data["messages"] = messages return data - def create_execute_code_tool(self, skill_modules: List[str]) -> Dict[str, Any]: + def create_execute_code_tool(self, skill_modules: list[str]) -> dict[str, Any]: """ Create the execute_code tool definition. @@ -224,7 +224,7 @@ class SkillPromptInjectionHandler: }, } - def convert_skill_to_tool(self, skill: LiteLLM_SkillsTable) -> Dict[str, Any]: + def convert_skill_to_tool(self, skill: LiteLLM_SkillsTable) -> dict[str, Any]: """ Convert a LiteLLM skill to an OpenAI-style tool. @@ -248,7 +248,7 @@ class SkillPromptInjectionHandler: if len(description) > max_desc_length: description = description[: max_desc_length - 3] + "..." - tool: Dict[str, Any] = { + tool: dict[str, Any] = { "type": "function", "function": { "name": func_name, @@ -269,7 +269,7 @@ class SkillPromptInjectionHandler: return tool - def convert_skill_to_anthropic_tool(self, skill: LiteLLM_SkillsTable) -> Dict[str, Any]: + def convert_skill_to_anthropic_tool(self, skill: LiteLLM_SkillsTable) -> dict[str, Any]: """ Convert a LiteLLM skill to an Anthropic-style tool (messages API format). @@ -287,7 +287,7 @@ class SkillPromptInjectionHandler: if len(description) > max_desc_length: description = description[: max_desc_length - 3] + "..." - input_schema: Dict[str, Any] = { + input_schema: dict[str, Any] = { "type": "object", "properties": {}, "required": [], diff --git a/litellm/llms/litellm_proxy/skills/sandbox_executor.py b/litellm/llms/litellm_proxy/skills/sandbox_executor.py index 5f1f129032c..e79b0c948c4 100644 --- a/litellm/llms/litellm_proxy/skills/sandbox_executor.py +++ b/litellm/llms/litellm_proxy/skills/sandbox_executor.py @@ -7,7 +7,7 @@ Supports Docker, Podman, and Kubernetes backends. import base64 import os -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger @@ -27,7 +27,7 @@ class SkillsSandboxExecutor: self, timeout: int = 60, backend: str = "docker", - image: Optional[str] = None, + image: str | None = None, ): """ Initialize the sandbox executor. @@ -45,9 +45,9 @@ class SkillsSandboxExecutor: def execute( self, code: str, - skill_files: Dict[str, bytes], - requirements: Optional[str] = None, - ) -> Dict[str, Any]: + skill_files: dict[str, bytes], + requirements: str | None = None, + ) -> dict[str, Any]: """ Execute code with skill files in sandbox. @@ -77,7 +77,7 @@ class SkillsSandboxExecutor: try: # Create sandbox session - session_kwargs: Dict[str, Any] = { + session_kwargs: dict[str, Any] = { "lang": "python", "verbose": False, } @@ -112,7 +112,7 @@ class SkillsSandboxExecutor: # requirements file inside the sandbox so standard syntax like # `-r`, `-e`, VCS URLs, and inline `#egg=` fragments continue to # work. - requirements_filename: Optional[str] = None + requirements_filename: str | None = None if requirements: with tempfile.NamedTemporaryFile( mode="w", @@ -198,8 +198,8 @@ sys.path.insert(0, '/sandbox') def _collect_generated_files( self, session: Any, - original_files: Dict[str, bytes], - ) -> List[Dict[str, Any]]: + original_files: dict[str, bytes], + ) -> list[dict[str, Any]]: """ Collect files generated during execution. @@ -213,7 +213,7 @@ sys.path.insert(0, '/sandbox') Returns: List of generated files with base64 content """ - generated_files: List[Dict[str, Any]] = [] + generated_files: list[dict[str, Any]] = [] try: import tempfile diff --git a/litellm/llms/litellm_proxy/skills/transformation.py b/litellm/llms/litellm_proxy/skills/transformation.py index eb1ac290807..56479b8bda2 100644 --- a/litellm/llms/litellm_proxy/skills/transformation.py +++ b/litellm/llms/litellm_proxy/skills/transformation.py @@ -8,7 +8,7 @@ Pattern follows litellm/llms/litellm_proxy/responses/transformation.py """ from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Optional from litellm.types.llms.anthropic_skills import ( DeleteSkillResponse, @@ -37,21 +37,21 @@ class LiteLLMSkillsTransformationHandler: def create_skill_handler( self, - display_title: Optional[str] = None, - description: Optional[str] = None, - instructions: Optional[str] = None, - files: Optional[List[Any]] = None, - file_content: Optional[bytes] = None, - file_name: Optional[str] = None, - file_type: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - user_id: Optional[str] = None, + display_title: str | None = None, + description: str | None = None, + instructions: str | None = None, + files: list[Any] | None = None, + file_content: bytes | None = None, + file_name: str | None = None, + file_type: str | None = None, + metadata: dict[str, Any] | None = None, + user_id: str | None = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, _is_async: bool = False, logging_obj: Optional["LiteLLMLoggingObj"] = None, - litellm_call_id: Optional[str] = None, + litellm_call_id: str | None = None, **kwargs, - ) -> Union[Skill, Coroutine[Any, Any, Skill]]: + ) -> Skill | Coroutine[Any, Any, Skill]: """ Create a skill in LiteLLM database. @@ -121,14 +121,14 @@ class LiteLLMSkillsTransformationHandler: async def _async_create_skill( self, - display_title: Optional[str] = None, - description: Optional[str] = None, - instructions: Optional[str] = None, - file_content: Optional[bytes] = None, - file_name: Optional[str] = None, - file_type: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - user_id: Optional[str] = None, + display_title: str | None = None, + description: str | None = None, + instructions: str | None = None, + file_content: bytes | None = None, + file_name: str | None = None, + file_type: str | None = None, + metadata: dict[str, Any] | None = None, + user_id: str | None = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, ) -> Skill: """Async implementation of create_skill.""" @@ -160,10 +160,10 @@ class LiteLLMSkillsTransformationHandler: offset: int = 0, _is_async: bool = False, logging_obj: Optional["LiteLLMLoggingObj"] = None, - litellm_call_id: Optional[str] = None, + litellm_call_id: str | None = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, **kwargs, - ) -> Union[ListSkillsResponse, Coroutine[Any, Any, ListSkillsResponse]]: + ) -> ListSkillsResponse | Coroutine[Any, Any, ListSkillsResponse]: """ List skills from LiteLLM database. @@ -232,10 +232,10 @@ class LiteLLMSkillsTransformationHandler: skill_id: str, _is_async: bool = False, logging_obj: Optional["LiteLLMLoggingObj"] = None, - litellm_call_id: Optional[str] = None, + litellm_call_id: str | None = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, **kwargs, - ) -> Union[Skill, Coroutine[Any, Any, Skill]]: + ) -> Skill | Coroutine[Any, Any, Skill]: """ Get a skill from LiteLLM database. @@ -293,10 +293,10 @@ class LiteLLMSkillsTransformationHandler: skill_id: str, _is_async: bool = False, logging_obj: Optional["LiteLLMLoggingObj"] = None, - litellm_call_id: Optional[str] = None, + litellm_call_id: str | None = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, **kwargs, - ) -> Union[DeleteSkillResponse, Coroutine[Any, Any, DeleteSkillResponse]]: + ) -> DeleteSkillResponse | Coroutine[Any, Any, DeleteSkillResponse]: """ Delete a skill from LiteLLM database. diff --git a/litellm/llms/llamafile/chat/transformation.py b/litellm/llms/llamafile/chat/transformation.py index 78cc58708ad..5fee8c588d2 100644 --- a/litellm/llms/llamafile/chat/transformation.py +++ b/litellm/llms/llamafile/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -9,7 +7,7 @@ class LlamafileChatConfig(OpenAIGPTConfig): """LlamafileChatConfig is used to provide configuration for the LlamaFile's chat API.""" @staticmethod - def _resolve_api_key(api_key: Optional[str] = None) -> str: + def _resolve_api_key(api_key: str | None = None) -> str: """Attempt to ensure that the API key is set, preferring the user-provided key over the secret manager key (``LLAMAFILE_API_KEY``). @@ -18,7 +16,7 @@ class LlamafileChatConfig(OpenAIGPTConfig): return api_key or get_secret_str("LLAMAFILE_API_KEY") or "fake-api-key" # llamafile does not require an API key @staticmethod - def _resolve_api_base(api_base: Optional[str] = None) -> Optional[str]: + def _resolve_api_base(api_base: str | None = None) -> str | None: """Attempt to ensure that the API base is set, preferring the user-provided key over the secret manager key (``LLAMAFILE_API_BASE``). @@ -28,8 +26,8 @@ class LlamafileChatConfig(OpenAIGPTConfig): return api_base or get_secret_str("LLAMAFILE_API_BASE") or "http://127.0.0.1:8080/v1" # type: ignore def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: """Attempts to ensure that the API base and key are set, preferring user-provided values, before falling back to secret manager values (``LLAMAFILE_API_BASE`` and ``LLAMAFILE_API_KEY`` respectively). diff --git a/litellm/llms/lm_studio/chat/transformation.py b/litellm/llms/lm_studio/chat/transformation.py index 64ed38467de..a40d6122ecd 100644 --- a/litellm/llms/lm_studio/chat/transformation.py +++ b/litellm/llms/lm_studio/chat/transformation.py @@ -2,8 +2,6 @@ Translate from OpenAI's `/v1/chat/completions` to LM Studio's `/chat/completions` """ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -11,8 +9,8 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class LMStudioChatConfig(OpenAIGPTConfig): def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("LM_STUDIO_API_BASE") # type: ignore dynamic_api_key = ( api_key or get_secret_str("LM_STUDIO_API_KEY") or "fake-api-key" diff --git a/litellm/llms/lm_studio/embed/transformation.py b/litellm/llms/lm_studio/embed/transformation.py index f0357b9428c..8c78be2c909 100644 --- a/litellm/llms/lm_studio/embed/transformation.py +++ b/litellm/llms/lm_studio/embed/transformation.py @@ -7,7 +7,6 @@ Docs - https://lmstudio.ai/docs/basics/server """ import types -from typing import List class LmStudioEmbeddingConfig: @@ -41,7 +40,7 @@ class LmStudioEmbeddingConfig: and v is not None } - def get_supported_openai_params(self) -> List[str]: + def get_supported_openai_params(self) -> list[str]: return [] def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict: diff --git a/litellm/llms/manus/files/transformation.py b/litellm/llms/manus/files/transformation.py index 4a65fac709b..cfa6d1cc722 100644 --- a/litellm/llms/manus/files/transformation.py +++ b/litellm/llms/manus/files/transformation.py @@ -11,15 +11,15 @@ Reference: https://open.manus.im/docs/openai-compatibility#file-management """ import time -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx from openai.types.file_deleted import FileDeleted import litellm from litellm._logging import verbose_logger -from litellm.litellm_core_utils.url_utils import encode_url_path_segment 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 from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, @@ -66,8 +66,8 @@ class ManusFilesConfig(BaseFilesConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Manus API. @@ -92,7 +92,7 @@ class ManusFilesConfig(BaseFilesConfig): ) return headers - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: """ Return supported OpenAI file creation parameters for Manus. Manus supports the standard 'purpose' parameter. @@ -114,12 +114,12 @@ class ManusFilesConfig(BaseFilesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Manus Files API endpoint. @@ -141,7 +141,7 @@ class ManusFilesConfig(BaseFilesConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: """ Return the appropriate error class for Manus API errors. @@ -216,7 +216,7 @@ class ManusFilesConfig(BaseFilesConfig): def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -279,8 +279,8 @@ class ManusFilesConfig(BaseFilesConfig): status_details=response_json.get("status_details"), ) except Exception as e: - verbose_logger.exception(f"Error parsing Manus file response: {str(e)}") - raise ValueError(f"Error parsing Manus file response: {str(e)}") + verbose_logger.exception(f"Error parsing Manus file response: {e!s}") + raise ValueError(f"Error parsing Manus file response: {e!s}") def transform_retrieve_file_request( self, @@ -342,7 +342,7 @@ class ManusFilesConfig(BaseFilesConfig): def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: @@ -364,13 +364,13 @@ class ManusFilesConfig(BaseFilesConfig): raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: """Transform list files response.""" response_json = raw_response.json() files_data = 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: dict[str, Any]) -> OpenAIFileObject: """Parse a file dict into OpenAIFileObject.""" created_at_str = file_dict.get("created_at", "") if created_at_str: diff --git a/litellm/llms/manus/responses/transformation.py b/litellm/llms/manus/responses/transformation.py index e6fbbac0563..25d4d0b8db6 100644 --- a/litellm/llms/manus/responses/transformation.py +++ b/litellm/llms/manus/responses/transformation.py @@ -1,15 +1,15 @@ import uuid -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers -from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.common_utils import OpenAIError from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str @@ -49,9 +49,9 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Manus API doesn't support real-time streaming. @@ -75,7 +75,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): # If no slash, assume the model name itself is the agent profile return model - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate environment and set up headers for Manus API. @@ -101,7 +101,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -123,11 +123,11 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """ Transform the request for Manus API. @@ -237,7 +237,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the get response API request into a URL and data. @@ -248,7 +248,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): """ encoded_response_id = encode_url_path_segment(response_id, field_name="response_id") url = f"{api_base}/{encoded_response_id}" - data: Dict = {} + data: dict = {} return url, data def transform_get_response_api_response( diff --git a/litellm/llms/maritalk.py b/litellm/llms/maritalk.py index 4b3a569357f..08300bbd16e 100644 --- a/litellm/llms/maritalk.py +++ b/litellm/llms/maritalk.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - from httpx._models import Headers from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -11,7 +9,7 @@ class MaritalkError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, Headers]] = None, + headers: dict | Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) @@ -19,18 +17,18 @@ class MaritalkError(BaseLLMException): class MaritalkConfig(OpenAIGPTConfig): def __init__( self, - frequency_penalty: Optional[float] = None, - presence_penalty: Optional[float] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - temperature: Optional[float] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - stop: Optional[List[str]] = None, - stream: Optional[bool] = None, - stream_options: Optional[dict] = None, - tools: Optional[List[dict]] = None, - tool_choice: Optional[Union[str, dict]] = None, + frequency_penalty: float | None = None, + presence_penalty: float | None = None, + top_p: float | None = None, + top_k: int | None = None, + temperature: float | None = None, + max_tokens: int | None = None, + n: int | None = None, + stop: list[str] | None = None, + stream: bool | None = None, + stream_options: dict | None = None, + tools: list[dict] | None = None, + tool_choice: str | dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -41,7 +39,7 @@ class MaritalkConfig(OpenAIGPTConfig): def get_config(cls): return super().get_config() - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "frequency_penalty", "presence_penalty", @@ -57,5 +55,5 @@ class MaritalkConfig(OpenAIGPTConfig): "tool_choice", ] - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return MaritalkError(status_code=status_code, message=error_message, headers=headers) diff --git a/litellm/llms/milvus/vector_stores/transformation.py b/litellm/llms/milvus/vector_stores/transformation.py index a53075ba1d6..48265d095a8 100644 --- a/litellm/llms/milvus/vector_stores/transformation.py +++ b/litellm/llms/milvus/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -47,8 +47,8 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): def __init__(self): super().__init__() - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: - api_key: Optional[str] = None + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: + api_key: str | None = None if litellm_params is not None: api_key = litellm_params.api_key or get_secret_str("MILVUS_API_KEY") @@ -94,7 +94,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -117,13 +117,13 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict[str, Any]]: """ Transform search request for Azure AI Search API @@ -158,14 +158,14 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): ) query_vector = embedding_response.data[0]["embedding"] except Exception as e: - raise Exception(f"Failed to generate embedding for query: {str(e)}") + raise Exception(f"Failed to generate embedding for query: {e!s}") # Azure AI Search endpoint for search index_name = vector_store_id # vector_store_id is the index name url = f"{api_base}/v2/vectordb/entities/search" # Build the request body for Azure AI Search with vector search - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "collectionName": index_name, "data": [query_vector], "annsField": "book_intro_vector", @@ -223,7 +223,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): ) # Transform results to standard format - search_results: List[VectorStoreSearchResult] = [] + search_results: list[VectorStoreSearchResult] = [] for result in results: # Extract text content text_content = result.get(text_field, "") @@ -271,7 +271,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: raise NotImplementedError def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: diff --git a/litellm/llms/minimax/__init__.py b/litellm/llms/minimax/__init__.py index e1b0e602e92..db884c27b99 100644 --- a/litellm/llms/minimax/__init__.py +++ b/litellm/llms/minimax/__init__.py @@ -8,6 +8,6 @@ from .text_to_speech.transformation import ( ) __all__ = [ - "MinimaxTextToSpeechConfig", "MinimaxException", + "MinimaxTextToSpeechConfig", ] diff --git a/litellm/llms/minimax/chat/transformation.py b/litellm/llms/minimax/chat/transformation.py index 512c162658c..df838d94b7d 100644 --- a/litellm/llms/minimax/chat/transformation.py +++ b/litellm/llms/minimax/chat/transformation.py @@ -2,8 +2,6 @@ MiniMax OpenAI transformation config - extends OpenAI chat config for MiniMax's OpenAI-compatible API """ -from typing import List, Optional, Tuple - import litellm from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.secret_managers.main import get_secret_str @@ -24,7 +22,7 @@ class MinimaxChatConfig(OpenAIGPTConfig): """ @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: """ Get MiniMax API key from environment or parameters. """ @@ -32,7 +30,7 @@ class MinimaxChatConfig(OpenAIGPTConfig): @staticmethod def get_api_base( - api_base: Optional[str] = None, + api_base: str | None = None, ) -> str: """ Get MiniMax API base URL. @@ -43,12 +41,12 @@ class MinimaxChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for MiniMax OpenAI API. @@ -70,9 +68,9 @@ class MinimaxChatConfig(OpenAIGPTConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, - messages: List[AllMessageValues], - tools: Optional[List[ChatCompletionToolParam]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List[ChatCompletionToolParam]]]: + messages: list[AllMessageValues], + tools: list[ChatCompletionToolParam] | None = None, + ) -> tuple[list[AllMessageValues], list[ChatCompletionToolParam] | None]: """ Override to preserve cache_control for MiniMax. MiniMax supports cache_control - don't strip it. diff --git a/litellm/llms/minimax/messages/transformation.py b/litellm/llms/minimax/messages/transformation.py index 3f46aae1aaa..4d6cbc4a4d7 100644 --- a/litellm/llms/minimax/messages/transformation.py +++ b/litellm/llms/minimax/messages/transformation.py @@ -2,8 +2,6 @@ MiniMax Anthropic transformation config - extends AnthropicConfig for MiniMax's Anthropic-compatible API """ -from typing import Optional - import litellm from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, @@ -25,14 +23,14 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "minimax" def should_strip_billing_metadata(self) -> bool: return True @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: """ Get MiniMax API key from environment or parameters. """ @@ -40,7 +38,7 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig): @staticmethod def get_api_base( - api_base: Optional[str] = None, + api_base: str | None = None, ) -> str: """ Get MiniMax API base URL. @@ -51,12 +49,12 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for MiniMax API. diff --git a/litellm/llms/minimax/text_to_speech/__init__.py b/litellm/llms/minimax/text_to_speech/__init__.py index bf4ac9010a4..16cfd98d3b4 100644 --- a/litellm/llms/minimax/text_to_speech/__init__.py +++ b/litellm/llms/minimax/text_to_speech/__init__.py @@ -4,4 +4,4 @@ MiniMax Text-to-Speech module from .transformation import MinimaxException, MinimaxTextToSpeechConfig -__all__ = ["MinimaxTextToSpeechConfig", "MinimaxException"] +__all__ = ["MinimaxException", "MinimaxTextToSpeechConfig"] diff --git a/litellm/llms/minimax/text_to_speech/transformation.py b/litellm/llms/minimax/text_to_speech/transformation.py index 70ce2e71731..93845d10789 100644 --- a/litellm/llms/minimax/text_to_speech/transformation.py +++ b/litellm/llms/minimax/text_to_speech/transformation.py @@ -5,7 +5,7 @@ Maps OpenAI TTS spec to MiniMax TTS API (WebSocket-based HTTP API) Reference: https://platform.minimax.io/docs """ -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from httpx import Headers @@ -33,7 +33,7 @@ class MinimaxException(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, Headers]] = None, + headers: dict | Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) @@ -86,13 +86,13 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): def _resolve_voice_id( self, - voice: Optional[Union[str, Dict[str, Any]]], - params: Dict[str, Any], + voice: str | dict[str, Any] | None, + params: dict[str, Any], ) -> str: """ Determine the MiniMax voice_id based on provided voice input or parameters. """ - mapped_voice: Optional[str] = None + mapped_voice: str | None = None if isinstance(voice, str) and voice.strip(): mapped_voice = self._extract_voice_id(voice) @@ -119,15 +119,15 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Optional[Dict[str, Any]] = None, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict[str, Any] | None = None, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to MiniMax TTS parameters """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} # Work on a copy so we don't mutate the caller's dictionary params = dict(optional_params) if optional_params else {} @@ -180,8 +180,8 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate MiniMax environment and set up authentication headers @@ -202,16 +202,16 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): return headers - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return MinimaxException(message=error_message, status_code=status_code, headers=headers) def transform_text_to_speech_request( self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ @@ -242,7 +242,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): # Output format: 'url' or 'hex' (default is 'hex') output_format = params.pop("output_format", "hex") - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "model": model, "text": input, "stream": False, # HTTP endpoint doesn't support streaming @@ -353,7 +353,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): except Exception as e: raise MinimaxException( status_code=500, - message=f"Failed to decode audio data: {str(e)}", + message=f"Failed to decode audio data: {e!s}", headers=dict(raw_response.headers), ) @@ -378,7 +378,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): except json.JSONDecodeError as e: raise MinimaxException( status_code=500, - message=f"Failed to parse MiniMax response: {str(e)}", + message=f"Failed to parse MiniMax response: {e!s}", headers=dict(raw_response.headers), ) except Exception as e: @@ -386,14 +386,14 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): raise raise MinimaxException( status_code=500, - message=f"Error processing MiniMax response: {str(e)}", + message=f"Error processing MiniMax response: {e!s}", headers=dict(raw_response.headers), ) def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/mistral/audio_transcription/transformation.py b/litellm/llms/mistral/audio_transcription/transformation.py index 53d1428e1f1..12950e9d61e 100644 --- a/litellm/llms/mistral/audio_transcription/transformation.py +++ b/litellm/llms/mistral/audio_transcription/transformation.py @@ -4,8 +4,6 @@ Support for Mistral Voxtral audio transcription via ``/v1/audio/transcriptions`` API reference: https://docs.mistral.ai/api/#tag/audio/operation/audio_transcriptions_v1_audio_transcriptions_post """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -27,7 +25,7 @@ class MistralAudioTranscriptionException(BaseLLMException): class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: return [ "language", "temperature", @@ -50,19 +48,17 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = "https://api.mistral.ai/v1" if api_base is None else api_base.rstrip("/") return f"{api_base}/audio/transcriptions" - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return MistralAudioTranscriptionException( message=error_message, status_code=status_code, @@ -73,11 +69,11 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("MISTRAL_API_KEY") diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 62b3ef92c48..d73435dbcfc 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -9,11 +9,7 @@ Docs - https://docs.mistral.ai/api/ from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( Any, - List, Literal, - Optional, - Tuple, - Union, cast, get_type_hints, overload, @@ -62,27 +58,27 @@ class MistralConfig(OpenAIGPTConfig): - `response_format` (object or null): An object specifying the format that the model must output. Setting to { "type": "json_object" } enables JSON mode, which guarantees the message the model generates is in JSON. When using JSON mode you MUST also instruct the model to produce JSON yourself with a system or a user message. """ - temperature: Optional[int] = None - top_p: Optional[int] = None - max_tokens: Optional[int] = None - tools: Optional[list] = None - tool_choice: Optional[Literal["auto", "any", "none"]] = None - random_seed: Optional[int] = None - safe_prompt: Optional[bool] = None - response_format: Optional[dict] = None - stop: Optional[Union[str, list]] = None + temperature: int | None = None + top_p: int | None = None + max_tokens: int | None = None + tools: list | None = None + tool_choice: Literal["auto", "any", "none"] | None = None + random_seed: int | None = None + safe_prompt: bool | None = None + response_format: dict | None = None + stop: str | list | None = None def __init__( self, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - max_tokens: Optional[int] = None, - tools: Optional[list] = None, - tool_choice: Optional[Literal["auto", "any", "none"]] = None, - random_seed: Optional[int] = None, - safe_prompt: Optional[bool] = None, - response_format: Optional[dict] = None, - stop: Optional[Union[str, list]] = None, + temperature: int | None = None, + top_p: int | None = None, + max_tokens: int | None = None, + tools: list | None = None, + tool_choice: Literal["auto", "any", "none"] | None = None, + random_seed: int | None = None, + safe_prompt: bool | None = None, + response_format: dict | None = None, + stop: str | list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -93,7 +89,7 @@ class MistralConfig(OpenAIGPTConfig): def get_config(cls): return super().get_config() - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: supported_params = [ "stream", "temperature", @@ -188,9 +184,7 @@ class MistralConfig(OpenAIGPTConfig): optional_params["parallel_tool_calls"] = value return optional_params - def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[str, Optional[str]]: + def _get_openai_compatible_provider_info(self, api_base: str | None, api_key: str | None) -> tuple[str, str | None]: # mistral is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.mistral.ai api_base = ( api_base @@ -212,23 +206,23 @@ class MistralConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: + ) -> list[AllMessageValues]: ... # fmt: on def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ - handles scenario where content is list and not string - content list is just text, and no images @@ -256,7 +250,7 @@ class MistralConfig(OpenAIGPTConfig): messages = handle_messages_with_content_list_to_str_conversion(messages) ## 3. Handle name in message - new_messages: List[AllMessageValues] = [] + new_messages: list[AllMessageValues] = [] for m in messages: m = MistralConfig._handle_name_in_message(m) m = MistralConfig._handle_tool_call_message(m) @@ -270,7 +264,7 @@ class MistralConfig(OpenAIGPTConfig): else: return super()._transform_messages(new_messages, model, False) - async def _transform_messages_async(self, messages: List[AllMessageValues], model: str) -> List[AllMessageValues]: + async def _transform_messages_async(self, messages: list[AllMessageValues], model: str) -> list[AllMessageValues]: """ Handle modification of messages for Mistral API in an async context. """ @@ -280,7 +274,7 @@ class MistralConfig(OpenAIGPTConfig): messages = self._handle_message_with_file(messages) return messages - def _transform_messages_sync(self, messages: List[AllMessageValues], model: str) -> List[AllMessageValues]: + def _transform_messages_sync(self, messages: list[AllMessageValues], model: str) -> list[AllMessageValues]: """Handle modification of messages for Mistral API in a sync context.""" # Call parent sync method to handle basic transformations # and then apply Mistral-specific handling for files @@ -289,7 +283,7 @@ class MistralConfig(OpenAIGPTConfig): messages = self._handle_message_with_file(messages) return messages - def _handle_message_with_file(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def _handle_message_with_file(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ Mistral API supports only 'file_id' in message content with type 'file'. """ @@ -309,8 +303,8 @@ class MistralConfig(OpenAIGPTConfig): return messages def _add_reasoning_system_prompt_if_needed( - self, messages: List[AllMessageValues], optional_params: dict - ) -> List[AllMessageValues]: + self, messages: list[AllMessageValues], optional_params: dict + ) -> list[AllMessageValues]: """ Add reasoning system prompt for Mistral magistral models when reasoning_effort is specified. """ @@ -330,13 +324,13 @@ class MistralConfig(OpenAIGPTConfig): # Handle both string and list content, preserving original format if isinstance(existing_content, str): # String content - prepend reasoning prompt - new_content: Union[str, list] = f"{reasoning_prompt}\n\n{existing_content}" + new_content: str | list = f"{reasoning_prompt}\n\n{existing_content}" elif isinstance(existing_content, list): # List content - prepend reasoning prompt as text block new_content = [{"type": "text", "text": reasoning_prompt + "\n\n"}] + existing_content else: # Fallback for any other type - convert to string - new_content = f"{reasoning_prompt}\n\n{str(existing_content)}" + new_content = f"{reasoning_prompt}\n\n{existing_content!s}" messages[i] = cast(AllMessageValues, {**msg, "content": new_content}) break @@ -414,10 +408,7 @@ class MistralConfig(OpenAIGPTConfig): if _name is not None: # Remove name if not a tool message - if message["role"] != "tool": - message.pop("name", None) # type: ignore - # For tool messages, remove name if it's an empty string - elif isinstance(_name, str) and len(_name.strip()) == 0: + if message["role"] != "tool" or isinstance(_name, str) and len(_name.strip()) == 0: message.pop("name", None) # type: ignore return message @@ -428,7 +419,7 @@ class MistralConfig(OpenAIGPTConfig): Mistral API only supports tool_calls in Messages in `MistralToolCallMessage` spec """ _tool_calls = message.get("tool_calls") - mistral_tool_calls: List[MistralToolCallMessage] = [] + mistral_tool_calls: list[MistralToolCallMessage] = [] if _tool_calls is not None and isinstance(_tool_calls, list): for _tool in _tool_calls: _tool_call_message = MistralToolCallMessage( @@ -530,7 +521,7 @@ class MistralConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -562,12 +553,12 @@ class MistralConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the raw response from Mistral API. @@ -596,9 +587,9 @@ class MistralConfig(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return MistralChatResponseIterator( streaming_response=streaming_response, @@ -634,20 +625,20 @@ class MistralChatResponseIterator(OpenAIChatCompletionStreamingHandler): @staticmethod def _normalize_content_blocks( - content_blocks: List[dict], - ) -> Tuple[Optional[str], List[dict], Optional[str]]: + content_blocks: list[dict], + ) -> tuple[str | None, list[dict], str | None]: """ Convert Mistral magistral content blocks into OpenAI-compatible content + thinking_blocks. """ - text_segments: List[str] = [] - thinking_blocks: List[dict] = [] - reasoning_segments: List[str] = [] + text_segments: list[str] = [] + thinking_blocks: list[dict] = [] + reasoning_segments: list[str] = [] for block in content_blocks: block_type = block.get("type") if block_type == "thinking": mistral_thinking = block.get("thinking", []) - thinking_text_parts: List[str] = [] + thinking_text_parts: list[str] = [] for thinking_block in mistral_thinking: if thinking_block.get("type") == "text": thinking_text_parts.append(thinking_block.get("text", "")) diff --git a/litellm/llms/mistral/ocr/guardrail_translation/__init__.py b/litellm/llms/mistral/ocr/guardrail_translation/__init__.py index da7b6ee6bf0..7595a0637b8 100644 --- a/litellm/llms/mistral/ocr/guardrail_translation/__init__.py +++ b/litellm/llms/mistral/ocr/guardrail_translation/__init__.py @@ -8,4 +8,4 @@ guardrail_translation_mappings = { CallTypes.aocr: OCRHandler, } -__all__ = ["guardrail_translation_mappings", "OCRHandler"] +__all__ = ["OCRHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/mistral/ocr/guardrail_translation/handler.py b/litellm/llms/mistral/ocr/guardrail_translation/handler.py index 9144c71f70a..729ac0adb6c 100644 --- a/litellm/llms/mistral/ocr/guardrail_translation/handler.py +++ b/litellm/llms/mistral/ocr/guardrail_translation/handler.py @@ -5,7 +5,7 @@ Provides guardrail translation support for the OCR endpoint. Processes the extracted markdown text from OCR pages. """ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -33,7 +33,7 @@ class OCRHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process OCR input by applying guardrails to the document reference. @@ -55,7 +55,7 @@ class OCRHandler(BaseTranslation): return data # Extract the document URL for guardrail checking - texts_to_check: List[str] = [] + texts_to_check: list[str] = [] doc_type = document.get("type") if doc_type == "document_url": url = document.get("document_url") @@ -87,9 +87,9 @@ class OCRHandler(BaseTranslation): self, response: "OCRResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process OCR output by applying guardrails to extracted page text. @@ -111,8 +111,8 @@ class OCRHandler(BaseTranslation): return response # Extract markdown text from all pages - texts_to_check: List[str] = [] - page_indices: List[int] = [] + texts_to_check: list[str] = [] + page_indices: list[int] = [] for i, page in enumerate(response.pages): if hasattr(page, "markdown") and page.markdown: texts_to_check.append(page.markdown) diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py index 07a67815f6e..e9d8280cc85 100644 --- a/litellm/llms/mistral/ocr/transformation.py +++ b/litellm/llms/mistral/ocr/transformation.py @@ -2,7 +2,7 @@ Mistral OCR transformation implementation. """ -from typing import Any, Dict +from typing import Any import httpx @@ -90,13 +90,13 @@ class MistralOCRConfig(BaseOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers for Mistral OCR. """ diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py index 4625b29d627..227b2e9e2c1 100644 --- a/litellm/llms/modelscope/chat/transformation.py +++ b/litellm/llms/modelscope/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/comple """ from collections.abc import Coroutine -from typing import Any, Literal, Optional, Tuple, Union, cast, overload +from typing import Any, Literal, cast, overload from typing_extensions import override @@ -39,7 +39,7 @@ class ModelScopeChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> Union[list[AllMessageValues], Coroutine[Any, Any, list[AllMessageValues]]]: + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Flatten text-only content lists to strings for ModelScope. @@ -60,8 +60,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL # type: ignore dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY") return api_base, dynamic_api_key @@ -69,12 +69,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig): @override def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ If api_base is not provided, use the default ModelScope /chat/completions endpoint. diff --git a/litellm/llms/modelscope/image_generation/transformation.py b/litellm/llms/modelscope/image_generation/transformation.py index a3d890734d1..55948561e72 100644 --- a/litellm/llms/modelscope/image_generation/transformation.py +++ b/litellm/llms/modelscope/image_generation/transformation.py @@ -6,7 +6,7 @@ Handles transformation between OpenAI-compatible format and ModelScope API forma API Reference: https://modelscope.cn/docs/model-service/API-Inference/intro """ -from typing import TYPE_CHECKING, Optional, Union +from typing import TYPE_CHECKING import httpx from typing_extensions import override @@ -75,12 +75,12 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): @override def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the ModelScope image generation API request. @@ -99,13 +99,13 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for ModelScope. """ - final_api_key: Optional[str] = api_key or get_secret_str("MODELSCOPE_API_KEY") + final_api_key: str | None = api_key or get_secret_str("MODELSCOPE_API_KEY") if not final_api_key: raise ValueError( @@ -158,8 +158,8 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: object, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform ModelScope response to OpenAI-compatible ImageResponse. @@ -204,7 +204,7 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: """Return the appropriate error class for ModelScope.""" from litellm.exceptions import ( diff --git a/litellm/llms/moonshot/chat/transformation.py b/litellm/llms/moonshot/chat/transformation.py index bcd63827324..5b339340657 100644 --- a/litellm/llms/moonshot/chat/transformation.py +++ b/litellm/llms/moonshot/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Moonshot AI's `/v1/chat/compl """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, cast, overload +from typing import Any, Literal, cast, overload import litellm from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -19,20 +19,20 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class MoonshotChatConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Moonshot text-only models don't support content in list format. Multimodal models (kimi-k2.5, kimi-latest, etc.) accept the @@ -59,20 +59,20 @@ class MoonshotChatConfig(OpenAIGPTConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("MOONSHOT_API_BASE") or "https://api.moonshot.ai/v1" # type: ignore dynamic_api_key = api_key or get_secret_str("MOONSHOT_API_KEY") return api_base, dynamic_api_key def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ If api_base is not provided, use the default Moonshot AI /chat/completions endpoint. @@ -94,14 +94,14 @@ class MoonshotChatConfig(OpenAIGPTConfig): - tool_choice doesn't support "required" value - kimi-thinking-preview doesn't support tool calls at all """ - excluded_params: List[str] = ["functions"] + excluded_params: list[str] = ["functions"] # kimi-thinking-preview has additional limitations if "kimi-thinking-preview" in model: excluded_params.extend(["tools", "tool_choice"]) base_openai_params = super().get_supported_openai_params(model=model) - final_params: List[str] = [] + final_params: list[str] = [] for param in base_openai_params: if param not in excluded_params: final_params.append(param) @@ -140,13 +140,12 @@ class MoonshotChatConfig(OpenAIGPTConfig): if supports_reasoning(model=model, custom_llm_provider="moonshot"): optional_params.pop("temperature", None) elif "temperature" in optional_params: - if optional_params["temperature"] > 1: - optional_params["temperature"] = 1 + optional_params["temperature"] = min(optional_params["temperature"], 1) if optional_params["temperature"] < 0.3 and optional_params.get("n", 1) > 1: optional_params["temperature"] = 0.3 return optional_params - def fill_reasoning_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def fill_reasoning_content(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ Moonshot reasoning models require `reasoning_content` on every assistant message that contains tool_calls (multi-turn tool-calling flows). @@ -160,7 +159,7 @@ class MoonshotChatConfig(OpenAIGPTConfig): Messages that already carry the field, or are not assistant/tool-call messages, are appended as-is (no copy made). """ - result: List[AllMessageValues] = [] + result: list[AllMessageValues] = [] for msg in messages: if ( msg.get("role") == "assistant" @@ -195,7 +194,7 @@ class MoonshotChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -226,8 +225,8 @@ class MoonshotChatConfig(OpenAIGPTConfig): ) def _add_tool_choice_required_message( - self, messages: List[AllMessageValues], optional_params: dict - ) -> List[AllMessageValues]: + self, messages: list[AllMessageValues], optional_params: dict + ) -> list[AllMessageValues]: """ Add a message to the messages list to indicate that the tool choice is required. diff --git a/litellm/llms/morph/chat/transformation.py b/litellm/llms/morph/chat/transformation.py index 97ddc12920f..8dc6e7fdd63 100644 --- a/litellm/llms/morph/chat/transformation.py +++ b/litellm/llms/morph/chat/transformation.py @@ -5,8 +5,6 @@ Transform request from OpenAI format to Morph format. https://docs.morphllm.com/quickstart """ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai_like.chat.transformation import OpenAILikeChatConfig @@ -18,12 +16,12 @@ class MorphChatConfig(OpenAILikeChatConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "morph" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = ( api_base or get_secret_str("MORPH_API_BASE") or "https://api.morphllm.com/v1" # default api base ) diff --git a/litellm/llms/nlp_cloud/chat/handler.py b/litellm/llms/nlp_cloud/chat/handler.py index a9f1c600492..ad14f032a3c 100644 --- a/litellm/llms/nlp_cloud/chat/handler.py +++ b/litellm/llms/nlp_cloud/chat/handler.py @@ -1,6 +1,5 @@ import json from collections.abc import Callable -from typing import Optional, Union import litellm from litellm.llms.custom_httpx.http_handler import ( @@ -28,7 +27,7 @@ def completion( litellm_params: dict, logger_fn=None, default_max_tokens_to_sample=None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, headers={}, ): headers = nlp_config.validate_environment( diff --git a/litellm/llms/nlp_cloud/chat/transformation.py b/litellm/llms/nlp_cloud/chat/transformation.py index 5aafc4cd45c..f092112825e 100644 --- a/litellm/llms/nlp_cloud/chat/transformation.py +++ b/litellm/llms/nlp_cloud/chat/transformation.py @@ -1,6 +1,6 @@ import json import time -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -50,33 +50,33 @@ class NLPCloudConfig(BaseConfig): - `num_return_sequences` (int): Optional. The number of independently computed returned sequences. """ - max_length: Optional[int] = None - length_no_input: Optional[bool] = None - end_sequence: Optional[str] = None - remove_end_sequence: Optional[bool] = None - remove_input: Optional[bool] = None - bad_words: Optional[list] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - top_k: Optional[int] = None - repetition_penalty: Optional[float] = None - num_beams: Optional[int] = None - num_return_sequences: Optional[int] = None + max_length: int | None = None + length_no_input: bool | None = None + end_sequence: str | None = None + remove_end_sequence: bool | None = None + remove_input: bool | None = None + bad_words: list | None = None + temperature: float | None = None + top_p: float | None = None + top_k: int | None = None + repetition_penalty: float | None = None + num_beams: int | None = None + num_return_sequences: int | None = None def __init__( self, - max_length: Optional[int] = None, - length_no_input: Optional[bool] = None, - end_sequence: Optional[str] = None, - remove_end_sequence: Optional[bool] = None, - remove_input: Optional[bool] = None, - bad_words: Optional[list] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - repetition_penalty: Optional[float] = None, - num_beams: Optional[int] = None, - num_return_sequences: Optional[int] = None, + max_length: int | None = None, + length_no_input: bool | None = None, + end_sequence: str | None = None, + remove_end_sequence: bool | None = None, + remove_input: bool | None = None, + bad_words: list | None = None, + temperature: float | None = None, + top_p: float | None = None, + top_k: int | None = None, + repetition_penalty: float | None = None, + num_beams: int | None = None, + num_return_sequences: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -91,11 +91,11 @@ class NLPCloudConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = { "accept": "application/json", @@ -105,7 +105,7 @@ class NLPCloudConfig(BaseConfig): headers["Authorization"] = f"Token {api_key}" return headers - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "max_tokens", "stream", @@ -143,15 +143,13 @@ class NLPCloudConfig(BaseConfig): optional_params["stop_sequences"] = value return optional_params - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return NLPCloudError(status_code=status_code, message=error_message, headers=headers) def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -172,12 +170,12 @@ class NLPCloudConfig(BaseConfig): model_response: ModelResponse, logging_obj: LoggingClass, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( diff --git a/litellm/llms/nlp_cloud/common_utils.py b/litellm/llms/nlp_cloud/common_utils.py index 232f56c9709..e64979d8359 100644 --- a/litellm/llms/nlp_cloud/common_utils.py +++ b/litellm/llms/nlp_cloud/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +8,6 @@ class NLPCloudError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) diff --git a/litellm/llms/novita/chat/transformation.py b/litellm/llms/novita/chat/transformation.py index 5a64a124ade..acdfa7e8790 100644 --- a/litellm/llms/novita/chat/transformation.py +++ b/litellm/llms/novita/chat/transformation.py @@ -6,8 +6,6 @@ Calls done in OpenAI/openai.py as Novita AI is openai-compatible. Docs: https://novita.ai/docs/guides/llm-api """ -from typing import List, Optional - from ....types.llms.openai import AllMessageValues from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -17,11 +15,11 @@ class NovitaConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: raise ValueError( diff --git a/litellm/llms/nscale/chat/transformation.py b/litellm/llms/nscale/chat/transformation.py index 1b032fab2ac..8404f862541 100644 --- a/litellm/llms/nscale/chat/transformation.py +++ b/litellm/llms/nscale/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Tuple - from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.secret_managers.main import get_secret_str @@ -14,20 +12,20 @@ class NscaleConfig(OpenAIGPTConfig): API_BASE_URL = "https://inference.api.nscale.com/v1" @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "nscale" @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("NSCALE_API_KEY") @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or get_secret_str("NSCALE_API_BASE") or NscaleConfig.API_BASE_URL def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # This method is called by get_llm_provider to resolve api_base and api_key resolved_api_base = NscaleConfig.get_api_base(api_base) resolved_api_key = NscaleConfig.get_api_key(api_key) diff --git a/litellm/llms/nvidia_nim/embed.py b/litellm/llms/nvidia_nim/embed.py index 61c8e8244e4..111111435a1 100644 --- a/litellm/llms/nvidia_nim/embed.py +++ b/litellm/llms/nvidia_nim/embed.py @@ -9,7 +9,6 @@ API calling is done using the OpenAI SDK with an api_base """ import types -from typing import Optional class NvidiaNimEmbeddingConfig: @@ -18,19 +17,19 @@ class NvidiaNimEmbeddingConfig: """ # OpenAI params - encoding_format: Optional[str] = None - user: Optional[str] = None + encoding_format: str | None = None + user: str | None = None # Nvidia NIM params - input_type: Optional[str] = None - truncate: Optional[str] = None + input_type: str | None = None + truncate: str | None = None def __init__( self, - encoding_format: Optional[str] = None, - user: Optional[str] = None, - input_type: Optional[str] = None, - truncate: Optional[str] = None, + encoding_format: str | None = None, + user: str | None = None, + input_type: str | None = None, + truncate: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -64,7 +63,7 @@ class NvidiaNimEmbeddingConfig: self, non_default_params: dict, optional_params: dict, - kwargs: Optional[dict] = None, + kwargs: dict | None = None, ): if "extra_body" not in optional_params: optional_params["extra_body"] = {} diff --git a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py index b9a46b8ac2b..34d1586e43c 100644 --- a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py +++ b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py @@ -6,8 +6,6 @@ Use this by passing "nvidia_nim/ranking/" to force the /v1/ranking endpoi Reference: https://build.nvidia.com/nvidia/llama-3_2-nv-rerankqa-1b-v2/deploy """ -from typing import Dict, Optional - from litellm.llms.nvidia_nim.rerank.transformation import NvidiaNimRerankConfig @@ -30,18 +28,16 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig): def _get_clean_model_name(self, model: str) -> str: """Strip 'nvidia_nim/' and 'ranking/' prefixes from model name.""" # First strip nvidia_nim/ prefix if present - if model.startswith("nvidia_nim/"): - model = model[len("nvidia_nim/") :] + model = model.removeprefix("nvidia_nim/") # Then strip ranking/ prefix if present - if model.startswith("ranking/"): - model = model[len("ranking/") :] + model = model.removeprefix("ranking/") return model def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, model: str, - optional_params: Optional[dict] = None, + optional_params: dict | None = None, ) -> str: """ Construct the Nvidia NIM ranking URL. @@ -56,17 +52,16 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig): if api_base.endswith("/ranking"): return api_base - if api_base.endswith("/v1"): - api_base = api_base[:-3] + api_base = api_base.removesuffix("/v1") return f"{api_base}/v1/ranking" def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> dict: """ Transform request, using clean model name without 'ranking/' prefix. diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index 2d72d52f991..9feaa9514dc 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Literal, Union +from typing import Any, Literal import httpx from typing_extensions import Required, TypedDict @@ -28,7 +28,7 @@ class NvidiaNimPassageObject(TypedDict): class NvidiaNimRerankRequest(TypedDict, total=False): model: Required[str] query: Required[NvidiaNimQueryObject] - passages: Required[List[NvidiaNimPassageObject]] + passages: Required[list[NvidiaNimPassageObject]] truncate: Literal["NONE", "END"] top_k: int @@ -39,7 +39,7 @@ class NvidiaNimRankingResult(TypedDict): class NvidiaNimRerankResponse(TypedDict): - rankings: Required[List[NvidiaNimRankingResult]] + rankings: Required[list[NvidiaNimRankingResult]] class NvidiaNimRerankConfig(BaseRerankConfig): @@ -86,8 +86,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): return api_base # Ensure we don't have duplicate /v1 - if api_base.endswith("/v1"): - api_base = api_base[:-3] + api_base = api_base.removesuffix("/v1") # Strip nvidia_nim/ prefix from model name if present clean_model = self._get_clean_model_name(model) @@ -110,15 +109,15 @@ class NvidiaNimRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map Cohere/OpenAI rerank params to Nvidia NIM format. @@ -128,7 +127,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: Dict[str, Any] = { + optional_nvidia_nim_rerank_params: dict[str, Any] = { "query": query, "documents": documents, } @@ -174,7 +173,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -202,7 +201,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): query_obj: NvidiaNimQueryObject = {"text": query} # Transform documents to passages format - passages: List[NvidiaNimPassageObject] = [] + passages: list[NvidiaNimPassageObject] = [] for doc in documents: if isinstance(doc, str): passages.append({"text": doc}) @@ -293,11 +292,11 @@ class NvidiaNimRerankConfig(BaseRerankConfig): nvidia_response: NvidiaNimRerankResponse = raw_response_json # Transform Nvidia NIM response to LiteLLM format - results: List[RerankResponseResult] = [] + results: list[RerankResponseResult] = [] rankings = nvidia_response.get("rankings", []) # Get original documents from request if we need to include them - original_passages: List[NvidiaNimPassageObject] = request_data.get("passages", []) + original_passages: list[NvidiaNimPassageObject] = request_data.get("passages", []) for ranking in rankings: result_item: RerankResponseResult = { @@ -327,9 +326,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): meta=meta, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return BaseLLMException( status_code=status_code, message=error_message, diff --git a/litellm/llms/nvidia_riva/audio_transcription/audio_utils.py b/litellm/llms/nvidia_riva/audio_transcription/audio_utils.py index 7ec679c858d..7eca93bfdcb 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/audio_utils.py +++ b/litellm/llms/nvidia_riva/audio_transcription/audio_utils.py @@ -16,7 +16,7 @@ import io import os import tempfile from dataclasses import dataclass -from typing import Any, Tuple, cast +from typing import Any, cast from litellm.llms.nvidia_riva.audio_transcription.transformation import ( RIVA_TARGET_NUM_CHANNELS, @@ -83,7 +83,7 @@ def resample_to_riva_pcm(file_bytes: bytes) -> ResampledAudio: ) -def _decode_to_float32(file_bytes: bytes) -> Tuple["FloatArray", int]: +def _decode_to_float32(file_bytes: bytes) -> tuple["FloatArray", int]: """ Decode arbitrary audio bytes into a float32 array shaped either ``(n_samples,)`` (mono) or ``(n_samples, n_channels)`` plus the source diff --git a/litellm/llms/nvidia_riva/audio_transcription/handler.py b/litellm/llms/nvidia_riva/audio_transcription/handler.py index eab5abd475b..ed99745cee4 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/handler.py +++ b/litellm/llms/nvidia_riva/audio_transcription/handler.py @@ -26,7 +26,7 @@ without the optional STT extras installed. import asyncio import inspect -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm.litellm_core_utils.audio_utils.utils import ( get_audio_file_name, @@ -36,9 +36,9 @@ from litellm.llms.nvidia_riva.audio_transcription.audio_utils import ( resample_to_riva_pcm, ) from litellm.llms.nvidia_riva.audio_transcription.transformation import ( - NvidiaRivaAudioTranscriptionConfig, RIVA_TARGET_NUM_CHANNELS, RIVA_TARGET_SAMPLE_RATE_HZ, + NvidiaRivaAudioTranscriptionConfig, ) from litellm.llms.nvidia_riva.common_utils import ( NvidiaRivaException, @@ -74,10 +74,10 @@ class NvidiaRivaAudioTranscription: model_response: TranscriptionResponse, timeout: float, logging_obj: "LiteLLMLoggingObj", - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, atranscription: bool = False, - provider_config: Optional[NvidiaRivaAudioTranscriptionConfig] = None, + provider_config: NvidiaRivaAudioTranscriptionConfig | None = None, ): if provider_config is None: provider_config = NvidiaRivaAudioTranscriptionConfig() @@ -119,9 +119,9 @@ class NvidiaRivaAudioTranscription: model_response: TranscriptionResponse, timeout: float, logging_obj: "LiteLLMLoggingObj", - api_key: Optional[str], - api_base: Optional[str], - provider_config: Optional[NvidiaRivaAudioTranscriptionConfig] = None, + api_key: str | None, + api_base: str | None, + provider_config: NvidiaRivaAudioTranscriptionConfig | None = None, ) -> TranscriptionResponse: # ``riva-client`` exposes a sync streaming generator, so we offload # the blocking call to a worker thread to keep the event loop free. @@ -149,8 +149,8 @@ class NvidiaRivaAudioTranscription: model_response: TranscriptionResponse, timeout: float, logging_obj: "LiteLLMLoggingObj", - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, provider_config: NvidiaRivaAudioTranscriptionConfig, atranscription: bool = False, ) -> TranscriptionResponse: @@ -184,7 +184,7 @@ class NvidiaRivaAudioTranscription: message="NvidiaRivaAudioTranscriptionConfig produced an unexpected request payload type.", ) - recognition_config_dict: Dict[str, Any] = request_payload["recognition_config"] + recognition_config_dict: dict[str, Any] = request_payload["recognition_config"] # The wire format is fixed by our resampler; override anything stale # the caller passed in so the gRPC config matches the bytes we send. recognition_config_dict["sample_rate_hertz"] = RIVA_TARGET_SAMPLE_RATE_HZ @@ -225,7 +225,7 @@ class NvidiaRivaAudioTranscription: try: asr_service = riva_module.ASRService(auth_obj) audio_chunks = self._iter_audio_chunks(resampled.pcm_bytes) - stream_kwargs: Dict[str, Any] = { + stream_kwargs: dict[str, Any] = { "audio_chunks": audio_chunks, "streaming_config": streaming_config, } @@ -276,7 +276,7 @@ class NvidiaRivaAudioTranscription: self, riva_module: Any, api_base: str, - api_key: Optional[str], + api_key: str | None, optional_params: dict, ) -> Any: """ @@ -293,7 +293,7 @@ class NvidiaRivaAudioTranscription: use_ssl_override = optional_params.get("use_ssl") use_ssl = bool(use_ssl_override) if use_ssl_override is not None else bool(nvcf_function_id) - metadata: List[Tuple[str, str]] = [] + metadata: list[tuple[str, str]] = [] if nvcf_function_id: metadata.append(("function-id", str(nvcf_function_id))) if api_key: @@ -305,7 +305,7 @@ class NvidiaRivaAudioTranscription: # Older riva-client signatures used positional-only args. return riva_module.Auth(None, use_ssl, api_base, metadata) - def _build_recognition_config_proto(self, riva_asr_module: Any, recognition_config_dict: Dict[str, Any]): + def _build_recognition_config_proto(self, riva_asr_module: Any, recognition_config_dict: dict[str, Any]): encoding_name = (recognition_config_dict.get("encoding") or "LINEAR_PCM").upper() encoding_enum = getattr( riva_asr_module.AudioEncoding, @@ -359,14 +359,14 @@ class NvidiaRivaAudioTranscription: yield chunk @staticmethod - def _collect_final_results(stream) -> List[Dict[str, Any]]: + def _collect_final_results(stream) -> list[dict[str, Any]]: """ Walk the gRPC stream, ignore empty / non-final chunks, and return a list of normalized final-result dicts. Matching the user's note: the ``id`` blocks with no ``results`` are streaming heartbeats and must be skipped. """ - final_results: List[Dict[str, Any]] = [] + final_results: list[dict[str, Any]] = [] for response in stream: results = getattr(response, "results", None) or [] for result in results: @@ -406,7 +406,7 @@ def _import_riva(): riva_asr_module = riva_client if not hasattr(riva_asr_module, "RecognitionConfig"): try: - import riva.client.proto.riva_asr_pb2 as riva_asr_pb2 # type: ignore + from riva.client.proto import riva_asr_pb2 # type: ignore riva_asr_module = riva_asr_pb2 except ImportError as e: diff --git a/litellm/llms/nvidia_riva/audio_transcription/transformation.py b/litellm/llms/nvidia_riva/audio_transcription/transformation.py index 43185cb2f7a..0c769a0497b 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/transformation.py +++ b/litellm/llms/nvidia_riva/audio_transcription/transformation.py @@ -11,7 +11,7 @@ dict at call time. Reference: https://docs.nvidia.com/deeplearning/riva/user-guide/docs/asr/asr-overview.html """ -from typing import Any, Dict, List, Optional, Union +from typing import Any from httpx import Headers, Response @@ -43,7 +43,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): optional TLS via ``use_ssl``). """ - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: # Riva natively understands language + word timestamps. # `response_format` is honored at response-shaping time in the handler. return ["language", "response_format", "timestamp_granularities"] @@ -77,7 +77,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return NvidiaRivaException(message=error_message, status_code=status_code, headers=headers) def transform_audio_transcription_request( @@ -102,7 +102,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): if endpointing_config is not None: recognition_config["endpointing_config"] = endpointing_config - request_payload: Dict[str, Any] = { + request_payload: dict[str, Any] = { "recognition_config": recognition_config, "response_format": optional_params.get("response_format") or "json", "timestamp_granularities": optional_params.get("timestamp_granularities"), @@ -126,16 +126,16 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: # 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, Any]: """ 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) -> Optional[Dict[str, Any]]: + def _build_endpointing_config_dict(self, optional_params: dict) -> dict[str, Any] | 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: Dict[str, Any] = {} + config: dict[str, Any] = {} if "threshold" in chunking: threshold = float(chunking["threshold"]) config["start_threshold"] = threshold @@ -219,10 +219,10 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): @staticmethod def build_transcription_response( - final_results: List[Dict[str, Any]], + final_results: list[dict[str, Any]], response_format: str, - duration_seconds: Optional[float], - timestamp_granularities: Optional[List[str]], + duration_seconds: float | None, + timestamp_granularities: list[str] | None, ) -> TranscriptionResponse: """ Aggregate a list of normalized "final result" dicts into a @@ -245,7 +245,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response["task"] = "transcribe" if response_format == "verbose_json": - words: List[Dict[str, Any]] = [] + words: list[dict[str, Any]] = [] 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/nvidia_riva/common_utils.py b/litellm/llms/nvidia_riva/common_utils.py index 4206fc91cc6..e4b43948b85 100644 --- a/litellm/llms/nvidia_riva/common_utils.py +++ b/litellm/llms/nvidia_riva/common_utils.py @@ -2,7 +2,7 @@ Common utilities and exceptions for the NVIDIA Riva STT provider """ -from typing import Any, Optional +from typing import Any from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -16,8 +16,6 @@ class NvidiaRivaException(BaseLLMException): classifiers (RateLimitError, AuthenticationError, etc.) keep working. """ - pass - # Mapping from grpc.StatusCode.name -> equivalent HTTP status code. # Kept as a plain dict (rather than importing grpc enums) so this module is @@ -43,7 +41,7 @@ _GRPC_STATUS_CODE_TO_HTTP: dict = { } -def _extract_grpc_status_name(error: Any) -> Optional[str]: +def _extract_grpc_status_name(error: Any) -> str | None: """ Best-effort extraction of a gRPC StatusCode name from an arbitrary error. @@ -62,7 +60,7 @@ def _extract_grpc_status_name(error: Any) -> Optional[str]: return None -def _extract_grpc_details(error: Any) -> Optional[str]: +def _extract_grpc_details(error: Any) -> str | None: """Best-effort extraction of a human-readable detail string from a gRPC error.""" details_fn = getattr(error, "details", None) if callable(details_fn): diff --git a/litellm/llms/oci/chat/cohere.py b/litellm/llms/oci/chat/cohere.py index ca85a6309d7..d3ffa926c46 100644 --- a/litellm/llms/oci/chat/cohere.py +++ b/litellm/llms/oci/chat/cohere.py @@ -8,7 +8,7 @@ response parsing, and streaming chunk parsing for models served with import datetime import json -from typing import Any, Dict, List, Optional +from typing import Any import httpx from pydantic import ValidationError @@ -42,8 +42,8 @@ from litellm.types.utils import ( ModelResponse, ModelResponseStream, StreamingChoices, + Usage, ) -from litellm.types.utils import Usage def _extract_text_content(content: Any) -> str: @@ -60,8 +60,8 @@ def _extract_text_content(content: Any) -> str: def adapt_messages_to_cohere_standard( - messages: List[AllMessageValues], -) -> List[CohereMessage]: + messages: list[AllMessageValues], +) -> list[CohereMessage]: """Build a Cohere ``chatHistory`` list from an OpenAI-format message array. - All messages except the *last user message* are included. The caller pulls @@ -78,7 +78,7 @@ def adapt_messages_to_cohere_standard( """ # First pass: build tool_call_id → CohereToolCall so tool-result messages can # reference the originating call by name and parameters. - tool_call_lookup: Dict[str, CohereToolCall] = {} + tool_call_lookup: dict[str, CohereToolCall] = {} for msg in messages: if msg.get("role") == "assistant": tool_calls_raw: Any = msg.get("tool_calls") or [] @@ -86,7 +86,7 @@ def adapt_messages_to_cohere_standard( tc_id = tc.get("id", "") raw_args: Any = tc.get("function", {}).get("arguments", "{}") try: - params: Dict[str, Any] = json.loads(raw_args) if isinstance(raw_args, str) else raw_args + params: dict[str, Any] = json.loads(raw_args) if isinstance(raw_args, str) else raw_args except json.JSONDecodeError: params = {} tool_call_lookup[tc_id] = CohereToolCall( @@ -102,19 +102,19 @@ def adapt_messages_to_cohere_standard( messages if last_user_index is None else [m for i, m in enumerate(messages) if i != last_user_index] ) - chat_history: List[CohereMessage] = [] + chat_history: list[CohereMessage] = [] for msg in history_source: role = msg.get("role") content = _extract_text_content(msg.get("content")) - tool_calls: Optional[List[CohereToolCall]] = None + tool_calls: list[CohereToolCall] | None = None if role == "assistant" and msg.get("tool_calls"): # type: ignore[union-attr,typeddict-item] tool_calls = [] for tc in msg["tool_calls"]: # pyright: ignore[reportOptionalIterable] # truthiness check above rules out None raw_arguments: Any = tc.get("function", {}).get("arguments", {}) if isinstance(raw_arguments, str): try: - arguments: Dict[str, Any] = json.loads(raw_arguments) + arguments: dict[str, Any] = json.loads(raw_arguments) except json.JSONDecodeError: arguments = {} else: @@ -151,8 +151,8 @@ def adapt_messages_to_cohere_standard( def adapt_tool_definitions_to_cohere_standard( - tools: List[Dict[str, Any]], -) -> List[CohereTool]: + tools: list[dict[str, Any]], +) -> list[CohereTool]: """Adapt OpenAI-format tool definitions to the OCI Cohere format. - Resolves ``$ref``/``$defs`` and ``anyOf`` patterns that OCI rejects. @@ -201,7 +201,7 @@ def handle_cohere_response( cohere_response = CohereChatResult(**json_response) except (TypeError, ValidationError) as e: raise OCIError( - message=f"Response cannot be casted to CohereChatResult: {str(e)}", + message=f"Response cannot be casted to CohereChatResult: {e!s}", status_code=raw_response.status_code, ) @@ -211,7 +211,7 @@ def handle_cohere_response( response_text = cohere_response.chatResponse.text finish_reason = _normalize_oci_finish_reason(cohere_response.chatResponse.finishReason) - tool_calls: Optional[List[Dict[str, Any]]] = None + tool_calls: list[dict[str, Any]] | None = None if cohere_response.chatResponse.toolCalls: tool_calls = [ { @@ -225,14 +225,14 @@ def handle_cohere_response( for i, tc in enumerate(cohere_response.chatResponse.toolCalls) ] - content: Optional[str] = response_text if response_text else None + content: str | None = response_text if response_text else None # Only include ``tool_calls`` in the message dict when actually present. # Passing an explicit ``None`` would let downstream consumers that key off # ``"tool_calls" in message`` (rather than truthiness) incorrectly conclude # that tool calls were attempted. Matches the generic handler's behaviour, # which only sets ``message.tool_calls`` when tool calls are present. - message: Dict[str, Any] = {"role": "assistant", "content": content} + message: dict[str, Any] = {"role": "assistant", "content": content} if tool_calls is not None: message["tool_calls"] = tool_calls @@ -283,7 +283,7 @@ def handle_cohere_stream_chunk( except (TypeError, ValidationError) as e: raise OCIError( status_code=500, - message=f"Chunk cannot be parsed as CohereStreamChunk: {str(e)}", + message=f"Chunk cannot be parsed as CohereStreamChunk: {e!s}", ) if typed_chunk.index is None: @@ -305,7 +305,7 @@ def handle_cohere_stream_chunk( # confirmed that text deltas were already emitted earlier — otherwise # (e.g. a degenerate stream that delivers the whole response in a # single SSE event), passing it through is the only chance to surface it. - text: Optional[str] = None if (is_terminal_consolidation and prior_text_emitted) else typed_chunk.text + text: str | None = None if (is_terminal_consolidation and prior_text_emitted) else typed_chunk.text # Tool calls on the terminal consolidation chunk (whether from # `typed_chunk.toolCalls` or from `chatHistory`) typically restate what @@ -317,7 +317,7 @@ def handle_cohere_stream_chunk( # passing them through is the only chance to surface them. cohere_tool_calls = None if (is_terminal_consolidation and prior_tool_calls_emitted) else typed_chunk.toolCalls - tool_calls: Optional[List[Dict[str, Any]]] = None + tool_calls: list[dict[str, Any]] | None = None if cohere_tool_calls: tool_calls = [ { diff --git a/litellm/llms/oci/chat/generic.py b/litellm/llms/oci/chat/generic.py index 02ec762488d..354bcbed3ba 100644 --- a/litellm/llms/oci/chat/generic.py +++ b/litellm/llms/oci/chat/generic.py @@ -8,7 +8,7 @@ parsing, and streaming chunk parsing for models served with import datetime import hashlib -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx from pydantic import ValidationError @@ -34,15 +34,16 @@ from litellm.types.llms.oci import ( ) from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ( + ChatCompletionMessageToolCall, Delta, ModelResponse, ModelResponseStream, StreamingChoices, + Usage, ) -from litellm.types.utils import ChatCompletionMessageToolCall, Usage # Maps OpenAI role names to OCI GENERIC role names. -open_ai_to_generic_oci_role_map: Dict[str, OCIRoles] = { +open_ai_to_generic_oci_role_map: dict[str, OCIRoles] = { "system": "SYSTEM", "user": "USER", "assistant": "ASSISTANT", @@ -55,9 +56,9 @@ open_ai_to_generic_oci_role_map: Dict[str, OCIRoles] = { # --------------------------------------------------------------------------- -def adapt_messages_to_generic_oci_standard_content_message(role: str, content: Union[str, list]) -> OCIMessage: +def adapt_messages_to_generic_oci_standard_content_message(role: str, content: str | list) -> OCIMessage: """Convert a plain-text or multipart content message to OCI format.""" - new_content: List[OCIContentPartUnion] = [] + new_content: list[OCIContentPartUnion] = [] if isinstance(content, str): return OCIMessage( role=open_ai_to_generic_oci_role_map[role], @@ -166,8 +167,8 @@ def adapt_messages_to_generic_oci_standard_tool_response(role: str, tool_call_id def adapt_messages_to_generic_oci_standard( - messages: List[AllMessageValues], -) -> List[OCIMessage]: + messages: list[AllMessageValues], +) -> list[OCIMessage]: """Convert an OpenAI-format message array to OCI GENERIC format.""" new_messages = [] for message in messages: @@ -210,7 +211,7 @@ def adapt_messages_to_generic_oci_standard( # --------------------------------------------------------------------------- -def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors) -> List[OCIToolDefinition]: +def adapt_tool_definition_to_oci_standard(tools: list[dict], vendor: OCIVendors) -> list[OCIToolDefinition]: """Convert OpenAI-format tool definitions to OCI GENERIC format. Resolves ``$ref``/``$defs`` and ``anyOf`` that the OCI endpoint rejects. @@ -239,7 +240,7 @@ def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors) return new_tools -def _normalize_oci_finish_reason(raw: Optional[str]) -> Optional[str]: +def _normalize_oci_finish_reason(raw: str | None) -> str | None: """Map an OCI-specific finish reason to its OpenAI-standard equivalent. OCI emits ``COMPLETE`` / ``MAX_TOKENS`` / ``TOOL_CALL(S)`` plus a long tail @@ -272,15 +273,15 @@ def _synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> st re-emissions while differing across truly distinct calls. """ digest = hashlib.sha256( - f"{position}|{name}|{arguments}".encode("utf-8"), + f"{position}|{name}|{arguments}".encode(), usedforsecurity=False, ).hexdigest()[:24] return f"call_{digest}" def adapt_tools_to_openai_standard( - tools: List[OCIToolCall], -) -> List[ChatCompletionMessageToolCall]: + tools: list[OCIToolCall], +) -> list[ChatCompletionMessageToolCall]: """Convert OCI tool-call objects in a response to the OpenAI format.""" return [ ChatCompletionMessageToolCall( @@ -308,7 +309,7 @@ def handle_generic_response( completion_response = OCICompletionResponse(**json_data) except (TypeError, ValidationError) as e: raise OCIError( - message=f"Response cannot be casted to OCICompletionResponse: {str(e)}", + message=f"Response cannot be casted to OCICompletionResponse: {e!s}", status_code=raw_response.status_code, ) @@ -331,7 +332,7 @@ def handle_generic_response( # Concatenate all text parts — matches the streaming handler, which # iterates the full content array. Skips non-text parts (e.g. image # parts) so a leading non-text part doesn't suppress trailing text. - text: Optional[str] = None + text: str | None = None for item in response_message.content: if isinstance(item, OCITextContentPart): text = (text or "") + item.text @@ -345,7 +346,7 @@ def handle_generic_response( ) oci_usage = completion_response.chatResponse.usage - reasoning_tokens: Optional[int] = None + reasoning_tokens: int | None = None if oci_usage.completionTokensDetails and oci_usage.completionTokensDetails.reasoningTokens is not None: reasoning_tokens = oci_usage.completionTokensDetails.reasoningTokens model_response.usage = Usage( # type: ignore[attr-defined] @@ -372,7 +373,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: except (TypeError, ValidationError) as e: raise OCIError( status_code=500, - message=f"Chunk cannot be parsed as OCIStreamChunk: {str(e)}", + message=f"Chunk cannot be parsed as OCIStreamChunk: {e!s}", ) if typed_chunk.index is None: @@ -382,7 +383,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: # parts (e.g. tool-call-only or keep-alive chunks) so downstream # stream-mergers that distinguish "no text in this delta" from "an # explicitly empty text delta" behave correctly. - text: Optional[str] = None + text: str | None = None if typed_chunk.message and typed_chunk.message.content: for item in typed_chunk.message.content: if isinstance(item, OCITextContentPart): @@ -405,7 +406,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: # same minimal ``{"id", "type", "function": {"name", "arguments"}}`` # shape keeps downstream stream-mergers behaving identically across # GENERIC and Cohere chunks. - tool_calls: Optional[List[Dict[str, Any]]] = None + tool_calls: list[dict[str, Any]] | None = None if typed_chunk.message and typed_chunk.message.toolCalls: tool_calls = [ { @@ -419,7 +420,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: for i, tc in enumerate(typed_chunk.message.toolCalls) ] - finish_reason: Optional[str] = _normalize_oci_finish_reason(typed_chunk.finishReason) + finish_reason: str | None = _normalize_oci_finish_reason(typed_chunk.finishReason) return ModelResponseStream( choices=[ diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index a90e19ad9eb..2d441cb4515 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -14,11 +14,6 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Union, ) import httpx @@ -142,7 +137,7 @@ async def _aiter_sse_events(stream: AsyncIterator[str]) -> AsyncIterator[str]: yield stripped -def _normalize_tool_choice(selected_params: Dict) -> None: +def _normalize_tool_choice(selected_params: dict) -> None: tc = selected_params.get("toolChoice") if tc is None: return @@ -189,7 +184,7 @@ def _normalize_tool_choice(selected_params: Dict) -> None: ) -def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> None: +def _normalize_response_format(selected_params: dict, vendor: OCIVendors) -> None: rf = selected_params.get("responseFormat") if not isinstance(rf, dict) or "type" not in rf: return @@ -204,7 +199,7 @@ def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> Non if vendor == OCIVendors.COHERE: # OCI Cohere has no JSON_SCHEMA type; a schema rides on JSON_OBJECT. - payload: Dict[str, Any] = {"type": "JSON_OBJECT"} + payload: dict[str, Any] = {"type": "JSON_OBJECT"} if json_schema is not None and json_schema.get("schema") is not None: payload["schema"] = json_schema["schema"] selected_params["responseFormat"] = payload @@ -219,7 +214,7 @@ def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> Non # OCI's ResponseJsonSchema accepts only name/description/schema/isStrict. # OpenAI sends `strict` instead of `isStrict`; forwarding it (or any # other extra key) makes OCI reject the whole request with HTTP 400. - oci_schema: Dict[str, Any] = {"name": json_schema.get("name") or "response"} + oci_schema: dict[str, Any] = {"name": json_schema.get("name") or "response"} if json_schema.get("description") is not None: oci_schema["description"] = json_schema["description"] if json_schema.get("schema") is not None: @@ -318,7 +313,7 @@ class OCIChatConfig(BaseConfig): self.openai_to_oci_cohere_param_map["logprobs"] = False self.openai_to_oci_cohere_param_map["logit_bias"] = False - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: param_map = ( self.openai_to_oci_cohere_param_map if get_vendor_from_model(model) == OCIVendors.COHERE @@ -385,11 +380,11 @@ class OCIChatConfig(BaseConfig): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, bytes]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes]: return sign_oci_request( headers=headers, optional_params=optional_params, @@ -405,11 +400,11 @@ class OCIChatConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if not messages: raise OCIError( @@ -443,21 +438,21 @@ class OCIChatConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base = get_oci_base_url(optional_params, api_base or litellm.api_base) return f"{base}/{OCI_API_VERSION}/actions/chat" - def _get_optional_params(self, vendor: OCIVendors, optional_params: dict, model: str = "") -> Dict: + def _get_optional_params(self, vendor: OCIVendors, optional_params: dict, model: str = "") -> dict: param_map = ( self.openai_to_oci_cohere_param_map if vendor == OCIVendors.COHERE else self.openai_to_oci_generic_param_map ) - selected_params: Dict = {} + selected_params: dict = {} # OpenAI reasoning models on OCI (e.g. GPT-5 family) reject "maxTokens" # and require "maxCompletionTokens" per OCI's /20231130/Chat schema. @@ -528,7 +523,7 @@ class OCIChatConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -608,12 +603,12 @@ class OCIChatConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: response_json = raw_response.json() @@ -648,9 +643,9 @@ class OCIChatConfig(BaseConfig): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "OCIStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -687,9 +682,9 @@ class OCIChatConfig(BaseConfig): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "OCIStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={}) @@ -716,9 +711,7 @@ class OCIChatConfig(BaseConfig): logging_obj=logging_obj, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OCIError(status_code=status_code, message=error_message) @@ -748,7 +741,7 @@ class OCIStreamWrapper(CustomStreamWrapper): except json.JSONDecodeError as e: raise OCIError( status_code=500, - message=f"Chunk cannot be parsed as JSON: {str(e)}", + message=f"Chunk cannot be parsed as JSON: {e!s}", ) if dict_chunk.get("apiFormat") == "COHERE": @@ -772,11 +765,11 @@ class OCIStreamWrapper(CustomStreamWrapper): __all__ = [ - "OCIChatConfig", - "OCIStreamWrapper", - "OCIRequestWrapper", "OCI_API_VERSION", "STREAMING_TIMEOUT", + "OCIChatConfig", + "OCIRequestWrapper", + "OCIStreamWrapper", "get_vendor_from_model", "version", ] diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 4ecbcbfb656..7277972f64a 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -5,7 +5,7 @@ import os import re from dataclasses import dataclass from email.utils import formatdate -from typing import Any, Dict, Optional, Protocol, Tuple +from typing import Any, Protocol from urllib.parse import urlparse import httpx @@ -42,7 +42,7 @@ class OCIError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[httpx.Headers] = None, + headers: httpx.Headers | None = None, ): super().__init__( status_code=status_code, @@ -177,7 +177,7 @@ _OCI_REGION_RE = re.compile(r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$") _OCI_ACTION_PATH_RE = re.compile(rf"/{OCI_API_VERSION}/actions/[^/?#]+/?$") -def get_oci_base_url(optional_params: dict, api_base: Optional[str] = None) -> str: +def get_oci_base_url(optional_params: dict, api_base: str | None = None) -> str: """Return the OCI inference base URL, respecting any explicit api_base override. If ``api_base`` already ends with a fully-formed OCI action path @@ -208,7 +208,7 @@ def sign_with_oci_signer( optional_params: dict, request_data: dict, api_base: str, -) -> Tuple[dict, bytes]: +) -> tuple[dict, bytes]: """Sign a request using an OCI SDK Signer object passed in optional_params.""" oci_signer = optional_params.get("oci_signer") body = json.dumps(request_data).encode("utf-8") @@ -232,7 +232,7 @@ def sign_with_oci_signer( raise OCIError( status_code=500, message=( - f"Failed to sign request with provided oci_signer: {str(e)}. " + f"Failed to sign request with provided oci_signer: {e!s}. " "The signer must implement the OCI SDK Signer interface with a " "do_request_sign(request, enforce_content_headers=True) method. " "See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html" @@ -248,7 +248,7 @@ def sign_with_manual_credentials( optional_params: dict, request_data: dict, api_base: str, -) -> Tuple[dict, bytes]: +) -> tuple[dict, bytes]: """Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key).""" creds = resolve_oci_credentials(optional_params) oci_user = creds["oci_user"] @@ -280,7 +280,7 @@ def sign_with_manual_credentials( content_length = str(len(body)) x_content_sha256 = sha256_base64(body) - headers_to_sign: Dict[str, str] = { + headers_to_sign: dict[str, str] = { "date": date, "host": host, "content-type": content_type, @@ -301,7 +301,7 @@ def sign_with_manual_credentials( _require_cryptography() # Resolve the private key — prefer inline PEM content over file path - oci_key_content: Optional[str] = None + oci_key_content: str | None = None if oci_key: if not isinstance(oci_key, str): raise OCIError( @@ -361,11 +361,11 @@ def sign_oci_request( optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, -) -> Tuple[dict, bytes]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, +) -> tuple[dict, bytes]: """ Route to the appropriate OCI signing method based on what credentials are present. @@ -384,7 +384,7 @@ def sign_oci_request( def validate_oci_environment( headers: dict, optional_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: """ Populate common OCI request headers (content-type, user-agent). @@ -410,7 +410,7 @@ def validate_oci_environment( # Mapping from JSON Schema type names to Python type names, as expected by # the OCI Cohere API's CohereParameterDefinition.type field. -OCI_JSON_TO_PYTHON_TYPES: Dict[str, str] = { +OCI_JSON_TO_PYTHON_TYPES: dict[str, str] = { "string": "str", "number": "float", "boolean": "bool", @@ -421,7 +421,7 @@ OCI_JSON_TO_PYTHON_TYPES: Dict[str, str] = { } -def resolve_oci_schema_refs(schema: Dict[str, Any]) -> Dict[str, Any]: +def resolve_oci_schema_refs(schema: dict[str, Any]) -> dict[str, Any]: """Inline all ``$ref``/``$defs`` references — OCI does not support JSON Schema ``$ref``.""" defs = schema.get("$defs", {}) resolving_stack: set = set() @@ -483,7 +483,7 @@ def sanitize_oci_schema(schema: Any) -> Any: if not isinstance(schema, dict): return schema - sanitized: Dict[str, Any] = {} + sanitized: dict[str, Any] = {} for key, value in schema.items(): if key == "title": continue @@ -513,7 +513,7 @@ def sanitize_oci_schema(schema: Any) -> Any: return sanitized -def enrich_cohere_param_description(description: str, param_schema: Dict[str, Any]) -> str: +def enrich_cohere_param_description(description: str, param_schema: dict[str, Any]) -> str: """Embed schema constraints into a Cohere parameter description. ``CohereParameterDefinition`` only has ``type``, ``description``, and diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index 44f5d941db4..834ffe1867e 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -22,7 +22,7 @@ Supported models: Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText """ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -87,7 +87,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): - ``dimensions``: output embedding dimensions (cohere.embed-v4.0+) """ - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["dimensions"] def map_openai_params( @@ -107,11 +107,11 @@ class OCIEmbedConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if optional_params.get("oci_signer") is None: creds = resolve_oci_credentials(optional_params) @@ -140,12 +140,12 @@ class OCIEmbedConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base = get_oci_base_url(optional_params, api_base or litellm.api_base) return f"{base}/{OCI_API_VERSION}/actions/embedText" @@ -156,11 +156,11 @@ class OCIEmbedConfig(BaseEmbeddingConfig): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, bytes]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes]: return sign_oci_request( headers=headers, optional_params=optional_params, @@ -251,7 +251,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -309,7 +309,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: return OCIError(status_code=status_code, message=error_message) diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index 2917e86b58e..e9e60106d2d 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -4,9 +4,6 @@ from collections.abc import AsyncIterator, Iterator from typing import ( TYPE_CHECKING, Any, - List, - Optional, - Union, cast, ) @@ -88,40 +85,40 @@ class OllamaChatConfig(BaseConfig): - `template` (string): the full prompt or prompt template (overrides what is defined in the Modelfile) """ - mirostat: Optional[int] = None - mirostat_eta: Optional[float] = None - mirostat_tau: Optional[float] = None - num_ctx: Optional[int] = None - num_gqa: Optional[int] = None - num_thread: Optional[int] = None - repeat_last_n: Optional[int] = None - repeat_penalty: Optional[float] = None - seed: Optional[int] = None - tfs_z: Optional[float] = None - num_predict: Optional[int] = None - top_k: Optional[int] = None - system: Optional[str] = None - template: Optional[str] = None + mirostat: int | None = None + mirostat_eta: float | None = None + mirostat_tau: float | None = None + num_ctx: int | None = None + num_gqa: int | None = None + num_thread: int | None = None + repeat_last_n: int | None = None + repeat_penalty: float | None = None + seed: int | None = None + tfs_z: float | None = None + num_predict: int | None = None + top_k: int | None = None + system: str | None = None + template: str | None = None def __init__( self, - mirostat: Optional[int] = None, - mirostat_eta: Optional[float] = None, - mirostat_tau: Optional[float] = None, - num_ctx: Optional[int] = None, - num_gqa: Optional[int] = None, - num_thread: Optional[int] = None, - repeat_last_n: Optional[int] = None, - repeat_penalty: Optional[float] = None, - temperature: Optional[float] = None, - seed: Optional[int] = None, - stop: Optional[list] = None, - tfs_z: Optional[float] = None, - num_predict: Optional[int] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - system: Optional[str] = None, - template: Optional[str] = None, + mirostat: int | None = None, + mirostat_eta: float | None = None, + mirostat_tau: float | None = None, + num_ctx: int | None = None, + num_gqa: int | None = None, + num_thread: int | None = None, + repeat_last_n: int | None = None, + repeat_penalty: float | None = None, + temperature: float | None = None, + seed: int | None = None, + stop: list | None = None, + tfs_z: float | None = None, + num_predict: int | None = None, + top_k: int | None = None, + top_p: float | None = None, + system: str | None = None, + template: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -198,11 +195,11 @@ class OllamaChatConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is not None and "Authorization" not in headers: headers["Authorization"] = f"Bearer {api_key}" @@ -210,12 +207,12 @@ class OllamaChatConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -236,7 +233,7 @@ class OllamaChatConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -256,7 +253,7 @@ class OllamaChatConfig(BaseConfig): ): # avoid message serialization issues - https://github.com/BerriAI/litellm/issues/5319 m = m.model_dump(exclude_none=True) tool_calls = m.get("tool_calls") - new_tools: Optional[List[OllamaToolCall]] = None + new_tools: list[OllamaToolCall] | None = None if tool_calls is not None and isinstance(tool_calls, list): new_tools = [] for tool in tool_calls: @@ -323,12 +320,12 @@ class OllamaChatConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -372,7 +369,7 @@ class OllamaChatConfig(BaseConfig): content=None, tool_calls=[ { - "id": f"call_{str(uuid.uuid4())}", + "id": f"call_{uuid.uuid4()!s}", "function": { "name": function_call.get("name", litellm_params.get("function_name")), "arguments": json.dumps(function_call.get("arguments", function_call)), @@ -409,14 +406,14 @@ class OllamaChatConfig(BaseConfig): ) return model_response - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return OllamaError(status_code=status_code, message=error_message, headers=headers) def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return OllamaChatCompletionResponseIterator( streaming_response=streaming_response, @@ -429,7 +426,7 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): started_reasoning_content: bool = False finished_reasoning_content: bool = False - def _is_function_call_complete(self, function_args: Union[str, dict]) -> bool: + def _is_function_call_complete(self, function_args: str | dict) -> bool: if isinstance(function_args, dict): return True try: @@ -481,8 +478,8 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): tool_call["id"] = str(uuid.uuid4()) # PROCESS REASONING CONTENT - reasoning_content: Optional[str] = None - content: Optional[str] = None + reasoning_content: str | None = None + content: str | None = None if chunk["message"].get("thinking"): reasoning_content = chunk["message"].get("thinking") self.started_reasoning_content = True diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 21ff3612a49..83c697d7cb5 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -1,4 +1,4 @@ -from typing import Any, List, Optional, Union +from typing import Any import httpx @@ -7,7 +7,7 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException class OllamaError(BaseLLMException): - def __init__(self, status_code: int, message: str, headers: Union[dict, httpx.Headers]): + def __init__(self, status_code: int, message: str, headers: dict | httpx.Headers): super().__init__(status_code=status_code, message=message, headers=headers) @@ -53,7 +53,7 @@ class OllamaModelInfo(BaseLLMModelInfo): """ @staticmethod - def get_api_key(api_key=None) -> Optional[str]: + def get_api_key(api_key=None) -> str | None: """Get API key from environment variables or litellm configuration""" import os @@ -69,14 +69,14 @@ class OllamaModelInfo(BaseLLMModelInfo): ) @staticmethod - def get_api_base(api_base: Optional[str] = None) -> str: + def get_api_base(api_base: str | None = None) -> str: from litellm.secret_managers.main import get_secret_str # env var OLLAMA_API_BASE or default return api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434" @classmethod - def get_server_api_base(cls, api_base: Optional[str] = None) -> str: + def get_server_api_base(cls, api_base: str | None = None) -> str: api_base = cls.get_api_base(api_base).rstrip("/") for suffix in ( "/api/generate", @@ -90,7 +90,7 @@ class OllamaModelInfo(BaseLLMModelInfo): return api_base[: -len(suffix)] return api_base - def get_models(self, api_key=None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key=None, api_base: str | None = None) -> list[str]: """ List all models available on the Ollama server via /api/tags endpoint. """ @@ -159,7 +159,7 @@ class OllamaModelInfo(BaseLLMModelInfo): return "tools" in _template.lower() @staticmethod - def _get_max_tokens(ollama_model_info: dict) -> Optional[int]: + def _get_max_tokens(ollama_model_info: dict) -> int | None: _model_info: dict = ollama_model_info.get("model_info", {}) for key, value in _model_info.items(): @@ -170,8 +170,8 @@ class OllamaModelInfo(BaseLLMModelInfo): def get_runtime_model_info( self, model: str, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, ) -> dict[str, Any]: from litellm import module_level_client @@ -219,9 +219,9 @@ class OllamaModelInfo(BaseLLMModelInfo): def get_model_info( self, model: str, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - ) -> Optional[dict[str, Any]]: + api_base: str | None = None, + api_key: str | None = None, + ) -> dict[str, Any] | None: if self._is_static_ollama_model(model): return None return self.get_runtime_model_info(model=model, api_base=api_base, api_key=api_key) diff --git a/litellm/llms/ollama/completion/handler.py b/litellm/llms/ollama/completion/handler.py index 7f229be53ae..f3c20b8075b 100644 --- a/litellm/llms/ollama/completion/handler.py +++ b/litellm/llms/ollama/completion/handler.py @@ -4,16 +4,16 @@ Ollama /chat/completion calls handled in llm_http_handler.py [TODO]: migrate embeddings to a base handler as well. """ -from typing import Any, Dict, List +from typing import Any import litellm from litellm.types.utils import EmbeddingResponse def _prepare_ollama_embedding_payload( - model: str, prompts: List[str], optional_params: Dict[str, Any] -) -> Dict[str, Any]: - data: Dict[str, Any] = {"model": model, "input": prompts} + model: str, prompts: list[str], optional_params: dict[str, Any] +) -> dict[str, Any]: + data: dict[str, Any] = {"model": model, "input": prompts} special_optional_params = ["truncate", "options", "keep_alive", "dimensions"] for k, v in optional_params.items(): @@ -28,14 +28,14 @@ def _prepare_ollama_embedding_payload( def _process_ollama_embedding_response( response_json: dict, - prompts: List[str], + prompts: list[str], model: str, model_response: EmbeddingResponse, logging_obj: Any, encoding: Any, ) -> EmbeddingResponse: output_data = [] - embeddings: List[List[float]] = response_json["embeddings"] + embeddings: list[list[float]] = response_json["embeddings"] for idx, emb in enumerate(embeddings): output_data.append({"object": "embedding", "index": idx, "embedding": emb}) @@ -68,7 +68,7 @@ def _process_ollama_embedding_response( async def ollama_aembeddings( api_base: str, model: str, - prompts: List[str], + prompts: list[str], model_response: EmbeddingResponse, optional_params: dict, logging_obj: Any, @@ -95,7 +95,7 @@ async def ollama_aembeddings( def ollama_embeddings( api_base: str, model: str, - prompts: List[str], + prompts: list[str], optional_params: dict, model_response: EmbeddingResponse, logging_obj: Any, diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 051e4e2b28a..0add66827f8 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -1,7 +1,7 @@ import json import time from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any from httpx._models import Headers, Response @@ -81,45 +81,45 @@ class OllamaConfig(BaseConfig): - `template` (string): the full prompt or prompt template (overrides what is defined in the Modelfile) """ - mirostat: Optional[int] = None - mirostat_eta: Optional[float] = None - mirostat_tau: Optional[float] = None - num_ctx: Optional[int] = None - num_gqa: Optional[int] = None - num_gpu: Optional[int] = None - num_thread: Optional[int] = None - repeat_last_n: Optional[int] = None - repeat_penalty: Optional[float] = None - temperature: Optional[float] = None - seed: Optional[int] = None - stop: Optional[list] = None # stop is a list based on this - https://github.com/ollama/ollama/pull/442 - tfs_z: Optional[float] = None - num_predict: Optional[int] = None - top_k: Optional[int] = None - top_p: Optional[float] = None - system: Optional[str] = None - template: Optional[str] = None + mirostat: int | None = None + mirostat_eta: float | None = None + mirostat_tau: float | None = None + num_ctx: int | None = None + num_gqa: int | None = None + num_gpu: int | None = None + num_thread: int | None = None + repeat_last_n: int | None = None + repeat_penalty: float | None = None + temperature: float | None = None + seed: int | None = None + stop: list | None = None # stop is a list based on this - https://github.com/ollama/ollama/pull/442 + tfs_z: float | None = None + num_predict: int | None = None + top_k: int | None = None + top_p: float | None = None + system: str | None = None + template: str | None = None def __init__( self, - mirostat: Optional[int] = None, - mirostat_eta: Optional[float] = None, - mirostat_tau: Optional[float] = None, - num_ctx: Optional[int] = None, - num_gqa: Optional[int] = None, - num_gpu: Optional[int] = None, - num_thread: Optional[int] = None, - repeat_last_n: Optional[int] = None, - repeat_penalty: Optional[float] = None, - temperature: Optional[float] = None, - seed: Optional[int] = None, - stop: Optional[list] = None, - tfs_z: Optional[float] = None, - num_predict: Optional[int] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - system: Optional[str] = None, - template: Optional[str] = None, + mirostat: int | None = None, + mirostat_eta: float | None = None, + mirostat_tau: float | None = None, + num_ctx: int | None = None, + num_gqa: int | None = None, + num_gpu: int | None = None, + num_thread: int | None = None, + repeat_last_n: int | None = None, + repeat_penalty: float | None = None, + temperature: float | None = None, + seed: int | None = None, + stop: list | None = None, + tfs_z: float | None = None, + num_predict: int | None = None, + top_k: int | None = None, + top_p: float | None = None, + system: str | None = None, + template: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -130,7 +130,7 @@ class OllamaConfig(BaseConfig): def get_config(cls): return super().get_config() - def get_required_params(self) -> List[ProviderField]: + def get_required_params(self) -> list[ProviderField]: """For a given provider, return it's required fields with a description""" return [ ProviderField( @@ -197,7 +197,7 @@ class OllamaConfig(BaseConfig): _template: str = str(ollama_model_info.get("template", "") or "") return "tools" in _template.lower() - def _get_max_tokens(self, ollama_model_info: dict) -> Optional[int]: + def _get_max_tokens(self, ollama_model_info: dict) -> int | None: _model_info: dict = ollama_model_info.get("model_info", {}) for k, v in _model_info.items(): @@ -206,7 +206,7 @@ class OllamaConfig(BaseConfig): return None @staticmethod - def get_api_key() -> Optional[str]: + def get_api_key() -> str | None: """Get API key from environment variables or litellm configuration""" import os @@ -223,8 +223,8 @@ class OllamaConfig(BaseConfig): def get_model_info( self, model: str, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, ) -> Any: """ curl http://localhost:11434/api/show -d '{ @@ -233,7 +233,7 @@ class OllamaConfig(BaseConfig): """ return OllamaModelInfo().get_model_info(model=model, api_base=api_base, api_key=api_key) - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return OllamaError(status_code=status_code, message=error_message, headers=headers) def transform_response( @@ -243,12 +243,12 @@ class OllamaConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: from litellm.litellm_core_utils.prompt_templates.common_utils import ( _parse_content_for_reasoning, @@ -282,7 +282,7 @@ class OllamaConfig(BaseConfig): content=None, tool_calls=[ { - "id": f"call_{str(uuid.uuid4())}", + "id": f"call_{uuid.uuid4()!s}", "function": { "name": function_call["name"], "arguments": json.dumps(function_call["arguments"]), @@ -303,8 +303,8 @@ class OllamaConfig(BaseConfig): except json.JSONDecodeError: # If JSON parsing fails, treat as regular text response ## output parse reasoning content from response_text - reasoning_content: Optional[str] = None - content: Optional[str] = None + reasoning_content: str | None = None + content: str | None = None if response_text is not None: reasoning_content, content = _parse_content_for_reasoning(response_text) message = litellm.Message(content=content, reasoning_content=reasoning_content) @@ -344,7 +344,7 @@ class OllamaConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -397,22 +397,22 @@ class OllamaConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return headers def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -432,9 +432,9 @@ class OllamaConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return OllamaTextCompletionResponseIterator( streaming_response=streaming_response, @@ -444,15 +444,15 @@ class OllamaConfig(BaseConfig): class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): - def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): + def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): super().__init__(streaming_response, sync_stream, json_mode) self.started_reasoning_content: bool = False self.finished_reasoning_content: bool = False - def _handle_string_chunk(self, str_line: str) -> Union[GenericStreamingChunk, ModelResponseStream]: + def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream: return self.chunk_parser(json.loads(str_line)) - def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]: + def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: try: if "error" in chunk: raise Exception(f"Ollama Error - {chunk}") @@ -464,10 +464,10 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): text = "" is_finished = True finish_reason = "stop" - prompt_eval_count: Optional[int] = chunk.get("prompt_eval_count", None) - eval_count: Optional[int] = chunk.get("eval_count", None) + prompt_eval_count: int | None = chunk.get("prompt_eval_count", None) + eval_count: int | None = chunk.get("eval_count", None) - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None if prompt_eval_count is not None and eval_count is not None: usage = ChatCompletionUsageBlock( prompt_tokens=prompt_eval_count, @@ -482,8 +482,8 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): ) elif chunk["response"]: text = chunk["response"] - reasoning_content: Optional[str] = None - content: Optional[str] = None + reasoning_content: str | None = None + content: str | None = None if text is not None: if "" in text: text = text.replace("", "") diff --git a/litellm/llms/oobabooga/chat/oobabooga.py b/litellm/llms/oobabooga/chat/oobabooga.py index 531baa9b43d..98a53c35b84 100644 --- a/litellm/llms/oobabooga/chat/oobabooga.py +++ b/litellm/llms/oobabooga/chat/oobabooga.py @@ -1,6 +1,6 @@ import json from collections.abc import Callable -from typing import Any, Optional +from typing import Any import litellm from litellm.llms.custom_httpx.http_handler import _get_httpx_client @@ -15,7 +15,7 @@ oobabooga_config = OobaboogaConfig() def completion( model: str, messages: list, - api_base: Optional[str], + api_base: str | None, model_response: ModelResponse, print_verbose: Callable, encoding, @@ -90,8 +90,8 @@ def embedding( model: str, input: list, model_response: EmbeddingResponse, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, logging_obj: Any, optional_params: dict, encoding=None, diff --git a/litellm/llms/oobabooga/chat/transformation.py b/litellm/llms/oobabooga/chat/transformation.py index 608fbc5cb35..808f115a307 100644 --- a/litellm/llms/oobabooga/chat/transformation.py +++ b/litellm/llms/oobabooga/chat/transformation.py @@ -1,5 +1,5 @@ import time -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -23,7 +23,7 @@ class OobaboogaConfig(OpenAIGPTConfig): self, error_message: str, status_code: int, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ) -> BaseLLMException: return OobaboogaError(status_code=status_code, message=error_message, headers=headers) @@ -34,12 +34,12 @@ class OobaboogaConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LoggingClass, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -82,11 +82,11 @@ class OobaboogaConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = { "accept": "application/json", diff --git a/litellm/llms/oobabooga/common_utils.py b/litellm/llms/oobabooga/common_utils.py index 82f8cda9511..962ecd1a926 100644 --- a/litellm/llms/oobabooga/common_utils.py +++ b/litellm/llms/oobabooga/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +8,6 @@ class OobaboogaError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index f0a859deba0..415e415b54f 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -1,7 +1,5 @@ """Support for OpenAI gpt-5 model family.""" -from typing import Optional, Union - import litellm from litellm.utils import ( _is_explicitly_disabled_factory, @@ -12,8 +10,8 @@ from .gpt_transformation import OpenAIGPTConfig def _normalize_reasoning_effort_for_chat_completion( - value: Union[str, dict, None], -) -> Optional[str]: + value: str | dict | None, +) -> str | None: """Convert reasoning_effort to the string format expected by OpenAI chat completion API. The chat completion API expects a simple string: 'none', 'low', 'medium', 'high', or 'xhigh'. @@ -28,7 +26,7 @@ def _normalize_reasoning_effort_for_chat_completion( return None -def _get_effort_level(value: Union[str, dict, None]) -> Optional[str]: +def _get_effort_level(value: str | dict | None) -> str | None: """Extract the effective effort level from reasoning_effort (string or dict). Use this for guards that compare effort level (e.g. xhigh validation, "none" checks). @@ -268,30 +266,28 @@ class OpenAIGPT5Config(OpenAIGPTConfig): raise litellm.utils.UnsupportedParamsError( message=( "gpt-5.1/5.2/5.4 only support logprobs, top_p, top_logprobs when " - "reasoning_effort='none'. Current reasoning_effort='{}'. " + f"reasoning_effort='none'. Current reasoning_effort='{effective_effort}'. " "To drop unsupported params set `litellm.drop_params = True`" - ).format(effective_effort), + ), status_code=400, ) if "temperature" in non_default_params: - temperature_value: Optional[float] = non_default_params.pop("temperature") + temperature_value: float | None = non_default_params.pop("temperature") if temperature_value is not None: # models supporting reasoning_effort="none" also support flexible temperature - if supports_none and (effective_effort == "none" or effective_effort is None): - optional_params["temperature"] = temperature_value - elif temperature_value == 1: + if supports_none and (effective_effort == "none" or effective_effort is None) or temperature_value == 1: optional_params["temperature"] = temperature_value elif litellm.drop_params or drop_params: pass else: raise litellm.utils.UnsupportedParamsError( message=( - "gpt-5 models (including gpt-5-codex) don't support temperature={}. " + f"gpt-5 models (including gpt-5-codex) don't support temperature={temperature_value}. " "Only temperature=1 is supported. " "For gpt-5.1, temperature is supported when reasoning_effort='none' (or not specified, as it defaults to 'none'). " "To drop unsupported params set `litellm.drop_params = True`" - ).format(temperature_value), + ), status_code=400, ) return super()._map_openai_params( diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index ff8bd446962..a37e15c1f86 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -8,12 +8,8 @@ from collections.abc import AsyncIterator, Coroutine, Iterator from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Union, cast, overload, ) @@ -98,31 +94,31 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): # Add a class variable to track if this is the base class _is_base_class = True - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -327,24 +323,24 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: + ) -> list[AllMessageValues]: ... # fmt: on def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """OpenAI no longer supports image_url as a string, so we need to convert it to a dict""" async def _async_transform(): @@ -353,7 +349,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): message_role = message.get("role") if message_role == "user" and message_content and isinstance(message_content, list): - message_content_types = cast(List[OpenAIMessageContentListBlock], message_content) + message_content_types = cast(list[OpenAIMessageContentListBlock], message_content) for i, content_item in enumerate(message_content_types): message_content_types[i] = await self._async_transform_content_item( cast(OpenAIMessageContentListBlock, content_item), @@ -367,7 +363,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): message_content = message.get("content") message_role = message.get("role") if message_role == "user" and message_content and isinstance(message_content, list): - message_content_types = cast(List[OpenAIMessageContentListBlock], message_content) + message_content_types = cast(list[OpenAIMessageContentListBlock], message_content) for i, content_item in enumerate(message_content): message_content_types[i] = self._transform_content_item( cast(OpenAIMessageContentListBlock, content_item) @@ -377,9 +373,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, # allows overrides to selectively run this - messages: List[AllMessageValues], - tools: Optional[List["ChatCompletionToolParam"]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + messages: list[AllMessageValues], + tools: list["ChatCompletionToolParam"] | None = None, + ) -> tuple[list[AllMessageValues], list["ChatCompletionToolParam"] | None]: from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, ) @@ -422,7 +418,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -454,7 +450,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -488,7 +484,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def _check_and_fix_if_content_is_tool_call( self, content: str, optional_params: dict - ) -> Optional[ChatCompletionMessageToolCall]: + ) -> ChatCompletionMessageToolCall | None: """ Check if the content is a tool call """ @@ -519,16 +515,16 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def _transform_choices( self, - choices: List[OpenAIChatCompletionChoices], - json_mode: Optional[bool] = None, - optional_params: Optional[dict] = None, - ) -> List[Choices]: + choices: list[OpenAIChatCompletionChoices], + json_mode: bool | None = None, + optional_params: dict | None = None, + ) -> list[Choices]: transformed_choices = [] for choice in choices: ## HANDLE JSON MODE - anthropic returns single function call] tool_calls = choice["message"].get("tool_calls", None) - new_tool_calls: Optional[List[ChatCompletionMessageToolCall]] = None + new_tool_calls: list[ChatCompletionMessageToolCall] | None = None message_content = choice["message"].get("content", None) if tool_calls is not None: _openai_tool_calls = [] @@ -545,14 +541,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): choice["message"]["content"] = None # remove the content new_tool_calls = [new_tool_call] - translated_message: Optional[Message] = None - finish_reason: Optional[str] = None + translated_message: Message | None = None + finish_reason: str | None = None if new_tool_calls and _should_convert_tool_call_to_json_mode( tool_calls=new_tool_calls, convert_tool_call_to_json_mode=json_mode, ): # to support response_format on claude models - json_mode_content_str: Optional[str] = str(new_tool_calls[0]["function"].get("arguments", "")) or None + json_mode_content_str: str | None = str(new_tool_calls[0]["function"].get("arguments", "")) or None if json_mode_content_str is not None: translated_message = Message(content=json_mode_content_str) finish_reason = "stop" @@ -597,12 +593,12 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the response from the API. @@ -625,7 +621,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): except Exception as e: response_headers = getattr(raw_response, "headers", None) raise OpenAIError( - message="Unable to get json response - {}, Original Response: {}".format(str(e), raw_response.text), + message=f"Unable to get json response - {e!s}, Original Response: {raw_response.text}", status_code=raw_response.status_code, headers=response_headers, ) @@ -639,9 +635,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): return cast(ModelResponse, final_response_obj) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OpenAIError( status_code=status_code, message=error_message, @@ -650,12 +644,12 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the API call. @@ -680,11 +674,11 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is not None: headers["Authorization"] = f"Bearer {api_key}" @@ -695,7 +689,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): return headers - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: """ Calls OpenAI's `/v1/models` endpoint and returns the list of models. """ @@ -723,11 +717,11 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): return [model["id"] for model in models] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return ( api_base or litellm.api_base @@ -737,7 +731,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ) @staticmethod - def get_base_model(model: Optional[str] = None) -> Optional[str]: + def get_base_model(model: str | None = None) -> str | None: return model def get_token_counter(self) -> Optional["BaseTokenCounter"]: @@ -749,9 +743,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return OpenAIChatCompletionStreamingHandler( streaming_response=streaming_response, @@ -781,7 +775,7 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): return choices @staticmethod - def _extract_error_from_chunk(chunk: dict) -> Optional[tuple[str, int]]: + def _extract_error_from_chunk(chunk: dict) -> tuple[str, int] | None: """OpenAI-compatible backends (vLLM, sglang) can return an HTTP 200 stream whose body carries an error payload, e.g. ``data: {"error": {"message": "...", "code": 400}}``.""" @@ -807,7 +801,7 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): choices = chunk.get("choices", []) choices = self._map_reasoning_to_reasoning_content(choices) - kwargs: Dict[str, Any] = { + kwargs: dict[str, Any] = { "id": chunk.get("id"), "object": "chat.completion.chunk", "created": chunk.get("created"), diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index fd2c9339248..836634ccc01 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -14,7 +14,7 @@ Pattern Overview: This pattern can be replicated for other message formats (e.g., Anthropic). """ -from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Union, cast import litellm from litellm._logging import verbose_proxy_logger @@ -56,7 +56,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): Methods can be overridden to customize behavior for different message formats. """ - def get_structured_messages(self, data: dict) -> List[AllMessageValues] | None: + def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None: """ Convert chat completions request data to OpenAI-spec structured messages. @@ -65,7 +65,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): messages = data.get("messages") if messages is None: return None - return cast(List[AllMessageValues], messages) + return cast(list[AllMessageValues], messages) async def process_input_messages( self, @@ -83,11 +83,11 @@ class OpenAIChatCompletionsHandler(BaseTranslation): skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply) skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply) - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - tool_calls_to_check: List[ChatCompletionToolParam] = [] - text_task_mappings: List[Tuple[int, int | None]] = [] - tool_call_task_mappings: List[Tuple[int, int]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + tool_calls_to_check: list[ChatCompletionToolParam] = [] + text_task_mappings: list[tuple[int, int | None]] = [] + tool_call_task_mappings: list[tuple[int, int]] = [] # Step 1: Extract all text content, images, and tool calls for msg_idx, message in enumerate(messages): @@ -170,9 +170,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): return data - def extract_request_tool_names(self, data: dict) -> List[str]: + def extract_request_tool_names(self, data: dict) -> list[str]: """Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name).""" - names: List[str] = [] + names: list[str] = [] for tool in data.get("tools") or []: if isinstance(tool, dict) and tool.get("type") == "function": fn = tool.get("function") @@ -185,13 +185,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def _extract_inputs( self, - message: Dict[str, Any], + message: dict[str, Any], msg_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - tool_calls_to_check: List[ChatCompletionToolParam], - text_task_mappings: List[Tuple[int, int | None]], - tool_call_task_mappings: List[Tuple[int, int]], + texts_to_check: list[str], + images_to_check: list[str], + tool_calls_to_check: list[ChatCompletionToolParam], + text_task_mappings: list[tuple[int, int | None]], + tool_call_task_mappings: list[tuple[int, int]], skip_system_message: bool = False, skip_tool_message: bool = False, ) -> None: @@ -243,9 +243,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_input_texts( self, - messages: List[Dict[str, Any]], - responses: List[str], - task_mappings: List[Tuple[int, int | None]], + messages: list[dict[str, Any]], + responses: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to input message text content. @@ -272,9 +272,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_input_tool_calls( self, - messages: List[Dict[str, Any]], - tool_calls: List[Dict[str, Any]], - task_mappings: List[Tuple[int, int]], + messages: list[dict[str, Any]], + tool_calls: list[dict[str, Any]], + task_mappings: list[tuple[int, int]], ) -> None: """ Apply guardrailed tool calls back to input messages. @@ -323,11 +323,11 @@ class OpenAIChatCompletionsHandler(BaseTranslation): verbose_proxy_logger.warning("OpenAI Chat Completions: No text content in response, skipping guardrail") return response - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - tool_calls_to_check: List[Dict[str, Any]] = [] - text_task_mappings: List[Tuple[int, int | None]] = [] - tool_call_task_mappings: List[Tuple[int, int]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + tool_calls_to_check: list[dict[str, Any]] = [] + text_task_mappings: list[tuple[int, int | None]] = [] + tool_call_task_mappings: list[tuple[int, int]] = [] # text_task_mappings: Track (choice_index, content_index) for each text # content_index is None for string content, int for list content # tool_call_task_mappings: Track (choice_index, tool_call_index) for each tool call @@ -378,8 +378,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): guardrailed_texts = guardrailed_inputs.get("texts", []) returned_tool_calls = guardrailed_inputs.get("tool_calls") - guardrailed_tool_calls: List[Dict[str, Any]] = ( - cast(List[Dict[str, Any]], returned_tool_calls) + guardrailed_tool_calls: list[dict[str, Any]] = ( + cast(list[dict[str, Any]], returned_tool_calls) if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check) else tool_calls_to_check ) @@ -406,13 +406,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def process_output_streaming_response( self, - responses_so_far: List["ModelResponseStream"], + responses_so_far: list["ModelResponseStream"], guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Any | None = None, user_api_key_dict: Any | None = None, request_data: dict | None = None, stream_transform_sink: StreamTransformSink | None = None, - ) -> List["ModelResponseStream"]: + ) -> list["ModelResponseStream"]: """ Process output streaming responses by applying guardrails to text content. @@ -508,9 +508,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): combined_texts = self._combine_streaming_texts(responses_so_far) # Step 2: Create lists for guardrail processing - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - task_mappings: List[Tuple[int, int | None]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + task_mappings: list[tuple[int, int | None]] = [] # Track (choice_index, content_index) for each combined text for (map_choice_idx, map_content_idx), combined_text in combined_texts.items(): @@ -664,8 +664,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): } def _combine_streaming_texts( - self, responses_so_far: List["ModelResponseStream"] - ) -> Dict[Tuple[int, int | None], str]: + self, responses_so_far: list["ModelResponseStream"] + ) -> dict[tuple[int, int | None], str]: """ Combine all streaming chunks into complete text per choice. @@ -677,7 +677,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): Returns: Dict mapping (choice_idx, content_idx) to combined text string """ - combined_texts: Dict[Tuple[int, int | None], str] = {} + combined_texts: dict[tuple[int, int | None], str] = {} for response_idx, response in enumerate(responses_so_far): for choice_idx, choice in enumerate(response.choices): @@ -693,7 +693,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): if isinstance(content, str): # String content - accumulate for this choice - str_key: Tuple[int, int | None] = (choice_idx, None) + str_key: tuple[int, int | None] = (choice_idx, None) if str_key not in combined_texts: combined_texts[str_key] = "" combined_texts[str_key] += content @@ -703,7 +703,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): for content_idx, content_item in enumerate(content): text_str = content_item.get("text") if text_str: - list_key: Tuple[int, int | None] = ( + list_key: tuple[int, int | None] = ( choice_idx, content_idx, ) @@ -745,13 +745,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def _extract_output_text_images_and_tool_calls( self, - choice: Union[Choices, StreamingChoices], + choice: Choices | StreamingChoices, choice_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - tool_calls_to_check: List[Dict[str, Any]], - text_task_mappings: List[Tuple[int, int | None]], - tool_call_task_mappings: List[Tuple[int, int]], + texts_to_check: list[str], + images_to_check: list[str], + tool_calls_to_check: list[dict[str, Any]], + text_task_mappings: list[tuple[int, int | None]], + tool_call_task_mappings: list[tuple[int, int]], ) -> None: """ Extract text content, images, and tool calls from a response choice. @@ -762,7 +762,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): # Determine content source and tool calls based on choice type content = None - tool_calls: List[Any] | None = None + tool_calls: list[Any] | None = None if isinstance(choice, litellm.Choices): content = choice.message.content tool_calls = choice.message.tool_calls @@ -805,7 +805,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): tool_calls_to_check.append(tool_call_dict) tool_call_task_mappings.append((choice_idx, int(tool_call_idx))) - def _convert_tool_call_to_dict(self, tool_call: Union[Dict[str, Any], Any]) -> Dict[str, Any] | None: + def _convert_tool_call_to_dict(self, tool_call: dict[str, Any] | Any) -> dict[str, Any] | None: """ Convert a tool call object to dictionary format. @@ -833,8 +833,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_output_texts( self, response: "ModelResponse", - responses: List[str], - task_mappings: List[Tuple[int, int | None]], + responses: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail text responses back to output response. @@ -864,8 +864,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_output_tool_calls( self, response: "ModelResponse", - tool_calls: List[Dict[str, Any]], - task_mappings: List[Tuple[int, int]], + tool_calls: list[dict[str, Any]], + task_mappings: list[tuple[int, int]], ) -> None: """ Apply guardrailed tool calls back to the output response. @@ -896,9 +896,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_output_streaming( self, - responses: List["ModelResponseStream"], - guardrailed_texts: List[str], - task_mappings: List[Tuple[int, int | None]], + responses: list["ModelResponseStream"], + guardrailed_texts: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to output streaming responses. @@ -914,7 +914,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): Override this method to customize how responses are applied to streaming responses. """ # Build a mapping of what guardrailed text to use for each (choice_idx, content_idx) - guardrail_map: Dict[Tuple[int, int | None], str] = {} + guardrail_map: dict[tuple[int, int | None], str] = {} for task_idx, guardrail_response in enumerate(guardrailed_texts): mapping = task_mappings[task_idx] choice_idx = cast(int, mapping[0]) @@ -923,7 +923,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): # Track which choices we've already set the guardrailed text for # Key: (choice_idx, content_idx), Value: boolean (True if already set) - already_set: Dict[Tuple[int, int | None], bool] = {} + already_set: dict[tuple[int, int | None], bool] = {} # Iterate through all responses and update content for response_idx, response in enumerate(responses): @@ -940,7 +940,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): if isinstance(content, str): # String content - str_key: Tuple[int, int | None] = (choice_idx_in_response, None) + str_key: tuple[int, int | None] = (choice_idx_in_response, None) if str_key in guardrail_map: if str_key not in already_set: # First chunk - set the complete guardrailed text @@ -960,7 +960,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): # List content - handle each content item for content_idx, content_item in enumerate(content): if "text" in content_item: - list_key: Tuple[int, int | None] = ( + list_key: tuple[int, int | None] = ( choice_idx_in_response, content_idx, ) diff --git a/litellm/llms/openai/chat/o_series_transformation.py b/litellm/llms/openai/chat/o_series_transformation.py index 6008ad332da..0aaf0315f2a 100644 --- a/litellm/llms/openai/chat/o_series_transformation.py +++ b/litellm/llms/openai/chat/o_series_transformation.py @@ -12,7 +12,7 @@ Translations handled by LiteLLM: """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Union, cast, overload +from typing import Any, Literal, cast, overload import litellm from litellm import verbose_logger @@ -37,7 +37,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): def get_config(cls): return super().get_config() - def translate_developer_role_to_system_role(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def translate_developer_role_to_system_role(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ O-series models support `developer` role. """ @@ -98,7 +98,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): if "max_tokens" in non_default_params: optional_params["max_completion_tokens"] = non_default_params.pop("max_tokens") if "temperature" in non_default_params: - temperature_value: Optional[float] = non_default_params.pop("temperature") + temperature_value: float | None = non_default_params.pop("temperature") if temperature_value is not None: if temperature_value == 1: optional_params["temperature"] = temperature_value @@ -108,9 +108,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): pass else: raise litellm.utils.UnsupportedParamsError( - message="O-series models don't support temperature={}. Only temperature=1 is supported. To drop unsupported openai params from the call, set `litellm.drop_params = True`".format( - temperature_value - ), + message=f"O-series models don't support temperature={temperature_value}. Only temperature=1 is supported. To drop unsupported openai params from the call, set `litellm.drop_params = True`", status_code=400, ) @@ -127,20 +125,20 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Handles limitations of O-1 model family. - modalities: image => drop param (if user opts in to dropping param) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 6731d4a6a4a..e72680f387d 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -10,13 +10,9 @@ import ssl from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, NamedTuple, Optional, - Tuple, - Union, ) import httpx @@ -35,13 +31,13 @@ from litellm.llms.custom_httpx.http_handler import ( ) -def _get_client_init_params(cls: type) -> Tuple[str, ...]: +def _get_client_init_params(cls: type) -> tuple[str, ...]: """Extract __init__ parameter names (excluding 'self') from a class.""" return tuple(p for p in inspect.signature(cls.__init__).parameters if p != "self") # type: ignore[misc] -_OPENAI_INIT_PARAMS: Tuple[str, ...] = _get_client_init_params(OpenAI) -_AZURE_OPENAI_INIT_PARAMS: Tuple[str, ...] = _get_client_init_params(AzureOpenAI) +_OPENAI_INIT_PARAMS: tuple[str, ...] = _get_client_init_params(OpenAI) +_AZURE_OPENAI_INIT_PARAMS: tuple[str, ...] = _get_client_init_params(AzureOpenAI) class OpenAIError(BaseLLMException): @@ -49,10 +45,10 @@ class OpenAIError(BaseLLMException): self, status_code: int, message: str, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[dict, httpx.Headers]] = None, - body: Optional[dict] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: dict | httpx.Headers | None = None, + body: dict | None = None, ): self.status_code = status_code self.message = message @@ -78,9 +74,9 @@ class OpenAIError(BaseLLMException): ####### Error Handling Utils for OpenAI API ####################### ################################################################### def drop_params_from_unprocessable_entity_error( - e: Union[openai.UnprocessableEntityError, httpx.HTTPStatusError], - data: Dict[str, Any], -) -> Dict[str, Any]: + e: openai.UnprocessableEntityError | httpx.HTTPStatusError, + data: dict[str, Any], +) -> dict[str, Any]: """ Helper function to read OpenAI UnprocessableEntityError and drop the params that raised an error from the error message. @@ -91,7 +87,7 @@ def drop_params_from_unprocessable_entity_error( Returns: Dict[str, Any]: A new dictionary with invalid parameters removed """ - invalid_params: List[str] = [] + invalid_params: list[str] = [] if isinstance(e, httpx.HTTPStatusError): error_json = e.response.json() error_message = error_json.get("error", {}) @@ -107,7 +103,7 @@ def drop_params_from_unprocessable_entity_error( message = {"detail": message} detail = message.get("detail") - if isinstance(detail, List) and len(detail) > 0 and isinstance(detail[0], dict): + if isinstance(detail, list) and len(detail) > 0 and isinstance(detail[0], dict): for error_dict in detail: if ( error_dict.get("loc") @@ -129,7 +125,7 @@ class BaseOpenAILLM: @staticmethod def get_cached_openai_client( client_initialization_params: dict, client_type: Literal["openai", "azure"] - ) -> Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]]: + ) -> OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None: """Retrieves the OpenAI client from the in-memory cache based on the client initialization parameters""" _cache_key = BaseOpenAILLM.get_openai_client_cache_key( client_initialization_params=client_initialization_params, @@ -140,7 +136,7 @@ class BaseOpenAILLM: @staticmethod def set_cached_openai_client( - openai_client: Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI], + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI, client_type: Literal["openai", "azure"], client_initialization_params: dict, ): @@ -190,7 +186,7 @@ class BaseOpenAILLM: @staticmethod def get_openai_client_initialization_param_fields( client_type: Literal["openai", "azure"], - ) -> Tuple[str, ...]: + ) -> tuple[str, ...]: """Returns a tuple of fields that are used to initialize the OpenAI client""" if client_type == "openai": return _OPENAI_INIT_PARAMS @@ -200,7 +196,7 @@ class BaseOpenAILLM: @staticmethod def _get_async_http_client( shared_session: Optional["ClientSession"] = None, - ) -> Optional[httpx.AsyncClient]: + ) -> httpx.AsyncClient | None: if litellm.aclient_session is not None: return litellm.aclient_session @@ -223,7 +219,7 @@ class BaseOpenAILLM: ) @staticmethod - def _get_sync_http_client() -> Optional[httpx.Client]: + def _get_sync_http_client() -> httpx.Client | None: if litellm.client_session is not None: return litellm.client_session @@ -243,14 +239,14 @@ class BaseOpenAILLM: class OpenAICredentials(NamedTuple): api_base: str - api_key: Optional[str] - organization: Optional[str] + api_key: str | None + organization: str | None def get_openai_credentials( - api_base: Optional[str] = None, - api_key: Optional[str] = None, - organization: Optional[str] = None, + api_base: str | None = None, + api_key: str | None = None, + organization: str | None = None, ) -> OpenAICredentials: """Resolve OpenAI credentials from params, litellm globals, and env vars.""" resolved_api_base = ( diff --git a/litellm/llms/openai/completion/guardrail_translation/__init__.py b/litellm/llms/openai/completion/guardrail_translation/__init__.py index 51e43c45937..fb6627d1594 100644 --- a/litellm/llms/openai/completion/guardrail_translation/__init__.py +++ b/litellm/llms/openai/completion/guardrail_translation/__init__.py @@ -10,4 +10,4 @@ guardrail_translation_mappings = { CallTypes.atext_completion: OpenAITextCompletionHandler, } -__all__ = ["guardrail_translation_mappings", "OpenAITextCompletionHandler"] +__all__ = ["OpenAITextCompletionHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/openai/completion/guardrail_translation/handler.py b/litellm/llms/openai/completion/guardrail_translation/handler.py index 8537fefe1e2..6f531644fd6 100644 --- a/litellm/llms/openai/completion/guardrail_translation/handler.py +++ b/litellm/llms/openai/completion/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for OpenAI's text completion The handler processes the 'prompt' parameter for guardrails. """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -33,7 +33,7 @@ class OpenAITextCompletionHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input prompt by applying guardrails to text content. @@ -120,9 +120,9 @@ class OpenAITextCompletionHandler(BaseTranslation): self, response: "TextCompletionResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response by applying guardrails to completion text. diff --git a/litellm/llms/openai/completion/handler.py b/litellm/llms/openai/completion/handler.py index d4c8139adc3..358cb08867d 100644 --- a/litellm/llms/openai/completion/handler.py +++ b/litellm/llms/openai/completion/handler.py @@ -1,6 +1,5 @@ import json from collections.abc import Callable -from typing import List, Optional, Union from openai import AsyncOpenAI, OpenAI @@ -35,19 +34,19 @@ class OpenAITextCompletion(BaseLLM): model_response: ModelResponse, api_key: str, model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], timeout: float, custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, optional_params: dict, - print_verbose: Optional[Callable] = None, - api_base: Optional[str] = None, + print_verbose: Callable | None = None, + api_base: str | None = None, acompletion: bool = False, litellm_params=None, logger_fn=None, client=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, + organization: str | None = None, + headers: dict | None = None, ): try: if headers: @@ -173,7 +172,7 @@ class OpenAITextCompletion(BaseLLM): model: str, timeout: float, max_retries: int, - organization: Optional[str] = None, + organization: str | None = None, client=None, ): try: @@ -224,7 +223,7 @@ class OpenAITextCompletion(BaseLLM): model_response: ModelResponse, model: str, timeout: float, - api_base: Optional[str] = None, + api_base: str | None = None, max_retries=None, client=None, organization=None, @@ -282,7 +281,7 @@ class OpenAITextCompletion(BaseLLM): model: str, timeout: float, max_retries: int, - api_base: Optional[str] = None, + api_base: str | None = None, client=None, organization=None, ): diff --git a/litellm/llms/openai/completion/transformation.py b/litellm/llms/openai/completion/transformation.py index 77dc0b54fe0..ff1af891a05 100644 --- a/litellm/llms/openai/completion/transformation.py +++ b/litellm/llms/openai/completion/transformation.py @@ -2,8 +2,6 @@ Support for gpt model family """ -from typing import List, Optional, Union - from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUserMessage from litellm.types.utils import Choices, Message, ModelResponse, TextCompletionResponse @@ -43,31 +41,31 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): - `top_p` (number or null): An alternative to sampling with temperature, used for nucleus sampling. """ - best_of: Optional[int] = None - echo: Optional[bool] = None - frequency_penalty: Optional[int] = None - logit_bias: Optional[dict] = None - logprobs: Optional[int] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - suffix: Optional[str] = None + best_of: int | None = None + echo: bool | None = None + frequency_penalty: int | None = None + logit_bias: dict | None = None + logprobs: int | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + suffix: str | None = None def __init__( self, - best_of: Optional[int] = None, - echo: Optional[bool] = None, - frequency_penalty: Optional[int] = None, - logit_bias: Optional[dict] = None, - logprobs: Optional[int] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - suffix: Optional[str] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, + best_of: int | None = None, + echo: bool | None = None, + frequency_penalty: int | None = None, + logit_bias: dict | None = None, + logprobs: int | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + suffix: str | None = None, + temperature: float | None = None, + top_p: float | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -80,14 +78,14 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): def convert_to_chat_model_response_object( self, - response_object: Optional[TextCompletionResponse] = None, - model_response_object: Optional[ModelResponse] = None, + response_object: TextCompletionResponse | None = None, + model_response_object: ModelResponse | None = None, ): try: ## RESPONSE OBJECT if response_object is None or model_response_object is None: raise ValueError("Error in response object format") - choice_list: List[Choices] = [] + choice_list: list[Choices] = [] for idx, choice in enumerate(response_object["choices"]): message = Message( content=choice["text"], @@ -118,7 +116,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): except Exception as e: raise e - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "functions", "function_call", @@ -146,7 +144,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): def transform_text_completion_request( self, model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], optional_params: dict, headers: dict, ) -> dict: diff --git a/litellm/llms/openai/completion/utils.py b/litellm/llms/openai/completion/utils.py index a7b7e7a67ce..e4ab74fe1f1 100644 --- a/litellm/llms/openai/completion/utils.py +++ b/litellm/llms/openai/completion/utils.py @@ -1,4 +1,4 @@ -from typing import List, Union, cast +from typing import cast from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, @@ -10,7 +10,7 @@ from litellm.types.llms.openai import ( ) -def is_tokens_or_list_of_tokens(value: List): +def is_tokens_or_list_of_tokens(value: list): # Check if it's a list of integers (tokens) if isinstance(value, list) and all(isinstance(item, int) for item in value): return True @@ -23,7 +23,7 @@ def is_tokens_or_list_of_tokens(value: List): def _transform_prompt( - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], ) -> AllPromptValues: if len(messages) == 1: # base case message_content = messages[0].get("content") @@ -34,7 +34,7 @@ def _transform_prompt( content = convert_content_list_to_str(cast(AllMessageValues, messages[0])) openai_prompt += content else: - prompt_str_list: List[str] = [] + prompt_str_list: list[str] = [] for m in messages: try: # expect list of int/list of list of int to be a 1 message array only. content = convert_content_list_to_str(cast(AllMessageValues, m)) diff --git a/litellm/llms/openai/containers/transformation.py b/litellm/llms/openai/containers/transformation.py index b5f4334af0a..6559b4b1d7b 100644 --- a/litellm/llms/openai/containers/transformation.py +++ b/litellm/llms/openai/containers/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -51,14 +51,14 @@ class OpenAIContainerConfig(BaseContainerConfig): self, container_create_optional_params: ContainerCreateOptionalRequestParams, drop_params: bool, - ) -> Dict: + ) -> dict: """No mapping applied since inputs are in OpenAI spec already""" return dict(container_create_optional_params) def validate_environment( self, headers: dict, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.update( @@ -70,7 +70,7 @@ class OpenAIContainerConfig(BaseContainerConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """Get the complete URL for OpenAI container API.""" @@ -87,10 +87,10 @@ class OpenAIContainerConfig(BaseContainerConfig): def transform_container_create_request( self, name: str, - container_create_optional_request_params: Dict, + container_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """Transform the container creation request for OpenAI API.""" # Remove extra_headers from optional params as they're handled separately container_create_optional_request_params = { @@ -137,11 +137,11 @@ class OpenAIContainerConfig(BaseContainerConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """Transform the container list request for OpenAI API. OpenAI API expects the following request: @@ -184,14 +184,14 @@ class OpenAIContainerConfig(BaseContainerConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform the OpenAI container retrieve request.""" # For container retrieve, we just need to construct the URL encoded_container_id = encode_url_path_segment(container_id, field_name="container_id") url = join_container_api_base_path(api_base, f"/{encoded_container_id}") # No additional data needed for GET request - data: Dict[str, Any] = {} + data: dict[str, Any] = {} return url, data @@ -213,7 +213,7 @@ class OpenAIContainerConfig(BaseContainerConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform the container delete request for OpenAI API. OpenAI API expects the following request: @@ -224,7 +224,7 @@ class OpenAIContainerConfig(BaseContainerConfig): url = join_container_api_base_path(api_base, f"/{encoded_container_id}") # No data needed for DELETE request - data: Dict[str, Any] = {} + data: dict[str, Any] = {} return url, data @@ -247,11 +247,11 @@ class OpenAIContainerConfig(BaseContainerConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """Transform the container file list request for OpenAI API. OpenAI API expects the following request: @@ -262,7 +262,7 @@ class OpenAIContainerConfig(BaseContainerConfig): url = join_container_api_base_path(api_base, f"/{encoded_container_id}/files") # Prepare query parameters - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if after is not None: params["after"] = after if limit is not None: @@ -296,7 +296,7 @@ class OpenAIContainerConfig(BaseContainerConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Transform the container file content request for OpenAI API. OpenAI API expects the following request: @@ -308,7 +308,7 @@ class OpenAIContainerConfig(BaseContainerConfig): url = join_container_api_base_path(api_base, f"/{encoded_container_id}/files/{encoded_file_id}/content") # No query parameters needed - params: Dict[str, Any] = {} + params: dict[str, Any] = {} return url, params @@ -327,7 +327,7 @@ class OpenAIContainerConfig(BaseContainerConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: from ...base_llm.chat.transformation import BaseLLMException diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index ffb8dbb9820..87568a4d399 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -4,7 +4,7 @@ Helper util for handling openai-specific cost calculation """ from collections.abc import Mapping -from typing import Any, Literal, Optional, Tuple +from typing import Any, Literal from litellm._logging import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token @@ -22,9 +22,9 @@ def cost_router(call_type: CallTypes) -> Literal["cost_per_token", "cost_per_sec def cost_per_token( model: str, usage: Usage, - service_tier: Optional[str] = None, - data_residency: Optional[str] = None, -) -> Tuple[float, float]: + service_tier: str | None = None, + data_residency: str | None = None, +) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -91,7 +91,7 @@ def cost_per_token( # return prompt_cost, completion_cost -def cost_per_second(model: str, custom_llm_provider: Optional[str], duration: float = 0.0) -> Tuple[float, float]: +def cost_per_second(model: str, custom_llm_provider: str | None, duration: float = 0.0) -> tuple[float, float]: """ Calculates the cost per second for a given model, prompt tokens, and completion tokens. @@ -126,7 +126,7 @@ def cost_per_second(model: str, custom_llm_provider: Optional[str], duration: fl return prompt_cost, completion_cost -def _video_resolution_to_cost_field_suffix(resolution: str) -> Optional[str]: +def _video_resolution_to_cost_field_suffix(resolution: str) -> str | None: """ Map usage resolution to a safe suffix for ``output_cost_per_second_`` keys. @@ -146,8 +146,8 @@ def _video_resolution_to_cost_field_suffix(resolution: str) -> Optional[str]: def _video_output_cost_per_second( model_info: Mapping[str, Any], - video_resolution: Optional[str], -) -> Optional[float]: + video_resolution: str | None, +) -> float | None: """ Per-second video output rate from model_info. @@ -172,9 +172,9 @@ def _video_output_cost_per_second( def video_generation_cost( model: str, duration_seconds: float, - custom_llm_provider: Optional[str] = None, - model_info: Optional[ModelInfo] = None, - video_resolution: Optional[str] = None, + custom_llm_provider: str | None = None, + model_info: ModelInfo | None = None, + video_resolution: str | None = None, ) -> float: """ Calculates the cost for video generation based on duration in seconds. diff --git a/litellm/llms/openai/data_residency.py b/litellm/llms/openai/data_residency.py index db3c49d7583..d84b0add468 100644 --- a/litellm/llms/openai/data_residency.py +++ b/litellm/llms/openai/data_residency.py @@ -7,20 +7,19 @@ enabled and rejects requests sent to the wrong host, so the api_base hostname is the authoritative signal of which region a request was processed in. """ -from typing import Dict, Optional from urllib.parse import urlparse # Mapping of OpenAI regional hostnames to the corresponding data-residency # value used by the cost calculator. See # https://developers.openai.com/api/docs/pricing for the regional-processing # uplift these hostnames trigger. -_OPENAI_REGIONAL_HOSTS: Dict[str, str] = { +_OPENAI_REGIONAL_HOSTS: dict[str, str] = { "eu.api.openai.com": "eu", "us.api.openai.com": "us", } -def infer_openai_data_residency(custom_llm_provider: Optional[str], api_base: Optional[str]) -> Optional[str]: +def infer_openai_data_residency(custom_llm_provider: str | None, api_base: str | None) -> str | None: """ Derive the OpenAI data-residency region from an api_base URL. diff --git a/litellm/llms/openai/embeddings/guardrail_translation/__init__.py b/litellm/llms/openai/embeddings/guardrail_translation/__init__.py index a60662282ca..d4f842d9dd3 100644 --- a/litellm/llms/openai/embeddings/guardrail_translation/__init__.py +++ b/litellm/llms/openai/embeddings/guardrail_translation/__init__.py @@ -10,4 +10,4 @@ guardrail_translation_mappings = { CallTypes.aembedding: OpenAIEmbeddingsHandler, } -__all__ = ["guardrail_translation_mappings", "OpenAIEmbeddingsHandler"] +__all__ = ["OpenAIEmbeddingsHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/openai/embeddings/guardrail_translation/handler.py b/litellm/llms/openai/embeddings/guardrail_translation/handler.py index d208c98b0e4..ab9bd4f2b25 100644 --- a/litellm/llms/openai/embeddings/guardrail_translation/handler.py +++ b/litellm/llms/openai/embeddings/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for OpenAI's embeddings endpo The handler processes the 'input' parameter for guardrails. """ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -35,7 +35,7 @@ class OpenAIEmbeddingsHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input text by applying guardrails to text content. @@ -70,7 +70,7 @@ class OpenAIEmbeddingsHandler(BaseTranslation): data: dict, input_data: str, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any], + litellm_logging_obj: Any | None, ) -> dict: """Process a single string input through the guardrail.""" inputs = GenericGuardrailAPIInputs(texts=[input_data]) @@ -97,9 +97,9 @@ class OpenAIEmbeddingsHandler(BaseTranslation): async def _process_list_input( self, data: dict, - input_data: List[Union[str, int, List[int]]], + input_data: list[str | int | list[int]], guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any], + litellm_logging_obj: Any | None, ) -> dict: """Process a list input through the guardrail (if it contains strings).""" if len(input_data) == 0: @@ -144,9 +144,9 @@ class OpenAIEmbeddingsHandler(BaseTranslation): self, response: "EmbeddingResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response - embeddings responses contain vectors, not text. diff --git a/litellm/llms/openai/fine_tuning/handler.py b/litellm/llms/openai/fine_tuning/handler.py index de9d7fa581a..d7b8dd80151 100644 --- a/litellm/llms/openai/fine_tuning/handler.py +++ b/litellm/llms/openai/fine_tuning/handler.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine -from typing import Any, Dict, Optional, Union, cast +from typing import Any, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -20,7 +20,7 @@ _AZURE_STATUS_MAP = { # because LiteLLMFineTuningJob schema has no intermediate cancellation state. -def _normalize_fine_tuning_job_dict(data: Dict[str, Any], is_azure: bool = False) -> Dict[str, Any]: +def _normalize_fine_tuning_job_dict(data: dict[str, Any], is_azure: bool = False) -> dict[str, Any]: """ Normalize Azure OpenAI FineTuningJob response to match OpenAI schema. @@ -61,25 +61,18 @@ class OpenAIFineTuningAPI: def get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, _is_async: bool = False, - api_version: Optional[str] = None, - litellm_params: Optional[dict] = None, - ) -> Optional[ - Union[ - OpenAI, - AsyncOpenAI, - AzureOpenAI, - AsyncAzureOpenAI, - ] - ]: + api_version: str | None = None, + litellm_params: dict | None = None, + ) -> OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None: received_args = locals() - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None if client is None: data = {} for k, v in received_args.items(): @@ -101,7 +94,7 @@ class OpenAIFineTuningAPI: async def acreate_fine_tuning_job( self, create_fine_tuning_job_data: dict, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.create(**create_fine_tuning_job_data) @@ -111,15 +104,15 @@ class OpenAIFineTuningAPI: self, _is_async: bool, create_fine_tuning_job_data: dict, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -150,7 +143,7 @@ class OpenAIFineTuningAPI: async def acancel_fine_tuning_job( self, fine_tuning_job_id: str, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) return _litellm_fine_tuning_job_from_response(response) @@ -159,15 +152,15 @@ class OpenAIFineTuningAPI: self, _is_async: bool, fine_tuning_job_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -197,9 +190,9 @@ class OpenAIFineTuningAPI: async def alist_fine_tuning_jobs( self, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], - after: Optional[str] = None, - limit: Optional[int] = None, + openai_client: AsyncOpenAI | AsyncAzureOpenAI, + after: str | None = None, + limit: int | None = None, ): response = await openai_client.fine_tuning.jobs.list(after=after, limit=limit) # type: ignore return response @@ -207,17 +200,17 @@ class OpenAIFineTuningAPI: def list_fine_tuning_jobs( self, _is_async: bool, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - after: Optional[str] = None, - limit: Optional[int] = None, + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + after: str | None = None, + limit: int | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -249,7 +242,7 @@ class OpenAIFineTuningAPI: async def aretrieve_fine_tuning_job( self, fine_tuning_job_id: str, - openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI], + openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response = await openai_client.fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) return _litellm_fine_tuning_job_from_response(response) @@ -258,15 +251,15 @@ class OpenAIFineTuningAPI: self, _is_async: bool, fine_tuning_job_id: str, - api_key: Optional[str], - api_base: Optional[str], - api_version: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = None, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + api_version: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, diff --git a/litellm/llms/openai/image_edit/__init__.py b/litellm/llms/openai/image_edit/__init__.py index 5d933b8186d..1bc6288f69c 100644 --- a/litellm/llms/openai/image_edit/__init__.py +++ b/litellm/llms/openai/image_edit/__init__.py @@ -4,8 +4,8 @@ from .dalle2_transformation import DallE2ImageEditConfig from .transformation import OpenAIImageEditConfig __all__ = [ - "OpenAIImageEditConfig", "DallE2ImageEditConfig", + "OpenAIImageEditConfig", "get_openai_image_edit_config", ] diff --git a/litellm/llms/openai/image_edit/dalle2_transformation.py b/litellm/llms/openai/image_edit/dalle2_transformation.py index ac08d056a34..63d244be676 100644 --- a/litellm/llms/openai/image_edit/dalle2_transformation.py +++ b/litellm/llms/openai/image_edit/dalle2_transformation.py @@ -1,5 +1,5 @@ from io import BufferedReader -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast from httpx._types import RequestFiles @@ -30,12 +30,12 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform image edit request for DALL-E-2. @@ -51,7 +51,7 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig): request_params["prompt"] = prompt request = ImageEditRequestParams(**request_params) - request_dict = cast(Dict, request) + request_dict = cast(dict, request) ######################################################### # Separate images and masks as `files` and send other parameters as `data` @@ -59,7 +59,7 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig): _image_list = request_dict.get("image") _mask = request_dict.get("mask") data_without_files = {k: v for k, v in request_dict.items() if k not in ["image", "mask"]} - files_list: List[Tuple[str, Any]] = [] + files_list: list[tuple[str, Any]] = [] # Handle image parameter - DALL-E-2 only supports single image if _image_list is not None: diff --git a/litellm/llms/openai/image_edit/transformation.py b/litellm/llms/openai/image_edit/transformation.py index f53c1731f58..7a08eeedddd 100644 --- a/litellm/llms/openai/image_edit/transformation.py +++ b/litellm/llms/openai/image_edit/transformation.py @@ -1,5 +1,5 @@ from io import BufferedReader -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -59,13 +59,13 @@ class OpenAIImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """No mapping applied since inputs are in OpenAI spec already""" return dict(image_edit_optional_params) def _add_image_to_files( self, - files_list: List[Tuple[str, Any]], + files_list: list[tuple[str, Any]], image: Any, field_name: str, ) -> None: @@ -80,12 +80,12 @@ class OpenAIImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform image edit request to OpenAI API format. @@ -103,7 +103,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig): request_params["prompt"] = prompt request = ImageEditRequestParams(**request_params) - request_dict = cast(Dict, request) + request_dict = cast(dict, request) ######################################################### # Separate images and masks as `files` and send other parameters as `data` @@ -111,7 +111,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig): _image_list = request_dict.get("image") _mask = request_dict.get("mask") data_without_files = {k: v for k, v in request_dict.items() if k not in ["image", "mask"]} - files_list: List[Tuple[str, Any]] = [] + files_list: list[tuple[str, Any]] = [] # Handle image parameter if _image_list is not None: @@ -156,9 +156,9 @@ class OpenAIImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.update( @@ -171,7 +171,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/openai/image_generation/cost_calculator.py b/litellm/llms/openai/image_generation/cost_calculator.py index effda2fa3ee..b134ecc9a24 100644 --- a/litellm/llms/openai/image_generation/cost_calculator.py +++ b/litellm/llms/openai/image_generation/cost_calculator.py @@ -4,8 +4,6 @@ Cost calculator for OpenAI image generation models (gpt-image family) These models use token-based pricing instead of pixel-based pricing like DALL-E. """ -from typing import Optional - from litellm import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import ( calculate_image_response_cost_from_usage, @@ -17,7 +15,7 @@ from litellm.types.utils import ImageResponse, Usage def cost_calculator( model: str, image_response: ImageResponse, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> float: """Calculate cost for OpenAI gpt-image models (token-based pricing).""" usage = getattr(image_response, "usage", None) diff --git a/litellm/llms/openai/image_generation/dall_e_2_transformation.py b/litellm/llms/openai/image_generation/dall_e_2_transformation.py index fbc2e8dec3d..78c6ef9f27b 100644 --- a/litellm/llms/openai/image_generation/dall_e_2_transformation.py +++ b/litellm/llms/openai/image_generation/dall_e_2_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -18,7 +18,7 @@ class DallE2ImageGenerationConfig(BaseImageGenerationConfig): OpenAI dall-e-2 image generation config """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "response_format", "quality", "size", "user"] def map_openai_params( @@ -29,8 +29,8 @@ class DallE2ImageGenerationConfig(BaseImageGenerationConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: @@ -52,8 +52,8 @@ class DallE2ImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: response = raw_response.json() diff --git a/litellm/llms/openai/image_generation/dall_e_3_transformation.py b/litellm/llms/openai/image_generation/dall_e_3_transformation.py index 3434c708113..e984c2dbb1f 100644 --- a/litellm/llms/openai/image_generation/dall_e_3_transformation.py +++ b/litellm/llms/openai/image_generation/dall_e_3_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -18,7 +18,7 @@ class DallE3ImageGenerationConfig(BaseImageGenerationConfig): OpenAI dall-e-3 image generation config """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "response_format", "quality", "size", "user", "style"] def map_openai_params( @@ -29,8 +29,8 @@ class DallE3ImageGenerationConfig(BaseImageGenerationConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: @@ -52,8 +52,8 @@ class DallE3ImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: response = raw_response.json() diff --git a/litellm/llms/openai/image_generation/gpt_transformation.py b/litellm/llms/openai/image_generation/gpt_transformation.py index b9c2368d4be..1ae700620b0 100644 --- a/litellm/llms/openai/image_generation/gpt_transformation.py +++ b/litellm/llms/openai/image_generation/gpt_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -18,7 +18,7 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig): OpenAI gpt-image image generation config """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return [ "background", "moderation", @@ -38,8 +38,8 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: @@ -61,8 +61,8 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: response = raw_response.json() diff --git a/litellm/llms/openai/image_generation/guardrail_translation/__init__.py b/litellm/llms/openai/image_generation/guardrail_translation/__init__.py index 1fba2a36927..5f7c6ef861a 100644 --- a/litellm/llms/openai/image_generation/guardrail_translation/__init__.py +++ b/litellm/llms/openai/image_generation/guardrail_translation/__init__.py @@ -10,4 +10,4 @@ guardrail_translation_mappings = { CallTypes.aimage_generation: OpenAIImageGenerationHandler, } -__all__ = ["guardrail_translation_mappings", "OpenAIImageGenerationHandler"] +__all__ = ["OpenAIImageGenerationHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/openai/image_generation/guardrail_translation/handler.py b/litellm/llms/openai/image_generation/guardrail_translation/handler.py index 56bc00f319c..394b6bfc199 100644 --- a/litellm/llms/openai/image_generation/guardrail_translation/handler.py +++ b/litellm/llms/openai/image_generation/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for OpenAI's image generation The handler processes the 'prompt' parameter for guardrails. """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -32,7 +32,7 @@ class OpenAIImageGenerationHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input prompt by applying guardrails to text content. @@ -82,9 +82,9 @@ class OpenAIImageGenerationHandler(BaseTranslation): self, response: "ImageResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response - typically not needed for image generation. diff --git a/litellm/llms/openai/image_variations/handler.py b/litellm/llms/openai/image_variations/handler.py index 37d3272cff5..2cdcebb456f 100644 --- a/litellm/llms/openai/image_variations/handler.py +++ b/litellm/llms/openai/image_variations/handler.py @@ -3,7 +3,6 @@ OpenAI Image Variations Handler """ from collections.abc import Callable -from typing import Optional import httpx from openai import AsyncOpenAI, OpenAI @@ -20,7 +19,7 @@ from ..common_utils import OpenAIError class OpenAIImageVariationsHandler: def get_sync_client( self, - client: Optional[OpenAI], + client: OpenAI | None, init_client_params: dict, ): if client is None: @@ -31,7 +30,7 @@ class OpenAIImageVariationsHandler: openai_client = client return openai_client - def get_async_client(self, client: Optional[AsyncOpenAI], init_client_params: dict) -> AsyncOpenAI: + def get_async_client(self, client: AsyncOpenAI | None, init_client_params: dict) -> AsyncOpenAI: if client is None: openai_client = AsyncOpenAI( **init_client_params, @@ -44,12 +43,12 @@ class OpenAIImageVariationsHandler: self, api_key: str, api_base: str, - organization: Optional[str], - client: Optional[AsyncOpenAI], + organization: str | None, + client: AsyncOpenAI | None, data: dict, headers: dict, - model: Optional[str], - timeout: Optional[float], + model: str | None, + timeout: float | None, max_retries: int, logging_obj: LiteLLMLoggingObj, model_response: ImageResponse, @@ -114,18 +113,18 @@ class OpenAIImageVariationsHandler: model_response: ImageResponse, api_key: str, api_base: str, - model: Optional[str], + model: str | None, image: FileTypes, - timeout: Optional[float], + timeout: float | None, custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, optional_params: dict, litellm_params: dict, - print_verbose: Optional[Callable] = None, + print_verbose: Callable | None = None, logger_fn=None, client=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, + organization: str | None = None, + headers: dict | None = None, ) -> ImageResponse: try: provider_config = ProviderConfigManager.get_provider_image_variation_config( diff --git a/litellm/llms/openai/image_variations/transformation.py b/litellm/llms/openai/image_variations/transformation.py index 2f16c6f3d23..be171bb3522 100644 --- a/litellm/llms/openai/image_variations/transformation.py +++ b/litellm/llms/openai/image_variations/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, List, Optional, Union +from typing import Any from aiohttp import ClientResponse from httpx import Headers, Response @@ -13,7 +13,7 @@ from ..common_utils import OpenAIError class OpenAIImageVariationConfig(BaseImageVariationConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIImageVariationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageVariationOptionalParams]: return ["n", "size", "response_format", "user"] def map_openai_params( @@ -28,7 +28,7 @@ class OpenAIImageVariationConfig(BaseImageVariationConfig): def transform_request_image_variation( self, - model: Optional[str], + model: str | None, image: FileTypes, optional_params: dict, headers: dict, @@ -42,7 +42,7 @@ class OpenAIImageVariationConfig(BaseImageVariationConfig): async def async_transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: ClientResponse, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -51,13 +51,13 @@ class OpenAIImageVariationConfig(BaseImageVariationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: return model_response def transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: Response, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -66,11 +66,11 @@ class OpenAIImageVariationConfig(BaseImageVariationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: return model_response - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return OpenAIError( status_code=status_code, message=error_message, diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index d67b765edd8..e4a13f0f526 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -4,10 +4,8 @@ from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterat from typing import ( TYPE_CHECKING, Any, - List, Literal, Optional, - Union, cast, ) from urllib.parse import urlparse @@ -133,33 +131,33 @@ class OpenAIConfig(BaseConfig): - `top_p` (number or null): An alternative to sampling with temperature, used for nucleus sampling. """ - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_completion_tokens: Optional[int] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_completion_tokens: int | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_completion_tokens: Optional[int] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_completion_tokens: int | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -201,7 +199,7 @@ class OpenAIConfig(BaseConfig): optional_params[param] = value return optional_params - def _transform_messages(self, messages: List[AllMessageValues], model: str) -> List[AllMessageValues]: + def _transform_messages(self, messages: list[AllMessageValues], model: str) -> list[AllMessageValues]: return messages def map_openai_params( @@ -241,9 +239,7 @@ class OpenAIConfig(BaseConfig): drop_params=drop_params, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OpenAIError( status_code=status_code, message=error_message, @@ -253,7 +249,7 @@ class OpenAIConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -268,12 +264,12 @@ class OpenAIConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: logging_obj.post_call(original_response=raw_response.text) logging_obj.model_call_details["response_headers"] = raw_response.headers @@ -293,11 +289,11 @@ class OpenAIConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return { "Authorization": f"Bearer {api_key}", @@ -306,9 +302,9 @@ class OpenAIConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return OpenAIChatCompletionResponseIterator( streaming_response=streaming_response, @@ -334,9 +330,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def _set_dynamic_params_on_client( self, - client: Union[OpenAI, AsyncOpenAI], - organization: Optional[str] = None, - max_retries: Optional[int] = None, + client: OpenAI | AsyncOpenAI, + organization: str | None = None, + max_retries: int | None = None, ): if organization is not None: client.organization = organization @@ -346,21 +342,21 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def _get_openai_client( self, is_async: bool, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - timeout: Union[float, httpx.Timeout] = httpx.Timeout(None), - max_retries: Optional[int] = DEFAULT_MAX_RETRIES, - organization: Optional[str] = None, - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + timeout: float | httpx.Timeout = httpx.Timeout(None), + max_retries: int | None = DEFAULT_MAX_RETRIES, + organization: str | None = None, + client: OpenAI | AsyncOpenAI | None = None, shared_session: Optional["ClientSession"] = None, - ) -> Optional[Union[OpenAI, AsyncOpenAI]]: + ) -> OpenAI | AsyncOpenAI | None: client_initialization_params: Dict = locals() if client is None: if not isinstance(max_retries, int): raise OpenAIError( status_code=422, - message="max retries must be an int. Passed in value: {}".format(max_retries), + message=f"max retries must be an int. Passed in value: {max_retries}", ) cached_client = self.get_cached_openai_client( client_initialization_params=client_initialization_params, @@ -371,7 +367,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): if isinstance(cached_client, OpenAI) or isinstance(cached_client, AsyncOpenAI): return cached_client if is_async: - _new_client: Union[OpenAI, AsyncOpenAI] = AsyncOpenAI( + _new_client: OpenAI | AsyncOpenAI = AsyncOpenAI( api_key=api_key, base_url=api_base, http_client=OpenAIChatCompletion._get_async_http_client(shared_session=shared_session), @@ -410,7 +406,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): self, openai_aclient: AsyncOpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, ) -> Tuple[dict, BaseModel]: """ @@ -447,7 +443,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): self, openai_client: OpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, ) -> Tuple[dict, BaseModel]: """ @@ -475,9 +471,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): except Exception as e: if raw_response is not None: raise Exception( - "error - {}, Received response - {}, Type of response - {}".format( - e, raw_response, type(raw_response) - ) + f"error - {e}, Received response - {raw_response}, Type of response - {type(raw_response)}" ) else: raise e @@ -486,12 +480,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): self, response: Any, model: str, - messages: List[Dict], + messages: list[Dict], optional_params: Dict, logging_obj: LiteLLMLoggingObj, stream: bool, litellm_params: Dict, - ) -> Optional[Any]: + ) -> Any | None: """ Call agentic completion hooks for all custom loggers (OpenAI Chat Completions API). @@ -557,7 +551,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): except Exception as e: verbose_logger.exception( - f"LiteLLM.AgenticHookError: Exception in agentic completion hooks for OpenAI: {str(e)}" + f"LiteLLM.AgenticHookError: Exception in agentic completion hooks for OpenAI: {e!s}" ) return None @@ -567,7 +561,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): response: ModelResponse, logging_obj: LiteLLMLoggingObj, model: str, - stream_options: Optional[dict] = None, + stream_options: dict | None = None, ) -> CustomStreamWrapper: completion_stream = MockResponseIterator(model_response=response) streaming_response = CustomStreamWrapper( @@ -583,35 +577,35 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def completion( # type: ignore self, model_response: ModelResponse, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, optional_params: dict, litellm_params: dict, logging_obj: Any, - model: Optional[str] = None, - messages: Optional[list] = None, - print_verbose: Optional[Callable] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - dynamic_params: Optional[bool] = None, - azure_ad_token: Optional[str] = None, + model: str | None = None, + messages: list | None = None, + print_verbose: Callable | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + dynamic_params: bool | None = None, + azure_ad_token: str | None = None, acompletion: bool = False, logger_fn=None, - headers: Optional[dict] = None, + headers: dict | None = None, custom_prompt_dict: dict = {}, client=None, - organization: Optional[str] = None, - custom_llm_provider: Optional[str] = None, - drop_params: Optional[bool] = None, + organization: str | None = None, + custom_llm_provider: str | None = None, + drop_params: bool | None = None, shared_session: Optional["ClientSession"] = None, ): super().completion(shared_session=shared_session) try: fake_stream: bool = False inference_params = optional_params.copy() - stream_options: Optional[dict] = inference_params.pop("stream_options", None) - stream: Optional[bool] = inference_params.pop("stream", False) - provider_config: Optional[BaseConfig] = None + stream_options: dict | None = inference_params.pop("stream_options", None) + stream: bool | None = inference_params.pop("stream", False) + provider_config: BaseConfig | None = None if custom_llm_provider is not None and model is not None: try: @@ -780,7 +774,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): # e.message except Exception as e: if print_verbose is not None: - print_verbose(f"openai.py: Received openai error - {str(e)}") + print_verbose(f"openai.py: Received openai error - {e!s}") if ( "Conversation roles must alternate user/assistant" in str(e) or "user and assistant roles should be alternating" in str(e) @@ -832,16 +826,16 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): model: str, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, - timeout: Union[float, httpx.Timeout], - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, + timeout: float | httpx.Timeout, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + organization: str | None = None, client=None, max_retries=None, headers=None, - drop_params: Optional[bool] = None, - stream_options: Optional[dict] = None, + drop_params: bool | None = None, + stream_options: dict | None = None, fake_stream: bool = False, shared_session: Optional["ClientSession"] = None, ): @@ -949,17 +943,17 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def streaming( self, logging_obj, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, data: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + organization: str | None = None, client=None, max_retries=None, headers=None, - stream_options: Optional[dict] = None, + stream_options: dict | None = None, ): data["stream"] = True data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base)) @@ -1005,22 +999,22 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): async def async_streaming( self, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, messages: list, optional_params: dict, litellm_params: dict, provider_config: BaseConfig, model: str, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + organization: str | None = None, client=None, max_retries=None, headers=None, - drop_params: Optional[bool] = None, - stream_options: Optional[dict] = None, + drop_params: bool | None = None, + stream_options: dict | None = None, shared_session: Optional["ClientSession"] = None, ): response = None @@ -1095,7 +1089,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): if response is not None and hasattr(response, "text"): raise OpenAIError( status_code=status_code, - message=f"{str(e)}\n\nOriginal Response: {response.text}", # type: ignore + message=f"{e!s}\n\nOriginal Response: {response.text}", # type: ignore headers=error_headers, body=exception_body, ) @@ -1117,12 +1111,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): else: raise OpenAIError( status_code=500, - message=f"{str(e)}", + message=f"{e!s}", headers=error_headers, body=exception_body, ) - def get_stream_options(self, stream_options: Optional[dict], api_base: Optional[str]) -> dict: + def get_stream_options(self, stream_options: dict | None, api_base: str | None) -> dict: """ Pass `stream_options` to the data dict for OpenAI requests """ @@ -1140,7 +1134,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): self, openai_aclient: AsyncOpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, ): """ @@ -1161,7 +1155,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): self, openai_client: OpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, logging_obj: LiteLLMLoggingObj, ): """ @@ -1185,9 +1179,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): model_response: EmbeddingResponse, timeout: float, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - client: Optional[AsyncOpenAI] = None, + api_key: str | None = None, + api_base: str | None = None, + client: AsyncOpenAI | None = None, max_retries=None, shared_session: Optional["ClientSession"] = None, ): @@ -1256,11 +1250,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): logging_obj, model_response: EmbeddingResponse, optional_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, client=None, aembedding=None, - max_retries: Optional[int] = None, + max_retries: int | None = None, shared_session: Optional["ClientSession"] = None, ) -> EmbeddingResponse: super().embedding() @@ -1300,7 +1294,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) ## embedding CALL - headers: Optional[Dict] = None + headers: Dict | None = None headers, sync_embedding_response = self.make_sync_openai_embedding_request( openai_client=openai_client, data=data, @@ -1341,12 +1335,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): model_response: ModelResponse, timeout: float, logging_obj: Any, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, client=None, max_retries=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, + organization: str | None = None, + headers: dict | None = None, ): response = None try: @@ -1387,18 +1381,18 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def image_generation( self, - model: Optional[str], + model: str | None, prompt: str, timeout: float, optional_params: dict, logging_obj: Any, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - model_response: Optional[ImageResponse] = None, + api_key: str | None = None, + api_base: str | None = None, + model_response: ImageResponse | None = None, client=None, aimg_generation=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, + organization: str | None = None, + headers: dict | None = None, ) -> ImageResponse: data = {} try: @@ -1490,13 +1484,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): input: str, voice: str, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - organization: Optional[str], - project: Optional[str], + api_key: str | None, + api_base: str | None, + organization: str | None, + project: str | None, max_retries: int, - timeout: Union[float, httpx.Timeout], - aspeech: Optional[bool] = None, + timeout: float | httpx.Timeout, + aspeech: bool | None = None, client=None, shared_session: Optional["ClientSession"] = None, ) -> HttpxBinaryResponseContent: @@ -1540,12 +1534,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): input: str, voice: str, optional_params: dict, - api_key: Optional[str], - api_base: Optional[str], - organization: Optional[str], - project: Optional[str], + api_key: str | None, + api_base: str | None, + organization: str | None, + project: str | None, max_retries: int, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, client=None, shared_session: Optional["ClientSession"] = None, ) -> HttpxBinaryResponseContent: @@ -1588,16 +1582,16 @@ class OpenAIFilesAPI(BaseLLM): def get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, _is_async: bool = False, - ) -> Optional[Union[OpenAI, AsyncOpenAI]]: + ) -> OpenAI | AsyncOpenAI | None: received_args = locals() - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = None + openai_client: OpenAI | AsyncOpenAI | None = None if client is None: data = {} for k, v in received_args.items(): @@ -1629,13 +1623,13 @@ class OpenAIFilesAPI(BaseLLM): _is_async: bool, create_file_data: CreateFileRequest, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, - ) -> Union[OpenAIFileObject, Coroutine[Any, Any, OpenAIFileObject]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, + ) -> OpenAIFileObject | Coroutine[Any, Any, OpenAIFileObject]: + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1673,13 +1667,13 @@ class OpenAIFilesAPI(BaseLLM): _is_async: bool, file_content_request: FileContentRequest, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, - ) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, + ) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1717,7 +1711,7 @@ class OpenAIFilesAPI(BaseLLM): headers = dict(response.headers) async def _stream() -> AsyncIterator[bytes]: - exc: Optional[BaseException] = None + exc: BaseException | None = None try: async for chunk in response.iter_bytes(chunk_size=chunk_size): yield chunk @@ -1737,14 +1731,14 @@ class OpenAIFilesAPI(BaseLLM): _is_async: bool, file_content_request: FileContentRequest, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, chunk_size: int = 1024 * 1024, - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + client: OpenAI | AsyncOpenAI | None = None, ) -> FileContentStreamingResult: - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1774,7 +1768,7 @@ class OpenAIFilesAPI(BaseLLM): headers = dict(response.headers) def _stream() -> Iterator[bytes]: - exc: Optional[BaseException] = None + exc: BaseException | None = None try: yield from response.iter_bytes(chunk_size=chunk_size) except BaseException as e: @@ -1801,13 +1795,13 @@ class OpenAIFilesAPI(BaseLLM): _is_async: bool, file_id: str, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1847,13 +1841,13 @@ class OpenAIFilesAPI(BaseLLM): _is_async: bool, file_id: str, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1883,7 +1877,7 @@ class OpenAIFilesAPI(BaseLLM): async def alist_files( self, openai_client: AsyncOpenAI, - purpose: Optional[str] = None, + purpose: str | None = None, ): if isinstance(purpose, str): response = await openai_client.files.list(purpose=purpose) @@ -1895,14 +1889,14 @@ class OpenAIFilesAPI(BaseLLM): self, _is_async: bool, api_base: str, - api_key: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - purpose: Optional[str] = None, - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + purpose: str | None = None, + client: OpenAI | AsyncOpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -1948,16 +1942,16 @@ class OpenAIBatchesAPI(BaseLLM): def get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, _is_async: bool = False, - ) -> Optional[Union[OpenAI, AsyncOpenAI]]: + ) -> OpenAI | AsyncOpenAI | None: received_args = locals() - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = None + openai_client: OpenAI | AsyncOpenAI | None = None if client is None: data = {} for k, v in received_args.items(): @@ -1988,14 +1982,14 @@ class OpenAIBatchesAPI(BaseLLM): self, _is_async: bool, create_batch_data: CreateBatchRequest, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[Union[OpenAI, AsyncOpenAI]] = None, - ) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | AsyncOpenAI | None = None, + ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -2034,14 +2028,14 @@ class OpenAIBatchesAPI(BaseLLM): self, _is_async: bool, retrieve_batch_data: RetrieveBatchRequest, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -2079,14 +2073,14 @@ class OpenAIBatchesAPI(BaseLLM): self, _is_async: bool, cancel_batch_data: CancelBatchRequest, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -2118,8 +2112,8 @@ class OpenAIBatchesAPI(BaseLLM): async def alist_batches( self, openai_client: AsyncOpenAI, - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, ): verbose_logger.debug("listing batches, after= %s, limit= %s", after, limit) response = await openai_client.batches.list(after=after, limit=limit) # type: ignore @@ -2128,16 +2122,16 @@ class OpenAIBatchesAPI(BaseLLM): def list_batches( self, _is_async: bool, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - after: Optional[str] = None, - limit: Optional[int] = None, - client: Optional[OpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + after: str | None = None, + limit: int | None = None, + client: OpenAI | None = None, ): - openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client( + openai_client: OpenAI | AsyncOpenAI | None = self.get_openai_client( api_key=api_key, api_base=api_base, timeout=timeout, @@ -2169,12 +2163,12 @@ class OpenAIAssistantsAPI(BaseLLM): def get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None = None, ) -> OpenAI: received_args = locals() if client is None: @@ -2194,12 +2188,12 @@ class OpenAIAssistantsAPI(BaseLLM): def async_get_openai_client( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None = None, ) -> AsyncOpenAI: received_args = locals() if client is None: @@ -2221,16 +2215,16 @@ class OpenAIAssistantsAPI(BaseLLM): async def async_get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], - order: Optional[str] = "desc", - limit: Optional[int] = 20, - before: Optional[str] = None, - after: Optional[str] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, + order: str | None = "desc", + limit: int | None = 20, + before: str | None = None, + after: str | None = None, ) -> AsyncCursorPage[Assistant]: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2258,12 +2252,12 @@ class OpenAIAssistantsAPI(BaseLLM): @overload def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, aget_assistants: Literal[True], ) -> Coroutine[None, None, AsyncCursorPage[Assistant]]: ... @@ -2271,13 +2265,13 @@ class OpenAIAssistantsAPI(BaseLLM): @overload def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI], - aget_assistants: Optional[Literal[False]], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None, + aget_assistants: Literal[False] | None, ) -> SyncCursorPage[Assistant]: ... @@ -2285,17 +2279,17 @@ class OpenAIAssistantsAPI(BaseLLM): def get_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client=None, aget_assistants=None, - order: Optional[str] = "desc", - limit: Optional[int] = 20, - before: Optional[str] = None, - after: Optional[str] = None, + order: str | None = "desc", + limit: int | None = 20, + before: str | None = None, + after: str | None = None, ): if aget_assistants is not None and aget_assistants is True: return self.async_get_assistants( @@ -2332,12 +2326,12 @@ class OpenAIAssistantsAPI(BaseLLM): # Create Assistant async def async_create_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, create_assistant_data: dict, ) -> Assistant: openai_client = self.async_get_openai_client( @@ -2355,11 +2349,11 @@ class OpenAIAssistantsAPI(BaseLLM): def create_assistants( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, create_assistant_data: dict, client=None, async_create_assistants=None, @@ -2389,12 +2383,12 @@ class OpenAIAssistantsAPI(BaseLLM): # Delete Assistant async def async_delete_assistant( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, assistant_id: str, ) -> AssistantDeleted: openai_client = self.async_get_openai_client( @@ -2412,11 +2406,11 @@ class OpenAIAssistantsAPI(BaseLLM): def delete_assistant( self, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, assistant_id: str, client=None, async_delete_assistants=None, @@ -2449,12 +2443,12 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None = None, ) -> OpenAIMessage: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2470,7 +2464,7 @@ class OpenAIAssistantsAPI(BaseLLM): **message_data, # type: ignore ) - response_obj: Optional[OpenAIMessage] = None + response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" response_obj = OpenAIMessage.model_validate(thread_message.dict()) @@ -2485,12 +2479,12 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, a_add_message: Literal[True], ) -> Coroutine[None, None, OpenAIMessage]: ... @@ -2500,13 +2494,13 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI], - a_add_message: Optional[Literal[False]], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None, + a_add_message: Literal[False] | None, ) -> OpenAIMessage: ... @@ -2516,13 +2510,13 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, message_data: dict, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client=None, - a_add_message: Optional[bool] = None, + a_add_message: bool | None = None, ): if a_add_message is not None and a_add_message is True: return self.a_add_message( @@ -2549,7 +2543,7 @@ class OpenAIAssistantsAPI(BaseLLM): **message_data, # type: ignore ) - response_obj: Optional[OpenAIMessage] = None + response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" response_obj = OpenAIMessage.model_validate(thread_message.dict()) @@ -2560,12 +2554,12 @@ class OpenAIAssistantsAPI(BaseLLM): async def async_get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI] = None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None = None, ) -> AsyncCursorPage[OpenAIMessage]: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2586,12 +2580,12 @@ class OpenAIAssistantsAPI(BaseLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, aget_messages: Literal[True], ) -> Coroutine[None, None, AsyncCursorPage[OpenAIMessage]]: ... @@ -2600,13 +2594,13 @@ class OpenAIAssistantsAPI(BaseLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI], - aget_messages: Optional[Literal[False]], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None, + aget_messages: Literal[False] | None, ) -> SyncCursorPage[OpenAIMessage]: ... @@ -2615,11 +2609,11 @@ class OpenAIAssistantsAPI(BaseLLM): def get_messages( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client=None, aget_messages=None, ): @@ -2650,14 +2644,14 @@ class OpenAIAssistantsAPI(BaseLLM): async def async_create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, ) -> Thread: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2683,14 +2677,14 @@ class OpenAIAssistantsAPI(BaseLLM): @overload def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], - client: Optional[AsyncOpenAI], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, + client: AsyncOpenAI | None, acreate_thread: Literal[True], ) -> Coroutine[None, None, Thread]: ... @@ -2698,15 +2692,15 @@ class OpenAIAssistantsAPI(BaseLLM): @overload def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], - client: Optional[OpenAI], - acreate_thread: Optional[Literal[False]], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, + client: OpenAI | None, + acreate_thread: Literal[False] | None, ) -> Thread: ... @@ -2714,13 +2708,13 @@ class OpenAIAssistantsAPI(BaseLLM): def create_thread( self, - metadata: Optional[dict], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], + metadata: dict | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + messages: Iterable[OpenAICreateThreadParamsMessage] | None, client=None, acreate_thread=None, ): @@ -2767,12 +2761,12 @@ class OpenAIAssistantsAPI(BaseLLM): async def async_get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, ) -> Thread: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2793,12 +2787,12 @@ class OpenAIAssistantsAPI(BaseLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, aget_thread: Literal[True], ) -> Coroutine[None, None, Thread]: ... @@ -2807,13 +2801,13 @@ class OpenAIAssistantsAPI(BaseLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[OpenAI], - aget_thread: Optional[Literal[False]], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: OpenAI | None, + aget_thread: Literal[False] | None, ) -> Thread: ... @@ -2822,11 +2816,11 @@ class OpenAIAssistantsAPI(BaseLLM): def get_thread( self, thread_id: str, - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client=None, aget_thread=None, ): @@ -2862,18 +2856,18 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], - client: Optional[AsyncOpenAI], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, + client: AsyncOpenAI | None, ) -> Run: openai_client = self.async_get_openai_client( api_key=api_key, @@ -2901,12 +2895,12 @@ class OpenAIAssistantsAPI(BaseLLM): client: AsyncOpenAI, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - tools: Optional[Iterable[AssistantToolParam]], - event_handler: Optional[AssistantEventHandler], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + tools: Iterable[AssistantToolParam] | None, + event_handler: AssistantEventHandler | None, ) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: data: Dict[str, Any] = { "thread_id": thread_id, @@ -2926,12 +2920,12 @@ class OpenAIAssistantsAPI(BaseLLM): client: OpenAI, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - tools: Optional[Iterable[AssistantToolParam]], - event_handler: Optional[AssistantEventHandler], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + tools: Iterable[AssistantToolParam] | None, + event_handler: AssistantEventHandler | None, ) -> AssistantStreamManager[AssistantEventHandler]: data: Dict[str, Any] = { "thread_id": thread_id, @@ -2953,20 +2947,20 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client, arun_thread: Literal[True], - event_handler: Optional[AssistantEventHandler], + event_handler: AssistantEventHandler | None, ) -> Coroutine[None, None, Run]: ... @@ -2975,20 +2969,20 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client, - arun_thread: Optional[Literal[False]], - event_handler: Optional[AssistantEventHandler], + arun_thread: Literal[False] | None, + event_handler: AssistantEventHandler | None, ) -> Run: ... @@ -2998,20 +2992,20 @@ class OpenAIAssistantsAPI(BaseLLM): self, thread_id: str, assistant_id: str, - additional_instructions: Optional[str], - instructions: Optional[str], - metadata: Optional[Dict], - model: Optional[str], - stream: Optional[bool], - tools: Optional[Iterable[AssistantToolParam]], - api_key: Optional[str], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - organization: Optional[str], + additional_instructions: str | None, + instructions: str | None, + metadata: Dict | None, + model: str | None, + stream: bool | None, + tools: Iterable[AssistantToolParam] | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + organization: str | None, client=None, arun_thread=None, - event_handler: Optional[AssistantEventHandler] = None, + event_handler: AssistantEventHandler | None = None, ): if arun_thread is not None and arun_thread is True: if stream is not None and stream is True: diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index 626d2f3a28e..14fa6dc9954 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -4,7 +4,7 @@ This file contains the calling OpenAI's `/v1/realtime` endpoint. This requires websockets, and is currently only supported on LiteLLM Proxy. """ -from typing import Any, Optional, cast +from typing import Any, cast from litellm._logging import _redact_string, verbose_logger from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES @@ -96,7 +96,7 @@ class OpenAIRealtime(OpenAIChatCompletion): url = url.copy_with(params=query_params) return str(url) - def _make_event_normalizer(self) -> Optional[RealtimeEventNormalizer]: + def _make_event_normalizer(self) -> RealtimeEventNormalizer | None: """Return a per-session GA event normalizer, or None for passthrough. Subclasses (e.g. XAIRealtime) override this to supply a provider-specific @@ -109,13 +109,13 @@ class OpenAIRealtime(OpenAIChatCompletion): model: str, websocket: Any, logging_obj: LiteLLMLogging, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - client: Optional[Any] = None, - timeout: Optional[float] = None, - query_params: Optional[RealtimeQueryParams] = None, - user_api_key_dict: Optional[Any] = None, - litellm_metadata: Optional[dict] = None, + api_base: str | None = None, + api_key: str | None = None, + client: Any | None = None, + timeout: float | None = None, + query_params: RealtimeQueryParams | None = None, + user_api_key_dict: Any | None = None, + litellm_metadata: dict | None = None, **kwargs: Any, ): import websockets @@ -178,7 +178,7 @@ class OpenAIRealtime(OpenAIChatCompletion): await websocket.close(code=e.status_code, reason=_redact_string(str(e))) except Exception as e: try: - await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {str(e)}")) + await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e!s}")) except RuntimeError as close_error: if "already completed" in str(close_error) or "websocket.close" in str(close_error): # The WebSocket is already closed or the response is completed, so we can ignore this error diff --git a/litellm/llms/openai/realtime/http_transformation.py b/litellm/llms/openai/realtime/http_transformation.py index 0a7e65dfea2..61dbf20397f 100644 --- a/litellm/llms/openai/realtime/http_transformation.py +++ b/litellm/llms/openai/realtime/http_transformation.py @@ -1,44 +1,37 @@ """OpenAI realtime HTTP transformation config (client_secrets + realtime_calls).""" -from typing import Optional - import litellm from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig from litellm.secret_managers.main import get_secret_str class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig): - def get_api_base(self, api_base: Optional[str], **kwargs) -> str: + def get_api_base(self, api_base: str | None, **kwargs) -> str: return api_base or litellm.api_base or get_secret_str("OPENAI_API_BASE") or "https://api.openai.com" - def get_api_key(self, api_key: Optional[str], **kwargs) -> str: + def get_api_key(self, api_key: str | None, **kwargs) -> str: return api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") or "" - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") - if base.endswith("/v1"): - base = base[:-3] + base = base.removesuffix("/v1") return f"{base}/v1/realtime/client_secrets" - def get_realtime_calls_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") - if base.endswith("/v1"): - base = base[:-3] + base = base.removesuffix("/v1") return f"{base}/v1/realtime/calls" - def get_transcription_session_url( - self, api_base: Optional[str], model: str, api_version: Optional[str] = None - ) -> str: + def get_transcription_session_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: base = self.get_api_base(api_base).rstrip("/") - if base.endswith("/v1"): - base = base[:-3] + base = base.removesuffix("/v1") return f"{base}/v1/realtime/transcription_sessions" def validate_environment( self, headers: dict, model: str, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> dict: return { **headers, diff --git a/litellm/llms/openai/responses/count_tokens/__init__.py b/litellm/llms/openai/responses/count_tokens/__init__.py index 8f129a6ff09..83985c92cd3 100644 --- a/litellm/llms/openai/responses/count_tokens/__init__.py +++ b/litellm/llms/openai/responses/count_tokens/__init__.py @@ -13,7 +13,7 @@ from litellm.llms.openai.responses.count_tokens.transformation import ( ) __all__ = [ - "OpenAICountTokensHandler", "OpenAICountTokensConfig", + "OpenAICountTokensHandler", "OpenAITokenCounter", ] diff --git a/litellm/llms/openai/responses/count_tokens/handler.py b/litellm/llms/openai/responses/count_tokens/handler.py index 3dded042de8..b7cc3b1673a 100644 --- a/litellm/llms/openai/responses/count_tokens/handler.py +++ b/litellm/llms/openai/responses/count_tokens/handler.py @@ -5,7 +5,7 @@ Uses httpx for HTTP requests to OpenAI's /v1/responses/input_tokens endpoint. """ import json -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -26,13 +26,13 @@ class OpenAICountTokensHandler(OpenAICountTokensConfig): async def handle_count_tokens_request( self, model: str, - input: Union[str, List[Any]], + input: str | list[Any], api_key: str, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tools: Optional[List[Dict[str, Any]]] = None, - instructions: Optional[str] = None, - ) -> Dict[str, Any]: + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + tools: list[dict[str, Any]] | None = None, + instructions: str | None = None, + ) -> dict[str, Any]: """ Handle a token counting request to OpenAI's Responses API. @@ -88,14 +88,14 @@ class OpenAICountTokensHandler(OpenAICountTokensConfig): except OpenAIError: raise except httpx.HTTPStatusError as e: - verbose_logger.error(f"HTTP error in CountTokens handler: {str(e)}") + verbose_logger.error(f"HTTP error in CountTokens handler: {e!s}") raise OpenAIError( status_code=e.response.status_code, message=e.response.text, ) except (httpx.RequestError, json.JSONDecodeError, ValueError) as e: - verbose_logger.error(f"Error in CountTokens handler: {str(e)}") + verbose_logger.error(f"Error in CountTokens handler: {e!s}") raise OpenAIError( status_code=500, - message=f"CountTokens processing error: {str(e)}", + message=f"CountTokens processing error: {e!s}", ) diff --git a/litellm/llms/openai/responses/count_tokens/token_counter.py b/litellm/llms/openai/responses/count_tokens/token_counter.py index 8e700ecafa1..d4494759f6c 100644 --- a/litellm/llms/openai/responses/count_tokens/token_counter.py +++ b/litellm/llms/openai/responses/count_tokens/token_counter.py @@ -3,7 +3,7 @@ OpenAI Token Counter implementation using the Responses API /input_tokens endpoi """ import os -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.llms.base_llm.base_utils import BaseTokenCounter @@ -25,20 +25,20 @@ class OpenAITokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: return custom_llm_provider == LlmProviders.OPENAI.value async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: """ Count tokens using OpenAI's Responses API /input_tokens endpoint. """ diff --git a/litellm/llms/openai/responses/count_tokens/transformation.py b/litellm/llms/openai/responses/count_tokens/transformation.py index 41d1a01ec66..d7ba49a7927 100644 --- a/litellm/llms/openai/responses/count_tokens/transformation.py +++ b/litellm/llms/openai/responses/count_tokens/transformation.py @@ -4,7 +4,7 @@ OpenAI Responses API token counting transformation logic. This module handles the transformation of requests to OpenAI's /v1/responses/input_tokens endpoint. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any class OpenAICountTokensConfig: @@ -16,7 +16,7 @@ class OpenAICountTokensConfig: - Response: {"input_tokens": } """ - def get_openai_count_tokens_endpoint(self, api_base: Optional[str] = None) -> str: + def get_openai_count_tokens_endpoint(self, api_base: str | None = None) -> str: base = api_base or "https://api.openai.com/v1" base = base.rstrip("/") return f"{base}/responses/input_tokens" @@ -24,16 +24,16 @@ class OpenAICountTokensConfig: def transform_request_to_count_tokens( self, model: str, - input: Union[str, List[Any]], - tools: Optional[List[Dict[str, Any]]] = None, - instructions: Optional[str] = None, - ) -> Dict[str, Any]: + input: str | list[Any], + tools: list[dict[str, Any]] | None = None, + instructions: str | None = None, + ) -> dict[str, Any]: """ Transform request to OpenAI Responses API token counting format. The Responses API uses `input` (not `messages`) and `instructions` (not `system`). """ - request: Dict[str, Any] = { + request: dict[str, Any] = { "model": model, "input": input, } @@ -46,13 +46,13 @@ class OpenAICountTokensConfig: return request - def get_required_headers(self, api_key: str) -> Dict[str, str]: + def get_required_headers(self, api_key: str) -> dict[str, str]: return { "Content-Type": "application/json", "Authorization": f"Bearer {api_key}", } - def validate_request(self, model: str, input: Union[str, List[Any]]) -> None: + def validate_request(self, model: str, input: str | list[Any]) -> None: if not model: raise ValueError("model parameter is required") @@ -61,8 +61,8 @@ class OpenAICountTokensConfig: @staticmethod def _transform_tools_for_responses_api( - tools: List[Dict[str, Any]], - ) -> List[Dict[str, Any]]: + tools: list[dict[str, Any]], + ) -> list[dict[str, Any]]: """ Transform OpenAI chat tools format to Responses API tools format. @@ -73,7 +73,7 @@ class OpenAICountTokensConfig: for tool in tools: if tool.get("type") == "function" and "function" in tool: func = tool["function"] - item: Dict[str, Any] = { + item: dict[str, Any] = { "type": "function", "name": func.get("name", ""), "description": func.get("description", ""), @@ -89,7 +89,7 @@ class OpenAICountTokensConfig: @staticmethod def messages_to_responses_input( - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], ) -> tuple: """ Convert standard chat messages format to OpenAI Responses API input format. @@ -98,8 +98,8 @@ class OpenAICountTokensConfig: (input_items, instructions) tuple where instructions is extracted from system/developer messages. """ - input_items: List[Dict[str, Any]] = [] - instructions_parts: List[str] = [] + input_items: list[dict[str, Any]] = [] + instructions_parts: list[str] = [] for msg in messages: role = msg.get("role", "") diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index d90703d1544..f2876b4f1bc 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -28,7 +28,7 @@ Output: response.output is List[GenericResponseOutputItem] where each has: - text: str """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Union, cast from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall from pydantic import BaseModel @@ -71,7 +71,7 @@ class OpenAIResponsesHandler(BaseTranslation): Methods can be overridden to customize behavior for different message formats. """ - def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]: + def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None: """ Convert Responses API request data to OpenAI-spec structured messages. @@ -85,21 +85,21 @@ class OpenAIResponsesHandler(BaseTranslation): input=input_data, responses_api_request=data, ) - return cast(List[AllMessageValues], messages) if messages else None + return cast(list[AllMessageValues], messages) if messages else None async def process_input_messages( self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input by applying guardrails to text content. Handles both string input and list of message objects. """ - input_data: Optional[Union[str, "ResponseInputParam"]] = data.get("input") - tools_to_check: List[ChatCompletionToolParam] = [] + input_data: str | ResponseInputParam | None = data.get("input") + tools_to_check: list[ChatCompletionToolParam] = [] if input_data is None: return data @@ -108,7 +108,7 @@ class OpenAIResponsesHandler(BaseTranslation): # Handle simple string input if isinstance(input_data, str): inputs = GenericGuardrailAPIInputs(texts=[input_data]) - original_tools: List[Dict[str, Any]] = [] + original_tools: list[dict[str, Any]] = [] # Extract and transform tools if present if "tools" in data and data["tools"]: @@ -139,10 +139,10 @@ class OpenAIResponsesHandler(BaseTranslation): if not isinstance(input_data, list): return data - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - task_mappings: List[Tuple[int, Optional[int]]] = [] - original_tools_list: List[Dict[str, Any]] = list(data.get("tools") or []) + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + task_mappings: list[tuple[int, int | None]] = [] + original_tools_list: list[dict[str, Any]] = list(data.get("tools") or []) # Step 1: Extract all text content, images, and tools for msg_idx, message in enumerate(input_data): @@ -196,10 +196,10 @@ class OpenAIResponsesHandler(BaseTranslation): return data - def extract_request_tool_names(self, data: dict) -> List[str]: + def extract_request_tool_names(self, data: dict) -> list[str]: """Extract tool names from Responses API request (tools[].name for function and custom, tools[].server_label for mcp).""" - names: List[str] = [] + names: list[str] = [] for tool in data.get("tools") or []: if not isinstance(tool, dict): continue @@ -211,8 +211,8 @@ class OpenAIResponsesHandler(BaseTranslation): def _extract_and_transform_tools( self, - tools: List[Dict[str, Any]], - tools_to_check: List[ChatCompletionToolParam], + tools: list[dict[str, Any]], + tools_to_check: list[ChatCompletionToolParam], ) -> None: """ Extract and transform tools from Responses API format to Chat Completion format. @@ -228,9 +228,9 @@ class OpenAIResponsesHandler(BaseTranslation): ) = LiteLLMCompletionResponsesConfig.transform_responses_api_tools_to_chat_completion_tools( tools # type: ignore ) - tools_to_check.extend(cast(List[ChatCompletionToolParam], transformed_tools)) + tools_to_check.extend(cast(list[ChatCompletionToolParam], transformed_tools)) - def _remap_tools_to_responses_api_format(self, guardrailed_tools: List[Any]) -> List[Dict[str, Any]]: + def _remap_tools_to_responses_api_format(self, guardrailed_tools: list[Any]) -> list[dict[str, Any]]: """ Remap guardrail-returned tools (Chat Completion format) back to Responses API request tool format. @@ -241,9 +241,9 @@ class OpenAIResponsesHandler(BaseTranslation): def _merge_tools_after_guardrail( self, - original_tools: List[Dict[str, Any]], - remapped: List[Dict[str, Any]], - ) -> List[Dict[str, Any]]: + original_tools: list[dict[str, Any]], + remapped: list[dict[str, Any]], + ) -> list[dict[str, Any]]: """ Merge remapped guardrailed tools with original tools that were not sent to the guardrail (e.g. web_search, web_search_preview), preserving order. @@ -252,7 +252,7 @@ class OpenAIResponsesHandler(BaseTranslation): """ if not original_tools: return remapped - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] j = 0 for tool in original_tools: if isinstance(tool, dict) and tool.get("type") in ( @@ -271,8 +271,8 @@ class OpenAIResponsesHandler(BaseTranslation): def _apply_guardrailed_tools_to_data( self, data: dict, - original_tools: List[Dict[str, Any]], - guardrailed_tools: Optional[List[Any]], + original_tools: list[dict[str, Any]], + guardrailed_tools: list[Any] | None, ) -> None: """Remap guardrailed tools to Responses API format and merge with original, then set data['tools'].""" if guardrailed_tools is not None: @@ -283,9 +283,9 @@ class OpenAIResponsesHandler(BaseTranslation): self, message: Any, # Can be Dict[str, Any] or ResponseInputParam msg_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - task_mappings: List[Tuple[int, Optional[int]]], + texts_to_check: list[str], + images_to_check: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Extract text content and images from an input message. @@ -322,8 +322,8 @@ class OpenAIResponsesHandler(BaseTranslation): async def _apply_guardrail_responses_to_input( self, messages: Any, # Can be List[Dict[str, Any]] or ResponseInputParam - responses: List[str], - task_mappings: List[Tuple[int, Optional[int]]], + responses: list[str], + task_mappings: list[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to input messages. @@ -333,7 +333,7 @@ class OpenAIResponsesHandler(BaseTranslation): for task_idx, guardrail_response in enumerate(responses): mapping = task_mappings[task_idx] msg_idx = cast(int, mapping[0]) - content_idx_optional = cast(Optional[int], mapping[1]) + content_idx_optional = cast(int | None, mapping[1]) content = messages[msg_idx].get("content", None) if content is None: @@ -352,9 +352,9 @@ class OpenAIResponsesHandler(BaseTranslation): self, response: "ResponsesAPIResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response by applying guardrails to text content and tool calls. @@ -376,10 +376,10 @@ class OpenAIResponsesHandler(BaseTranslation): - Each OutputText object has a text field """ - texts_to_check: List[str] = [] - images_to_check: List[str] = [] - tool_calls_to_check: List[ChatCompletionToolCallChunk] = [] - task_mappings: List[Tuple[int, int]] = [] + texts_to_check: list[str] = [] + images_to_check: list[str] = [] + tool_calls_to_check: list[ChatCompletionToolCallChunk] = [] + task_mappings: list[tuple[int, int]] = [] # Track (output_item_index, content_index) for each text # Handle both dict and Pydantic object responses @@ -458,12 +458,12 @@ class OpenAIResponsesHandler(BaseTranslation): async def process_output_streaming_response( self, - responses_so_far: List[Any], + responses_so_far: list[Any], guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, - ) -> List[Any]: + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, + ) -> list[Any]: """ Process output streaming response by applying guardrails to text content. @@ -493,11 +493,11 @@ class OpenAIResponsesHandler(BaseTranslation): response_obj = final_chunk.get("response") or {} if not hasattr(response_obj, "get"): return responses_so_far - outputs: List[Any] = response_obj.get("output") or [] + outputs: list[Any] = response_obj.get("output") or [] - texts_to_check: List[str] = [] - tool_calls_to_check: List[ChatCompletionToolCallChunk] = [] - task_mappings: List[Tuple[int, int]] = [] + texts_to_check: list[str] = [] + tool_calls_to_check: list[ChatCompletionToolCallChunk] = [] + task_mappings: list[tuple[int, int]] = [] for output_idx, output_item in enumerate(outputs): self._extract_output_text_and_images( @@ -521,7 +521,7 @@ class OpenAIResponsesHandler(BaseTranslation): inputs = GenericGuardrailAPIInputs(texts=texts_to_check) if tool_calls_to_check: - inputs["tool_calls"] = cast(List[ChatCompletionToolCallChunk], tool_calls_to_check) + inputs["tool_calls"] = cast(list[ChatCompletionToolCallChunk], tool_calls_to_check) response_model = response_obj.get("model") if response_model: inputs["model"] = response_model @@ -556,7 +556,7 @@ class OpenAIResponsesHandler(BaseTranslation): tool_calls = model_response_stream.choices[0].delta.tool_calls if tool_calls: inputs = GenericGuardrailAPIInputs() - inputs["tool_calls"] = cast(List[ChatCompletionToolCallChunk], tool_calls) + inputs["tool_calls"] = cast(list[ChatCompletionToolCallChunk], tool_calls) if hasattr(model_response_stream, "model") and model_response_stream.model: inputs["model"] = model_response_stream.model await guardrail_to_apply.apply_guardrail( @@ -588,7 +588,7 @@ class OpenAIResponsesHandler(BaseTranslation): ) return responses_so_far - def _check_streaming_has_ended(self, responses_so_far: List[Any]) -> bool: + def _check_streaming_has_ended(self, responses_so_far: list[Any]) -> bool: """ Check if the streaming has ended. """ @@ -601,7 +601,7 @@ class OpenAIResponsesHandler(BaseTranslation): } return responses_so_far[-1].get("type") in terminal_types - def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str: + def get_streaming_string_so_far(self, responses_so_far: list[Any]) -> str: """ Get the string so far from the responses so far. """ @@ -645,10 +645,10 @@ class OpenAIResponsesHandler(BaseTranslation): self, output_item: Any, output_idx: int, - texts_to_check: List[str], - images_to_check: List[str], - task_mappings: List[Tuple[int, int]], - tool_calls_to_check: Optional[List[ChatCompletionToolCallChunk]] = None, + texts_to_check: list[str], + images_to_check: list[str], + task_mappings: list[tuple[int, int]], + tool_calls_to_check: list[ChatCompletionToolCallChunk] | None = None, ) -> None: """ Extract text content, images, and tool calls from a response output item. @@ -657,17 +657,7 @@ class OpenAIResponsesHandler(BaseTranslation): """ # Check if this is a tool call (OutputFunctionToolCall) - if isinstance(output_item, OutputFunctionToolCall): - if tool_calls_to_check is not None: - tool_call_dict = ( - LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( - tool_call_item=output_item, - index=output_idx, - ) - ) - tool_calls_to_check.append(cast(ChatCompletionToolCallChunk, tool_call_dict)) - return - elif ( + if isinstance(output_item, OutputFunctionToolCall) or ( isinstance(output_item, BaseModel) and hasattr(output_item, "type") and getattr(output_item, "type") == "function_call" @@ -697,7 +687,7 @@ class OpenAIResponsesHandler(BaseTranslation): return # Handle both GenericResponseOutputItem and dict - content: Optional[Union[List[OutputText], List[dict]]] = None + content: list[OutputText] | list[dict] | None = None if isinstance(output_item, BaseModel): try: output_item_dump = output_item.model_dump() @@ -736,9 +726,9 @@ class OpenAIResponsesHandler(BaseTranslation): async def _apply_guardrail_responses_to_output( self, - response: Union["ResponsesAPIResponse", Dict[Any, Any]], - responses: List[str], - task_mappings: List[Tuple[int, int]], + response: Union["ResponsesAPIResponse", dict[Any, Any]], + responses: list[str], + task_mappings: list[tuple[int, int]], ) -> None: """ Apply guardrail responses back to output response. diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index dc4e98e6216..2d0ce47e595 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional, Union, cast, get_type_hints +from typing import TYPE_CHECKING, Any, cast, get_type_hints import httpx from openai.types.responses import ResponseReasoningItem @@ -7,10 +7,10 @@ from pydantic import BaseModel, ValidationError import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers -from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import * @@ -98,7 +98,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """No mapping applied since inputs are in OpenAI spec already. GPT-5 models have restrictions on temperature (only temperature=1 @@ -123,12 +123,12 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): else: raise litellm.UnsupportedParamsError( message=( - "gpt-5 models don't support temperature={}. " + f"gpt-5 models don't support temperature={temperature}. " "Only temperature=1 is supported. " "For models like gpt-5.1/5.4, temperature is supported " "when reasoning.effort='none' (or not specified). " "To drop unsupported params set `litellm.drop_params = True`" - ).format(temperature), + ), status_code=400, ) @@ -137,11 +137,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """Strip Anthropic-only `cache_control` markers before sending to OpenAI. OpenAI's Responses API rejects unknown fields on input content blocks @@ -164,11 +164,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def remove_cache_control_flag_from_input_and_tools( self, model: str, # allows overrides to selectively run this - input: Union[str, ResponseInputParam], - tools: Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]] = None, + input: str | ResponseInputParam, + tools: List[ALL_RESPONSES_API_TOOL_PARAMS] | None = None, ) -> Tuple[ - Union[str, ResponseInputParam], - Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]], + str | ResponseInputParam, + List[ALL_RESPONSES_API_TOOL_PARAMS] | None, ]: """Sibling of `remove_cache_control_flag_from_messages_and_tools` on the chat path. Strips Anthropic-only `cache_control` markers from @@ -193,7 +193,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): return input, tools - def _validate_input_param(self, input: Union[str, ResponseInputParam]) -> Union[str, ResponseInputParam]: + def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam: """ Ensure all input fields if pydantic are converted to dict @@ -211,11 +211,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): if item.get("type") == "reasoning": verbose_logger.debug(f"Handling reasoning item: {item}") # Type assertion since we know it's a dict at this point - dict_item = cast(Dict[str, Any], item) + dict_item = cast(dict[str, Any], item) filtered_item = self._handle_reasoning_item(dict_item) else: # For other dict items, just pass through - filtered_item = cast(Dict[str, Any], item) + filtered_item = cast(dict[str, Any], item) validated_input.append(filtered_item) else: validated_input.append(item) @@ -223,7 +223,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): # Input is expected to be either str or List, no single BaseModel expected return input - def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]: + def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]: """ Handle reasoning items specifically to filter out status=None using OpenAI's model. Issue: https://github.com/BerriAI/litellm/issues/13484 @@ -290,7 +290,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): response._hidden_params["headers"] = raw_response_headers return response - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.setdefault("Content-Type", "application/json") @@ -299,7 +299,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -412,9 +412,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: if stream is not True: return False @@ -445,7 +445,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> Tuple[str, dict]: """ Transform the delete response API request into a URL and data @@ -454,7 +454,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): """ encoded_response_id = encode_url_path_segment(response_id, field_name="response_id") url = f"{api_base}/{encoded_response_id}" - data: Dict = {} + data: dict = {} return url, data def transform_delete_response_api_response( @@ -480,7 +480,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> Tuple[str, dict]: """ Transform the get response API request into a URL and data @@ -489,7 +489,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): """ encoded_response_id = encode_url_path_segment(response_id, field_name="response_id") url = f"{api_base}/{encoded_response_id}" - data: Dict = {} + data: dict = {} return url, data def transform_get_response_api_response( @@ -521,15 +521,15 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + after: str | None = None, + before: str | None = None, + include: List[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - ) -> Tuple[str, Dict]: + ) -> Tuple[str, dict]: encoded_response_id = encode_url_path_segment(response_id, field_name="response_id") url = f"{api_base}/{encoded_response_id}/input_items" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if after is not None: params["after"] = after if before is not None: @@ -546,7 +546,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - ) -> Dict: + ) -> dict: try: return raw_response.json() except Exception: @@ -561,7 +561,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> Tuple[str, dict]: """ Transform the cancel response API request into a URL and data @@ -570,7 +570,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): """ encoded_response_id = encode_url_path_segment(response_id, field_name="response_id") url = f"{api_base}/{encoded_response_id}/cancel" - data: Dict = {} + data: dict = {} return url, data def transform_cancel_response_api_response( @@ -600,12 +600,12 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def transform_compact_response_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> Tuple[str, dict]: """ Transform the compact response API request into a URL and data diff --git a/litellm/llms/openai/speech/guardrail_translation/__init__.py b/litellm/llms/openai/speech/guardrail_translation/__init__.py index ef7d50f861a..7f2a0988468 100644 --- a/litellm/llms/openai/speech/guardrail_translation/__init__.py +++ b/litellm/llms/openai/speech/guardrail_translation/__init__.py @@ -10,4 +10,4 @@ guardrail_translation_mappings = { CallTypes.aspeech: OpenAITextToSpeechHandler, } -__all__ = ["guardrail_translation_mappings", "OpenAITextToSpeechHandler"] +__all__ = ["OpenAITextToSpeechHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/openai/speech/guardrail_translation/handler.py b/litellm/llms/openai/speech/guardrail_translation/handler.py index 3f29a8055d8..5e7c5a481e6 100644 --- a/litellm/llms/openai/speech/guardrail_translation/handler.py +++ b/litellm/llms/openai/speech/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for OpenAI's text-to-speech e The handler processes the 'input' text parameter (output is audio, so no text to guardrail). """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -31,7 +31,7 @@ class OpenAITextToSpeechHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input text by applying guardrails. @@ -80,9 +80,9 @@ class OpenAITextToSpeechHandler(BaseTranslation): self, response: "HttpxBinaryResponseContent", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output - not applicable for text-to-speech. diff --git a/litellm/llms/openai/transcriptions/gpt_transformation.py b/litellm/llms/openai/transcriptions/gpt_transformation.py index 56a1e39ecef..41112cf921f 100644 --- a/litellm/llms/openai/transcriptions/gpt_transformation.py +++ b/litellm/llms/openai/transcriptions/gpt_transformation.py @@ -1,5 +1,3 @@ -from typing import List - from litellm.llms.base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, ) @@ -10,7 +8,7 @@ from .whisper_transformation import OpenAIWhisperAudioTranscriptionConfig class OpenAIGPTAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: """ Get the supported OpenAI params for the `gpt-4o-transcribe` models """ diff --git a/litellm/llms/openai/transcriptions/guardrail_translation/__init__.py b/litellm/llms/openai/transcriptions/guardrail_translation/__init__.py index a6a1a8c2ccf..a9d401d6f49 100644 --- a/litellm/llms/openai/transcriptions/guardrail_translation/__init__.py +++ b/litellm/llms/openai/transcriptions/guardrail_translation/__init__.py @@ -10,4 +10,4 @@ guardrail_translation_mappings = { CallTypes.atranscription: OpenAIAudioTranscriptionHandler, } -__all__ = ["guardrail_translation_mappings", "OpenAIAudioTranscriptionHandler"] +__all__ = ["OpenAIAudioTranscriptionHandler", "guardrail_translation_mappings"] diff --git a/litellm/llms/openai/transcriptions/guardrail_translation/handler.py b/litellm/llms/openai/transcriptions/guardrail_translation/handler.py index fc1cae75b80..7b45cd6d594 100644 --- a/litellm/llms/openai/transcriptions/guardrail_translation/handler.py +++ b/litellm/llms/openai/transcriptions/guardrail_translation/handler.py @@ -5,7 +5,7 @@ This module provides guardrail translation support for OpenAI's audio transcript The handler processes the output transcribed text (input is audio, so no text to guardrail). """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -31,7 +31,7 @@ class OpenAIAudioTranscriptionHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, + litellm_logging_obj: Any | None = None, ) -> Any: """ Process input - not applicable for audio transcription. @@ -55,9 +55,9 @@ class OpenAIAudioTranscriptionHandler(BaseTranslation): self, response: "TranscriptionResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output transcription by applying guardrails to transcribed text. diff --git a/litellm/llms/openai/transcriptions/handler.py b/litellm/llms/openai/transcriptions/handler.py index 76178051ca1..ecfc6d121e0 100644 --- a/litellm/llms/openai/transcriptions/handler.py +++ b/litellm/llms/openai/transcriptions/handler.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Optional, Union, cast +from typing import TYPE_CHECKING, Optional, cast import httpx from openai import AsyncOpenAI, OpenAI @@ -29,7 +29,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): self, openai_aclient: AsyncOpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, ): """ Helper to: @@ -49,7 +49,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): self, openai_client: OpenAI, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, ): """ Helper to: @@ -78,11 +78,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): timeout: float, max_retries: int, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, client=None, atranscription: bool = False, - provider_config: Optional[BaseAudioTranscriptionConfig] = None, + provider_config: BaseAudioTranscriptionConfig | None = None, shared_session: Optional["ClientSession"] = None, ) -> TranscriptionResponse: """ @@ -167,8 +167,8 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): model_response: TranscriptionResponse, timeout: float, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, client=None, max_retries=None, shared_session: Optional["ClientSession"] = None, diff --git a/litellm/llms/openai/transcriptions/whisper_transformation.py b/litellm/llms/openai/transcriptions/whisper_transformation.py index ae7d0bb30b2..84590171b80 100644 --- a/litellm/llms/openai/transcriptions/whisper_transformation.py +++ b/litellm/llms/openai/transcriptions/whisper_transformation.py @@ -1,5 +1,4 @@ import json -from typing import List, Optional, Union from httpx import Headers, Response @@ -21,12 +20,12 @@ from ..common_utils import OpenAIError class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ OPTIONAL @@ -47,7 +46,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return api_base or "" - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: """ Get the supported OpenAI params for the `whisper-1` models """ @@ -79,11 +78,11 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or get_secret_str("OPENAI_API_KEY") @@ -113,7 +112,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): data=data, ) - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return OpenAIError( status_code=status_code, message=error_message, diff --git a/litellm/llms/openai/vector_store_files/transformation.py b/litellm/llms/openai/vector_store_files/transformation.py index 653a31f2e80..1f0971a5917 100644 --- a/litellm/llms/openai/vector_store_files/transformation.py +++ b/litellm/llms/openai/vector_store_files/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional, Tuple, cast +from typing import Any, cast import httpx @@ -22,7 +22,7 @@ from litellm.types.vector_store_files import ( from litellm.utils import add_openai_metadata -def _clean_dict(source: Dict[str, Any]) -> Dict[str, Any]: +def _clean_dict(source: dict[str, Any]) -> dict[str, Any]: return {k: v for k, v in source.items() if v is not None} @@ -30,7 +30,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): ASSISTANTS_HEADER_KEY = "OpenAI-Beta" ASSISTANTS_HEADER_VALUE = "assistants=v2" - def get_auth_credentials(self, litellm_params: Dict[str, Any]) -> VectorStoreFileAuthCredentials: + def get_auth_credentials(self, litellm_params: dict[str, Any]) -> VectorStoreFileAuthCredentials: api_key = litellm_params.get("api_key") if api_key is None: raise ValueError("api_key is required") @@ -42,7 +42,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): def get_vector_store_file_endpoints_by_type( self, - ) -> Dict[str, Tuple[Tuple[str, str], ...]]: + ) -> dict[str, tuple[tuple[str, str], ...]]: return { "read": ( ("GET", "/vector_stores/{vector_store_id}/files"), @@ -62,9 +62,9 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): def validate_environment( self, *, - headers: Dict[str, str], - litellm_params: Optional[GenericLiteLLMParams], - ) -> Dict[str, str]: + headers: dict[str, str], + litellm_params: GenericLiteLLMParams | None, + ) -> dict[str, str]: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.update( @@ -80,9 +80,9 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): def get_complete_url( self, *, - api_base: Optional[str], + api_base: str | None, vector_store_id: str, - litellm_params: Dict[str, Any], + litellm_params: dict[str, Any], ) -> str: base_url = ( api_base @@ -101,8 +101,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): vector_store_id: str, create_request: VectorStoreFileCreateRequest, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: - payload: Dict[str, Any] = _clean_dict(dict(create_request)) + ) -> tuple[str, dict[str, Any]]: + payload: dict[str, Any] = _clean_dict(dict(create_request)) attributes = payload.get("attributes") if isinstance(attributes, dict): filtered_attributes = add_openai_metadata(attributes) @@ -133,7 +133,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): vector_store_id: str, query_params: VectorStoreFileListQueryParams, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: + ) -> tuple[str, dict[str, Any]]: params = _clean_dict(dict(query_params)) return api_base, params @@ -157,7 +157,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: + ) -> tuple[str, dict[str, Any]]: encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") return f"{api_base}/{encoded_file_id}", {} @@ -181,7 +181,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: + ) -> tuple[str, dict[str, Any]]: encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") return f"{api_base}/{encoded_file_id}/content", {} @@ -206,8 +206,8 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): file_id: str, update_request: VectorStoreFileUpdateRequest, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: - payload: Dict[str, Any] = dict(update_request) + ) -> tuple[str, dict[str, Any]]: + payload: dict[str, Any] = dict(update_request) attributes = payload.get("attributes") if isinstance(attributes, dict): filtered_attributes = add_openai_metadata(attributes) @@ -238,7 +238,7 @@ class OpenAIVectorStoreFilesConfig(BaseVectorStoreFilesConfig): vector_store_id: str, file_id: str, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: + ) -> tuple[str, dict[str, Any]]: encoded_file_id = encode_url_path_segment(file_id, field_name="file_id") return f"{api_base}/{encoded_file_id}", {} diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 6ccf8e271e5..2e314b3a429 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx @@ -47,7 +47,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): "write": [("POST", "/vector_stores")], } - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.update( @@ -71,7 +71,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -93,13 +93,13 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url = f"{api_base}/{encoded_vector_store_id}/search" typed_request_body = VectorStoreSearchRequest( @@ -130,7 +130,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: url = api_base # Base URL for creating vector stores metadata = vector_store_create_optional_params.get("metadata", None) metadata_payload = add_openai_metadata(metadata) diff --git a/litellm/llms/openai/videos/transformation.py b/litellm/llms/openai/videos/transformation.py index 855bc410cea..726b4441f49 100644 --- a/litellm/llms/openai/videos/transformation.py +++ b/litellm/llms/openai/videos/transformation.py @@ -1,6 +1,6 @@ import mimetypes from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast from urllib.parse import quote import httpx @@ -64,7 +64,7 @@ class OpenAIVideoConfig(BaseVideoConfig): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """No mapping applied since inputs are in OpenAI spec already""" return dict(video_create_optional_params) @@ -72,8 +72,8 @@ class OpenAIVideoConfig(BaseVideoConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: # Use api_key from litellm_params if available, otherwise fall back to other sources if litellm_params and litellm_params.api_key: @@ -90,7 +90,7 @@ class OpenAIVideoConfig(BaseVideoConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -106,10 +106,10 @@ class OpenAIVideoConfig(BaseVideoConfig): model: str, prompt: str, api_base: str, - video_create_optional_request_params: Dict, + video_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles, str]: + ) -> tuple[dict, RequestFiles, str]: """ Transform the video creation request for OpenAI API. """ @@ -122,13 +122,13 @@ class OpenAIVideoConfig(BaseVideoConfig): # Create the request data video_create_request = CreateVideoRequest(model=model, prompt=prompt, **video_create_optional_request_params) - request_dict = cast(Dict, video_create_request) + request_dict = cast(dict, video_create_request) request_dict = self._decode_character_ids_in_create_video_request(request_dict) # Handle input_reference parameter if provided _input_reference = video_create_optional_request_params.get("input_reference") data_without_files = {k: v for k, v in request_dict.items() if k not in ["input_reference"]} - files_list: List[Tuple[str, FileTypes]] = [] + files_list: list[tuple[str, FileTypes]] = [] # Handle input_reference parameter if _input_reference is not None: @@ -139,7 +139,7 @@ class OpenAIVideoConfig(BaseVideoConfig): ) return data_without_files, files_list, api_base - def _decode_character_ids_in_create_video_request(self, request_dict: Dict) -> Dict: + def _decode_character_ids_in_create_video_request(self, request_dict: dict) -> dict: """ Decode LiteLLM-managed encoded character ids for provider requests. @@ -151,7 +151,7 @@ class OpenAIVideoConfig(BaseVideoConfig): if not isinstance(raw_characters, list): return request_dict - decoded_characters: List[Any] = [] + decoded_characters: list[Any] = [] for character in raw_characters: if not isinstance(character, dict): decoded_characters.append(character) @@ -173,8 +173,8 @@ class OpenAIVideoConfig(BaseVideoConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: """Transform the OpenAI video creation response.""" video_obj = VideoObject.model_validate(raw_response.json()) @@ -199,8 +199,8 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - variant: Optional[str] = None, - ) -> Tuple[str, Dict]: + variant: str | None = None, + ) -> tuple[str, dict]: """ Transform the video content request for OpenAI API. @@ -221,7 +221,7 @@ class OpenAIVideoConfig(BaseVideoConfig): url = f"{url}?variant={quote(variant, safe='')}" # No additional data needed for GET content request - data: Dict[str, object] = {} + data: dict[str, object] = {} return url, data @@ -232,8 +232,8 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video remix request for OpenAI API. @@ -267,7 +267,7 @@ class OpenAIVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """ Transform the OpenAI video remix response. @@ -297,11 +297,11 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video list request for OpenAI API. @@ -331,8 +331,8 @@ class OpenAIVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - ) -> Dict[str, str]: + custom_llm_provider: str | None = None, + ) -> dict[str, str]: response_data = raw_response.json() if custom_llm_provider and "data" in response_data: @@ -374,7 +374,7 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video delete request for OpenAI API. @@ -388,7 +388,7 @@ class OpenAIVideoConfig(BaseVideoConfig): url = f"{api_base.rstrip('/')}/{encoded_video_id}" # No data needed for DELETE request - data: Dict[str, object] = {} + data: dict[str, object] = {} return url, data @@ -411,7 +411,7 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the OpenAI video retrieve request. """ @@ -423,7 +423,7 @@ class OpenAIVideoConfig(BaseVideoConfig): url = f"{api_base.rstrip('/')}/{encoded_video_id}" # No additional data needed for GET request - data: Dict[str, object] = {} + data: dict[str, object] = {} return url, data @@ -431,7 +431,7 @@ class OpenAIVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """ Transform the OpenAI video retrieve response. @@ -444,9 +444,7 @@ class OpenAIVideoConfig(BaseVideoConfig): return video_obj - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ...base_llm.chat.transformation import BaseLLMException raise BaseLLMException( @@ -462,9 +460,9 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, list]: + ) -> tuple[str, list]: url = f"{api_base.rstrip('/')}/characters" - files_list: List[Tuple[str, FileTypes]] = [("name", (None, name))] + files_list: list[tuple[str, FileTypes]] = [("name", (None, name))] self._add_video_to_files(files_list, video, "video") return url, files_list @@ -481,7 +479,7 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: original_character_id = extract_original_character_id(character_id) encoded_character_id = encode_url_path_segment(original_character_id, field_name="character_id") url = f"{api_base.rstrip('/')}/characters/{encoded_character_id}" @@ -501,12 +499,12 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, object]] = None, - prefetched_source_data: Optional[Dict[str, object]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, object] | None = None, + prefetched_source_data: dict[str, object] | None = None, + ) -> tuple[str, dict]: original_video_id = extract_original_video_id(video_id) url = f"{api_base.rstrip('/')}/edits" - data: Dict[str, object] = {"prompt": prompt, "video": {"id": original_video_id}} + data: dict[str, object] = {"prompt": prompt, "video": {"id": original_video_id}} if extra_body: data.update(extra_body) return url, data @@ -515,8 +513,8 @@ class OpenAIVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: video_obj = VideoObject.model_validate(raw_response.json()) if custom_llm_provider and video_obj.id: @@ -531,11 +529,11 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, object]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, object] | None = None, + ) -> tuple[str, dict]: original_video_id = extract_original_video_id(video_id) url = f"{api_base.rstrip('/')}/extensions" - data: Dict[str, object] = { + data: dict[str, object] = { "prompt": prompt, "seconds": seconds, "video": {"id": original_video_id}, @@ -548,7 +546,7 @@ class OpenAIVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: video_obj = VideoObject.model_validate(raw_response.json()) if custom_llm_provider and video_obj.id: @@ -557,7 +555,7 @@ class OpenAIVideoConfig(BaseVideoConfig): def _add_image_to_files( self, - files_list: List[Tuple[str, Any]], + files_list: list[tuple[str, Any]], image: Any, field_name: str, ) -> None: @@ -571,7 +569,7 @@ class OpenAIVideoConfig(BaseVideoConfig): def _add_video_to_files( self, - files_list: List[Tuple[str, FileTypes]], + files_list: list[tuple[str, FileTypes]], video: FileContent, field_name: str, ) -> None: @@ -593,12 +591,7 @@ class OpenAIVideoConfig(BaseVideoConfig): # Fast-path detection for common MP4 signatures when filename is missing/incorrect. try: header_bytes = b"" - if isinstance(video, BytesIO): - current_pos = video.tell() - video.seek(0) - header_bytes = video.read(64) - video.seek(current_pos) - elif isinstance(video, BufferedReader): + if isinstance(video, BytesIO) or isinstance(video, BufferedReader): current_pos = video.tell() video.seek(0) header_bytes = video.read(64) diff --git a/litellm/llms/openai_like/chat/handler.py b/litellm/llms/openai_like/chat/handler.py index 866cb7a9531..cdf4ab7abb8 100644 --- a/litellm/llms/openai_like/chat/handler.py +++ b/litellm/llms/openai_like/chat/handler.py @@ -6,7 +6,7 @@ For handling OpenAI-like chat completions, like IBM WatsonX, etc. import json from collections.abc import Callable -from typing import Any, Optional, Union +from typing import Any import httpx @@ -25,14 +25,14 @@ from .transformation import OpenAILikeChatConfig async def make_call( - client: Optional[AsyncHTTPHandler], + client: AsyncHTTPHandler | None, api_base: str, headers: dict, data: str, model: str, messages: list, logging_obj, - streaming_decoder: Optional[CustomStreamingDecoder] = None, + streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ): if client is None: @@ -59,16 +59,16 @@ async def make_call( def make_sync_call( - client: Optional[HTTPHandler], + client: HTTPHandler | None, api_base: str, headers: dict, data: str, model: str, messages: list, logging_obj, - streaming_decoder: Optional[CustomStreamingDecoder] = None, + streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, ): if client is None: client = litellm.module_level_client # Create a new client if none provided @@ -119,8 +119,8 @@ class OpenAILikeChatHandler(OpenAILikeBase): litellm_params=None, logger_fn=None, headers={}, - client: Optional[AsyncHTTPHandler] = None, - streaming_decoder: Optional[CustomStreamingDecoder] = None, + client: AsyncHTTPHandler | None = None, + streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ) -> CustomStreamWrapper: data["stream"] = True @@ -152,18 +152,18 @@ class OpenAILikeChatHandler(OpenAILikeBase): model_response: ModelResponse, custom_llm_provider: str, print_verbose: Callable, - client: Optional[AsyncHTTPHandler], + client: AsyncHTTPHandler | None, encoding, api_key, logging_obj, stream, data: dict, - base_model: Optional[str], + base_model: str | None, optional_params: dict, litellm_params=None, logger_fn=None, headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, json_mode: bool = False, ) -> ModelResponse: if timeout is None: @@ -213,23 +213,22 @@ class OpenAILikeChatHandler(OpenAILikeBase): model_response: ModelResponse, print_verbose: Callable, encoding, - api_key: Optional[str], + api_key: str | None, logging_obj, optional_params: dict, acompletion=None, litellm_params: dict = {}, logger_fn=None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - custom_endpoint: Optional[bool] = None, - streaming_decoder: Optional[ - CustomStreamingDecoder - ] = None, # if openai-compatible api needs custom stream decoder - e.g. sagemaker + headers: dict | None = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + custom_endpoint: bool | None = None, + streaming_decoder: CustomStreamingDecoder + | None = None, # if openai-compatible api needs custom stream decoder - e.g. sagemaker fake_stream: bool = False, ): custom_endpoint = custom_endpoint or optional_params.pop("custom_endpoint", None) - base_model: Optional[str] = optional_params.pop("base_model", None) + base_model: str | None = optional_params.pop("base_model", None) api_base, headers = self._validate_environment( api_base=api_base, api_key=api_key, diff --git a/litellm/llms/openai_like/chat/transformation.py b/litellm/llms/openai_like/chat/transformation.py index a2c847a410f..895d5a99971 100644 --- a/litellm/llms/openai_like/chat/transformation.py +++ b/litellm/llms/openai_like/chat/transformation.py @@ -2,7 +2,7 @@ OpenAI-like chat completion transformation """ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -23,9 +23,9 @@ else: class OpenAILikeChatConfig(OpenAIGPTConfig): def _get_openai_compatible_provider_info( self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[Optional[str], Optional[str]]: + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("OPENAI_LIKE_API_BASE") # type: ignore dynamic_api_key = api_key or get_secret_str("OPENAI_LIKE_API_KEY") or "" # vllm does not require an api key return api_base, dynamic_api_key @@ -83,14 +83,14 @@ class OpenAILikeChatConfig(OpenAIGPTConfig): stream: bool, logging_obj: LiteLLMLoggingObj, optional_params: dict, - api_key: Optional[str], - data: Union[dict, str], - messages: List, + api_key: str | None, + data: dict | str, + messages: list, print_verbose, encoding, - json_mode: Optional[bool], - custom_llm_provider: Optional[str], - base_model: Optional[str], + json_mode: bool | None, + custom_llm_provider: str | None, + base_model: str | None, ) -> ModelResponse: response_json = response.json() logging_obj.post_call( @@ -126,12 +126,12 @@ class OpenAILikeChatConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: return OpenAILikeChatConfig._transform_response( model=model, diff --git a/litellm/llms/openai_like/common_utils.py b/litellm/llms/openai_like/common_utils.py index 40f2e5c3f5c..11b85e52af5 100644 --- a/litellm/llms/openai_like/common_utils.py +++ b/litellm/llms/openai_like/common_utils.py @@ -1,4 +1,4 @@ -from typing import Literal, Optional, Tuple +from typing import Literal import httpx @@ -18,12 +18,12 @@ class OpenAILikeBase: def _validate_environment( self, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, endpoint_type: Literal["chat_completions", "embeddings"], - headers: Optional[dict], - custom_endpoint: Optional[bool], - ) -> Tuple[str, dict]: + headers: dict | None, + custom_endpoint: bool | None, + ) -> tuple[str, dict]: if api_key is None and headers is None: raise OpenAILikeError( status_code=400, @@ -44,11 +44,11 @@ class OpenAILikeBase: if ( api_key is not None and "Authorization" not in headers ): # [TODO] remove 'validate_environment' from OpenAI base. should use llm providers config for this only. - headers.update({"Authorization": "Bearer {}".format(api_key)}) + headers.update({"Authorization": f"Bearer {api_key}"}) if not custom_endpoint: if endpoint_type == "chat_completions": - api_base = "{}/chat/completions".format(api_base) + api_base = f"{api_base}/chat/completions" elif endpoint_type == "embeddings": - api_base = "{}/embeddings".format(api_base) + api_base = f"{api_base}/embeddings" return api_base, headers diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index 6e6a40bed38..40c3e2a07a7 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -3,7 +3,7 @@ Dynamic configuration class generator for JSON-based providers. """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Tuple, Union, overload +from typing import Any, Literal, overload from litellm._logging import verbose_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -26,20 +26,20 @@ def create_config_class(provider: SimpleProviderConfig): class JSONProviderConfig(base_class): # type: ignore[valid-type,misc] @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """Transform messages based on special_handling config""" # Handle content list to string conversion if configured @@ -52,8 +52,8 @@ def create_config_class(provider: SimpleProviderConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: """Get API base and key from JSON config""" # Resolve base URL @@ -70,12 +70,12 @@ def create_config_class(provider: SimpleProviderConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """Build complete URL for the API endpoint""" if not api_base: @@ -163,7 +163,7 @@ def create_config_class(provider: SimpleProviderConfig): return optional_params @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return provider.slug return JSONProviderConfig @@ -196,7 +196,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or get_secret_str(provider.api_key_env) @@ -206,7 +206,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: if not api_base: @@ -224,7 +224,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, diff --git a/litellm/llms/openai_like/embedding/handler.py b/litellm/llms/openai_like/embedding/handler.py index 52eafc05b2c..8e82cc8f3e9 100644 --- a/litellm/llms/openai_like/embedding/handler.py +++ b/litellm/llms/openai_like/embedding/handler.py @@ -3,7 +3,6 @@ ## Allows jina ai embedding calls - which don't allow 'encoding_format' in payload. import json -from typing import Optional import httpx @@ -86,14 +85,14 @@ class OpenAILikeEmbeddingHandler(OpenAILikeBase): input: list, timeout: float, logging_obj, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, optional_params: dict, - model_response: Optional[EmbeddingResponse] = None, + model_response: EmbeddingResponse | None = None, client=None, aembedding=None, - custom_endpoint: Optional[bool] = None, - headers: Optional[dict] = None, + custom_endpoint: bool | None = None, + headers: dict | None = None, ) -> EmbeddingResponse: api_base, headers = self._validate_environment( api_base=api_base, diff --git a/litellm/llms/openai_like/json_loader.py b/litellm/llms/openai_like/json_loader.py index 4640bb8a422..bc10b7bd62f 100644 --- a/litellm/llms/openai_like/json_loader.py +++ b/litellm/llms/openai_like/json_loader.py @@ -4,7 +4,6 @@ JSON-based provider configuration loader for OpenAI-compatible providers. import json from pathlib import Path -from typing import Dict, Optional from litellm._logging import verbose_logger @@ -27,7 +26,7 @@ class SimpleProviderConfig: class JSONProviderRegistry: """Load providers from JSON once on import""" - _providers: Dict[str, SimpleProviderConfig] = {} + _providers: dict[str, SimpleProviderConfig] = {} _loaded = False @classmethod @@ -56,7 +55,7 @@ class JSONProviderRegistry: cls._loaded = True @classmethod - def get(cls, slug: str) -> Optional[SimpleProviderConfig]: + def get(cls, slug: str) -> SimpleProviderConfig | None: """Get a provider configuration by slug""" return cls._providers.get(slug) diff --git a/litellm/llms/openai_like/messages/transformation.py b/litellm/llms/openai_like/messages/transformation.py index 0d593d8d0f4..ca7602961ff 100644 --- a/litellm/llms/openai_like/messages/transformation.py +++ b/litellm/llms/openai_like/messages/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Optional +from typing import Any import litellm from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -30,9 +30,9 @@ class OpenAILikeAnthropicMessagesConfig(AnthropicMessagesConfig): messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> tuple[dict[str, str], Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict[str, str], str | None]: present = {key.lower() for key in headers} needs_auth = bool(api_key) and "authorization" not in present and "x-api-key" not in present defaults: dict[str, str] = { @@ -55,20 +55,19 @@ class OpenAILikeAnthropicMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if not api_base: raise ValueError("api_base is required to forward Anthropic /v1/messages to a native endpoint") base = api_base.rstrip("/") if base.endswith("/v1/messages"): return base - if base.endswith("/v1"): - base = base[: -len("/v1")] + base = base.removesuffix("/v1") return f"{base}/v1/messages" @@ -86,16 +85,16 @@ class JSONProviderAnthropicMessagesConfig(OpenAILikeAnthropicMessagesConfig): self._provider = provider @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return self._provider.slug def should_strip_billing_metadata(self) -> bool: return True - def _resolve_api_key(self, api_key: Optional[str]) -> Optional[str]: + def _resolve_api_key(self, api_key: str | None) -> str | None: return api_key or get_secret_str(self._provider.api_key_env) or litellm.api_key - def _resolve_api_base(self, api_base: Optional[str]) -> str: + def _resolve_api_base(self, api_base: str | None) -> str: env_api_base = get_secret_str(self._provider.api_base_env) if self._provider.api_base_env else None return api_base or env_api_base or self._provider.base_url @@ -106,9 +105,9 @@ class JSONProviderAnthropicMessagesConfig(OpenAILikeAnthropicMessagesConfig): messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> tuple[dict[str, str], Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict[str, str], str | None]: return super().validate_anthropic_messages_environment( headers=headers, model=model, @@ -121,12 +120,12 @@ class JSONProviderAnthropicMessagesConfig(OpenAILikeAnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: return super().get_complete_url( api_base=self._resolve_api_base(api_base), diff --git a/litellm/llms/openai_like/responses/transformation.py b/litellm/llms/openai_like/responses/transformation.py index ff496901363..ea8830feb04 100644 --- a/litellm/llms/openai_like/responses/transformation.py +++ b/litellm/llms/openai_like/responses/transformation.py @@ -6,8 +6,6 @@ Inherits everything from OpenAIResponsesAPIConfig; subclasses only override provider-specific resolution (slug, API key env var, base URL). """ -from typing import Optional, Union - from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -24,14 +22,14 @@ class OpenAILikeResponsesConfig(OpenAIResponsesAPIConfig): """ @property - def custom_llm_provider(self) -> Union[str, LlmProviders]: # type: ignore[override] + def custom_llm_provider(self) -> str | LlmProviders: # type: ignore[override] return "openai_like" def validate_environment( self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or get_secret_str("OPENAI_LIKE_API_KEY") @@ -41,7 +39,7 @@ class OpenAILikeResponsesConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = api_base or get_secret_str("OPENAI_LIKE_API_BASE") diff --git a/litellm/llms/openrouter/chat/transformation.py b/litellm/llms/openrouter/chat/transformation.py index 10a3e208089..da88ac81cbb 100644 --- a/litellm/llms/openrouter/chat/transformation.py +++ b/litellm/llms/openrouter/chat/transformation.py @@ -8,7 +8,7 @@ Docs: https://openrouter.ai/docs/parameters from collections.abc import AsyncIterator, Iterator from enum import Enum -from typing import Any, List, Optional, Tuple, Union, cast +from typing import Any, cast import httpx @@ -89,15 +89,15 @@ class OpenrouterConfig(OpenAIGPTConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, - messages: List[AllMessageValues], - tools: Optional[List["ChatCompletionToolParam"]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + messages: list[AllMessageValues], + tools: list["ChatCompletionToolParam"] | None = None, + ) -> tuple[list[AllMessageValues], list["ChatCompletionToolParam"] | None]: if self._supports_cache_control_in_content(model): return messages, tools else: return super().remove_cache_control_flag_from_messages_and_tools(model, messages, tools) - def _move_cache_control_to_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + def _move_cache_control_to_content(self, messages: list[AllMessageValues]) -> list[AllMessageValues]: """ Move cache_control from message level to content blocks. OpenRouter requires cache_control to be inside content blocks, not at message level. @@ -105,7 +105,7 @@ class OpenrouterConfig(OpenAIGPTConfig): To avoid exceeding Anthropic's limit of 4 cache breakpoints, cache_control is only added to the LAST content block in each message. """ - transformed_messages: List[AllMessageValues] = [] + transformed_messages: list[AllMessageValues] = [] for message in messages: message_dict = dict(message) cache_control = message_dict.pop("cache_control", None) @@ -142,7 +142,7 @@ class OpenrouterConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -174,12 +174,12 @@ class OpenrouterConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: Any, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the response from OpenRouter API. @@ -225,9 +225,7 @@ class OpenrouterConfig(OpenAIGPTConfig): return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OpenRouterException( message=error_message, status_code=status_code, @@ -236,9 +234,9 @@ class OpenrouterConfig(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return OpenRouterChatCompletionStreamingHandler( streaming_response=streaming_response, diff --git a/litellm/llms/openrouter/embedding/transformation.py b/litellm/llms/openrouter/embedding/transformation.py index c6c3df083a1..1d74504f0e7 100644 --- a/litellm/llms/openrouter/embedding/transformation.py +++ b/litellm/llms/openrouter/embedding/transformation.py @@ -7,7 +7,7 @@ OpenRouter is OpenAI-compatible and supports embeddings via the /v1/embeddings e Docs: https://openrouter.ai/docs """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -40,8 +40,8 @@ class OpenrouterEmbeddingConfig(BaseEmbeddingConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for OpenRouter API. @@ -74,12 +74,12 @@ class OpenrouterEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for OpenRouter Embedding API endpoint. @@ -125,7 +125,7 @@ class OpenrouterEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, diff --git a/litellm/llms/openrouter/image_edit/transformation.py b/litellm/llms/openrouter/image_edit/transformation.py index f4531932f96..fad7d53577c 100644 --- a/litellm/llms/openrouter/image_edit/transformation.py +++ b/litellm/llms/openrouter/image_edit/transformation.py @@ -42,7 +42,7 @@ Response format: import base64 from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -88,9 +88,9 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: supported_params = self.get_supported_openai_params(model) - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} for key, value in image_edit_optional_params.items(): if key in supported_params: @@ -113,9 +113,9 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or litellm.api_key or get_secret_str("OPENROUTER_API_KEY") if not api_key: @@ -134,7 +134,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: base_url = api_base or get_secret_str("OPENROUTER_API_BASE") or "https://openrouter.ai/api/v1" @@ -146,13 +146,13 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: - content_parts: List[Dict[str, Any]] = [] + ) -> tuple[dict, RequestFiles]: + content_parts: list[dict[str, Any]] = [] # Add source image(s) as base64 data URLs if image is not None: @@ -174,7 +174,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): if prompt: content_parts.append({"type": "text", "text": prompt}) - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "model": model, "messages": [ { @@ -203,7 +203,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): response_json = raw_response.json() except Exception as e: raise OpenRouterException( - message=f"Error parsing OpenRouter response: {str(e)}", + message=f"Error parsing OpenRouter response: {e!s}", status_code=raw_response.status_code, headers=raw_response.headers, ) @@ -246,7 +246,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): except Exception as e: raise OpenRouterException( - message=f"Error transforming OpenRouter image edit response: {str(e)}", + message=f"Error transforming OpenRouter image edit response: {e!s}", status_code=500, headers={}, ) @@ -254,9 +254,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): self._set_usage_and_cost(model_response, response_json, model) return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OpenRouterException( message=error_message, status_code=status_code, @@ -284,7 +282,7 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): } return size_to_aspect_ratio.get(size, "1:1") - def _map_quality_to_image_size(self, quality: str) -> Optional[str]: + def _map_quality_to_image_size(self, quality: str) -> str | None: """ Map OpenAI quality to OpenRouter image_size format. diff --git a/litellm/llms/openrouter/image_generation/transformation.py b/litellm/llms/openrouter/image_generation/transformation.py index eabb76f00c0..1114bb41275 100644 --- a/litellm/llms/openrouter/image_generation/transformation.py +++ b/litellm/llms/openrouter/image_generation/transformation.py @@ -27,7 +27,7 @@ Response format: } """ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -36,10 +36,11 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) +from litellm.llms.openrouter.common_utils import OpenRouterException from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( - OpenAIImageGenerationOptionalParams, AllMessageValues, + OpenAIImageGenerationOptionalParams, ) from litellm.types.utils import ( ImageObject, @@ -47,7 +48,6 @@ from litellm.types.utils import ( ImageUsage, ImageUsageInputTokensDetails, ) -from litellm.llms.openrouter.common_utils import OpenRouterException if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -64,7 +64,7 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): and extract images from chat responses. """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for OpenRouter image generation. @@ -158,7 +158,7 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): } return size_to_aspect_ratio.get(size, "1:1") - def _map_quality_to_image_size(self, quality: str) -> Optional[str]: + def _map_quality_to_image_size(self, quality: str) -> str | None: """ Map OpenAI quality to OpenRouter image_size format. @@ -236,12 +236,12 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for OpenRouter image generation. @@ -261,11 +261,11 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: api_key = api_key or litellm.api_key or get_secret_str("OPENROUTER_API_KEY") headers.update( @@ -318,8 +318,8 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform OpenRouter chat completion response to ImageResponse format. @@ -345,7 +345,7 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): response_json = raw_response.json() except Exception as e: raise OpenRouterException( - message=f"Error parsing OpenRouter response: {str(e)}", + message=f"Error parsing OpenRouter response: {e!s}", status_code=raw_response.status_code, headers=raw_response.headers, ) @@ -394,14 +394,12 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): except Exception as e: raise OpenRouterException( - message=f"Error transforming OpenRouter image generation response: {str(e)}", + message=f"Error transforming OpenRouter image generation response: {e!s}", status_code=500, headers={}, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """Get the appropriate error class for OpenRouter errors.""" return OpenRouterException( message=error_message, diff --git a/litellm/llms/openrouter/responses/transformation.py b/litellm/llms/openrouter/responses/transformation.py index 217a419ed22..7fdb9e896e2 100644 --- a/litellm/llms/openrouter/responses/transformation.py +++ b/litellm/llms/openrouter/responses/transformation.py @@ -8,8 +8,6 @@ encrypted_content for multi-turn stateless workflows. Docs: https://openrouter.ai/docs/api/reference/responses/overview """ -from typing import Optional - import litellm from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str @@ -37,7 +35,7 @@ class OpenRouterResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, - litellm_params: Optional[GenericLiteLLMParams], + litellm_params: GenericLiteLLMParams | None, ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = ( @@ -61,7 +59,7 @@ class OpenRouterResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: api_base = ( diff --git a/litellm/llms/opensandbox/sandbox/transformation.py b/litellm/llms/opensandbox/sandbox/transformation.py index 60266c988df..6b573b35b02 100644 --- a/litellm/llms/opensandbox/sandbox/transformation.py +++ b/litellm/llms/opensandbox/sandbox/transformation.py @@ -1,7 +1,7 @@ import asyncio import json import time -from typing import Union, cast +from typing import cast import httpx @@ -19,10 +19,10 @@ from litellm.constants import ( OPEN_SANDBOX_READY_TIMEOUT, ) from litellm.llms.base_llm.sandbox.transformation import ( + SANDBOX_MAX_OUTPUT_BYTES, BaseSandboxConfig, CodeExecutionResult, ContainerHandle, - SANDBOX_MAX_OUTPUT_BYTES, ) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -130,7 +130,7 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig): async def arun_code( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, code: str, api_key: str | None = None, api_base: str | None = None, @@ -172,7 +172,7 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig): async def adelete_sandbox( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, api_key: str | None = None, api_base: str | None = None, client: AsyncHTTPHandler | None = None, @@ -198,7 +198,7 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig): async def _ensure_handle( self, *, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, api_key: str | None, api_base: str | None, use_server_proxy: bool, @@ -428,7 +428,7 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig): return f"{protocol}://{normalized_endpoint}" @staticmethod - def _as_handle(container: Union[ContainerHandle, str], *, api_base: str | None) -> ContainerHandle: + def _as_handle(container: ContainerHandle | str, *, api_base: str | None) -> ContainerHandle: if isinstance(container, ContainerHandle): return container handle = ContainerHandle( diff --git a/litellm/llms/ovhcloud/audio_transcription/transformation.py b/litellm/llms/ovhcloud/audio_transcription/transformation.py index 43b68c6503d..1fd2174a5ba 100644 --- a/litellm/llms/ovhcloud/audio_transcription/transformation.py +++ b/litellm/llms/ovhcloud/audio_transcription/transformation.py @@ -5,8 +5,6 @@ Our unified API follows the OpenAI standard. More information on our website: https://endpoints.ai.cloud.ovh.net """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -26,7 +24,7 @@ from ..utils import OVHCloudException class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: # OVHCloud implements the OpenAI-compatible Whisper interface. # We pass through the same optional params as the OpenAI Whisper API. return [ @@ -52,20 +50,18 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1" if api_base is None else api_base.rstrip("/") complete_url = f"{api_base}/audio/transcriptions" return complete_url - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OVHCloudException( message=error_message, status_code=status_code, @@ -76,11 +72,11 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("OVHCLOUD_API_KEY") diff --git a/litellm/llms/ovhcloud/chat/transformation.py b/litellm/llms/ovhcloud/chat/transformation.py index 0090ae168f7..0b4f8e6168b 100644 --- a/litellm/llms/ovhcloud/chat/transformation.py +++ b/litellm/llms/ovhcloud/chat/transformation.py @@ -5,39 +5,35 @@ Our unified API follows the OpenAI standard. More information on our website: https://endpoints.ai.cloud.ovh.net """ -from typing import Optional, Union, List - import httpx -from litellm.utils import ModelResponseStream -from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig -from litellm.llms.ovhcloud.utils import OVHCloudException + from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseLLMException - +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.llms.ovhcloud.utils import OVHCloudException from litellm.types.llms.openai import AllMessageValues +from litellm.utils import ModelResponseStream class OVHCloudChatConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "ovhcloud" def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1" if api_base is None else api_base.rstrip("/") complete_url = f"{api_base}/chat/completions" return complete_url - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OVHCloudException( message=error_message, status_code=status_code, @@ -57,7 +53,7 @@ class OVHCloudChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/ovhcloud/embedding/transformation.py b/litellm/llms/ovhcloud/embedding/transformation.py index 006f2a2349b..93d761ef408 100644 --- a/litellm/llms/ovhcloud/embedding/transformation.py +++ b/litellm/llms/ovhcloud/embedding/transformation.py @@ -3,8 +3,6 @@ This is OpenAI compatible - no transformation is applied """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -23,12 +21,12 @@ class OVHCloudEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1" if api_base is None else api_base.rstrip("/") complete_url = f"{api_base}/embeddings" @@ -38,11 +36,11 @@ class OVHCloudEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("OVHCLOUD_API_KEY") @@ -89,7 +87,7 @@ class OVHCloudEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -115,7 +113,5 @@ class OVHCloudEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return OVHCloudException(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/ovhcloud/utils.py b/litellm/llms/ovhcloud/utils.py index 046df4bca1b..d5e8a34f655 100644 --- a/litellm/llms/ovhcloud/utils.py +++ b/litellm/llms/ovhcloud/utils.py @@ -3,5 +3,3 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException class OVHCloudException(BaseLLMException): """OVHCloud AI Endpoints exception handling class""" - - pass diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index 56566aea0b1..3d90212208d 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -4,7 +4,7 @@ Calls Parallel AI's /v1/search endpoint to search the web. Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -18,8 +18,8 @@ from litellm.secret_managers.main import get_secret_str class _ParallelAISourcePolicy(TypedDict, total=False): - include_domains: List[str] - exclude_domains: List[str] + include_domains: list[str] + exclude_domains: list[str] after_date: str @@ -30,7 +30,7 @@ class _ParallelAIExcerptSettings(TypedDict, total=False): class _ParallelAIAdvancedSettings(TypedDict, total=False): source_policy: _ParallelAISourcePolicy excerpt_settings: _ParallelAIExcerptSettings - fetch_policy: Dict + fetch_policy: dict location: str max_results: int @@ -41,7 +41,7 @@ class ParallelAISearchRequest(TypedDict, total=False): Based on: https://docs.parallel.ai/api-reference/search/search """ - search_queries: List[str] # Required - at least one keyword search query + search_queries: list[str] # Required - at least one keyword search query objective: str # Optional - natural-language description of search goal mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced') max_chars_total: int # Optional - upper bound on total excerpt characters @@ -62,11 +62,11 @@ class ParallelAISearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: api_key = self.resolve_server_api_key( caller_api_key=api_key, caller_api_base=api_base, @@ -82,9 +82,9 @@ class ParallelAISearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: api_base = api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE @@ -97,10 +97,10 @@ class ParallelAISearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Parallel AI v1 API format. @@ -170,7 +170,7 @@ class ParallelAISearchConfig(BaseSearchConfig): # unified-spec param with no v1 equivalent params.pop("max_tokens_per_page", None) - result_data: Dict = dict(request_data) + result_data: dict = dict(request_data) result_data.update(params) return result_data diff --git a/litellm/llms/pass_through/guardrail_translation/__init__.py b/litellm/llms/pass_through/guardrail_translation/__init__.py index 46fea242c13..deffaa3aec6 100644 --- a/litellm/llms/pass_through/guardrail_translation/__init__.py +++ b/litellm/llms/pass_through/guardrail_translation/__init__.py @@ -12,7 +12,7 @@ guardrail_translation_mappings = { } __all__ = [ - "guardrail_translation_mappings", "LlmPassthroughRouteHandler", "PassThroughEndpointHandler", + "guardrail_translation_mappings", ] diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 8ca600b0bcf..11d2cbfb65b 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -6,7 +6,7 @@ It uses the field targeting configuration from litellm_logging_obj to extract specific fields for guardrail processing. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -31,8 +31,8 @@ class PassThroughEndpointHandler(BaseTranslation): def _get_guardrail_settings( self, litellm_logging_obj: Optional["LiteLLMLoggingObj"], - guardrail_name: Optional[str], - ) -> Optional[PassThroughGuardrailSettings]: + guardrail_name: str | None, + ) -> PassThroughGuardrailSettings | None: """ Get the guardrail settings for a specific guardrail from logging_obj. """ @@ -52,7 +52,7 @@ class PassThroughEndpointHandler(BaseTranslation): def _extract_text_for_guardrail( self, data: dict, - field_expressions: Optional[List[str]], + field_expressions: list[str] | None, ) -> str: """ Extract text from data for guardrail processing. @@ -130,8 +130,8 @@ class PassThroughEndpointHandler(BaseTranslation): response: Any, guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """ Process output response by applying guardrails to targeted fields. @@ -192,10 +192,10 @@ class PassThroughEndpointHandler(BaseTranslation): return response -_PROVIDER_HANDLERS: Dict[str, Type[BaseTranslation]] = {} +_PROVIDER_HANDLERS: dict[str, type[BaseTranslation]] = {} -def _get_provider_handlers() -> Dict[str, Type[BaseTranslation]]: +def _get_provider_handlers() -> dict[str, type[BaseTranslation]]: global _PROVIDER_HANDLERS if not _PROVIDER_HANDLERS: from litellm.llms.bedrock.passthrough.guardrail_translation.handler import ( @@ -239,8 +239,8 @@ class LlmPassthroughRouteHandler(BaseTranslation): response: Any, guardrail_to_apply: "CustomGuardrail", litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: provider = (request_data or {}).get("custom_llm_provider") handler_cls = _get_provider_handlers().get(provider or "") @@ -259,7 +259,7 @@ class LlmPassthroughRouteHandler(BaseTranslation): ) @staticmethod - def is_event_stream_response(provider: Optional[str], content_type: str) -> bool: + def is_event_stream_response(provider: str | None, content_type: str) -> bool: handler_cls = _get_provider_handlers().get(provider or "") detector = getattr(handler_cls, "is_event_stream_content_type", None) if detector is None: @@ -267,7 +267,7 @@ class LlmPassthroughRouteHandler(BaseTranslation): return detector(content_type) @staticmethod - def event_stream_media_type(provider: Optional[str]) -> Optional[str]: + def event_stream_media_type(provider: str | None) -> str | None: handler_cls = _get_provider_handlers().get(provider or "") getter = getattr(handler_cls, "event_stream_media_type", None) if getter is None: @@ -275,12 +275,12 @@ class LlmPassthroughRouteHandler(BaseTranslation): return getter() @staticmethod - def _resolve_event_stream_de_anonymizer(provider: Optional[str]): + def _resolve_event_stream_de_anonymizer(provider: str | None): handler_cls = _get_provider_handlers().get(provider or "") return getattr(handler_cls, "de_anonymize_event_stream", None) @staticmethod - def supports_event_stream_de_anonymization(provider: Optional[str], endpoint: Optional[str]) -> bool: + def supports_event_stream_de_anonymization(provider: str | None, endpoint: str | None) -> bool: handler_cls = _get_provider_handlers().get(provider or "") endpoint_check = getattr(handler_cls, "event_stream_endpoint_is_de_anonymizable", None) if endpoint_check is None: diff --git a/litellm/llms/perplexity/chat/transformation.py b/litellm/llms/perplexity/chat/transformation.py index 93afccd5c9d..c6fb750ec1b 100644 --- a/litellm/llms/perplexity/chat/transformation.py +++ b/litellm/llms/perplexity/chat/transformation.py @@ -2,29 +2,27 @@ Translate from OpenAI's `/v1/chat/completions` to Perplexity's `/v1/chat/completions` """ -from typing import Any, List, Optional, Tuple +from typing import Any import httpx + import litellm from litellm._logging import verbose_logger -from litellm.secret_managers.main import get_secret_str -from litellm.types.llms.openai import AllMessageValues -from litellm.types.utils import Usage, PromptTokensDetailsWrapper from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig -from litellm.types.utils import ModelResponse -from litellm.types.llms.openai import ChatCompletionAnnotation -from litellm.types.llms.openai import ChatCompletionAnnotationURLCitation +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues, ChatCompletionAnnotation, ChatCompletionAnnotationURLCitation +from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage class PerplexityChatConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "perplexity" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("PERPLEXITY_API_BASE") or "https://api.perplexity.ai" # type: ignore dynamic_api_key = api_key or get_secret_str("PERPLEXITYAI_API_KEY") or get_secret_str("PERPLEXITY_API_KEY") return api_base, dynamic_api_key @@ -71,12 +69,12 @@ class PerplexityChatConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: # Call the parent transform_response first to handle the standard transformation model_response = super().transform_response( diff --git a/litellm/llms/perplexity/cost_calculator.py b/litellm/llms/perplexity/cost_calculator.py index c9574f3be80..3e6520a1896 100644 --- a/litellm/llms/perplexity/cost_calculator.py +++ b/litellm/llms/perplexity/cost_calculator.py @@ -3,13 +3,11 @@ Helper util for handling perplexity-specific cost calculation - e.g.: citation tokens, search queries """ -from typing import Tuple, Union - from litellm.types.utils import Usage from litellm.utils import get_model_info -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -34,7 +32,7 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: ## GET MODEL INFO model_info = get_model_info(model=model, custom_llm_provider="perplexity") - def _safe_float_cast(value: Union[str, int, float, None, object], default: float = 0.0) -> float: + def _safe_float_cast(value: str | float | None | object, default: float = 0.0) -> float: """Safely cast a value to float with proper type handling for mypy.""" if value is None: return default diff --git a/litellm/llms/perplexity/embedding/transformation.py b/litellm/llms/perplexity/embedding/transformation.py index a52eab34c08..812f5240b94 100644 --- a/litellm/llms/perplexity/embedding/transformation.py +++ b/litellm/llms/perplexity/embedding/transformation.py @@ -13,7 +13,7 @@ This module decodes them into float arrays for OpenAI-compatible responses. import base64 import struct -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -30,7 +30,7 @@ class PerplexityEmbeddingError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] = {}, + headers: dict | httpx.Headers = {}, ): self.status_code = status_code self.message = message @@ -53,12 +53,12 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: if not api_base.endswith("/embeddings"): @@ -90,11 +90,11 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("PERPLEXITYAI_API_KEY") or get_secret_str("PERPLEXITY_API_KEY") @@ -117,7 +117,7 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): } @staticmethod - def _decode_base64_embedding(embedding_value: Any) -> List[float]: + def _decode_base64_embedding(embedding_value: Any) -> list[float]: """ Decode a Perplexity embedding into a list of floats. @@ -140,7 +140,7 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -154,7 +154,7 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): model_response.object = raw_response_json.get("object", "list") raw_data = raw_response_json.get("data", []) - decoded_data: List[Dict[str, Any]] = [] + decoded_data: list[dict[str, Any]] = [] for item in raw_data: decoded_item = dict(item) decoded_item["embedding"] = self._decode_base64_embedding(item.get("embedding")) @@ -173,6 +173,6 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): self, error_message: str, status_code: int, - headers: Union[dict, httpx.Headers], + headers: dict | httpx.Headers, ) -> BaseLLMException: return PerplexityEmbeddingError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/perplexity/responses/transformation.py b/litellm/llms/perplexity/responses/transformation.py index dd5517f6c33..09b4275bd85 100644 --- a/litellm/llms/perplexity/responses/transformation.py +++ b/litellm/llms/perplexity/responses/transformation.py @@ -9,7 +9,7 @@ The only provider quirks: Ref: https://docs.perplexity.ai/api-reference/responses-post """ -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -40,7 +40,7 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.PERPLEXITY - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() api_key = ( litellm_params.api_key or get_secret_str("PERPLEXITYAI_API_KEY") or get_secret_str("PERPLEXITY_API_KEY") @@ -49,16 +49,16 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): headers["Authorization"] = f"Bearer {api_key}" return headers - def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str: + def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str: api_base = api_base or get_secret_str("PERPLEXITY_API_BASE") or "https://api.perplexity.ai" return f"{api_base.rstrip('/')}/v1/responses" - def _ensure_message_type(self, input: Union[str, ResponseInputParam]) -> Union[str, ResponseInputParam]: + def _ensure_message_type(self, input: str | ResponseInputParam) -> str | ResponseInputParam: """Ensure list input items have type='message' (required by Perplexity).""" if isinstance(input, str): return input if isinstance(input, list): - result: List[Any] = [] + result: list[Any] = [] for item in input: if isinstance(item, dict) and "type" not in item: new_item = dict(item) # convert to plain dict to avoid TypedDict checking @@ -72,16 +72,16 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, + input: str | ResponseInputParam, + response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: """Handle preset/ model prefix: send as {"preset": name} instead of {"model": name}.""" input = self._ensure_message_type(input) if model.startswith("preset/"): input = self._validate_input_param(input) - data: Dict = { + data: dict = { "preset": model[len("preset/") :], "input": input, } diff --git a/litellm/llms/perplexity/search/transformation.py b/litellm/llms/perplexity/search/transformation.py index 8ed165de742..1d65bffa822 100644 --- a/litellm/llms/perplexity/search/transformation.py +++ b/litellm/llms/perplexity/search/transformation.py @@ -2,7 +2,7 @@ Calls Perplexity's /search endpoint to search the web. """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -18,7 +18,7 @@ from litellm.secret_managers.main import get_secret_str class _PerplexitySearchRequestRequired(TypedDict): """Required fields for Perplexity Search API request.""" - query: Union[str, List[str]] # Required - search query or queries + query: str | list[str] # Required - search query or queries class PerplexitySearchRequest(_PerplexitySearchRequestRequired, total=False): @@ -28,7 +28,7 @@ class PerplexitySearchRequest(_PerplexitySearchRequestRequired, total=False): """ max_results: int # Optional - maximum number of results (1-20), default 10 - search_domain_filter: List[str] # Optional - list of domains to filter (max 20) + search_domain_filter: list[str] # Optional - list of domains to filter (max 20) max_tokens_per_page: int # Optional - max tokens per page, default 1024 country: str # Optional - country code filter (e.g., 'US', 'GB', 'DE') @@ -42,11 +42,11 @@ class PerplexitySearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -65,9 +65,9 @@ class PerplexitySearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -83,10 +83,10 @@ class PerplexitySearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Perplexity API format. diff --git a/litellm/llms/petals/common_utils.py b/litellm/llms/petals/common_utils.py index bffee338f2b..973a1714320 100644 --- a/litellm/llms/petals/common_utils.py +++ b/litellm/llms/petals/common_utils.py @@ -1,10 +1,8 @@ -from typing import Union - from httpx import Headers from litellm.llms.base_llm.chat.transformation import BaseLLMException class PetalsError(BaseLLMException): - def __init__(self, status_code: int, message: str, headers: Union[dict, Headers]): + def __init__(self, status_code: int, message: str, headers: dict | Headers): super().__init__(status_code=status_code, message=message, headers=headers) diff --git a/litellm/llms/petals/completion/handler.py b/litellm/llms/petals/completion/handler.py index 8d0f02ddd0f..ca932833dcb 100644 --- a/litellm/llms/petals/completion/handler.py +++ b/litellm/llms/petals/completion/handler.py @@ -1,6 +1,5 @@ import time from collections.abc import Callable -from typing import Optional, Union import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -20,7 +19,7 @@ from ..common_utils import PetalsError def completion( model: str, messages: list, - api_base: Optional[str], + api_base: str | None, model_response: ModelResponse, print_verbose: Callable, encoding, @@ -29,7 +28,7 @@ def completion( stream=False, litellm_params=None, logger_fn=None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ): ## Load Config config = litellm.PetalsConfig.get_config() @@ -51,7 +50,7 @@ def completion( else: prompt = prompt_factory(model=model, messages=messages) - output_text: Optional[str] = None + output_text: str | None = None if api_base: ## LOGGING logging_obj.pre_call( diff --git a/litellm/llms/petals/completion/transformation.py b/litellm/llms/petals/completion/transformation.py index ae6415680b1..85fd1bd267b 100644 --- a/litellm/llms/petals/completion/transformation.py +++ b/litellm/llms/petals/completion/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, List, Optional, Union +from typing import Any from httpx import Headers, Response @@ -36,23 +36,23 @@ class PetalsConfig(BaseConfig): - `repetition_penalty` (float, optional): This helps apply the repetition penalty during text generation, as discussed in this paper. """ - max_length: Optional[int] = None - max_new_tokens: Optional[int] = litellm.max_tokens # petals requires max tokens to be set - do_sample: Optional[bool] = None - temperature: Optional[float] = None - top_k: Optional[int] = None - top_p: Optional[float] = None - repetition_penalty: Optional[float] = None + max_length: int | None = None + max_new_tokens: int | None = litellm.max_tokens # petals requires max tokens to be set + do_sample: bool | None = None + temperature: float | None = None + top_k: int | None = None + top_p: float | None = None + repetition_penalty: float | None = None def __init__( self, - max_length: Optional[int] = None, - max_new_tokens: Optional[int] = litellm.max_tokens, # petals requires max tokens to be set - do_sample: Optional[bool] = None, - temperature: Optional[float] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - repetition_penalty: Optional[float] = None, + max_length: int | None = None, + max_new_tokens: int | None = litellm.max_tokens, # petals requires max tokens to be set + do_sample: bool | None = None, + temperature: float | None = None, + top_k: int | None = None, + top_p: float | None = None, + repetition_penalty: float | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -63,10 +63,10 @@ class PetalsConfig(BaseConfig): def get_config(cls): return super().get_config() - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return PetalsError(status_code=status_code, message=error_message, headers=headers) - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return ["max_tokens", "temperature", "top_p", "stream"] def map_openai_params( @@ -90,7 +90,7 @@ class PetalsConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -106,12 +106,12 @@ class PetalsConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: raise NotImplementedError( "Petals transformation currently done in handler.py. [TODO] Move to the transformation.py" @@ -121,10 +121,10 @@ class PetalsConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {} diff --git a/litellm/llms/pg_vector/vector_stores/transformation.py b/litellm/llms/pg_vector/vector_stores/transformation.py index b58b6e7f498..e30591d99c5 100644 --- a/litellm/llms/pg_vector/vector_stores/transformation.py +++ b/litellm/llms/pg_vector/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig @@ -27,7 +27,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): - api_key: API key for authentication with the PG vector service """ - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate environment and set headers for PG vector service authentication """ @@ -52,7 +52,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -74,13 +74,13 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url = f"{api_base}/{encoded_vector_store_id}/search" _, request_body = super().transform_search_vector_store_request( diff --git a/litellm/llms/predibase/chat/handler.py b/litellm/llms/predibase/chat/handler.py index d1eb4e590d3..36537562638 100644 --- a/litellm/llms/predibase/chat/handler.py +++ b/litellm/llms/predibase/chat/handler.py @@ -4,7 +4,6 @@ import json from collections.abc import Callable from functools import partial -from typing import Optional, Union import httpx # type: ignore @@ -26,7 +25,7 @@ async def make_call( model: str, messages: list, logging_obj, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, ): response = await client.post(api_base, headers=headers, data=data, stream=True, timeout=timeout) @@ -63,11 +62,11 @@ class PredibaseChatCompletion: optional_params: dict, litellm_params: dict, tenant_id: str, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, acompletion=None, logger_fn=None, headers: dict = {}, - ) -> Union[ModelResponse, CustomStreamWrapper]: + ) -> ModelResponse | CustomStreamWrapper: predibase_config = litellm.PredibaseConfig() headers = predibase_config.validate_environment( api_key=api_key, @@ -202,7 +201,7 @@ class PredibaseChatCompletion: stream, data: dict, optional_params: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, litellm_params=None, logger_fn=None, headers={}, @@ -219,16 +218,14 @@ class PredibaseChatCompletion: except httpx.HTTPStatusError as e: raise PredibaseError( status_code=e.response.status_code, - message="HTTPStatusError - received status_code={}, error_message={}".format( - e.response.status_code, e.response.text - ), + message=f"HTTPStatusError - received status_code={e.response.status_code}, error_message={e.response.text}", ) except Exception as e: for exception in litellm.LITELLM_EXCEPTION_TYPES: if isinstance(e, exception): raise e raise PredibaseError( - status_code=500, message="{}".format(str(e)) + status_code=500, message=f"{e!s}" ) # don't use verbose_logger.exception, if exception is raised return predibase_config.transform_response( model=model, @@ -254,7 +251,7 @@ class PredibaseChatCompletion: api_key, logging_obj, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: float | httpx.Timeout, optional_params=None, litellm_params=None, logger_fn=None, diff --git a/litellm/llms/predibase/chat/transformation.py b/litellm/llms/predibase/chat/transformation.py index fcb21272be2..942e36ce9fc 100644 --- a/litellm/llms/predibase/chat/transformation.py +++ b/litellm/llms/predibase/chat/transformation.py @@ -1,6 +1,6 @@ import os import time -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Literal from httpx import Headers, Response @@ -30,39 +30,39 @@ class PredibaseConfig(BaseConfig): Reference: https://docs.predibase.com/user-guide/inference/rest_api """ - adapter_id: Optional[str] = None - adapter_source: Optional[Literal["pbase", "hub", "s3"]] = None - best_of: Optional[int] = None - decoder_input_details: Optional[bool] = None + adapter_id: str | None = None + adapter_source: Literal["pbase", "hub", "s3"] | None = None + best_of: int | None = None + decoder_input_details: bool | None = None details: bool = True # enables returning logprobs + best of max_new_tokens: int = DEFAULT_MAX_TOKENS # openai default - requests hang if max_new_tokens not given - repetition_penalty: Optional[float] = None - return_full_text: Optional[bool] = False # by default don't return the input as part of the output - seed: Optional[int] = None - stop: Optional[List[str]] = None - temperature: Optional[float] = None - top_k: Optional[int] = None - top_p: Optional[int] = None - truncate: Optional[int] = None - typical_p: Optional[float] = None - watermark: Optional[bool] = None + repetition_penalty: float | None = None + return_full_text: bool | None = False # by default don't return the input as part of the output + seed: int | None = None + stop: list[str] | None = None + temperature: float | None = None + top_k: int | None = None + top_p: int | None = None + truncate: int | None = None + typical_p: float | None = None + watermark: bool | None = None def __init__( self, - best_of: Optional[int] = None, - decoder_input_details: Optional[bool] = None, - details: Optional[bool] = None, - max_new_tokens: Optional[int] = None, - repetition_penalty: Optional[float] = None, - return_full_text: Optional[bool] = None, - seed: Optional[int] = None, - stop: Optional[List[str]] = None, - temperature: Optional[float] = None, - top_k: Optional[int] = None, - top_p: Optional[int] = None, - truncate: Optional[int] = None, - typical_p: Optional[float] = None, - watermark: Optional[bool] = None, + best_of: int | None = None, + decoder_input_details: bool | None = None, + details: bool | None = None, + max_new_tokens: int | None = None, + repetition_penalty: float | None = None, + return_full_text: bool | None = None, + seed: int | None = None, + stop: list[str] | None = None, + temperature: float | None = None, + top_k: int | None = None, + top_p: int | None = None, + truncate: int | None = None, + typical_p: float | None = None, + watermark: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -130,12 +130,12 @@ class PredibaseConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: logging_obj.post_call( input=messages, @@ -253,7 +253,7 @@ class PredibaseConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -305,12 +305,12 @@ class PredibaseConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: tenant_id = litellm_params.get("predibase_tenant_id") or litellm_params.get("tenant_id") if tenant_id is None: @@ -332,18 +332,18 @@ class PredibaseConfig(BaseConfig): completion_url += "/generate" return completion_url - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return PredibaseError(status_code=status_code, message=error_message, headers=headers) def validate_environment( self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: raise ValueError( @@ -352,7 +352,7 @@ class PredibaseConfig(BaseConfig): default_headers = { "content-type": "application/json", - "Authorization": "Bearer {}".format(api_key), + "Authorization": f"Bearer {api_key}", } if headers is not None and isinstance(headers, dict): headers = {**default_headers, **headers} diff --git a/litellm/llms/predibase/common_utils.py b/litellm/llms/predibase/common_utils.py index 2dad5861208..36b47a3a4b2 100644 --- a/litellm/llms/predibase/common_utils.py +++ b/litellm/llms/predibase/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,9 +8,9 @@ class PredibaseError(BaseLLMException): self, status_code: int, message: str, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, ): super().__init__( status_code=status_code, diff --git a/litellm/llms/ragflow/chat/transformation.py b/litellm/llms/ragflow/chat/transformation.py index be3417d1aad..9b997c9772a 100644 --- a/litellm/llms/ragflow/chat/transformation.py +++ b/litellm/llms/ragflow/chat/transformation.py @@ -10,8 +10,6 @@ Model name format: - Agent: ragflow/agent/{agent_id}/{model_name} """ -from typing import List, Optional, Tuple - import litellm from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.openai import OpenAIConfig @@ -28,7 +26,7 @@ class RAGFlowConfig(OpenAIConfig): - ragflow/agent/{agent_id}/{model_name} for agent endpoints """ - def _parse_ragflow_model(self, model: str) -> Tuple[str, str, str]: + def _parse_ragflow_model(self, model: str) -> tuple[str, str, str]: """ Parse RAGFlow model name format: ragflow/{endpoint_type}/{id}/{model_name} @@ -62,12 +60,12 @@ class RAGFlowConfig(OpenAIConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the RAGFlow API call. @@ -127,10 +125,10 @@ class RAGFlowConfig(OpenAIConfig): def _get_openai_compatible_provider_info( self, model: str, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, custom_llm_provider: str, - ) -> Tuple[Optional[str], Optional[str], str]: + ) -> tuple[str | None, str | None, str]: """ Get OpenAI-compatible provider information for RAGFlow. @@ -161,11 +159,11 @@ class RAGFlowConfig(OpenAIConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for RAGFlow API. @@ -211,7 +209,7 @@ class RAGFlowConfig(OpenAIConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/ragflow/vector_stores/transformation.py b/litellm/llms/ragflow/vector_stores/transformation.py index d8bdd981425..bcb54689c93 100644 --- a/litellm/llms/ragflow/vector_stores/transformation.py +++ b/litellm/llms/ragflow/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -44,7 +44,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): "write": [], } - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """Validate environment and set headers for RAGFlow API.""" litellm_params = litellm_params or GenericLiteLLMParams() api_key = litellm_params.api_key or get_secret_str("RAGFLOW_API_KEY") @@ -62,7 +62,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -86,13 +86,13 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """RAGFlow vector stores are management-only, search is not supported.""" raise NotImplementedError("RAGFlow vector stores support dataset management only, not search/retrieval") @@ -106,7 +106,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform create request to RAGFlow POST /api/v1/datasets format. @@ -121,7 +121,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): raise ValueError("name is required for RAGFlow dataset creation") # Build request body - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "name": name, } diff --git a/litellm/llms/recraft/image_edit/transformation.py b/litellm/llms/recraft/image_edit/transformation.py index 61c669b50c0..554a1d918d5 100644 --- a/litellm/llms/recraft/image_edit/transformation.py +++ b/litellm/llms/recraft/image_edit/transformation.py @@ -1,5 +1,5 @@ from io import BufferedReader -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -26,7 +26,7 @@ class RecraftImageEditConfig(BaseImageEditConfig): IMAGE_EDIT_ENDPOINT: str = "v1/images/imageToImage" DEFAULT_STRENGTH: float = 0.2 - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: """ Supported OpenAI parameters that can be mapped to Recraft image edit API. @@ -43,7 +43,7 @@ class RecraftImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI image edit parameters to Recraft parameters. Reuses OpenAI logic but filters to supported params only. @@ -60,7 +60,7 @@ class RecraftImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -78,11 +78,11 @@ class RecraftImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("RECRAFT_API_KEY") + final_api_key: str | None = api_key or get_secret_str("RECRAFT_API_KEY") if not final_api_key: raise ValueError("RECRAFT_API_KEY is not set") @@ -92,12 +92,12 @@ class RecraftImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform the image edit request to Recraft's multipart form format. Reuses OpenAI file handling logic but adapts for Recraft API structure. @@ -114,7 +114,7 @@ class RecraftImageEditConfig(BaseImageEditConfig): request_params["prompt"] = prompt request_body = RecraftImageEditRequestParams(**request_params) - request_dict = cast(Dict, request_body) + request_dict = cast(dict, request_body) ######################################################### # Reuse OpenAI logic: Separate images as `files` and send other parameters as `data` ######################################################### @@ -125,9 +125,9 @@ class RecraftImageEditConfig(BaseImageEditConfig): def _get_image_files_for_request( self, - image: Optional[FileTypes], - ) -> List[Tuple[str, Any]]: - files_list: List[Tuple[str, Any]] = [] + image: FileTypes | None, + ) -> list[tuple[str, Any]]: + files_list: list[tuple[str, Any]] = [] # Handle single image (Recraft expects single image, not array) if image: diff --git a/litellm/llms/recraft/image_generation/transformation.py b/litellm/llms/recraft/image_generation/transformation.py index 9f48273c306..b3542e80f59 100644 --- a/litellm/llms/recraft/image_generation/transformation.py +++ b/litellm/llms/recraft/image_generation/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -25,7 +25,7 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://external.api.recraft.ai" IMAGE_GENERATION_ENDPOINT: str = "v1/images/generations" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ https://www.recraft.ai/docs#generate-image """ @@ -39,8 +39,8 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: @@ -54,12 +54,12 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -76,13 +76,13 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = api_key or get_secret_str("RECRAFT_API_KEY") + final_api_key: str | None = api_key or get_secret_str("RECRAFT_API_KEY") if not final_api_key: raise ValueError("RECRAFT_API_KEY is not set") @@ -121,8 +121,8 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the image generation response to the litellm image response diff --git a/litellm/llms/reducto/common.py b/litellm/llms/reducto/common.py index 364f269feb1..4eebce59973 100644 --- a/litellm/llms/reducto/common.py +++ b/litellm/llms/reducto/common.py @@ -1,7 +1,7 @@ import base64 import binascii from collections import defaultdict -from typing import TYPE_CHECKING, Any, Dict, List, NoReturn, Optional, Tuple +from typing import TYPE_CHECKING, Any, NoReturn from litellm.constants import request_timeout @@ -12,7 +12,7 @@ if TYPE_CHECKING: from litellm.llms.base_llm.ocr.transformation import OCRPage -def _normalize_api_base(api_base: Optional[str]) -> str: +def _normalize_api_base(api_base: str | None) -> str: return (api_base or REDUCTO_API_BASE).rstrip("/") @@ -29,7 +29,7 @@ def _raise_bad_request(message: str, model: str) -> NoReturn: def extract_file_id_or_bytes( source_url: str, model: str, -) -> Tuple[Optional[str], Optional[bytes], Optional[str]]: +) -> tuple[str | None, bytes | None, str | None]: if source_url.startswith(REDUCTO_ID_PREFIX): return source_url, None, None @@ -66,18 +66,18 @@ def _extract_file_id_from_upload_response(response: Any) -> str: try: payload = response.json() except ValueError as exc: - raise ValueError("Reducto /upload returned a non-JSON 200 response: {}".format(response.text)) from exc + raise ValueError(f"Reducto /upload returned a non-JSON 200 response: {response.text}") from exc file_id = (payload or {}).get("file_id") if isinstance(payload, dict) else None if not isinstance(file_id, str) or not file_id: - raise ValueError("Reducto /upload returned 200 without a file_id; got payload={}".format(payload)) + raise ValueError(f"Reducto /upload returned 200 without a file_id; got payload={payload}") return file_id def upload_bytes_sync( raw_bytes: bytes, - mime: Optional[str], + mime: str | None, api_key: str, - api_base: Optional[str], + api_base: str | None, ) -> str: import litellm @@ -93,9 +93,9 @@ def upload_bytes_sync( async def upload_bytes_async( raw_bytes: bytes, - mime: Optional[str], + mime: str | None, api_key: str, - api_base: Optional[str], + api_base: str | None, ) -> str: import litellm @@ -109,11 +109,11 @@ async def upload_bytes_async( return _extract_file_id_from_upload_response(response) -def build_pages_from_reducto(result: Dict[str, Any]) -> List["OCRPage"]: +def build_pages_from_reducto(result: dict[str, Any]) -> list["OCRPage"]: from litellm.llms.base_llm.ocr.transformation import OCRPage chunks = result.get("chunks", []) or [] - blocks_by_page: Dict[int, List[Dict[str, Any]]] = defaultdict(list) + blocks_by_page: dict[int, list[dict[str, Any]]] = defaultdict(list) for chunk in chunks: for block in chunk.get("blocks", []) or []: @@ -132,7 +132,7 @@ def build_pages_from_reducto(result: Dict[str, Any]) -> List["OCRPage"]: return [] return [OCRPage(index=0, markdown=fallback_markdown)] - pages: List["OCRPage"] = [] + pages: list[OCRPage] = [] for page_no, blocks in sorted(blocks_by_page.items()): markdown = "\n\n".join(block.get("content", "") for block in blocks if block.get("content")) page_index = max(page_no - 1, 0) diff --git a/litellm/llms/reducto/ocr/transformation.py b/litellm/llms/reducto/ocr/transformation.py index e8bfcceea2a..635a4165cf3 100644 --- a/litellm/llms/reducto/ocr/transformation.py +++ b/litellm/llms/reducto/ocr/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional, Tuple +from typing import Any import httpx @@ -34,13 +34,13 @@ class _BaseReductoOCRConfig(BaseOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - litellm_params: Optional[dict] = None, + api_key: str | None = None, + api_base: str | None = None, + litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: from litellm.secret_managers.main import get_secret_str resolved_key = api_key or get_secret_str("REDUCTO_API_KEY") @@ -57,10 +57,10 @@ class _BaseReductoOCRConfig(BaseOCRConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, model: str, optional_params: dict, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, **kwargs, ) -> str: return "{}/parse".format((api_base or REDUCTO_API_BASE).rstrip("/")) @@ -69,12 +69,12 @@ class _BaseReductoOCRConfig(BaseOCRConfig): source_url = document.get("document_url") or document.get("image_url") if source_url is None: raise ValueError( - "Reducto expected OCR preprocessing to produce document_url or image_url for model={}".format(model) + f"Reducto expected OCR preprocessing to produce document_url or image_url for model={model}" ) return source_url @staticmethod - def _resolve_credentials(api_key: Optional[str], api_base: Optional[str]) -> Tuple[str, str]: + def _resolve_credentials(api_key: str | None, api_base: str | None) -> tuple[str, str]: from litellm.secret_managers.main import get_secret_str resolved_key = api_key or get_secret_str("REDUCTO_API_KEY") @@ -89,8 +89,8 @@ class _BaseReductoOCRConfig(BaseOCRConfig): self, model: str, document: DocumentType, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, ) -> str: source_url = self._get_source_url(document=document, model=model) file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model) @@ -108,8 +108,8 @@ class _BaseReductoOCRConfig(BaseOCRConfig): self, model: str, document: DocumentType, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, ) -> str: source_url = self._get_source_url(document=document, model=model) file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model) @@ -187,8 +187,8 @@ class ReductoParseLegacyConfig(_BaseReductoOCRConfig): def get_supported_ocr_params(self, model: str) -> list: return ["enhance"] - def _build_legacy_body(self, file_id: str, optional_params: dict) -> Dict[str, Any]: - body: Dict[str, Any] = {"document_url": file_id} + def _build_legacy_body(self, file_id: str, optional_params: dict) -> dict[str, Any]: + body: dict[str, Any] = {"document_url": file_id} enhance = optional_params.get("enhance") if enhance is not None: body["options"] = {"enhance": enhance} diff --git a/litellm/llms/replicate/chat/handler.py b/litellm/llms/replicate/chat/handler.py index a9123d47007..9f4d3238421 100644 --- a/litellm/llms/replicate/chat/handler.py +++ b/litellm/llms/replicate/chat/handler.py @@ -2,7 +2,6 @@ import asyncio import json import time from collections.abc import Callable -from typing import List, Union import litellm from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS @@ -135,7 +134,7 @@ def completion( logger_fn=None, acompletion=None, headers={}, -) -> Union[ModelResponse, CustomStreamWrapper]: +) -> ModelResponse | CustomStreamWrapper: headers = replicate_config.validate_environment( api_key=api_key, headers=headers, @@ -238,7 +237,7 @@ def completion( async def async_completion( model_response: ModelResponse, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], encoding, optional_params: dict, litellm_params: dict, @@ -249,7 +248,7 @@ async def async_completion( logging_obj, print_verbose, headers: dict, -) -> Union[ModelResponse, CustomStreamWrapper]: +) -> ModelResponse | CustomStreamWrapper: prediction_url = replicate_config.get_complete_url( api_base=api_base, api_key=api_key, diff --git a/litellm/llms/replicate/chat/transformation.py b/litellm/llms/replicate/chat/transformation.py index 6da26b966f3..605a8ca6d80 100644 --- a/litellm/llms/replicate/chat/transformation.py +++ b/litellm/llms/replicate/chat/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -52,27 +52,27 @@ class ReplicateConfig(BaseConfig): Please note that Replicate's mapping of these parameters can be inconsistent across different models, indicating that not all of these parameters may be available for use with all models. """ - system_prompt: Optional[str] = None - max_new_tokens: Optional[int] = None - min_new_tokens: Optional[int] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - stop_sequences: Optional[str] = None - seed: Optional[int] = None - debug: Optional[bool] = None + system_prompt: str | None = None + max_new_tokens: int | None = None + min_new_tokens: int | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + stop_sequences: str | None = None + seed: int | None = None + debug: bool | None = None def __init__( self, - system_prompt: Optional[str] = None, - max_new_tokens: Optional[int] = None, - min_new_tokens: Optional[int] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - top_k: Optional[int] = None, - stop_sequences: Optional[str] = None, - seed: Optional[int] = None, - debug: Optional[bool] = None, + system_prompt: str | None = None, + max_new_tokens: int | None = None, + min_new_tokens: int | None = None, + temperature: int | None = None, + top_p: int | None = None, + top_k: int | None = None, + stop_sequences: str | None = None, + seed: int | None = None, + debug: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -130,19 +130,17 @@ class ReplicateConfig(BaseConfig): return split_model[1] return model - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return ReplicateError(status_code=status_code, message=error_message, headers=headers) def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: version_id = self.model_to_version_id(model) base_url = api_base @@ -158,7 +156,7 @@ class ReplicateConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -201,7 +199,7 @@ class ReplicateConfig(BaseConfig): if prompt is None or not isinstance(prompt, str): raise ReplicateError( status_code=400, - message="LiteLLM Error - prompt is not a string - {}".format(prompt), + message=f"LiteLLM Error - prompt is not a string - {prompt}", headers={}, ) @@ -234,12 +232,12 @@ class ReplicateConfig(BaseConfig): model_response: ModelResponse, logging_obj: LoggingClass, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: logging_obj.post_call( input=messages, @@ -251,7 +249,7 @@ class ReplicateConfig(BaseConfig): if raw_response_json.get("status") != "succeeded": raise ReplicateError( status_code=422, - message="LiteLLM Error - prediction not succeeded - {}".format(raw_response_json), + message=f"LiteLLM Error - prediction not succeeded - {raw_response_json}", headers=raw_response.headers, ) outputs = raw_response_json.get("output", []) @@ -292,7 +290,7 @@ class ReplicateConfig(BaseConfig): if prediction_url is None: raise ReplicateError( status_code=400, - message="LiteLLM Error - prediction url is None - {}".format(response_json), + message=f"LiteLLM Error - prediction url is None - {response_json}", headers=response.headers, ) return prediction_url @@ -301,11 +299,11 @@ class ReplicateConfig(BaseConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = { "Authorization": f"Token {api_key}", diff --git a/litellm/llms/replicate/common_utils.py b/litellm/llms/replicate/common_utils.py index c52b47a46aa..e1fa3828379 100644 --- a/litellm/llms/replicate/common_utils.py +++ b/litellm/llms/replicate/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +8,6 @@ class ReplicateError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]], + headers: dict | httpx.Headers | None, ): super().__init__(status_code=status_code, message=message, headers=headers) diff --git a/litellm/llms/runwayml/image_generation/transformation.py b/litellm/llms/runwayml/image_generation/transformation.py index fddd0b1350b..ba51a7ba093 100644 --- a/litellm/llms/runwayml/image_generation/transformation.py +++ b/litellm/llms/runwayml/image_generation/transformation.py @@ -1,6 +1,6 @@ import asyncio import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -37,12 +37,12 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete url for the request @@ -60,13 +60,13 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: - final_api_key: Optional[str] = ( + final_api_key: str | None = ( api_key or get_secret_str("RUNWAYML_API_SECRET") or get_secret_str("RUNWAYML_API_KEY") ) if not final_api_key: @@ -78,7 +78,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): @staticmethod def _transform_runwayml_response_to_openai( - response_data: Dict[str, Any], + response_data: dict[str, Any], model_response: ImageResponse, ) -> ImageResponse: """ @@ -153,7 +153,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): raise TimeoutError(f"RunwayML task polling timed out after {timeout_secs} seconds") @staticmethod - def _check_task_status(response_data: Dict[str, Any]) -> str: + def _check_task_status(response_data: dict[str, Any]) -> str: """ Check RunwayML task status from response. @@ -189,7 +189,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): self, task_id: str, api_base: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: float = 600, ) -> httpx.Response: """ @@ -240,7 +240,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): self, task_id: str, api_base: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: float = 600, ) -> httpx.Response: """ @@ -295,8 +295,8 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform the image generation response to the litellm image response. @@ -370,8 +370,8 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Async transform the image generation response to the litellm image response. @@ -420,7 +420,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): model_response=model_response, ) - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Get supported OpenAI parameters for RunwayML image generation """ @@ -450,8 +450,8 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): } optional_params["ratio"] = size_to_ratio_map.get(size, "1920:1080") - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: diff --git a/litellm/llms/runwayml/text_to_speech/transformation.py b/litellm/llms/runwayml/text_to_speech/transformation.py index d27e28c23e0..46a5f606853 100644 --- a/litellm/llms/runwayml/text_to_speech/transformation.py +++ b/litellm/llms/runwayml/text_to_speech/transformation.py @@ -7,7 +7,7 @@ Maps OpenAI TTS spec to RunwayML Text-to-Speech API import asyncio import time from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -59,16 +59,16 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): self, model: str, input: str, - voice: Optional[Union[str, Dict]], - optional_params: Dict, - litellm_params_dict: Dict, + voice: str | dict | None, + optional_params: dict, + litellm_params_dict: dict, logging_obj: "LiteLLMLoggingObj", - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, Any]], + timeout: float | httpx.Timeout, + extra_headers: dict[str, Any] | None, base_llm_http_handler: Any, aspeech: bool, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, **kwargs: Any, ) -> Union[ "HttpxBinaryResponseContent", @@ -101,7 +101,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): ) # Convert voice to appropriate format - voice_param: Optional[Union[str, Dict]] = voice + voice_param: str | dict | None = voice if isinstance(voice, str): # Keep as string, will be processed in map_openai_params voice_param = voice @@ -143,11 +143,11 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Dict = {}, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict = {}, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to RunwayML TTS parameters @@ -160,7 +160,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): mapped_params = {} # Map voice parameter to RunwayML format dict - voice_dict: Optional[Dict] = None + voice_dict: dict | None = None if isinstance(voice, str): # Check if it's an OpenAI voice name that needs mapping if voice in self.VOICE_MAPPINGS: @@ -193,8 +193,8 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate RunwayML environment and set up authentication headers @@ -215,7 +215,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -242,7 +242,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): raise TimeoutError(f"RunwayML TTS task polling timed out after {timeout_secs} seconds") @staticmethod - def _check_task_status(response_data: Dict[str, Any]) -> str: + def _check_task_status(response_data: dict[str, Any]) -> str: """ Check RunwayML task status from response. @@ -278,7 +278,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): self, task_id: str, api_base: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: float = 600, ) -> httpx.Response: """ @@ -329,7 +329,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): self, task_id: str, api_base: str, - headers: Dict[str, str], + headers: dict[str, str], timeout_secs: float = 600, ) -> httpx.Response: """ @@ -377,9 +377,9 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): self, model: str, input: str, - voice: Optional[Union[str, Dict]], - optional_params: Dict, - litellm_params: Dict, + voice: str | dict | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ diff --git a/litellm/llms/runwayml/videos/transformation.py b/litellm/llms/runwayml/videos/transformation.py index b11671c9431..8a89f53f13a 100644 --- a/litellm/llms/runwayml/videos/transformation.py +++ b/litellm/llms/runwayml/videos/transformation.py @@ -1,5 +1,5 @@ from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -68,7 +68,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI parameters to RunwayML format. @@ -78,7 +78,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): - size -> ratio (convert "WIDTHxHEIGHT" to "WIDTH:HEIGHT") - seconds -> duration (convert to integer) """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} # Handle input_reference parameter - map to promptImage if "input_reference" in video_create_optional_params: @@ -114,8 +114,8 @@ class RunwayMLVideoConfig(BaseVideoConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, ) -> dict: """ Validate environment and set up authentication headers. @@ -146,7 +146,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -163,10 +163,10 @@ class RunwayMLVideoConfig(BaseVideoConfig): model: str, prompt: str, api_base: str, - video_create_optional_request_params: Dict, + video_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles, str]: + ) -> tuple[dict, RequestFiles, str]: """ Transform the video creation request for RunwayML API. @@ -180,7 +180,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): } """ # Build the request data - request_data: Dict[str, Any] = { + request_data: dict[str, Any] = { "model": model, "promptText": prompt, } @@ -189,7 +189,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): request_data.update(video_create_optional_request_params) # RunwayML uses JSON body, no files multipart - files_list: List[Tuple[str, Any]] = [] + files_list: list[tuple[str, Any]] = [] # Append the specific endpoint for video generation full_api_base = f"{api_base}/image_to_video" @@ -201,8 +201,8 @@ class RunwayMLVideoConfig(BaseVideoConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: """ Transform the RunwayML video creation response. @@ -219,7 +219,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): response_data = raw_response.json() # Map RunwayML task response to VideoObject format - video_data: Dict[str, Any] = { + video_data: dict[str, Any] = { "id": response_data.get("id", ""), "object": "video", "status": self._map_runway_status(response_data.get("status", "pending")), @@ -287,7 +287,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): } return status_map.get(runway_status.upper(), "queued") - def _parse_runway_timestamp(self, timestamp_str: Optional[str]) -> int: + def _parse_runway_timestamp(self, timestamp_str: str | None) -> int: """ Convert RunwayML ISO 8601 timestamp to Unix timestamp. @@ -311,8 +311,8 @@ class RunwayMLVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - variant: Optional[str] = None, - ) -> Tuple[str, Dict]: + variant: str | None = None, + ) -> tuple[str, dict]: """ Transform the video content request for RunwayML API. @@ -326,11 +326,11 @@ class RunwayMLVideoConfig(BaseVideoConfig): # Get task status to retrieve video URL url = f"{api_base}/tasks/{encoded_video_id}" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} return url, params - def _extract_video_url_from_response(self, response_data: Dict[str, Any]) -> str: + def _extract_video_url_from_response(self, response_data: dict[str, Any]) -> str: """ Helper method to extract video URL from RunwayML response. Shared between sync and async transforms. @@ -421,8 +421,8 @@ class RunwayMLVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video remix request for RunwayML API. @@ -435,7 +435,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """Transform the RunwayML video remix response.""" raise NotImplementedError("Video remix is not yet supported by RunwayML API") @@ -445,11 +445,11 @@ class RunwayMLVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Transform the video list request for RunwayML API. @@ -461,8 +461,8 @@ class RunwayMLVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - ) -> Dict[str, str]: + custom_llm_provider: str | None = None, + ) -> dict[str, str]: """Transform the RunwayML video list response.""" raise NotImplementedError("Video listing is not yet supported by RunwayML API") @@ -472,7 +472,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video delete request for RunwayML API. @@ -484,7 +484,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): # Construct the URL for task cancellation url = f"{api_base}/tasks/{encoded_video_id}/cancel" - data: Dict[str, Any] = {} + data: dict[str, Any] = {} return url, data @@ -511,7 +511,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the RunwayML video status retrieve request. @@ -524,7 +524,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): url = f"{api_base}/tasks/{encoded_video_id}" # Empty dict for GET request (no body) - data: Dict[str, Any] = {} + data: dict[str, Any] = {} return url, data @@ -532,7 +532,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """ Transform the RunwayML video status retrieve response. @@ -540,7 +540,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): response_data = raw_response.json() # Map RunwayML task response to VideoObject format - video_data: Dict[str, Any] = { + video_data: dict[str, Any] = { "id": response_data.get("id", ""), "object": "video", "status": self._map_runway_status(response_data.get("status", "pending")), @@ -620,9 +620,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None): raise NotImplementedError("video extension is not supported for RunwayML") - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from ...base_llm.chat.transformation import BaseLLMException raise BaseLLMException( diff --git a/litellm/llms/s3_vectors/vector_stores/transformation.py b/litellm/llms/s3_vectors/vector_stores/transformation.py index b31e6f4511a..fe57e8d3964 100644 --- a/litellm/llms/s3_vectors/vector_stores/transformation.py +++ b/litellm/llms/s3_vectors/vector_stores/transformation.py @@ -1,5 +1,5 @@ import re -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -38,7 +38,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): "write": [], } - def get_supported_openai_params(self, model: str) -> List[VECTOR_STORE_OPENAI_PARAMS]: + def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]: return ["max_num_results"] def map_openai_params( @@ -52,12 +52,12 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): optional_params["maxResults"] = value return optional_params - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: headers = headers or {} headers.setdefault("Content-Type", "application/json") return headers - def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str: + def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str: aws_region_name = litellm_params.get("aws_region_name") if not aws_region_name: raise ValueError("aws_region_name is required for S3 Vectors") @@ -68,13 +68,13 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """Sync version - generates embedding synchronously.""" # For S3 Vectors, vector_store_id should be in format: bucket_name:index_name # If not in that format, try to construct it from litellm_params @@ -107,7 +107,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): url = f"{api_base}/QueryVectors" - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "vectorBucketName": bucket_name, "indexName": index_name, "queryVector": {"float32": query_embedding}, @@ -122,13 +122,13 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): async def atransform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """Async version - generates embedding asynchronously.""" # For S3 Vectors, vector_store_id should be in format: bucket_name:index_name # If not in that format, try to construct it from litellm_params @@ -161,7 +161,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): url = f"{api_base}/QueryVectors" - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "vectorBucketName": bucket_name, "indexName": index_name, "queryVector": {"float32": query_embedding}, @@ -176,11 +176,11 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): def sign_request( self, headers: dict, - optional_params: Dict, - request_data: Dict, + optional_params: dict, + request_data: dict, api_base: str, - api_key: Optional[str] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + ) -> tuple[dict, bytes | None]: return self._sign_request( service_name="s3vectors", headers=headers, @@ -195,7 +195,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): ) -> VectorStoreSearchResponse: try: response_data = response.json() - results: List[VectorStoreSearchResult] = [] + results: list[VectorStoreSearchResult] = [] for item in response_data.get("vectors", []) or []: metadata = item.get("metadata", {}) or {} @@ -248,7 +248,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): self, vector_store_create_optional_params, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: raise NotImplementedError def transform_create_vector_store_response(self, response: httpx.Response): diff --git a/litellm/llms/sagemaker/chat/handler.py b/litellm/llms/sagemaker/chat/handler.py index 4278ef34b31..daa930794dc 100644 --- a/litellm/llms/sagemaker/chat/handler.py +++ b/litellm/llms/sagemaker/chat/handler.py @@ -1,7 +1,6 @@ import json from collections.abc import Callable from copy import deepcopy -from typing import Optional, Union import httpx @@ -70,7 +69,7 @@ class SagemakerChatHandler(BaseAWSLLM): data: dict, optional_params: dict, aws_region_name: str, - extra_headers: Optional[dict] = None, + extra_headers: dict | None = None, ): try: from botocore.auth import SigV4Auth @@ -113,12 +112,12 @@ class SagemakerChatHandler(BaseAWSLLM): logging_obj, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, custom_prompt_dict={}, logger_fn=None, acompletion: bool = False, headers: dict = {}, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, ): # pop streaming if it's in the optional params as 'stream' raises an error with sagemaker credentials, aws_region_name = self._load_credentials(optional_params) diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 4447d63e5a1..183f55b295b 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -7,7 +7,7 @@ LiteLLM Docs: https://docs.litellm.ai/docs/providers/aws_sagemaker#sagemaker-mes Huggingface Docs: https://huggingface.co/docs/text-generation-inference/en/messages_api """ -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._models import Headers @@ -41,29 +41,29 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): OpenAIGPTConfig.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return SagemakerError(status_code=status_code, message=error_message, headers=headers) def validate_environment( self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return headers def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: aws_region_name = self._get_aws_region_name( optional_params=optional_params, @@ -75,7 +75,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): else: api_base = f"https://runtime.sagemaker.{aws_region_name}.amazonaws.com/endpoints/{model}/invocations" - sagemaker_base_url = cast(Optional[str], optional_params.get("sagemaker_base_url")) + sagemaker_base_url = cast(str | None, optional_params.get("sagemaker_base_url")) if sagemaker_base_url is not None: api_base = sagemaker_base_url @@ -87,11 +87,11 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> tuple[dict, bytes | None]: return self._sign_request( service_name="sagemaker", headers=headers, @@ -121,9 +121,9 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -163,9 +163,9 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, HTTPHandler): try: diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index da9572267ac..5f5c0250273 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -1,7 +1,6 @@ import functools import json from collections.abc import AsyncIterator, Iterator -from typing import List, Optional, Union import httpx @@ -45,18 +44,18 @@ class SagemakerError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) class AWSEventStreamDecoder: - def __init__(self, model: str, is_messages_api: Optional[bool] = None) -> None: + def __init__(self, model: str, is_messages_api: bool | None = None) -> None: from botocore.parsers import EventStreamJSONParser self.model = model self.parser = EventStreamJSONParser() - self.content_blocks: List = [] + self.content_blocks: list = [] self.is_messages_api = is_messages_api def _chunk_parser_messages_api(self, chunk_data: dict) -> StreamingChatCompletionChunk: @@ -89,7 +88,7 @@ class AWSEventStreamDecoder: usage=None, ) - def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[Optional[Union[GChunk, StreamingChatCompletionChunk]]]: + def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | StreamingChatCompletionChunk | None]: """Given an iterator that yields lines, iterate over it & yield every event encountered""" from botocore.eventstream import EventStreamBuffer @@ -136,7 +135,7 @@ class AWSEventStreamDecoder: async def aiter_bytes( self, iterator: AsyncIterator[bytes] - ) -> AsyncIterator[Optional[Union[GChunk, StreamingChatCompletionChunk]]]: + ) -> AsyncIterator[GChunk | StreamingChatCompletionChunk | None]: """Given an async iterator that yields lines, iterate over it & yield every event encountered""" from botocore.eventstream import EventStreamBuffer @@ -190,7 +189,7 @@ class AWSEventStreamDecoder: except Exception as e: verbose_logger.error(f"Final error parsing accumulated JSON: {e}") - def _parse_message_from_event(self, event) -> Optional[str]: + def _parse_message_from_event(self, event) -> str | None: response_stream_shape = get_sagemaker_response_stream_shape() if response_stream_shape is None: raise SagemakerError( diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index 216ce08678f..868f4f9696e 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -1,7 +1,7 @@ import json from collections.abc import Callable from copy import deepcopy -from typing import Any, List, Optional, Union, cast +from typing import Any, cast import httpx @@ -90,11 +90,11 @@ class SagemakerLLM(BaseAWSLLM): credentials, model: str, data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], litellm_params: dict, optional_params: dict, aws_region_name: str, - extra_headers: Optional[dict] = None, + extra_headers: dict | None = None, ): try: from botocore.auth import SigV4Auth @@ -141,7 +141,7 @@ class SagemakerLLM(BaseAWSLLM): logging_obj, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]] = None, + timeout: float | httpx.Timeout | None = None, custom_prompt_dict={}, hf_model_name=None, logger_fn=None, @@ -393,16 +393,16 @@ class SagemakerLLM(BaseAWSLLM): async def async_streaming( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, custom_prompt_dict: dict, - hf_model_name: Optional[str], + hf_model_name: str | None, credentials, aws_region_name: str, optional_params, encoding, model_response: ModelResponse, - model_id: Optional[str], + model_id: str | None, logging_obj: Any, litellm_params: dict, headers: dict, @@ -456,17 +456,17 @@ class SagemakerLLM(BaseAWSLLM): async def async_completion( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, custom_prompt_dict: dict, - hf_model_name: Optional[str], + hf_model_name: str | None, credentials, aws_region_name: str, encoding, model_response: ModelResponse, optional_params: dict, logging_obj: Any, - model_id: Optional[str], + model_id: str | None, headers: dict, litellm_params: dict, ): @@ -529,7 +529,7 @@ class SagemakerLLM(BaseAWSLLM): ) raise e except Exception as e: - error_message = f"{str(e)}" + error_message = f"{e!s}" if "Inference Component Name header is required" in error_message: error_message += "\n pass in via `litellm.completion(..., model_id={InferenceComponentName})`" raise SagemakerError(status_code=500, message=error_message) diff --git a/litellm/llms/sagemaker/completion/transformation.py b/litellm/llms/sagemaker/completion/transformation.py index 918af7f586d..5769f13885e 100644 --- a/litellm/llms/sagemaker/completion/transformation.py +++ b/litellm/llms/sagemaker/completion/transformation.py @@ -6,8 +6,7 @@ In the Huggingface TGI format. import json import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union - +from typing import TYPE_CHECKING, Any from httpx._models import Headers, Response @@ -37,19 +36,19 @@ class SagemakerConfig(BaseConfig): Reference: https://d-uuwbxj1u4cnu.studio.us-west-2.sagemaker.aws/jupyter/default/lab/workspaces/auto-q/tree/DemoNotebooks/meta-textgeneration-llama-2-7b-SDK_1.ipynb """ - max_new_tokens: Optional[int] = None - max_completion_tokens: Optional[int] = None - top_p: Optional[float] = None - temperature: Optional[float] = None - return_full_text: Optional[bool] = None + max_new_tokens: int | None = None + max_completion_tokens: int | None = None + top_p: float | None = None + temperature: float | None = None + return_full_text: bool | None = None def __init__( self, - max_new_tokens: Optional[int] = None, - max_completion_tokens: Optional[int] = None, - top_p: Optional[float] = None, - temperature: Optional[float] = None, - return_full_text: Optional[bool] = None, + max_new_tokens: int | None = None, + max_completion_tokens: int | None = None, + top_p: float | None = None, + temperature: float | None = None, + return_full_text: bool | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -60,10 +59,10 @@ class SagemakerConfig(BaseConfig): def get_config(cls): return super().get_config() - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return SagemakerError(message=error_message, status_code=status_code, headers=headers) - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "stream", "temperature", @@ -113,9 +112,9 @@ class SagemakerConfig(BaseConfig): def _transform_prompt( self, model: str, - messages: List, + messages: list, custom_prompt_dict: dict, - hf_model_name: Optional[str], + hf_model_name: str | None, ) -> str: if model in custom_prompt_dict: # check if the model has a registered custom prompt @@ -152,14 +151,14 @@ class SagemakerConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, ) -> dict: inference_params = optional_params.copy() stream = inference_params.pop("stream", False) - data: Dict = {"parameters": inference_params} + data: dict = {"parameters": inference_params} if stream is True: data["stream"] = True @@ -180,7 +179,7 @@ class SagemakerConfig(BaseConfig): async def async_transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -194,12 +193,12 @@ class SagemakerConfig(BaseConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: completion_response = raw_response.json() ## LOGGING @@ -256,13 +255,13 @@ class SagemakerConfig(BaseConfig): def validate_environment( self, - headers: Optional[dict], + headers: dict | None, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = {"Content-Type": "application/json"} diff --git a/litellm/llms/sagemaker/embedding/cohere_transformation.py b/litellm/llms/sagemaker/embedding/cohere_transformation.py index 126f153222d..b05e146a966 100644 --- a/litellm/llms/sagemaker/embedding/cohere_transformation.py +++ b/litellm/llms/sagemaker/embedding/cohere_transformation.py @@ -10,7 +10,7 @@ be of type Object`. Reference: https://docs.cohere.com/v2/reference/embed """ -from typing import TYPE_CHECKING, Any, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast if TYPE_CHECKING: from litellm.types.llms.openai import AllEmbeddingInputValues @@ -37,7 +37,7 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): def __init__(self) -> None: pass - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return ["encoding_format", "dimensions", "input_type"] def map_openai_params( @@ -55,7 +55,7 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): optional_params["input_type"] = non_default_params["input_type"] return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return SagemakerError(message=error_message, status_code=status_code, headers=headers) def transform_embedding_request( @@ -69,11 +69,11 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): Transform embedding request for Cohere models on SageMaker """ if isinstance(input, str): - input_list: List[str] = [input] + input_list: list[str] = [input] elif isinstance(input, list): if input and (isinstance(input[0], list) or isinstance(input[0], int)): raise ValueError("Input must be a list of strings") - input_list = cast(List[str], input) + input_list = cast(list[str], input) else: input_list = [str(input)] @@ -91,7 +91,7 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): raw_response: Response, model_response: "EmbeddingResponse", logging_obj: Any, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -122,11 +122,11 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment for SageMaker Cohere embeddings diff --git a/litellm/llms/sagemaker/embedding/transformation.py b/litellm/llms/sagemaker/embedding/transformation.py index fce2bfd22e7..7221e030d97 100644 --- a/litellm/llms/sagemaker/embedding/transformation.py +++ b/litellm/llms/sagemaker/embedding/transformation.py @@ -4,7 +4,7 @@ Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke` In the Huggingface TGI format. """ -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from litellm.types.llms.openai import AllEmbeddingInputValues @@ -46,7 +46,7 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): return SagemakerCohereEmbeddingConfig() return cls() - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: model_lower = model.lower() if "voyage" in model_lower: return VoyageEmbeddingConfig().get_supported_openai_params(model) @@ -63,7 +63,7 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): ) -> dict: return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return SagemakerError(message=error_message, status_code=status_code, headers=headers) def transform_embedding_request( @@ -85,7 +85,7 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): raw_response: Response, model_response: "EmbeddingResponse", logging_obj: Any, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -97,7 +97,7 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): response_data = raw_response.json() except Exception as e: raise SagemakerError( - message=f"Failed to parse response: {str(e)}", + message=f"Failed to parse response: {e!s}", status_code=raw_response.status_code, ) @@ -146,11 +146,11 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment for SageMaker embeddings diff --git a/litellm/llms/sagemaker/nova/transformation.py b/litellm/llms/sagemaker/nova/transformation.py index 41c20847b53..0b37e4920df 100644 --- a/litellm/llms/sagemaker/nova/transformation.py +++ b/litellm/llms/sagemaker/nova/transformation.py @@ -7,8 +7,6 @@ additional Nova-specific parameters (top_k, reasoning_effort, etc.). Docs: https://docs.aws.amazon.com/nova/latest/nova2-userguide/nova-sagemaker-inference-api-reference.html """ -from typing import List - from litellm.types.llms.openai import AllMessageValues from ..chat.transformation import SagemakerChatConfig @@ -31,7 +29,7 @@ class SagemakerNovaConfig(SagemakerChatConfig): """Nova expects `stream: true` in the request body for streaming.""" return True - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: """Extend parent params with Nova-specific parameters.""" params = super().get_supported_openai_params(model) nova_params = [ @@ -48,7 +46,7 @@ class SagemakerNovaConfig(SagemakerChatConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, diff --git a/litellm/llms/sambanova/chat.py b/litellm/llms/sambanova/chat.py index be5f8b3dffb..26c53267a93 100644 --- a/litellm/llms/sambanova/chat.py +++ b/litellm/llms/sambanova/chat.py @@ -5,7 +5,7 @@ this is OpenAI compatible - no translation needed / occurs """ from collections.abc import Coroutine -from typing import Any, List, Literal, Optional, Union, overload +from typing import Any, Literal, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -21,29 +21,29 @@ class SambanovaConfig(OpenAIGPTConfig): Below are the parameters: """ - max_tokens: Optional[int] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - top_k: Optional[int] = None - stop: Optional[Union[str, list]] = None - stream: Optional[bool] = None - stream_options: Optional[dict] = None - tool_choice: Optional[str] = None - response_format: Optional[dict] = None - tools: Optional[list] = None + max_tokens: int | None = None + temperature: int | None = None + top_p: int | None = None + top_k: int | None = None + stop: str | list | None = None + stream: bool | None = None + stream_options: dict | None = None + tool_choice: str | None = None + response_format: dict | None = None + tools: list | None = None def __init__( self, - max_tokens: Optional[int] = None, - response_format: Optional[dict] = None, - stop: Optional[str] = None, - stream: Optional[bool] = None, - stream_options: Optional[dict] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - tool_choice: Optional[str] = None, - tools: Optional[list] = None, + max_tokens: int | None = None, + response_format: dict | None = None, + stop: str | None = None, + stream: bool | None = None, + stream_options: dict | None = None, + temperature: float | None = None, + top_p: float | None = None, + top_k: int | None = None, + tool_choice: str | None = None, + tools: list | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -100,20 +100,20 @@ class SambanovaConfig(OpenAIGPTConfig): @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... @overload def _transform_messages( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, is_async: Literal[False] = False, - ) -> List[AllMessageValues]: ... + ) -> list[AllMessageValues]: ... def _transform_messages( - self, messages: List[AllMessageValues], model: str, is_async: bool = False - ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: """ Transform messages to handle content list conversion. diff --git a/litellm/llms/sambanova/embedding/transformation.py b/litellm/llms/sambanova/embedding/transformation.py index 611507bcf0d..c6730aa1ddd 100644 --- a/litellm/llms/sambanova/embedding/transformation.py +++ b/litellm/llms/sambanova/embedding/transformation.py @@ -3,8 +3,6 @@ This is OpenAI compatible - no transformation is applied """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -23,12 +21,12 @@ class SambaNovaEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: raise ValueError("api_base is required for SambaNova embeddings") @@ -42,11 +40,11 @@ class SambaNovaEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("SAMBANOVA_API_KEY") @@ -106,7 +104,7 @@ class SambaNovaEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -132,7 +130,5 @@ class SambaNovaEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return SambaNovaError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/sap/chat/handler.py b/litellm/llms/sap/chat/handler.py index 28aecaa38a0..dd806600a9b 100755 --- a/litellm/llms/sap/chat/handler.py +++ b/litellm/llms/sap/chat/handler.py @@ -3,7 +3,6 @@ from __future__ import annotations import json import time from collections.abc import AsyncIterator, Iterator -from typing import Optional import httpx @@ -47,7 +46,7 @@ class _StreamParser: """Normalize orchestration streaming events into OpenAI-like chunks.""" @staticmethod - def _from_orchestration_result(evt: dict) -> Optional[OpenAIChatCompletionChunk]: + def _from_orchestration_result(evt: dict) -> OpenAIChatCompletionChunk | None: """ Accepts orchestration_result shape and maps it to an OpenAI-like *chunk*. """ @@ -73,7 +72,7 @@ class _StreamParser: ) @staticmethod - def to_openai_chunk(event_obj: dict) -> Optional[OpenAIChatCompletionChunk]: + def to_openai_chunk(event_obj: dict) -> OpenAIChatCompletionChunk | None: """ Accepts: - {"final_result": } (IMPORTANT: this is just another chunk, NOT terminal) @@ -140,7 +139,7 @@ class SAPStreamIterator: if not line: continue - payload = line[len(self._prefix) :] if line.startswith(self._prefix) else line + payload = line.removeprefix(self._prefix) if payload == self._final: self._safe_close() raise StopIteration @@ -212,7 +211,7 @@ class AsyncSAPStreamIterator: continue # now = lambda: int(time.time() * 1000) - payload = line[len(self._prefix) :] if line.startswith(self._prefix) else line + payload = line.removeprefix(self._prefix) if payload == self._final: await self._aclose() raise StopAsyncIteration diff --git a/litellm/llms/sap/chat/models.py b/litellm/llms/sap/chat/models.py index 2756dd0e67e..ff089e8b680 100644 --- a/litellm/llms/sap/chat/models.py +++ b/litellm/llms/sap/chat/models.py @@ -1,11 +1,11 @@ -from typing import Union, Literal, Optional -from enum import Enum import warnings +from enum import Enum +from typing import Literal, Union from pydantic import BaseModel, Field, field_validator, model_validator -def validate_different_content(v: Union[str, dict, list]) -> str: +def validate_different_content(v: str | dict | list) -> str: if v in ((), {}, []): return "" elif isinstance(v, dict) and "text" in v: @@ -95,7 +95,7 @@ class SAPMessage(BaseModel): class SAPUserMessage(BaseModel): role: Literal["user"] = "user" - content: Union[str, TextContent, ImageContent, list[Union[TextContent, ImageContent]]] + content: str | TextContent | ImageContent | list[TextContent | ImageContent] class SAPAssistantMessage(BaseModel): @@ -140,12 +140,12 @@ class KeyValueListPair(BaseModel): class DocumentMetadataKeyValueListPairs(KeyValueListPair): - select_mode: Optional[list[Literal["ignoreIfKeyAbsent"]]] = None + select_mode: list[Literal["ignoreIfKeyAbsent"]] | None = None class GroundingSearchConfig(BaseModel): - max_chunk_count: Optional[int] = Field(default=None, ge=0) - max_document_count: Optional[int] = Field(default=None, ge=0) + max_chunk_count: int | None = Field(default=None, ge=0) + max_document_count: int | None = Field(default=None, ge=0) @model_validator(mode="after") def validate_max_chunk_count_and_max_document_count(self): @@ -155,13 +155,13 @@ class GroundingSearchConfig(BaseModel): class DocumentGroundingFilter(BaseModel): - id_: Optional[str] = Field(default=None, alias="id") + id_: str | None = Field(default=None, alias="id") data_repository_type: Literal["vector", "help.sap.com"] - search_config: Optional[GroundingSearchConfig] = None - data_repositories: Optional[list[str]] = None - data_repository_metadata: Optional[list[KeyValueListPair]] = None - document_metadata: Optional[list[DocumentMetadataKeyValueListPairs]] = None - chunk_metadata: Optional[list[KeyValueListPair]] = None + search_config: GroundingSearchConfig | None = None + data_repositories: list[str] | None = None + data_repository_metadata: list[KeyValueListPair] | None = None + document_metadata: list[DocumentMetadataKeyValueListPairs] | None = None + chunk_metadata: list[KeyValueListPair] | None = None class DocumentGroundingPlaceholders(BaseModel): @@ -170,9 +170,9 @@ class DocumentGroundingPlaceholders(BaseModel): class DocumentGroundingConfig(BaseModel): - filters: Optional[list[DocumentGroundingFilter]] = None + filters: list[DocumentGroundingFilter] | None = None placeholders: DocumentGroundingPlaceholders - metadata_params: Optional[list[str]] = None + metadata_params: list[str] | None = None class GroundingModuleConfig(BaseModel): @@ -182,15 +182,15 @@ class GroundingModuleConfig(BaseModel): class Template(BaseModel): template: list[ChatMessage] - defaults: Optional[dict[str, str]] = None - response_format: Optional[Union[ResponseFormat, ResponseFormatJSONSchema]] = None - tools: Optional[list[ChatCompletionTool]] = None + defaults: dict[str, str] | None = None + response_format: ResponseFormat | ResponseFormatJSONSchema | None = None + tools: list[ChatCompletionTool] | None = None class LLMModelDetails(BaseModel): name: str version: str = "latest" - params: Optional[dict] = None + params: dict | None = None class PromptTemplatingModuleConfig(BaseModel): @@ -319,7 +319,7 @@ class DPIStandardEntity(BaseModel): """ type_: SAPMaskingProfileEntity = Field(..., alias="type") - replacement_strategy: Optional[Union[DPIMethodConstant, DPIMethodFabricatedData]] = None + replacement_strategy: DPIMethodConstant | DPIMethodFabricatedData | None = None class MaskGroundingInput(BaseModel): @@ -351,9 +351,9 @@ class MaskingProviderConfig(BaseModel): type_: Literal["sap_data_privacy_integration"] = Field(default="sap_data_privacy_integration", alias="type") method: Literal["anonymization", "pseudonymization"] - entities: list[Union[DPIStandardEntity, DPICustomEntity]] - allowlist: Optional[list[str]] = None - mask_grounding_input: Optional[MaskGroundingInput] = None + entities: list[DPIStandardEntity | DPICustomEntity] + allowlist: list[str] | None = None + mask_grounding_input: MaskGroundingInput | None = None class MaskingModuleConfig(BaseModel): @@ -367,8 +367,8 @@ class MaskingModuleConfig(BaseModel): DEPRECATED: parameter 'masking_providers' will be removed Sept 15, 2026. Use 'providers' instead. """ - providers: Optional[list[MaskingProviderConfig]] = Field(min_length=1, default=None) - masking_providers: Optional[list[MaskingProviderConfig]] = Field(min_length=1, default=None) + providers: list[MaskingProviderConfig] | None = Field(min_length=1, default=None) + masking_providers: list[MaskingProviderConfig] | None = Field(min_length=1, default=None) @model_validator(mode="after") def enforce_exactly_one_provider_list(self): @@ -435,10 +435,10 @@ class AzureContentFilter(BaseModel): self_harm: Threshold for self-harm content. """ - hate: Optional[Union[AzureThreshold, Literal[0, 2, 4, 6]]] = None - sexual: Optional[Union[AzureThreshold, Literal[0, 2, 4, 6]]] = None - violence: Optional[Union[AzureThreshold, Literal[0, 2, 4, 6]]] = None - self_harm: Optional[Union[AzureThreshold, Literal[0, 2, 4, 6]]] = None + hate: AzureThreshold | Literal[0, 2, 4, 6] | None = None + sexual: AzureThreshold | Literal[0, 2, 4, 6] | None = None + violence: AzureThreshold | Literal[0, 2, 4, 6] | None = None + self_harm: AzureThreshold | Literal[0, 2, 4, 6] | None = None class AzureContentSafetyInput(AzureContentFilter): @@ -457,7 +457,7 @@ class AzureContentSafetyInput(AzureContentFilter): prompt_shield: A flag to use prompt shield """ - prompt_shield: Optional[bool] = False + prompt_shield: bool | None = False class AzureContentSafetyOutput(AzureContentFilter): @@ -478,7 +478,7 @@ class AzureContentSafetyOutput(AzureContentFilter): and other proprietary programming content. """ - protected_material_code: Optional[bool] = False + protected_material_code: bool | None = False class LlamaGuard38bFilter(BaseModel): @@ -539,12 +539,12 @@ class LlamaGuard38bFilterConfig(BaseModel): class AzureContentSafetyInputFilterConfig(BaseModel): type_: Literal["azure_content_safety"] = Field(default="azure_content_safety", alias="type") - config: Optional[AzureContentSafetyInput] = None + config: AzureContentSafetyInput | None = None class AzureContentSafetyOutputFilterConfig(BaseModel): type_: Literal["azure_content_safety"] = Field(default="azure_content_safety", alias="type") - config: Optional[AzureContentSafetyOutput] = None + config: AzureContentSafetyOutput | None = None class FilteringStreamOptions(BaseModel): @@ -553,7 +553,7 @@ class FilteringStreamOptions(BaseModel): from previous chunks as additional context. """ - overlap: Optional[int] = Field(default=0, ge=0, le=10000) + overlap: int | None = Field(default=0, ge=0, le=10000) class InputFiltering(BaseModel): @@ -563,7 +563,7 @@ class InputFiltering(BaseModel): filters: List of ContentFilter objects to be applied to input content. """ - filters: list[Union[AzureContentSafetyInputFilterConfig, LlamaGuard38bFilterConfig]] = Field(min_length=1) + filters: list[AzureContentSafetyInputFilterConfig | LlamaGuard38bFilterConfig] = Field(min_length=1) class OutputFiltering(BaseModel): @@ -575,8 +575,8 @@ class OutputFiltering(BaseModel): stream_options: Module-specific streaming options. """ - filters: list[Union[AzureContentSafetyOutputFilterConfig, LlamaGuard38bFilterConfig]] = Field(min_length=1) - stream_options: Optional[FilteringStreamOptions] = None + filters: list[AzureContentSafetyOutputFilterConfig | LlamaGuard38bFilterConfig] = Field(min_length=1) + stream_options: FilteringStreamOptions | None = None class FilteringModuleConfig(BaseModel): @@ -588,8 +588,8 @@ class FilteringModuleConfig(BaseModel): output: Module for filtering and validating output content after generation. """ - input: Optional[InputFiltering] = None - output: Optional[OutputFiltering] = None + input: InputFiltering | None = None + output: OutputFiltering | None = None @model_validator(mode="after") def enforce_min_properties(self) -> "FilteringModuleConfig": @@ -629,14 +629,14 @@ class InputTranslationConfig(BaseModel): apply_to: List of selectors that define the scope of translation. """ - source_language: Optional[str] = None + source_language: str | None = None target_language: str - apply_to: Optional[list[SAPDocumentTranslationApplyToSelector]] = None + apply_to: list[SAPDocumentTranslationApplyToSelector] | None = None class OutputTranslationConfig(BaseModel): - source_language: Optional[str] = None - target_language: Union[str, SAPDocumentTranslationApplyToSelector] + source_language: str | None = None + target_language: str | SAPDocumentTranslationApplyToSelector class SAPDocumentTranslationInput(BaseModel): @@ -652,7 +652,7 @@ class SAPDocumentTranslationInput(BaseModel): """ type_: Literal["sap_document_translation"] = Field(default="sap_document_translation", alias="type") - translate_messages_history: Optional[bool] = None + translate_messages_history: bool | None = None config: InputTranslationConfig @@ -680,8 +680,8 @@ class TranslationModuleConfig(BaseModel): output: Configuration for output translation """ - input: Optional[SAPDocumentTranslationInput] = None - output: Optional[SAPDocumentTranslationOutput] = None + input: SAPDocumentTranslationInput | None = None + output: SAPDocumentTranslationOutput | None = None @model_validator(mode="after") def enforce_min_properties(self) -> "TranslationModuleConfig": @@ -692,23 +692,23 @@ class TranslationModuleConfig(BaseModel): class ModuleConfig(BaseModel): prompt_templating: PromptTemplatingModuleConfig - filtering: Optional[FilteringModuleConfig] = None - masking: Optional[MaskingModuleConfig] = None - grounding: Optional[GroundingModuleConfig] = None - translation: Optional[TranslationModuleConfig] = None + filtering: FilteringModuleConfig | None = None + masking: MaskingModuleConfig | None = None + grounding: GroundingModuleConfig | None = None + translation: TranslationModuleConfig | None = None class GlobalStreamOptions(BaseModel): enabled: bool = False - chunk_size: Optional[int] = Field(default=None, ge=1) - delimiters: Optional[list[str]] = None + chunk_size: int | None = Field(default=None, ge=1) + delimiters: list[str] | None = None class OrchestrationConfig(BaseModel): - modules: Union[ModuleConfig, list[ModuleConfig]] - stream: Optional[GlobalStreamOptions] = None + modules: ModuleConfig | list[ModuleConfig] + stream: GlobalStreamOptions | None = None class OrchestrationRequest(BaseModel): config: OrchestrationConfig - placeholder_values: Optional[dict[str, str]] = None + placeholder_values: dict[str, str] | None = None diff --git a/litellm/llms/sap/chat/transformation.py b/litellm/llms/sap/chat/transformation.py index 0ec95c15be0..753dae4c783 100755 --- a/litellm/llms/sap/chat/transformation.py +++ b/litellm/llms/sap/chat/transformation.py @@ -7,11 +7,6 @@ from functools import cached_property from typing import ( TYPE_CHECKING, Any, - Dict, - FrozenSet, - List, - Optional, - Tuple, Union, ) @@ -48,7 +43,7 @@ from .models import ( ) # Keys routed outside SAP orchestration `model.params` (prompt, stream, fallbacks, etc.) -_SAP_MODEL_PARAMS_EXCLUDED_KEYS: FrozenSet[str] = frozenset( +_SAP_MODEL_PARAMS_EXCLUDED_KEYS: frozenset[str] = frozenset( { "tools", "tool_choice", @@ -64,7 +59,7 @@ def validate_dict(data: dict, model) -> dict: return model(**data).model_dump(by_alias=True, exclude_unset=True) -def _messages_to_sap_template(messages: List[Dict[str, str]]) -> list: # type: ignore[type-arg] +def _messages_to_sap_template(messages: list[dict[str, str]]) -> list: # type: ignore[type-arg] template = [] for message in messages: if message["role"] == "user": @@ -78,7 +73,7 @@ def _messages_to_sap_template(messages: List[Dict[str, str]]) -> list: # type: return template -def _tools_response_format_and_stream(optional_params: dict, model_params: dict) -> Tuple[dict, dict, dict]: +def _tools_response_format_and_stream(optional_params: dict, model_params: dict) -> tuple[dict, dict, dict]: tools_ = optional_params.pop("tools", []) tools_ = [validate_dict(tool, ChatCompletionTool) for tool in tools_] tools: dict = {"tools": tools_} if tools_ else {} @@ -105,36 +100,36 @@ def _tools_response_format_and_stream(optional_params: dict, model_params: dict) class GenAIHubOrchestrationConfig(OpenAIGPTConfig): - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None - tools: Optional[list] = None - tool_choice: Optional[Union[str, dict]] = None # + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None + tools: list | None = None + tool_choice: str | dict | None = None model_version: str = "latest" def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, - tools: Optional[list] = None, - tool_choice: Optional[Union[str, dict]] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -144,14 +139,14 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): self._base_url = None self._resource_group = None - def run_env_setup(self, service_key: Optional[str] = None) -> None: + def run_env_setup(self, service_key: str | None = None) -> None: try: self.token_creator, self._base_url, self._resource_group = get_token_creator(service_key) # type: ignore except ValueError as err: raise GenAIHubOrchestrationError(status_code=400, message=err.args[0]) @property - def headers(self) -> Dict[str, str]: + def headers(self) -> dict[str, str]: if self.token_creator is None: self.run_env_setup() access_token = self.token_creator() # pyright: ignore[reportOptionalCall] # run_env_setup set it or raised @@ -180,7 +175,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): client = litellm.module_level_client # with httpx.Client(timeout=30) as client: deployments = client.get(f"{self.base_url}/lm/deployments", headers=self.headers).json() - valid: List[Tuple[str, str]] = [] + valid: list[tuple[str, str]] = [] for dep in deployments.get("resources", []): if dep.get("scenarioId") == "orchestration": cfg = client.get( @@ -238,11 +233,11 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key: self.run_env_setup(api_key) @@ -250,12 +245,12 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ): api_base_ = f"{self.deployment_url}/v2/completion" return api_base_ @@ -263,7 +258,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): def _build_prompt_module( self, model_name: str, - template_messages: List[Dict[str, str]], + template_messages: list[dict[str, str]], params: dict, ) -> dict: # Filter strict for GPT models only - SAP AI Core doesn't accept it as a model param @@ -318,7 +313,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[Dict[str, str]], # type: ignore + messages: list[dict[str, str]], # type: ignore optional_params: dict, litellm_params: dict, headers: dict, @@ -355,8 +350,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): fallback_model = modules_dict.pop("model", None) if fallback_model is None: raise ValueError("Each entry in `fallback_sap_modules` must include a 'model' key.") - if fallback_model.startswith("sap/"): - fallback_model = fallback_model[4:] + fallback_model = fallback_model.removeprefix("sap/") fallback_template = modules_dict.pop("messages", []) modules.append( @@ -367,13 +361,13 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): ) ) - config_payload: Dict[str, Any] = { + config_payload: dict[str, Any] = { "modules": modules if len(modules) > 1 else modules[0], } if stream_config: config_payload["stream"] = stream_config - request_body: Dict[str, Any] = {"config": config_payload} + request_body: dict[str, Any] = {"config": config_payload} if placeholder_values is not None: request_body["placeholder_values"] = placeholder_values @@ -388,12 +382,12 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: logging_obj.post_call( input=messages, @@ -437,7 +431,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): self, streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse"], sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): if sync_stream: return SAPStreamIterator(response=streaming_response) # type: ignore diff --git a/litellm/llms/sap/credentials.py b/litellm/llms/sap/credentials.py index 34bd6eda2ff..5b2e02875b8 100644 --- a/litellm/llms/sap/credentials.py +++ b/litellm/llms/sap/credentials.py @@ -8,7 +8,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path from threading import Lock -from typing import Any, Dict, Final, List, Optional, Tuple, Union +from typing import Any, Final import httpx @@ -33,7 +33,7 @@ def _get_home() -> str: return os.getenv(HOME_PATH_ENV_VAR, DEFAULT_HOME_PATH) -def _get_nested(d: Union[Dict[str, Any], str], path: Sequence[str]) -> Any: +def _get_nested(d: dict[str, Any] | str, path: Sequence[str]) -> Any: cur: Any = d if isinstance(cur, str): # This shouldn't happen if service keys are pre-parsed correctly @@ -54,7 +54,7 @@ def _get_nested(d: Union[Dict[str, Any], str], path: Sequence[str]) -> Any: return cur -def _load_json_env(var_name: str) -> Optional[Dict[str, Any]]: +def _load_json_env(var_name: str) -> dict[str, Any] | None: raw = os.environ.get(var_name) if not raw: return None @@ -64,18 +64,18 @@ def _load_json_env(var_name: str) -> Optional[Dict[str, Any]]: return None -def _str_or_none(value) -> Optional[str]: +def _str_or_none(value) -> str | None: try: return str(value) if value is not None else None except Exception: return None -def _load_vcap() -> Dict[str, Any]: +def _load_vcap() -> dict[str, Any]: return _load_json_env(VCAP_SERVICES_ENV_VAR) or {} -def _get_vcap_service(label: str) -> Optional[Dict[str, Any]]: +def _get_vcap_service(label: str) -> dict[str, Any] | None: for services in _load_vcap().values(): for svc in services: if svc.get("label") == label: @@ -86,18 +86,18 @@ def _get_vcap_service(label: str) -> Optional[Dict[str, Any]]: @dataclass class Source: name: str - get: Callable[[CredentialsValue], Optional[str]] + get: Callable[[CredentialsValue], str | None] @dataclass(frozen=True) class CredentialsValue: name: str - vcap_key: Optional[Tuple[str, ...]] = None - default: Optional[str] = None - transform_fn: Optional[Callable[[str], str]] = None + vcap_key: tuple[str, ...] | None = None + default: str | None = None + transform_fn: Callable[[str], str] | None = None -CREDENTIAL_VALUES: Final[List[CredentialsValue]] = [ +CREDENTIAL_VALUES: Final[list[CredentialsValue]] = [ CredentialsValue("client_id", ("clientid",)), CredentialsValue("client_secret", ("clientsecret",)), CredentialsValue( @@ -124,7 +124,7 @@ CREDENTIAL_VALUES: Final[List[CredentialsValue]] = [ ] -def init_conf(profile: Optional[str] = None) -> Dict[str, Any]: +def init_conf(profile: str | None = None) -> dict[str, Any]: """ Loads config JSON from: 1) $AICORE_CONFIG if set, otherwise @@ -158,7 +158,7 @@ def _env_name(name: str) -> str: return f"AICORE_{name.upper()}" -def extract_credentials(source: Source) -> Dict[str, str]: +def extract_credentials(source: Source) -> dict[str, str]: """Extract all credentials from a source.""" credentials = {} for cv in CREDENTIAL_VALUES: @@ -168,7 +168,7 @@ def extract_credentials(source: Source) -> Dict[str, str]: return credentials -def resolve_credentials(sources: List[Source]) -> Dict[str, str]: +def resolve_credentials(sources: list[Source]) -> dict[str, str]: """Extract credentials from the first source that has any defined.""" for source in sources: credentials = extract_credentials(source) @@ -178,7 +178,7 @@ def resolve_credentials(sources: List[Source]) -> Dict[str, str]: raise ValueError("No credentials found in any source") -def resolve_resource_group(sources: List[Source]) -> Optional[str]: +def resolve_resource_group(sources: list[Source]) -> str | None: """Find resource_group from the first source that defines it.""" rg_cred = CredentialsValue("resource_group", default="default") for source in sources: @@ -190,8 +190,8 @@ def resolve_resource_group(sources: List[Source]) -> Optional[str]: def _parse_service_key_once( - service_key: Optional[Union[str, dict]], -) -> Optional[Dict[str, Any]]: + service_key: str | dict | None, +) -> dict[str, Any] | None: """ Pre-parse service_key if it's a string to avoid repeated JSON parsing. @@ -213,9 +213,7 @@ def _parse_service_key_once( return None -def _resolve_credential_from_service_key( - service_key: Optional[Union[str, dict]], cv: CredentialsValue -) -> Optional[str]: +def _resolve_credential_from_service_key(service_key: str | dict | None, cv: CredentialsValue) -> str | None: if service_key is None: return None val = _str_or_none(_get_nested(service_key, (("credentials",) + cv.vcap_key) if cv.vcap_key else (cv.name,))) @@ -225,10 +223,10 @@ def _resolve_credential_from_service_key( def fetch_credentials( - service_key: Optional[Union[str, dict]] = None, - profile: Optional[str] = None, + service_key: str | dict | None = None, + profile: str | None = None, **kwargs, -) -> Dict[str, str]: +) -> dict[str, str]: """ Resolution order (first-source-wins): @@ -298,14 +296,14 @@ def fetch_credentials( def validate_credentials( - auth_url: Optional[str] = None, - base_url: Optional[str] = None, - client_id: Optional[str] = None, - client_secret: Optional[str] = None, - cert_str: Optional[str] = None, - key_str: Optional[str] = None, - cert_file_path: Optional[str] = None, - key_file_path: Optional[str] = None, + auth_url: str | None = None, + base_url: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + cert_str: str | None = None, + key_str: str | None = None, + cert_file_path: str | None = None, + key_file_path: str | None = None, ) -> None: """ Validate SAP AI Core credentials for completeness and consistency. @@ -357,7 +355,7 @@ def _request_token( if client_secret: data["client_secret"] = client_secret - resp: Optional[httpx.Response] = None + resp: httpx.Response | None = None try: if cert_pair: with httpx.Client(cert=cert_pair) as raw_client: @@ -378,13 +376,13 @@ def _request_token( def get_token_creator( - service_key: Optional[Union[str, dict]] = None, - profile: Optional[str] = None, + service_key: str | dict | None = None, + profile: str | None = None, *, timeout: float = 30.0, expiry_buffer_minutes: int = 60, **overrides, -) -> Tuple[Callable[[], str], str, str]: +) -> tuple[Callable[[], str], str, str]: """ Creates a callable that fetches and caches an OAuth2 bearer token using credentials from `fetch_credentials()`. @@ -405,7 +403,7 @@ def get_token_creator( """ # Resolve credentials using your helper - credentials: Dict[str, str] = fetch_credentials(service_key=service_key, profile=profile, **overrides) + credentials: dict[str, str] = fetch_credentials(service_key=service_key, profile=profile, **overrides) auth_url = credentials.get("auth_url") base_url = credentials.get("base_url") @@ -429,8 +427,8 @@ def get_token_creator( ) lock = Lock() - token: Optional[str] = None - token_expiry: Optional[datetime] = None + token: str | None = None + token_expiry: datetime | None = None def _fetch_token() -> tuple[str, datetime]: # Case 1: secret-based auth diff --git a/litellm/llms/sap/embed/transformation.py b/litellm/llms/sap/embed/transformation.py index 8368be718ad..3ac200e778e 100644 --- a/litellm/llms/sap/embed/transformation.py +++ b/litellm/llms/sap/embed/transformation.py @@ -2,17 +2,17 @@ Translates from OpenAI's `/v1/embeddings` to IBM's `/text/embeddings` route. """ -from typing import Optional, List, Dict, Literal, Union -from pydantic import BaseModel, Field from functools import cached_property -from litellm.llms.sap.chat.models import MaskingModuleConfig +from typing import Literal import httpx +from pydantic import BaseModel, Field from litellm.llms.base_llm.embedding.transformation import ( BaseEmbeddingConfig, LiteLLMLoggingObj, ) +from litellm.llms.sap.chat.models import MaskingModuleConfig from litellm.types.llms.openai import AllEmbeddingInputValues from litellm.types.utils import EmbeddingResponse @@ -27,13 +27,13 @@ class Usage(BaseModel): class EmbeddingItem(BaseModel): object: Literal["embedding"] - embedding: List[float] = Field(..., description="Vector of floats (length varies by model).") + embedding: list[float] = Field(..., description="Vector of floats (length varies by model).") index: int class FinalResult(BaseModel): object: Literal["list"] - data: List[EmbeddingItem] + data: list[EmbeddingItem] model: str usage: Usage @@ -47,8 +47,8 @@ class EmbeddingModel(BaseModel): name: str version: str = "latest" params: dict = Field(default_factory=dict) - timeout: Optional[int] = Field(default=None, ge=1, le=600) - max_retries: Optional[int] = Field(default=None, ge=0, le=5) + timeout: int | None = Field(default=None, ge=1, le=600) + max_retries: int | None = Field(default=None, ge=0, le=5) class EmbeddingsModelConfig(BaseModel): @@ -57,12 +57,12 @@ class EmbeddingsModelConfig(BaseModel): class EmbeddingsModules(BaseModel): embeddings: EmbeddingsModelConfig - masking: Optional[MaskingModuleConfig] = None + masking: MaskingModuleConfig | None = None class EmbeddingInput(BaseModel): - text: Union[str, List[str]] - type: Optional[Literal["text", "document", "query"]] = None + text: str | list[str] + type: Literal["text", "document", "query"] | None = None class EmbeddingConfig(BaseModel): @@ -85,7 +85,7 @@ class GenAIHubEmbeddingConfig(BaseEmbeddingConfig): self.token_creator, self.base_url, self.resource_group = get_token_creator() @property - def headers(self) -> Dict: + def headers(self) -> dict: access_token = self.token_creator() # headers for completions and embeddings requests headers = { @@ -136,12 +136,12 @@ class GenAIHubEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: url = self.deployment_url.rstrip("/") + "/v2/embeddings" return url @@ -182,7 +182,7 @@ class GenAIHubEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, diff --git a/litellm/llms/scaleway/audio_transcription/transformation.py b/litellm/llms/scaleway/audio_transcription/transformation.py index d5438cbf930..fbd278a70b7 100644 --- a/litellm/llms/scaleway/audio_transcription/transformation.py +++ b/litellm/llms/scaleway/audio_transcription/transformation.py @@ -4,8 +4,6 @@ Support for Scaleway's OpenAI-compatible `/v1/audio/transcriptions` endpoint. API reference: https://www.scaleway.com/en/developers/api/generative-apis/#path-audio-create-an-audio-transcription """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -27,7 +25,7 @@ class ScalewayAudioTranscriptionException(BaseLLMException): class ScalewayAudioTranscriptionConfig(BaseAudioTranscriptionConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: return [ "language", "prompt", @@ -51,19 +49,17 @@ class ScalewayAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = "https://api.scaleway.ai/v1" if api_base is None else api_base.rstrip("/") return f"{api_base}/audio/transcriptions" - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return ScalewayAudioTranscriptionException( message=error_message, status_code=status_code, @@ -74,11 +70,11 @@ class ScalewayAudioTranscriptionConfig(BaseAudioTranscriptionConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = get_secret_str("SCW_SECRET_KEY") diff --git a/litellm/llms/searchapi/search/transformation.py b/litellm/llms/searchapi/search/transformation.py index 5f3e535d7fd..54296255e64 100644 --- a/litellm/llms/searchapi/search/transformation.py +++ b/litellm/llms/searchapi/search/transformation.py @@ -4,7 +4,7 @@ Calls SearchAPI.io's Google Search API endpoint. SearchAPI.io API Reference: https://www.searchapi.io/docs/google """ -from typing import Dict, List, Literal, Optional, TypedDict, Union, cast +from typing import Literal, TypedDict, cast from urllib.parse import urlencode import httpx @@ -66,11 +66,11 @@ class SearchAPIConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -91,9 +91,9 @@ class SearchAPIConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -113,13 +113,13 @@ class SearchAPIConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, - api_key: Optional[str] = None, + api_key: str | None = None, api_base: str | None = None, - search_engine_id: Optional[str] = None, + search_engine_id: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to SearchAPI.io format. @@ -189,7 +189,7 @@ class SearchAPIConfig(BaseSearchConfig): } @staticmethod - def _append_domain_filters(query: str, domains: List[str]) -> str: + def _append_domain_filters(query: str, domains: list[str]) -> str: """ Add site: filters to restrict search to specific domains. """ @@ -201,7 +201,7 @@ class SearchAPIConfig(BaseSearchConfig): def transform_search_response( self, raw_response: httpx.Response, - logging_obj: Optional[LiteLLMLoggingObj], + logging_obj: LiteLLMLoggingObj | None, **kwargs, ) -> SearchResponse: """ @@ -216,7 +216,7 @@ class SearchAPIConfig(BaseSearchConfig): response_json = raw_response.json() # Transform results to SearchResult objects - results: List[SearchResult] = [] + results: list[SearchResult] = [] # Process organic results for result in response_json.get("organic_results", []): diff --git a/litellm/llms/searxng/search/transformation.py b/litellm/llms/searxng/search/transformation.py index b5f41015112..750b1a628a1 100644 --- a/litellm/llms/searxng/search/transformation.py +++ b/litellm/llms/searxng/search/transformation.py @@ -4,7 +4,7 @@ Calls SearXNG's /search endpoint to search the web. SearXNG API Reference: https://docs.searxng.org/dev/search_api.html """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -50,11 +50,11 @@ class SearXNGSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. SearXNG is open-source and doesn't require an API key by default. @@ -75,9 +75,9 @@ class SearXNGSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -113,10 +113,10 @@ class SearXNGSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to SearXNG API format. diff --git a/litellm/llms/serper/search/transformation.py b/litellm/llms/serper/search/transformation.py index 31a0d3f2bac..e56cc7ceb2d 100644 --- a/litellm/llms/serper/search/transformation.py +++ b/litellm/llms/serper/search/transformation.py @@ -4,7 +4,7 @@ Calls Serper's /search endpoint to search Google. Serper API Reference: https://serper.dev """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -47,11 +47,11 @@ class SerperSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -70,9 +70,9 @@ class SerperSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -88,10 +88,10 @@ class SerperSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Serper API format. diff --git a/litellm/llms/snowflake/chat/transformation.py b/litellm/llms/snowflake/chat/transformation.py index 8b23ae135b5..e3ad3cd4400 100644 --- a/litellm/llms/snowflake/chat/transformation.py +++ b/litellm/llms/snowflake/chat/transformation.py @@ -9,7 +9,7 @@ Ref: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -67,7 +67,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): def get_config(cls): return super().get_config() - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: params = [ "temperature", "max_tokens", @@ -83,12 +83,12 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = self._get_api_base(api_base, optional_params) if _is_claude_model(model): @@ -99,11 +99,11 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = super().validate_environment( headers=headers, @@ -118,7 +118,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): headers["anthropic-version"] = ANTHROPIC_VERSION return headers - def _transform_tools_to_anthropic(self, tools: List[Dict]) -> List[Dict]: + def _transform_tools_to_anthropic(self, tools: list[dict]) -> list[dict]: """ Convert tools from OpenAI format to Anthropic format. @@ -129,7 +129,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): for tool in tools: if tool.get("type") == "function" and "function" in tool: func = tool["function"] - anthropic_tool: Dict[str, Any] = { + anthropic_tool: dict[str, Any] = { "name": func.get("name", ""), } if "description" in func: @@ -146,7 +146,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): anthropic_tools.append(tool) return anthropic_tools - def _extract_system_and_messages(self, messages: List[AllMessageValues]) -> tuple[Optional[str], List[Dict]]: + def _extract_system_and_messages(self, messages: list[AllMessageValues]) -> tuple[str | None, list[dict]]: """ Split messages into system prompt and conversation turns for Anthropic format. @@ -154,8 +154,8 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): - assistant messages with tool_calls → tool_use content blocks - tool role messages → user role with tool_result content blocks """ - system_parts: List[str] = [] - conversation: List[Dict] = [] + system_parts: list[str] = [] + conversation: list[dict] = [] for msg in messages: if isinstance(msg, dict): @@ -173,7 +173,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): elif role == "assistant": tool_calls = msg.get("tool_calls") if isinstance(msg, dict) else getattr(msg, "tool_calls", None) if tool_calls: # type: ignore[truthy-bool] - content_blocks: List[Dict[str, Any]] = [] + content_blocks: list[dict[str, Any]] = [] if content: content_blocks.append({"type": "text", "text": content}) for tc in tool_calls: # type: ignore[attr-defined] @@ -221,13 +221,13 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): else: conversation.append({"role": role, "content": content}) - system: Optional[str] = "\n\n".join(system_parts) if system_parts else None + system: str | None = "\n\n".join(system_parts) if system_parts else None return system, conversation def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -242,7 +242,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): def _transform_request_openai( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, stream: bool, extra_body: dict, @@ -265,7 +265,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): return body - def _transform_tool_choice_to_anthropic(self, tool_choice: Any) -> Dict[str, Any]: + def _transform_tool_choice_to_anthropic(self, tool_choice: Any) -> dict[str, Any]: """ Convert tool_choice from OpenAI format to Anthropic format. @@ -290,7 +290,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): def _transform_request_anthropic( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, stream: bool, extra_body: dict, @@ -310,7 +310,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): model_name = model.removeprefix("snowflake/") - body: Dict[str, Any] = { + body: dict[str, Any] = { "model": model_name, "messages": conversation, "stream": stream, @@ -333,12 +333,12 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: if _is_claude_model(model): return self._transform_response_anthropic( @@ -353,7 +353,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], ) -> ModelResponse: """Parse standard OpenAI chat completions response.""" response_json = raw_response.json() @@ -380,7 +380,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], ) -> ModelResponse: """Parse Anthropic Messages response into OpenAI format.""" response_json = raw_response.json() @@ -449,7 +449,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): self, streaming_response: Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return SnowflakeStreamingHandler( streaming_response=streaming_response, @@ -470,7 +470,7 @@ class SnowflakeStreamingHandler(BaseModelResponseIterator): self, streaming_response: Any, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): super().__init__(streaming_response=streaming_response, sync_stream=sync_stream) self._tool_index = 0 diff --git a/litellm/llms/snowflake/common_utils.py b/litellm/llms/snowflake/common_utils.py index 40c8270f95f..d8c4aec73f7 100644 --- a/litellm/llms/snowflake/common_utils.py +++ b/litellm/llms/snowflake/common_utils.py @@ -1,11 +1,8 @@ -from typing import Optional - - class SnowflakeBase: def validate_environment( self, headers: dict, - JWT: Optional[str] = None, + JWT: str | None = None, ) -> dict: """ Return headers to use for Snowflake completion request diff --git a/litellm/llms/snowflake/embedding/transformation.py b/litellm/llms/snowflake/embedding/transformation.py index 44abb66b900..8470db109b7 100644 --- a/litellm/llms/snowflake/embedding/transformation.py +++ b/litellm/llms/snowflake/embedding/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -8,7 +6,7 @@ from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.types.llms.openai import AllEmbeddingInputValues from litellm.types.utils import EmbeddingResponse -from ..utils import SnowflakeException, SnowflakeBaseConfig +from ..utils import SnowflakeBaseConfig, SnowflakeException class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): @@ -18,12 +16,12 @@ class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = self._get_api_base(api_base, optional_params) @@ -44,7 +42,7 @@ class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -61,7 +59,5 @@ class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): returned_response._hidden_params["model"] = model return returned_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return SnowflakeException(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/snowflake/utils.py b/litellm/llms/snowflake/utils.py index 4f79006f6f8..3a1657b01e1 100644 --- a/litellm/llms/snowflake/utils.py +++ b/litellm/llms/snowflake/utils.py @@ -1,9 +1,9 @@ import re -from typing import TYPE_CHECKING, Any, List, Optional, Tuple +from typing import TYPE_CHECKING, Any +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues -from litellm.llms.base_llm.chat.transformation import BaseLLMException if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -16,11 +16,9 @@ else: class SnowflakeException(BaseLLMException): """Snowflake AI Endpoints exception handling class""" - pass - class SnowflakeBaseConfig: - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return [ "temperature", "max_tokens", @@ -76,11 +74,11 @@ class SnowflakeBaseConfig: self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Return headers to use for Snowflake completion request @@ -116,7 +114,7 @@ class SnowflakeBaseConfig: return headers def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: dynamic_api_key = api_key or get_secret_str("SNOWFLAKE_JWT") return api_base, dynamic_api_key diff --git a/litellm/llms/soniox/audio_transcription/handler.py b/litellm/llms/soniox/audio_transcription/handler.py index d25b3022a1b..f2b83e30251 100644 --- a/litellm/llms/soniox/audio_transcription/handler.py +++ b/litellm/llms/soniox/audio_transcription/handler.py @@ -22,11 +22,6 @@ from collections.abc import Coroutine from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Union, ) import httpx @@ -75,20 +70,20 @@ class SonioxAudioTranscriptionHandler: def audio_transcriptions( self, model: str, - audio_file: Optional[FileTypes], + audio_file: FileTypes | None, optional_params: dict, litellm_params: dict, model_response: TranscriptionResponse, timeout: float, max_retries: int, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: str | None, + api_base: str | None, + client: HTTPHandler | AsyncHTTPHandler | None = None, atranscription: bool = False, - headers: Optional[Dict[str, Any]] = None, - provider_config: Optional[SonioxAudioTranscriptionConfig] = None, - ) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: + headers: dict[str, Any] | None = None, + provider_config: SonioxAudioTranscriptionConfig | None = None, + ) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]: """Sync/async dispatch for Soniox transcription requests. Note: ``max_retries`` is accepted for signature compatibility with @@ -136,18 +131,18 @@ class SonioxAudioTranscriptionHandler: def _prepare( self, - audio_file: Optional[FileTypes], + audio_file: FileTypes | None, optional_params: dict, litellm_params: dict, - api_key: Optional[str], - api_base: Optional[str], + api_key: str | None, + api_base: str | None, provider_config: SonioxAudioTranscriptionConfig, - headers: Dict[str, Any], - ) -> Tuple[ - Dict[str, str], # auth headers + headers: dict[str, Any], + ) -> tuple[ + dict[str, str], # auth headers str, # api_base (no trailing slash) - Dict[str, Any], # body for POST /v1/transcriptions (without file_id/audio_url) - Dict[str, Any], # handler-only options (poll interval, cleanup, ...) + dict[str, Any], # body for POST /v1/transcriptions (without file_id/audio_url) + dict[str, Any], # handler-only options (poll interval, cleanup, ...) ]: # Validate env -> auth headers. auth_headers = provider_config.validate_environment( @@ -175,7 +170,7 @@ class SonioxAudioTranscriptionHandler: max_attempts = SONIOX_DEFAULT_MAX_POLL_ATTEMPTS cleanup_raw = params.pop("soniox_cleanup", SONIOX_DEFAULT_CLEANUP) if cleanup_raw is None: - cleanup: List[str] = [] + cleanup: list[str] = [] elif isinstance(cleanup_raw, str): cleanup = [cleanup_raw] else: @@ -192,7 +187,7 @@ class SonioxAudioTranscriptionHandler: clamped_poll_interval = max(SONIOX_MIN_POLL_INTERVAL, min(poll_interval, SONIOX_MAX_POLL_INTERVAL)) clamped_max_attempts = max(1, min(max_attempts, SONIOX_MAX_POLL_ATTEMPTS)) - handler_opts: Dict[str, Any] = { + handler_opts: dict[str, Any] = { "poll_interval": clamped_poll_interval, "max_attempts": clamped_max_attempts, "cleanup": cleanup, @@ -214,10 +209,10 @@ class SonioxAudioTranscriptionHandler: self, model: str, optional_params: dict, - handler_opts: Dict[str, Any], - file_id: Optional[str], - ) -> Dict[str, Any]: - body: Dict[str, Any] = {"model": model} + handler_opts: dict[str, Any], + file_id: str | None, + ) -> dict[str, Any]: + body: dict[str, Any] = {"model": model} # Soniox-native passthrough fields for key, value in optional_params.items(): if value is None: @@ -232,7 +227,7 @@ class SonioxAudioTranscriptionHandler: return body @staticmethod - def _redact_body_for_logging(body: Dict[str, Any]) -> Dict[str, Any]: + def _redact_body_for_logging(body: dict[str, Any]) -> dict[str, Any]: """Return a shallow copy of ``body`` with secret fields redacted. Soniox's create-transcription body can include @@ -254,9 +249,9 @@ class SonioxAudioTranscriptionHandler: @staticmethod def _safe_log_pre_call( logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, api_base: str, - body: Dict[str, Any], + body: dict[str, Any], ) -> None: try: logging_obj.pre_call( @@ -276,9 +271,9 @@ class SonioxAudioTranscriptionHandler: @staticmethod def _safe_log_post_call( logging_obj: LiteLLMLoggingObj, - audio_file: Optional[FileTypes], - api_key: Optional[str], - body: Dict[str, Any], + audio_file: FileTypes | None, + api_key: str | None, + body: dict[str, Any], original_response: Any, ) -> None: try: @@ -318,16 +313,16 @@ class SonioxAudioTranscriptionHandler: def _sync_audio_transcriptions( self, model: str, - audio_file: Optional[FileTypes], + audio_file: FileTypes | None, optional_params: dict, litellm_params: dict, model_response: TranscriptionResponse, timeout: float, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - client: Optional[HTTPHandler], - headers: Dict[str, Any], + api_key: str | None, + api_base: str | None, + client: HTTPHandler | None, + headers: dict[str, Any], provider_config: SonioxAudioTranscriptionConfig, ) -> TranscriptionResponse: auth_headers, base_url, opt_params, handler_opts = self._prepare( @@ -351,8 +346,8 @@ class SonioxAudioTranscriptionHandler: ) file_id = handler_opts.get("file_id") - uploaded_file_id: Optional[str] = None - transcription_id: Optional[str] = None + uploaded_file_id: str | None = None + transcription_id: str | None = None try: if not file_id and not handler_opts.get("audio_url"): @@ -442,9 +437,9 @@ class SonioxAudioTranscriptionHandler: self, http_client: HTTPHandler, base_url: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], audio_file: FileTypes, - filename_override: Optional[str], + filename_override: str | None, timeout: float, provider_config: SonioxAudioTranscriptionConfig, ) -> str: @@ -468,13 +463,13 @@ class SonioxAudioTranscriptionHandler: self, http_client: HTTPHandler, base_url: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], transcription_id: str, poll_interval: float, max_attempts: int, timeout: float, provider_config: SonioxAudioTranscriptionConfig, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: for _ in range(max_attempts): resp = http_client.get( url=f"{base_url}/v1/transcriptions/{transcription_id}", @@ -509,10 +504,10 @@ class SonioxAudioTranscriptionHandler: self, http_client: HTTPHandler, base_url: str, - auth_headers: Dict[str, str], - cleanup: List[str], - file_id_to_cleanup: Optional[str], - transcription_id: Optional[str], + auth_headers: dict[str, str], + cleanup: list[str], + file_id_to_cleanup: str | None, + transcription_id: str | None, timeout: float, ) -> None: if not cleanup: @@ -547,16 +542,16 @@ class SonioxAudioTranscriptionHandler: async def _async_audio_transcriptions( self, model: str, - audio_file: Optional[FileTypes], + audio_file: FileTypes | None, optional_params: dict, litellm_params: dict, model_response: TranscriptionResponse, timeout: float, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - api_base: Optional[str], - client: Optional[AsyncHTTPHandler], - headers: Dict[str, Any], + api_key: str | None, + api_base: str | None, + client: AsyncHTTPHandler | None, + headers: dict[str, Any], provider_config: SonioxAudioTranscriptionConfig, ) -> TranscriptionResponse: import litellm @@ -583,8 +578,8 @@ class SonioxAudioTranscriptionHandler: ) file_id = handler_opts.get("file_id") - uploaded_file_id: Optional[str] = None - transcription_id: Optional[str] = None + uploaded_file_id: str | None = None + transcription_id: str | None = None try: if not file_id and not handler_opts.get("audio_url"): @@ -674,9 +669,9 @@ class SonioxAudioTranscriptionHandler: self, http_client: AsyncHTTPHandler, base_url: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], audio_file: FileTypes, - filename_override: Optional[str], + filename_override: str | None, timeout: float, provider_config: SonioxAudioTranscriptionConfig, ) -> str: @@ -699,13 +694,13 @@ class SonioxAudioTranscriptionHandler: self, http_client: AsyncHTTPHandler, base_url: str, - auth_headers: Dict[str, str], + auth_headers: dict[str, str], transcription_id: str, poll_interval: float, max_attempts: int, timeout: float, provider_config: SonioxAudioTranscriptionConfig, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: for _ in range(max_attempts): resp = await http_client.get( url=f"{base_url}/v1/transcriptions/{transcription_id}", @@ -740,10 +735,10 @@ class SonioxAudioTranscriptionHandler: self, http_client: AsyncHTTPHandler, base_url: str, - auth_headers: Dict[str, str], - cleanup: List[str], - file_id_to_cleanup: Optional[str], - transcription_id: Optional[str], + auth_headers: dict[str, str], + cleanup: list[str], + file_id_to_cleanup: str | None, + transcription_id: str | None, timeout: float, ) -> None: if not cleanup: diff --git a/litellm/llms/soniox/audio_transcription/transformation.py b/litellm/llms/soniox/audio_transcription/transformation.py index 7160d2548df..0a42528fb75 100644 --- a/litellm/llms/soniox/audio_transcription/transformation.py +++ b/litellm/llms/soniox/audio_transcription/transformation.py @@ -9,7 +9,7 @@ async API requires multiple HTTP calls and does not fit the single-request contract of `base_llm_http_handler.audio_transcriptions`. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any from httpx import Headers, Response @@ -34,7 +34,7 @@ from litellm.types.utils import FileTypes, TranscriptionResponse # Soniox-native kwargs the user can pass through `litellm.transcription(..., **kwargs)` # in addition to the standard OpenAI params. -SONIOX_PASSTHROUGH_PARAMS: List[str] = [ +SONIOX_PASSTHROUGH_PARAMS: list[str] = [ "language_hints", "language_hints_strict", "enable_language_identification", @@ -50,7 +50,7 @@ SONIOX_PASSTHROUGH_PARAMS: List[str] = [ ] # Handler-only kwargs (consumed by the handler, not sent to Soniox). -SONIOX_HANDLER_ONLY_PARAMS: List[str] = [ +SONIOX_HANDLER_ONLY_PARAMS: list[str] = [ "soniox_polling_interval", "soniox_max_polling_attempts", "soniox_cleanup", @@ -61,7 +61,7 @@ SONIOX_HANDLER_ONLY_PARAMS: List[str] = [ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): """Configuration for Soniox async speech-to-text transcription.""" - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: # `language` is mapped onto Soniox's `language_hints`. # `response_format` is handled by LiteLLM (Soniox doesn't support # SRT/VTT natively but we synthesize them from token timestamps). @@ -94,18 +94,18 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return SonioxException(message=error_message, status_code=status_code, headers=headers) def validate_environment( self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: resolved_key = get_soniox_api_key(api_key) if not resolved_key: @@ -118,7 +118,7 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): headers=None, ) - merged_headers: Dict[str, str] = { + merged_headers: dict[str, str] = { "Authorization": f"Bearer {resolved_key}", } if headers: @@ -127,12 +127,12 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: # The handler builds per-call URLs (uploads, create, poll, fetch, delete); # we just return the resolved base. @@ -152,7 +152,7 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): and for filling in `file_id`/`audio_url`. This method exists so the config can be exercised in isolation by unit tests. """ - body: Dict[str, Any] = {"model": model} + body: dict[str, Any] = {"model": model} for key in SONIOX_PASSTHROUGH_PARAMS: value = optional_params.get(key) @@ -164,7 +164,7 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def transform_audio_transcription_response( self, raw_response: Response, - model_response: Optional[TranscriptionResponse] = None, + model_response: TranscriptionResponse | None = None, ) -> TranscriptionResponse: """ Build a TranscriptionResponse from a Soniox transcript payload. @@ -187,13 +187,13 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def _build_response_from_payload( self, - payload: Dict[str, Any], - model_response: Optional[TranscriptionResponse] = None, - response_format: Optional[str] = None, + payload: dict[str, Any], + model_response: TranscriptionResponse | None = None, + response_format: str | None = None, ) -> TranscriptionResponse: """Shared response-building logic (also used by the handler).""" - transcription_meta: Dict[str, Any] = {} - transcript: Dict[str, Any] + transcription_meta: dict[str, Any] = {} + transcript: dict[str, Any] if isinstance(payload, dict) and "transcript" in payload: transcription_meta = payload.get("transcription") or {} @@ -201,7 +201,7 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): else: transcript = payload if isinstance(payload, dict) else {} - tokens: List[Dict[str, Any]] = transcript.get("tokens") or [] + tokens: list[dict[str, Any]] = transcript.get("tokens") or [] # Decide what to put in `text` based on response_format: # - "srt": render tokens as SRT subtitles (synthesized from timestamps) @@ -247,9 +247,9 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): # For verbose_json, include word-level timing from tokens. if response_format == "verbose_json" and tokens: - words: List[Dict[str, Any]] = [] + words: list[dict[str, Any]] = [] for token in tokens: - word_entry: Dict[str, Any] = {"word": token.get("text", "")} + word_entry: dict[str, Any] = {"word": token.get("text", "")} if token.get("start_ms") is not None: word_entry["start"] = float(token["start_ms"]) / 1000.0 if token.get("end_ms") is not None: diff --git a/litellm/llms/soniox/common_utils.py b/litellm/llms/soniox/common_utils.py index 76aa25522d0..2f951b352b8 100644 --- a/litellm/llms/soniox/common_utils.py +++ b/litellm/llms/soniox/common_utils.py @@ -2,7 +2,7 @@ Shared utilities for the Soniox provider (https://soniox.com). """ -from typing import Any, Dict, List, Optional +from typing import Any from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -33,22 +33,20 @@ SONIOX_MAX_POLL_ATTEMPTS: int = 6000 # Default cleanup behaviour: delete both the uploaded file (if any) and the # transcription record after the transcript has been fetched. -SONIOX_DEFAULT_CLEANUP: List[str] = ["file", "transcription"] +SONIOX_DEFAULT_CLEANUP: list[str] = ["file", "transcription"] # Body fields that may carry secrets and must be redacted before being # forwarded to logging callbacks. Soniox accepts a webhook auth header value # alongside the create-transcription request; that value lets the recipient # authenticate webhook callbacks and must not leak into observability sinks. -SONIOX_SECRET_FIELDS: List[str] = ["webhook_auth_header_value"] +SONIOX_SECRET_FIELDS: list[str] = ["webhook_auth_header_value"] class SonioxException(BaseLLMException): """Provider-specific exception class for Soniox.""" - pass - -def get_soniox_api_key(api_key: Optional[str] = None) -> Optional[str]: +def get_soniox_api_key(api_key: str | None = None) -> str | None: """Resolve the Soniox API key from arg or env var.""" # Local import to avoid a circular import: litellm.secret_managers.main # imports from litellm at top-level. @@ -57,7 +55,7 @@ def get_soniox_api_key(api_key: Optional[str] = None) -> Optional[str]: return api_key or get_secret_str("SONIOX_API_KEY") -def get_soniox_api_base(api_base: Optional[str] = None) -> str: +def get_soniox_api_base(api_base: str | None = None) -> str: """Resolve the Soniox API base URL from arg or env var (defaults to public API).""" from litellm.secret_managers.main import get_secret_str @@ -65,7 +63,7 @@ def get_soniox_api_base(api_base: Optional[str] = None) -> str: return base.rstrip("/") -def render_soniox_tokens(tokens: List[Dict[str, Any]]) -> str: +def render_soniox_tokens(tokens: list[dict[str, Any]]) -> str: """ Render a list of Soniox tokens to a readable transcript string. @@ -81,9 +79,9 @@ def render_soniox_tokens(tokens: List[Dict[str, Any]]) -> str: if not tokens: return "" - text_parts: List[str] = [] - current_speaker: Optional[Any] = None - current_language: Optional[Any] = None + text_parts: list[str] = [] + current_speaker: Any | None = None + current_language: Any | None = None for token in tokens: text = token.get("text", "") @@ -124,8 +122,7 @@ _CUE_MAX_DURATION_MS: int = 5000 def _format_timestamp_srt(ms: int) -> str: """Format milliseconds as SRT timestamp: HH:MM:SS,mmm""" - if ms < 0: - ms = 0 + ms = max(ms, 0) hours = ms // 3_600_000 ms %= 3_600_000 minutes = ms // 60_000 @@ -137,8 +134,7 @@ def _format_timestamp_srt(ms: int) -> str: def _format_timestamp_vtt(ms: int) -> str: """Format milliseconds as VTT timestamp: HH:MM:SS.mmm""" - if ms < 0: - ms = 0 + ms = max(ms, 0) hours = ms // 3_600_000 ms %= 3_600_000 minutes = ms // 60_000 @@ -149,8 +145,8 @@ def _format_timestamp_vtt(ms: int) -> str: def _group_tokens_into_cues( - tokens: List[Dict[str, Any]], -) -> List[Dict[str, Any]]: + tokens: list[dict[str, Any]], +) -> list[dict[str, Any]]: """ Group Soniox tokens into subtitle cues. @@ -165,11 +161,11 @@ def _group_tokens_into_cues( - A new cue starts when the speaker changes (if diarization is on). - Tokens without timestamps are appended to the current cue. """ - cues: List[Dict[str, Any]] = [] - current_tokens: List[str] = [] - current_start: Optional[int] = None - current_end: Optional[int] = None - current_speaker: Optional[Any] = None + cues: list[dict[str, Any]] = [] + current_tokens: list[str] = [] + current_start: int | None = None + current_end: int | None = None + current_speaker: Any | None = None def _flush() -> None: if current_tokens and current_start is not None: @@ -205,9 +201,12 @@ def _group_tokens_into_cues( # Duration or token count exceeded -> flush should_break = False - if len(current_tokens) >= _CUE_MAX_TOKENS: - should_break = True - elif current_start is not None and start_ms is not None and (start_ms - current_start) >= _CUE_MAX_DURATION_MS: + if ( + len(current_tokens) >= _CUE_MAX_TOKENS + or current_start is not None + and start_ms is not None + and (start_ms - current_start) >= _CUE_MAX_DURATION_MS + ): should_break = True if should_break: @@ -227,7 +226,7 @@ def _group_tokens_into_cues( return cues -def render_soniox_tokens_as_srt(tokens: List[Dict[str, Any]]) -> str: +def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str: """ Render Soniox tokens as SRT (SubRip) subtitle format. @@ -237,7 +236,7 @@ def render_soniox_tokens_as_srt(tokens: List[Dict[str, Any]]) -> str: if not cues: return "" - lines: List[str] = [] + lines: list[str] = [] for idx, cue in enumerate(cues, start=1): start = _format_timestamp_srt(cue["start_ms"]) end = _format_timestamp_srt(cue["end_ms"]) @@ -249,7 +248,7 @@ def render_soniox_tokens_as_srt(tokens: List[Dict[str, Any]]) -> str: return "\n".join(lines) -def render_soniox_tokens_as_vtt(tokens: List[Dict[str, Any]]) -> str: +def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str: """ Render Soniox tokens as WebVTT subtitle format. @@ -257,7 +256,7 @@ def render_soniox_tokens_as_vtt(tokens: List[Dict[str, Any]]) -> str: """ cues = _group_tokens_into_cues(tokens) - lines: List[str] = ["WEBVTT", ""] + lines: list[str] = ["WEBVTT", ""] for cue in cues: start = _format_timestamp_vtt(cue["start_ms"]) end = _format_timestamp_vtt(cue["end_ms"]) diff --git a/litellm/llms/stability/image_edit/transformations.py b/litellm/llms/stability/image_edit/transformations.py index 05a200246a1..b6e5f213e4c 100644 --- a/litellm/llms/stability/image_edit/transformations.py +++ b/litellm/llms/stability/image_edit/transformations.py @@ -6,7 +6,7 @@ Handles transformation between OpenAI-compatible format and Stability AI API for API Reference: https://platform.stability.ai/docs/api-reference """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import httpx from httpx._types import RequestFiles @@ -40,7 +40,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): DEFAULT_BASE_URL: str = "https://api.stability.ai" - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Return list of OpenAI params supported by Stability AI. @@ -58,7 +58,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map OpenAI parameters to Stability AI parameters. @@ -74,7 +74,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): } # Create a copy to not mutate original - convert TypedDict to regular dict - mapped_params: Dict[str, Any] = dict(image_edit_optional_params) + mapped_params: dict[str, Any] = dict(image_edit_optional_params) for k, v in image_edit_optional_params.items(): if k in param_mapping: @@ -102,8 +102,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): # Remove OpenAI params that have been mapped unless they're in stability for mapped in ["size", "n", "response_format"]: - if mapped in mapped_params: - del mapped_params[mapped] + mapped_params.pop(mapped, None) return mapped_params @@ -113,8 +112,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): """ # Remove "stability/" prefix if present model_name = model.lower() - if model_name.startswith("stability/"): - model_name = model_name[10:] # Remove "stability/" prefix + model_name = model_name.removeprefix("stability/") # Remove "stability/" prefix # Check if model is in our mapping for key, endpoint in STABILITY_EDIT_ENDPOINTS.items(): @@ -127,7 +125,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -148,14 +146,14 @@ class StabilityImageEditConfig(BaseImageEditConfig): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Stability AI. """ - final_api_key: Optional[str] = api_key or get_secret_str("STABILITY_API_KEY") + final_api_key: str | None = api_key or get_secret_str("STABILITY_API_KEY") if not final_api_key: raise ValueError( @@ -169,12 +167,12 @@ class StabilityImageEditConfig(BaseImageEditConfig): def transform_image_edit_request( self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles]: + ) -> tuple[dict, RequestFiles]: """ Transform OpenAI-style request to Stability AI request format. @@ -184,7 +182,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): # Build Stability request # Populate multipart form-data as separate text fields (data) and files. # Stability expects prompt/output_format/etc. as normal form fields, not file parts. - data: Dict[str, Any] = { + data: dict[str, Any] = { "output_format": "png", # Default to PNG } @@ -193,7 +191,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): data["prompt"] = prompt # Handle image parameter - could be a single file or list image_file = image[0] if isinstance(image, list) else image # type: ignore - files: Dict[str, Any] = {} + files: dict[str, Any] = {} if image is not None: image_file = image[0] if isinstance(image, list) else image # type: ignore files["image"] = image_file @@ -251,8 +249,8 @@ class StabilityImageEditConfig(BaseImageEditConfig): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Stability AI response to OpenAI-compatible ImageResponse. diff --git a/litellm/llms/stability/image_generation/transformation.py b/litellm/llms/stability/image_generation/transformation.py index a5b18b0f325..2faf781f7ea 100644 --- a/litellm/llms/stability/image_generation/transformation.py +++ b/litellm/llms/stability/image_generation/transformation.py @@ -6,7 +6,7 @@ Handles transformation between OpenAI-compatible format and Stability AI API for API Reference: https://platform.stability.ai/docs/api-reference """ -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -45,7 +45,7 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://api.stability.ai" - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Return list of OpenAI params supported by Stability AI. @@ -104,8 +104,7 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): """ # Remove "stability/" prefix if present model_name = model.lower() - if model_name.startswith("stability/"): - model_name = model_name[10:] # Remove "stability/" prefix + model_name = model_name.removeprefix("stability/") # Remove "stability/" prefix # Check if model is in our mapping for key, endpoint in STABILITY_GENERATION_MODELS.items(): @@ -117,12 +116,12 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the Stability AI API request. @@ -137,16 +136,16 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Stability AI. """ - final_api_key: Optional[str] = api_key or get_secret_str("STABILITY_API_KEY") + final_api_key: str | None = api_key or get_secret_str("STABILITY_API_KEY") if not final_api_key: raise ValueError( @@ -207,8 +206,8 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Stability AI response to OpenAI-compatible ImageResponse. diff --git a/litellm/llms/tavily/search/transformation.py b/litellm/llms/tavily/search/transformation.py index 51b897d93b2..cf72744b504 100644 --- a/litellm/llms/tavily/search/transformation.py +++ b/litellm/llms/tavily/search/transformation.py @@ -4,7 +4,7 @@ Calls Tavily's /search endpoint to search the web. Tavily API Reference: https://docs.tavily.com/documentation/api-reference/endpoint/search """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -30,12 +30,12 @@ class TavilySearchRequest(_TavilySearchRequestRequired, total=False): """ max_results: int # Optional - maximum number of results (0-20), default 5 - include_domains: List[str] # Optional - list of domains to include (max 300) - exclude_domains: List[str] # Optional - list of domains to exclude (max 150) + include_domains: list[str] # Optional - list of domains to include (max 300) + exclude_domains: list[str] # Optional - list of domains to exclude (max 150) topic: str # Optional - category of search ('general', 'news', 'finance'), default 'general' search_depth: str # Optional - depth of search ('basic', 'advanced'), default 'basic' - include_answer: Union[bool, str] # Optional - include LLM-generated answer - include_raw_content: Union[bool, str] # Optional - include raw HTML content + include_answer: bool | str # Optional - include LLM-generated answer + include_raw_content: bool | str # Optional - include raw HTML content include_images: bool # Optional - perform image search include_image_descriptions: bool # Optional - add descriptions for images include_favicon: bool # Optional - include favicon URL @@ -54,11 +54,11 @@ class TavilySearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers. """ @@ -77,9 +77,9 @@ class TavilySearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -95,10 +95,10 @@ class TavilySearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to Tavily API format. diff --git a/litellm/llms/tencent/chat/transformation.py b/litellm/llms/tencent/chat/transformation.py index 4dea0c4b8c7..ff0dc40d85e 100644 --- a/litellm/llms/tencent/chat/transformation.py +++ b/litellm/llms/tencent/chat/transformation.py @@ -3,8 +3,6 @@ Translates from OpenAI's `/v1/chat/completions` to Tencent TokenHub's OpenAI-compatible endpoint. """ -from typing import Optional - from litellm.secret_managers.main import get_secret_str from litellm.utils import supports_reasoning @@ -39,20 +37,20 @@ class TencentChatConfig(OpenAIGPTConfig): return optional_params def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("TENCENT_API_BASE") or "https://tokenhub-intl.tencentcloudmaas.com/v1" dynamic_api_key = api_key or get_secret_str("TENCENT_API_KEY") return api_base, dynamic_api_key def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if not api_base: api_base = "https://tokenhub-intl.tencentcloudmaas.com/v1" diff --git a/litellm/llms/tencent/messages/transformation.py b/litellm/llms/tencent/messages/transformation.py index e0f13aa9ca4..f1d9ee966ff 100644 --- a/litellm/llms/tencent/messages/transformation.py +++ b/litellm/llms/tencent/messages/transformation.py @@ -5,7 +5,7 @@ Tencent TokenHub exposes an Anthropic-compatible Messages API endpoint alongside its standard OpenAI-compatible chat completions endpoint. """ -from typing import Any, Optional +from typing import Any import litellm from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -24,18 +24,18 @@ class TencentAnthropicMessagesConfig(AnthropicMessagesConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "tencent" def should_strip_billing_metadata(self) -> bool: return True @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("TENCENT_API_KEY") or litellm.api_key @staticmethod - def get_api_base(api_base: Optional[str] = None) -> str: + def get_api_base(api_base: str | None = None) -> str: return ( api_base or get_secret_str("TENCENT_ANTHROPIC_API_BASE") @@ -50,9 +50,9 @@ class TencentAnthropicMessagesConfig(AnthropicMessagesConfig): messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: return super().validate_anthropic_messages_environment( headers=headers, model=model, @@ -65,12 +65,12 @@ class TencentAnthropicMessagesConfig(AnthropicMessagesConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: base_url = self.get_api_base(api_base=api_base).rstrip("/") diff --git a/litellm/llms/together_ai/chat.py b/litellm/llms/together_ai/chat.py index a78b023f287..5920f02a44e 100644 --- a/litellm/llms/together_ai/chat.py +++ b/litellm/llms/together_ai/chat.py @@ -6,10 +6,8 @@ Calls done in OpenAI/openai.py as TogetherAI is openai-compatible. Docs: https://docs.together.ai/reference/completions-1 """ -from typing import Optional - -from litellm.utils import supports_function_calling from litellm._logging import verbose_logger +from litellm.utils import supports_function_calling from ..openai.chat.gpt_transformation import OpenAIGPTConfig @@ -27,12 +25,11 @@ class TogetherAIConfig(OpenAIGPTConfig): # into this method for together_ai models, creating a recursion that # only terminates when Python's recursion limit or the "not mapped" # exception in _get_model_info_helper is hit (~332 deep calls). - supports_fc: Optional[bool] = None + supports_fc: bool | None = None try: supports_fc = supports_function_calling(model, custom_llm_provider="together_ai") except Exception as e: verbose_logger.debug(f"Error getting supported openai params: {e}") - pass optional_params = super().get_supported_openai_params(model) if supports_fc is not True: diff --git a/litellm/llms/together_ai/completion/transformation.py b/litellm/llms/together_ai/completion/transformation.py index 6e0b862c183..d2017def06d 100644 --- a/litellm/llms/together_ai/completion/transformation.py +++ b/litellm/llms/together_ai/completion/transformation.py @@ -6,7 +6,7 @@ Calls done in OpenAI/openai.py as TogetherAI is openai-compatible. Docs: https://docs.together.ai/reference/completions-1 """ -from typing import List, Union, cast +from typing import cast from litellm.llms.openai.completion.utils import is_tokens_or_list_of_tokens from litellm.types.llms.openai import ( @@ -22,7 +22,7 @@ from ...openai.completion.utils import _transform_prompt class TogetherAITextCompletionConfig(OpenAITextCompletionConfig): def _transform_prompt( self, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], ) -> AllPromptValues: """ TogetherAI expects a string prompt. @@ -43,7 +43,7 @@ class TogetherAITextCompletionConfig(OpenAITextCompletionConfig): def transform_text_completion_request( self, model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], optional_params: dict, headers: dict, ) -> dict: diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index 08acdead386..0a18db317ac 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -4,7 +4,7 @@ Re rank api LiteLLM supports the re rank API format, no paramter transformation occurs """ -from typing import Any, Dict, List, Optional, Union +from typing import Any import litellm from litellm.llms.base import BaseLLM @@ -22,12 +22,12 @@ class TogetherAIRerank(BaseLLM): model: str, api_key: str, query: str, - documents: List[Union[str, Dict[str, Any]]], - top_n: Optional[int] = None, - rank_fields: Optional[List[str]] = None, - return_documents: Optional[bool] = True, - max_chunks_per_doc: Optional[int] = None, - _is_async: Optional[bool] = False, + documents: list[str | dict[str, Any]], + top_n: int | None = None, + rank_fields: list[str] | None = None, + return_documents: bool | None = True, + max_chunks_per_doc: int | None = None, + _is_async: bool | None = False, ) -> RerankResponse: client = _get_httpx_client() @@ -67,7 +67,7 @@ class TogetherAIRerank(BaseLLM): async def async_rerank( # New async method self, - request_data_dict: Dict[str, Any], + request_data_dict: dict[str, Any], api_key: str, ) -> RerankResponse: client = get_async_httpx_client(llm_provider=litellm.LlmProviders.TOGETHER_AI) # Use async client diff --git a/litellm/llms/together_ai/rerank/transformation.py b/litellm/llms/together_ai/rerank/transformation.py index 3610a5853ac..1876b809265 100644 --- a/litellm/llms/together_ai/rerank/transformation.py +++ b/litellm/llms/together_ai/rerank/transformation.py @@ -5,8 +5,6 @@ Why separate file? Make it easy to see how transformation works """ from litellm._uuid import uuid -from typing import List, Optional - from litellm.types.rerank import ( RerankBilledUnits, RerankResponse, @@ -23,12 +21,12 @@ class TogetherAIRerankConfig: _tokens = RerankTokens(**response.get("usage", {})) rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) - _results: Optional[List[dict]] = response.get("results") + _results: list[dict] | None = response.get("results") if _results is None: raise ValueError(f"No results found in the response={response}") - rerank_results: List[RerankResponseResult] = [] + rerank_results: list[RerankResponseResult] = [] for result in _results: # Validate required fields exist diff --git a/litellm/llms/topaz/common_utils.py b/litellm/llms/topaz/common_utils.py index 27603b3b401..784987c3adc 100644 --- a/litellm/llms/topaz/common_utils.py +++ b/litellm/llms/topaz/common_utils.py @@ -1,5 +1,3 @@ -from typing import List, Optional - from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues @@ -16,11 +14,11 @@ class TopazModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: raise ValueError("API key is required for Topaz image variations. Set via `TOPAZ_API_KEY` or `api_key=..`") @@ -30,7 +28,7 @@ class TopazModelInfo(BaseLLMModelInfo): "X-API-Key": api_key, } - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: return [ "topaz/Standard V2", "topaz/Low Resolution V2", @@ -40,11 +38,11 @@ class TopazModelInfo(BaseLLMModelInfo): ] @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return api_key or get_secret_str("TOPAZ_API_KEY") @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or get_secret_str("TOPAZ_API_BASE") or "https://api.topazlabs.com" @staticmethod diff --git a/litellm/llms/topaz/image_variations/transformation.py b/litellm/llms/topaz/image_variations/transformation.py index d495f1c28ab..bc2fffe88aa 100644 --- a/litellm/llms/topaz/image_variations/transformation.py +++ b/litellm/llms/topaz/image_variations/transformation.py @@ -2,7 +2,7 @@ import base64 import time from collections.abc import Mapping from io import BytesIO -from typing import Any, List, Optional, Tuple, Union +from typing import Any from aiohttp import ClientResponse from httpx import Headers, Response @@ -24,17 +24,17 @@ from ..common_utils import TopazException, TopazModelInfo class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): - def get_supported_openai_params(self, model: str) -> List[OpenAIImageVariationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageVariationOptionalParams]: return ["response_format", "size"] def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: api_base = api_base or "https://api.topazlabs.com" return f"{api_base}/image/v1/enhance" @@ -59,14 +59,14 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): def prepare_file_tuple( self, file_data: FileTypes, - ) -> Tuple[str, Optional[FileTypes], str, Mapping[str, str]]: + ) -> tuple[str, FileTypes | None, str, Mapping[str, str]]: """ Convert various file input formats to a consistent tuple format for HTTPX Returns: (filename, file_content, content_type, headers) """ # Default values filename = "image.png" - content: Optional[FileTypes] = None + content: FileTypes | None = None content_type = "image/png" headers: Mapping[str, str] = {} @@ -94,7 +94,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): def transform_request_image_variation( self, - model: Optional[str], + model: str | None, image: FileTypes, optional_params: dict, headers: dict, @@ -128,7 +128,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): async def async_transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: ClientResponse, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -137,7 +137,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: image_content = await raw_response.read() @@ -147,7 +147,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): def transform_response_image_variation( self, - model: Optional[str], + model: str | None, raw_response: Response, model_response: ImageResponse, logging_obj: LiteLLMLoggingObj, @@ -156,7 +156,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, + api_key: str | None = None, ) -> ImageResponse: image_content = raw_response.content @@ -164,7 +164,7 @@ class TopazImageVariationConfig(TopazModelInfo, BaseImageVariationConfig): return self._common_transform_response_image_variation(image_content, response_ms) - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return TopazException( status_code=status_code, message=error_message, diff --git a/litellm/llms/triton/common_utils.py b/litellm/llms/triton/common_utils.py index d5372eee00c..01702e3fccf 100644 --- a/litellm/llms/triton/common_utils.py +++ b/litellm/llms/triton/common_utils.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +8,6 @@ class TritonError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ) -> None: super().__init__(status_code=status_code, message=message, headers=headers) diff --git a/litellm/llms/triton/completion/transformation.py b/litellm/llms/triton/completion/transformation.py index dc1ed427c2d..7d977ff402c 100644 --- a/litellm/llms/triton/completion/transformation.py +++ b/litellm/llms/triton/completion/transformation.py @@ -4,7 +4,7 @@ Translates from OpenAI's `/v1/chat/completions` endpoint to Triton's `/generate` import json from collections.abc import AsyncIterator, Iterator -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal from httpx import Headers, Response @@ -36,31 +36,31 @@ class TritonConfig(BaseConfig): Handles routing between /infer and /generate triton completion llms """ - def get_error_class(self, error_message: str, status_code: int, headers: Union[Dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return TritonError(status_code=status_code, message=error_message, headers=headers) def validate_environment( self, - headers: Dict, + headers: dict, model: str, - messages: List[AllMessageValues], - optional_params: Dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: return {"Content-Type": "application/json"} - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return ["max_tokens", "max_completion_tokens"] def map_openai_params( self, - non_default_params: Dict, - optional_params: Dict, + non_default_params: dict, + optional_params: dict, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: for param, value in non_default_params.items(): if param == "max_tokens" or param == "max_completion_tokens": optional_params[param] = value @@ -68,12 +68,12 @@ class TritonConfig(BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: raise ValueError("api_base is required") @@ -88,13 +88,13 @@ class TritonConfig(BaseConfig): raw_response: Response, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: api_base = litellm_params.get("api_base", "") llm_type = self._get_triton_llm_type(api_base) @@ -131,7 +131,7 @@ class TritonConfig(BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -166,9 +166,9 @@ class TritonConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return TritonResponseIterator( streaming_response=streaming_response, @@ -185,14 +185,14 @@ class TritonGenerateConfig(TritonConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, ) -> dict: inference_params = optional_params.copy() stream = inference_params.pop("stream", False) - data_for_triton: Dict[str, Any] = { + data_for_triton: dict[str, Any] = { "text_input": prompt_factory(model=model, messages=messages), "parameters": { "max_tokens": int(optional_params.get("max_tokens", DEFAULT_MAX_TOKENS_FOR_TRITON)), @@ -208,13 +208,13 @@ class TritonGenerateConfig(TritonConfig): raw_response: Response, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: raw_response_json = raw_response.json() @@ -233,7 +233,7 @@ class TritonInferConfig(TritonConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -273,13 +273,13 @@ class TritonInferConfig(TritonConfig): raw_response: Response, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: try: raw_response_json = raw_response.json() @@ -287,7 +287,7 @@ class TritonInferConfig(TritonConfig): raise TritonError(message=raw_response.text, status_code=raw_response.status_code) _triton_response_data = raw_response_json["outputs"][0]["data"] - triton_response_data: Optional[str] = None + triton_response_data: str | None = None if isinstance(_triton_response_data, list): triton_response_data = "".join(_triton_response_data) else: @@ -307,10 +307,10 @@ class TritonResponseIterator(BaseModelResponseIterator): def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: try: text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None + tool_use: ChatCompletionToolCallChunk | None = None is_finished = False finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None + usage: ChatCompletionUsageBlock | None = None provider_specific_fields = None index = int(chunk.get("index", 0)) diff --git a/litellm/llms/triton/embedding/transformation.py b/litellm/llms/triton/embedding/transformation.py index 2426520e630..9969e7d78d5 100644 --- a/litellm/llms/triton/embedding/transformation.py +++ b/litellm/llms/triton/embedding/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - import httpx from litellm.llms.base_llm.chat.transformation import AllMessageValues, BaseLLMException @@ -41,11 +39,11 @@ class TritonEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: return {} @@ -73,7 +71,7 @@ class TritonEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -107,7 +105,7 @@ class TritonEmbeddingConfig(BaseEmbeddingConfig): def _build_embedding_usage(self, model: str, request_data: dict) -> Usage: input_data = request_data.get("inputs", []) - input_text_values: List[str] = [] + input_text_values: list[str] = [] for item in input_data: if isinstance(item, dict) and item.get("name") == "input_text": data_values = item.get("data", []) @@ -130,13 +128,11 @@ class TritonEmbeddingConfig(BaseEmbeddingConfig): total_tokens=prompt_tokens, ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return TritonError(message=error_message, status_code=status_code, headers=headers) @staticmethod - def split_embedding_by_shape(data: List[float], shape: List[int]) -> List[List[float]]: + def split_embedding_by_shape(data: list[float], shape: list[int]) -> list[list[float]]: if len(shape) != 2: raise ValueError("Shape must be of length 2.") embedding_size = shape[1] diff --git a/litellm/llms/v0/chat/transformation.py b/litellm/llms/v0/chat/transformation.py index 5e029512471..c0683837dfc 100644 --- a/litellm/llms/v0/chat/transformation.py +++ b/litellm/llms/v0/chat/transformation.py @@ -2,8 +2,6 @@ Translate from OpenAI's `/v1/chat/completions` to v0's `/v1/chat/completions` """ -from typing import Optional, Tuple - from litellm.secret_managers.main import get_secret_str from ...openai_like.chat.transformation import OpenAILikeChatConfig @@ -15,12 +13,12 @@ class V0ChatConfig(OpenAILikeChatConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "v0" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: # v0 is openai compatible, we just need to set the api_base api_base = ( api_base or get_secret_str("V0_API_BASE") or "https://api.v0.dev/v1" # Default v0 API base URL diff --git a/litellm/llms/vercel_ai_gateway/chat/transformation.py b/litellm/llms/vercel_ai_gateway/chat/transformation.py index 1c2e29234e6..f87c4e18068 100644 --- a/litellm/llms/vercel_ai_gateway/chat/transformation.py +++ b/litellm/llms/vercel_ai_gateway/chat/transformation.py @@ -6,14 +6,12 @@ Calls done in OpenAI/openai.py as Vercel AI Gateway is openai-compatible. Docs: https://vercel.com/docs/ai-gateway """ -from typing import List, Optional, Tuple, Union - import httpx -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.types.llms.openai import AllMessageValues -from litellm.secret_managers.main import get_secret_str import litellm +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues from ...openai.chat.gpt_transformation import OpenAIGPTConfig from ..common_utils import VercelAIGatewayException @@ -21,7 +19,7 @@ from ..common_utils import VercelAIGatewayException class VercelAIGatewayConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "vercel_ai_gateway" def get_supported_openai_params(self, model: str) -> list: @@ -31,8 +29,8 @@ class VercelAIGatewayConfig(OpenAIGPTConfig): return base_params def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("VERCEL_AI_GATEWAY_API_BASE") or "https://ai-gateway.vercel.sh/v1" user_api_key = api_key or get_secret_str("VERCEL_AI_GATEWAY_API_KEY") or get_secret_str("VERCEL_OIDC_TOKEN") return api_base, user_api_key @@ -59,7 +57,7 @@ class VercelAIGatewayConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -72,16 +70,14 @@ class VercelAIGatewayConfig(OpenAIGPTConfig): """ return super().transform_request(model, messages, optional_params, litellm_params, headers) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VercelAIGatewayException( message=error_message, status_code=status_code, headers=headers, ) - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base, _ = self._get_openai_compatible_provider_info(api_base, api_key) if api_base is None: diff --git a/litellm/llms/vercel_ai_gateway/embedding/transformation.py b/litellm/llms/vercel_ai_gateway/embedding/transformation.py index e4036f415a9..a8a01fc0ed8 100644 --- a/litellm/llms/vercel_ai_gateway/embedding/transformation.py +++ b/litellm/llms/vercel_ai_gateway/embedding/transformation.py @@ -7,7 +7,7 @@ Vercel AI Gateway is OpenAI-compatible and supports embeddings via the /v1/embed Docs: https://vercel.com/docs/ai-gateway/openai-compat/embeddings """ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -41,8 +41,8 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig): messages: list, optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate environment and set up headers for Vercel AI Gateway API. @@ -65,12 +65,12 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Vercel AI Gateway Embedding API endpoint. @@ -112,7 +112,7 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, diff --git a/litellm/llms/vertex_ai/agent_engine/sse_iterator.py b/litellm/llms/vertex_ai/agent_engine/sse_iterator.py index d3e95f46be9..5f802e3a6e1 100644 --- a/litellm/llms/vertex_ai/agent_engine/sse_iterator.py +++ b/litellm/llms/vertex_ai/agent_engine/sse_iterator.py @@ -4,7 +4,7 @@ SSE Stream Iterator for Vertex AI Agent Engine. Handles Server-Sent Events (SSE) streaming responses from Vertex AI Reasoning Engines. """ -from typing import Any, Union +from typing import Any from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.types.llms.openai import ChatCompletionUsageBlock @@ -27,7 +27,7 @@ class VertexAgentEngineResponseIterator(BaseModelResponseIterator): def __init__(self, streaming_response: Any, sync_stream: bool) -> None: super().__init__(streaming_response=streaming_response, sync_stream=sync_stream) - def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]: + def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: """ Parse a Vertex Agent Engine response chunk into ModelResponseStream. diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index 20c86a25f82..e785e7ec28e 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -10,7 +10,7 @@ API Reference: """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Optional, Union, cast import httpx @@ -62,7 +62,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): BaseConfig.__init__(self, **kwargs) VertexBase.__init__(self) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """Vertex Agent Engine has limited OpenAI compatible params.""" return ["user"] @@ -79,7 +79,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): optional_params["user_id"] = non_default_params["user"] return optional_params - def _parse_model_string(self, model: str) -> Tuple[str, str]: + def _parse_model_string(self, model: str) -> tuple[str, str]: """ Parse model string to extract resource ID. @@ -89,8 +89,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): Returns: (resource_path, engine_id) """ # Remove 'agent_engine/' prefix if present - if model.startswith("agent_engine/"): - model = model[len("agent_engine/") :] + model = model.removeprefix("agent_engine/") # Check if it's a full resource path if model.startswith("projects/"): @@ -102,12 +101,12 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for the request. @@ -145,7 +144,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): self, optional_params: dict, litellm_params: dict, - ) -> Dict[str, str]: + ) -> dict[str, str]: """Get authentication headers using Google Cloud credentials.""" vertex_credentials = self.safe_get_vertex_ai_credentials(litellm_params) vertex_project = self.safe_get_vertex_ai_project(litellm_params) @@ -171,14 +170,14 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): # Generate a user ID return f"litellm-user-{str(uuid.uuid4())[:8]}" - def _get_session_id(self, optional_params: dict) -> Optional[str]: + def _get_session_id(self, optional_params: dict) -> str | None: """Get session ID if provided.""" return optional_params.get("session_id") def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -204,7 +203,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): session_id = self._get_session_id(optional_params) # Build the input - input_data: Dict[str, Any] = { + input_data: dict[str, Any] = { "message": prompt, "user_id": user_id, } @@ -227,11 +226,11 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """Validate environment and set up authentication headers.""" auth_headers = self._get_auth_headers(optional_params, litellm_params) @@ -256,7 +255,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): return "" - def _calculate_usage(self, model: str, messages: List[AllMessageValues], content: str) -> Optional[Usage]: + def _calculate_usage(self, model: str, messages: list[AllMessageValues], content: str) -> Usage | None: """Calculate token usage using LiteLLM's token counter.""" try: from litellm.utils import token_counter @@ -271,7 +270,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): total_tokens=total_tokens, ) except Exception as e: - verbose_logger.warning(f"Failed to calculate token usage: {str(e)}") + verbose_logger.warning(f"Failed to calculate token usage: {e!s}") return None def transform_response( @@ -281,12 +280,12 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform Vertex Agent Engine response to LiteLLM ModelResponse format. @@ -336,9 +335,9 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): return model_response except Exception as e: - verbose_logger.error(f"Error processing Vertex Agent Engine response: {str(e)}") + verbose_logger.error(f"Error processing Vertex Agent Engine response: {e!s}") raise VertexAgentEngineError( - message=f"Error processing response: {str(e)}", + message=f"Error processing response: {e!s}", status_code=raw_response.status_code, ) @@ -362,9 +361,9 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): headers: dict, data: dict, messages: list, - client: Optional[Union[HTTPHandler, "AsyncHTTPHandler"]] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for synchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( @@ -421,8 +420,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): data: dict, messages: list, client: Optional["AsyncHTTPHandler"] = None, - json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + json_mode: bool | None = None, + signed_json_body: bytes | None = None, ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for asynchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( @@ -482,16 +481,14 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): """Agent Engine does not allow passing `stream` in the request body.""" return False - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VertexAgentEngineError(status_code=status_code, message=error_message) def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """Agent Engine always returns SSE streams, so we use real streaming.""" return False diff --git a/litellm/llms/vertex_ai/batches/handler.py b/litellm/llms/vertex_ai/batches/handler.py index c8c951e1d3b..5e2d8d778f4 100644 --- a/litellm/llms/vertex_ai/batches/handler.py +++ b/litellm/llms/vertex_ai/batches/handler.py @@ -1,6 +1,6 @@ import json from collections.abc import Coroutine -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -35,13 +35,13 @@ class VertexAIBatchPrediction(VertexLLM): self, _is_async: bool, create_batch_data: CreateBatchRequest, - api_base: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - ) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: + api_base: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: sync_handler = _get_httpx_client() access_token, project_id = self._ensure_access_token( @@ -111,7 +111,7 @@ class VertexAIBatchPrediction(VertexLLM): self, vertex_batch_request: VertexAIBatchPredictionJob, api_base: str, - headers: Dict[str, str], + headers: dict[str, str], ) -> LiteLLMBatch: client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, @@ -153,14 +153,14 @@ class VertexAIBatchPrediction(VertexLLM): self, _is_async: bool, batch_id: str, - api_base: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - logging_obj: Optional[Any] = None, - ) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: + api_base: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + logging_obj: Any | None = None, + ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: sync_handler = _get_httpx_client() access_token, project_id = self._ensure_access_token( @@ -254,8 +254,8 @@ class VertexAIBatchPrediction(VertexLLM): async def _async_retrieve_batch( self, api_base: str, - headers: Dict[str, str], - logging_obj: Optional[Any] = None, + headers: dict[str, str], + logging_obj: Any | None = None, ) -> LiteLLMBatch: client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, @@ -304,14 +304,14 @@ class VertexAIBatchPrediction(VertexLLM): def list_batches( self, _is_async: bool, - after: Optional[str], - limit: Optional[int], - api_base: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], + after: str | None, + limit: int | None, + api_base: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, ): sync_handler = _get_httpx_client() @@ -346,7 +346,7 @@ class VertexAIBatchPrediction(VertexLLM): "Authorization": f"Bearer {access_token}", } - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if limit is not None: params["pageSize"] = str(limit) if after is not None: @@ -379,8 +379,8 @@ class VertexAIBatchPrediction(VertexLLM): async def _async_list_batches( self, api_base: str, - headers: Dict[str, str], - params: Dict[str, Any], + headers: dict[str, str], + params: dict[str, Any], ): client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, @@ -405,13 +405,13 @@ class VertexAIBatchPrediction(VertexLLM): self, _is_async: bool, batch_id: str, - api_base: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - ) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]: + api_base: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, @@ -501,8 +501,8 @@ class VertexAIBatchPrediction(VertexLLM): self, api_base: str, retrieve_api_base: str, - headers: Dict[str, str], - timeout: Union[float, httpx.Timeout] = 600.0, + headers: dict[str, str], + timeout: float | httpx.Timeout = 600.0, ) -> LiteLLMBatch: client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, diff --git a/litellm/llms/vertex_ai/batches/transformation.py b/litellm/llms/vertex_ai/batches/transformation.py index df903ba7ef0..8329e0881b7 100644 --- a/litellm/llms/vertex_ai/batches/transformation.py +++ b/litellm/llms/vertex_ai/batches/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional +from typing import Any from litellm._uuid import uuid from litellm.llms.vertex_ai.common_utils import ( @@ -59,8 +59,8 @@ class VertexAIBatchTransformation: @classmethod def transform_vertex_ai_batch_list_response_to_openai_list_response( - cls, response: Dict[str, Any] - ) -> Dict[str, Any]: + cls, response: dict[str, Any] + ) -> dict[str, Any]: """ Transforms Vertex AI batch list response into OpenAI-compatible list response. """ @@ -152,7 +152,7 @@ class VertexAIBatchTransformation: ref: https://cloud.google.com/vertex-ai/docs/reference/rest/v1/JobState """ - state_mapping: Dict[str, BatchJobStatus] = { + state_mapping: dict[str, BatchJobStatus] = { "JOB_STATE_UNSPECIFIED": "failed", "JOB_STATE_QUEUED": "validating", "JOB_STATE_PENDING": "validating", @@ -210,7 +210,7 @@ class VertexAIBatchTransformation: return model @classmethod - def is_unmanaged_gcs_batch_input_file_id(cls, input_file_id: Optional[str]) -> bool: + def is_unmanaged_gcs_batch_input_file_id(cls, input_file_id: str | None) -> bool: """ Returns True if `input_file_id` is a raw gs:// Vertex batch input file (i.e. not a LiteLLM-managed unified file id) with a `publishers/` model path that diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 7dcb4dcf2e8..b627444b181 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -1,7 +1,7 @@ import re from copy import deepcopy from enum import Enum -from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, get_type_hints +from typing import Any, Literal, get_type_hints import httpx @@ -26,7 +26,7 @@ class VertexAIError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[Dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(message=message, status_code=status_code, headers=headers) @@ -73,8 +73,8 @@ def redact_vertex_ai_metadata_from_litellm_params(model_call_details: dict) -> N def vertex_request_labels_from_litellm_params( - litellm_params: Optional[dict], -) -> Optional[Dict[str, str]]: + litellm_params: dict | None, +) -> dict[str, str] | None: """ Build Vertex/GCP billing labels from LiteLLM user metadata on ``litellm_params``: ``metadata`` (``completion(..., metadata=...)``) or ``litellm_metadata``, @@ -101,15 +101,15 @@ def vertex_request_labels_from_litellm_params( def pop_vertex_request_labels( - optional_params: Optional[dict], - litellm_params: Optional[dict], -) -> Optional[Dict[str, str]]: + optional_params: dict | None, + litellm_params: dict | None, +) -> dict[str, str] | None: """ Resolve labels from optional ``labels`` (Gemini-style) and/or ``litellm_params["metadata"]`` / ``litellm_params["litellm_metadata"]`` (``requester_metadata``). Pops ``labels`` from optional_params when present. """ - labels: Optional[Dict[str, str]] = None + labels: dict[str, str] | None = None if optional_params is not None and "labels" in optional_params: raw = optional_params.pop("labels") if isinstance(raw, dict): @@ -135,7 +135,7 @@ class VertexAIModelRoute(str, Enum): VERTEX_AI_MODEL_ROUTES = [f"{route.value}/" for route in VertexAIModelRoute] -def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) -> VertexAIModelRoute: +def get_vertex_ai_model_route(model: str, litellm_params: dict | None = None) -> VertexAIModelRoute: """ Determine which handler to use for a Vertex AI model based on the model name. @@ -221,9 +221,7 @@ def get_supports_system_message( supports_system_message = True except Exception as e: verbose_logger.warning( - "Unable to identify if system message supported. Defaulting to 'False'. Received error message - {}\nAdd it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json".format( - str(e) - ) + f"Unable to identify if system message supported. Defaulting to 'False'. Received error message - {e!s}\nAdd it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) supports_system_message = False @@ -270,7 +268,7 @@ def supports_response_json_schema(model: str) -> bool: return bool(gemini_2_plus_pattern.search(model_lower)) -from typing import Literal, Optional +from typing import Literal all_gemini_url_modes = Literal["chat", "embedding", "batch_embedding", "image_generation", "count_tokens"] @@ -311,7 +309,7 @@ def get_vertex_base_model_name(model: str) -> str: return model -def validate_vertex_location(vertex_location: Optional[str]) -> str: +def validate_vertex_location(vertex_location: str | None) -> str: """ Validate a Vertex AI location before interpolating it into a request host or URL path. @@ -334,7 +332,7 @@ def validate_vertex_location(vertex_location: Optional[str]) -> str: def get_vertex_base_url( - vertex_location: Optional[str], + vertex_location: str | None, ) -> str: """ Get the base URL for Vertex AI API calls. @@ -353,10 +351,10 @@ def get_vertex_base_url( def _get_embedding_url( model: str, - vertex_project: Optional[str], - vertex_location: Optional[str], + vertex_project: str | None, + vertex_location: str | None, vertex_api_version: Literal["v1", "v1beta1"], -) -> Tuple[str, str]: +) -> tuple[str, str]: """ Get URL for embedding models. @@ -393,13 +391,13 @@ def _get_embedding_url( def _get_vertex_url( mode: all_gemini_url_modes, model: str, - stream: Optional[bool], - vertex_project: Optional[str], - vertex_location: Optional[str], + stream: bool | None, + vertex_project: str | None, + vertex_location: str | None, vertex_api_version: Literal["v1", "v1beta1"], -) -> Tuple[str, str]: - url: Optional[str] = None - endpoint: Optional[str] = None +) -> tuple[str, str]: + url: str | None = None + endpoint: str | None = None model = litellm.VertexGeminiConfig.get_model_for_vertex_ai_url(model=model) @@ -451,8 +449,8 @@ def _get_vertex_url( def _get_gemini_url( mode: all_gemini_url_modes, model: str, - stream: Optional[bool], -) -> Tuple[str, str]: + stream: bool | None, +) -> tuple[str, str]: """Build the Gemini API URL for the given mode. The API key is NOT included in the URL. Callers must pass it via the @@ -463,27 +461,25 @@ def _get_gemini_url( VertexGeminiConfig, ) - _gemini_model_name = "models/{}".format(model) + _gemini_model_name = f"models/{model}" api_version = "v1alpha" if VertexGeminiConfig._is_gemini_3_or_newer(model) else "v1beta" if mode == "chat": endpoint = "generateContent" if stream is True: endpoint = "streamGenerateContent" - url = "https://generativelanguage.googleapis.com/{}/{}:{}?alt=sse".format( - api_version, _gemini_model_name, endpoint - ) + url = f"https://generativelanguage.googleapis.com/{api_version}/{_gemini_model_name}:{endpoint}?alt=sse" else: - url = "https://generativelanguage.googleapis.com/{}/{}:{}".format(api_version, _gemini_model_name, endpoint) + url = f"https://generativelanguage.googleapis.com/{api_version}/{_gemini_model_name}:{endpoint}" elif mode == "embedding": endpoint = "embedContent" - url = "https://generativelanguage.googleapis.com/v1beta/{}:{}".format(_gemini_model_name, endpoint) + url = f"https://generativelanguage.googleapis.com/v1beta/{_gemini_model_name}:{endpoint}" elif mode == "batch_embedding": endpoint = "batchEmbedContents" - url = "https://generativelanguage.googleapis.com/v1beta/{}:{}".format(_gemini_model_name, endpoint) + url = f"https://generativelanguage.googleapis.com/v1beta/{_gemini_model_name}:{endpoint}" elif mode == "count_tokens": endpoint = "countTokens" - url = "https://generativelanguage.googleapis.com/v1beta/{}:{}".format(_gemini_model_name, endpoint) + url = f"https://generativelanguage.googleapis.com/v1beta/{_gemini_model_name}:{endpoint}" elif mode == "image_generation": raise ValueError( "LiteLLM's `gemini/` route does not support image generation yet. Let us know if you need this feature by opening an issue at https://github.com/BerriAI/litellm/issues" @@ -494,7 +490,7 @@ def _get_gemini_url( return url, endpoint -def _check_text_in_content(parts: List[PartType]) -> bool: +def _check_text_in_content(parts: list[PartType]) -> bool: """ check that user_content has 'text' parameter. - Known Vertex Error: Unable to submit request because it must have a text parameter. @@ -654,7 +650,7 @@ def _build_json_schema(parameters: dict) -> dict: return parameters -def _filter_anyof_fields(schema_dict: Dict[str, Any]) -> Dict[str, Any]: +def _filter_anyof_fields(schema_dict: dict[str, Any]) -> dict[str, Any]: """ When anyof is present, only keep the anyof field and its contents - otherwise VertexAI will throw an error - https://github.com/BerriAI/litellm/issues/11164 Filter out other fields in the same dict. @@ -689,13 +685,15 @@ def process_items(schema, depth=0): # Normalize: empty `items: {}` and missing-items both become {"type": "object"}. type_val = schema.get("type") if ( - isinstance(type_val, str) - and type_val.lower() == "array" - and ("items" not in schema or schema.get("items") == {}) + ( + isinstance(type_val, str) + and type_val.lower() == "array" + and ("items" not in schema or schema.get("items") == {}) + ) + or schema.get("type") == "array" + and "items" not in schema ): schema["items"] = {"type": "object"} - elif schema.get("type") == "array" and "items" not in schema: - schema["items"] = {"type": "object"} for key, value in schema.items(): if isinstance(value, dict): process_items(value, depth + 1) @@ -705,7 +703,7 @@ def process_items(schema, depth=0): process_items(item, depth + 1) -def set_schema_property_ordering(schema: Dict[str, Any], depth: int = 0) -> Dict[str, Any]: +def set_schema_property_ordering(schema: dict[str, Any], depth: int = 0) -> dict[str, Any]: """ vertex ai and generativeai apis order output of fields alphabetically, unless you specify the order. python dicts retain order, so we just use that. Note that this field only applies to structured outputs, and not tools. @@ -732,7 +730,7 @@ def set_schema_property_ordering(schema: Dict[str, Any], depth: int = 0) -> Dict return schema -def filter_schema_fields(schema_dict: Dict[str, Any], valid_fields: Set[str], processed=None) -> Dict[str, Any]: +def filter_schema_fields(schema_dict: dict[str, Any], valid_fields: set[str], processed=None) -> dict[str, Any]: """ Recursively filter a schema dictionary to keep only valid fields. """ @@ -909,7 +907,7 @@ def _convert_schema_types(schema, depth=0): "maxProperties", } - any_of: List[Dict[str, Any]] = [] + any_of: list[dict[str, Any]] = [] for t in type_val: if not isinstance(t, str): continue @@ -957,7 +955,7 @@ def _convert_schema_types(schema, depth=0): _convert_schema_types(anyof_schema, depth + 1) -def get_vertex_project_id_from_url(url: str) -> Optional[str]: +def get_vertex_project_id_from_url(url: str) -> str | None: """ Get the vertex project id from the url @@ -967,7 +965,7 @@ def get_vertex_project_id_from_url(url: str) -> Optional[str]: return match.group(1) if match else None -def get_vertex_location_from_url(url: str) -> Optional[str]: +def get_vertex_location_from_url(url: str) -> str | None: """ Get the vertex location from the url @@ -977,7 +975,7 @@ def get_vertex_location_from_url(url: str) -> Optional[str]: return match.group(1) if match else None -def get_vertex_model_id_from_url(url: str) -> Optional[str]: +def get_vertex_model_id_from_url(url: str) -> str | None: """ Get the vertex model id from the url @@ -1003,8 +1001,8 @@ def replace_project_and_location_in_route(requested_route: str, vertex_project: def construct_target_url( base_url: str, requested_route: str, - vertex_location: Optional[str], - vertex_project: Optional[str], + vertex_location: str | None, + vertex_project: str | None, ) -> httpx.URL: """ Allow user to specify their own project id / location. @@ -1041,7 +1039,7 @@ def construct_target_url( vertex_version = "v1beta1" requested_route = requested_route.replace("/v1beta1/", "/", 1) - base_requested_route = "{}/projects/{}/locations/{}".format(vertex_version, vertex_project, vertex_location) + base_requested_route = f"{vertex_version}/projects/{vertex_project}/locations/{vertex_location}" updated_requested_route = "/" + base_requested_route + requested_route @@ -1050,7 +1048,7 @@ def construct_target_url( class VertexAIModelInfo(BaseLLMModelInfo): - def get_token_counter(self) -> Optional[BaseTokenCounter]: + def get_token_counter(self) -> BaseTokenCounter | None: """ Factory method to create a token counter for this provider. @@ -1064,32 +1062,32 @@ class VertexAIModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: raise NotImplementedError("Vertex AI models are not supported yet") - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: """ Returns a list of models supported by this provider. """ raise NotImplementedError("Vertex AI models are not supported yet") @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: raise NotImplementedError("Vertex AI models are not supported yet") @staticmethod def get_api_base( - api_base: Optional[str] = None, - ) -> Optional[str]: + api_base: str | None = None, + ) -> str | None: raise NotImplementedError("Vertex AI models are not supported yet") @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: """ Returns the base model name from the given model name. @@ -1104,7 +1102,7 @@ class VertexAITokenCounter(BaseTokenCounter): def should_use_token_counting_api( self, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> bool: from litellm.types.utils import LlmProviders @@ -1113,13 +1111,13 @@ class VertexAITokenCounter(BaseTokenCounter): async def count_tokens( self, model_to_use: str, - messages: Optional[List[Dict[str, Any]]], - contents: Optional[List[Dict[str, Any]]], - deployment: Optional[Dict[str, Any]] = None, + messages: list[dict[str, Any]] | None, + contents: list[dict[str, Any]] | None, + deployment: dict[str, Any] | None = None, request_model: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[Any] = None, - ) -> Optional[TokenCountResponse]: + tools: list[dict[str, Any]] | None = None, + system: Any | None = None, + ) -> TokenCountResponse | None: import copy from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index c60574ce934..1cd04cfb5e7 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -6,7 +6,7 @@ Why separate file? Make it easy to see how transformation works import re from collections.abc import Sequence -from typing import List, Literal, Optional, Tuple +from typing import Literal from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import CachedContentRequestBody @@ -20,7 +20,7 @@ from ..gemini.transformation import ( def get_first_continuous_block_idx( - filtered_messages: List[Tuple[int, AllMessageValues]], # (idx, message) + filtered_messages: list[tuple[int, AllMessageValues]], # (idx, message) ) -> int: """ Find the array index that ends the first continuous sequence of message blocks. @@ -49,7 +49,7 @@ def get_first_continuous_block_idx( return len(filtered_messages) - 1 -def extract_ttl_from_cached_messages(messages: List[AllMessageValues]) -> Optional[str]: +def extract_ttl_from_cached_messages(messages: list[AllMessageValues]) -> str | None: """ Extract TTL from cached messages. Returns the first valid TTL found. @@ -116,8 +116,8 @@ def _is_valid_ttl_format(ttl: str) -> bool: def separate_cached_messages( - messages: List[AllMessageValues], -) -> Tuple[List[AllMessageValues], List[AllMessageValues]]: + messages: list[AllMessageValues], +) -> tuple[list[AllMessageValues], list[AllMessageValues]]: """ Returns separated cached and non-cached messages. @@ -129,11 +129,11 @@ def separate_cached_messages( - cached_messages: List of cached messages. - non_cached_messages: List of non-cached messages. """ - cached_messages: List[AllMessageValues] = [] - non_cached_messages: List[AllMessageValues] = [] + cached_messages: list[AllMessageValues] = [] + non_cached_messages: list[AllMessageValues] = [] # Extract cached messages and their indices - filtered_messages: List[Tuple[int, AllMessageValues]] = [] + filtered_messages: list[tuple[int, AllMessageValues]] = [] for idx, message in enumerate(messages): if is_cached_message(message=message): filtered_messages.append((idx, message)) @@ -169,11 +169,11 @@ def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageVa def transform_openai_messages_to_gemini_context_caching( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], cache_key: str, - vertex_project: Optional[str], - vertex_location: Optional[str], + vertex_project: str | None, + vertex_location: str | None, ) -> CachedContentRequestBody: # Extract TTL from cached messages BEFORE system message transformation ttl = extract_ttl_from_cached_messages(messages) @@ -190,7 +190,7 @@ def transform_openai_messages_to_gemini_context_caching( custom_llm_provider=custom_llm_provider, ) - model_name = "models/{}".format(model) + model_name = f"models/{model}" if custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta": model_name = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/{model_name}" diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index f8774e33ca4..71a73f1981b 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -1,8 +1,9 @@ -from typing import List, Literal, Optional, Tuple, Union +from typing import Literal import httpx import litellm +from litellm._logging import verbose_logger from litellm.caching.caching import Cache, LiteLLMCacheType from litellm.constants import MINIMUM_PROMPT_CACHE_TOKEN_COUNT from litellm.litellm_core_utils.litellm_logging import Logging @@ -11,13 +12,12 @@ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, get_async_httpx_client, ) -from litellm._logging import verbose_logger from litellm.llms.openai.openai import AllMessageValues -from litellm.utils import is_prompt_caching_valid_prompt from litellm.types.llms.vertex_ai import ( CachedContentListAllResponseBody, VertexAICachedContentResponseObject, ) +from litellm.utils import is_prompt_caching_valid_prompt from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase @@ -44,14 +44,14 @@ class ContextCachingEndpoints(VertexBase): def _get_token_and_url_context_caching( self, - gemini_api_key: Optional[str], + gemini_api_key: str | None, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - api_base: Optional[str], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], - model: Optional[str] = None, - ) -> Tuple[Optional[str], str]: + api_base: str | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, + model: str | None = None, + ) -> tuple[str | None, str]: """ Internal function. Returns the token and url for the call. @@ -60,11 +60,11 @@ class ContextCachingEndpoints(VertexBase): Returns token, url """ - auth_header: Optional[str] + auth_header: str | None if custom_llm_provider == "gemini": auth_header = {"x-goog-api-key": gemini_api_key} # type: ignore[assignment] endpoint = "cachedContents" - url = "https://generativelanguage.googleapis.com/v1beta/{}".format(endpoint) + url = f"https://generativelanguage.googleapis.com/v1beta/{endpoint}" elif custom_llm_provider == "vertex_ai": auth_header = vertex_auth_header endpoint = "cachedContents" @@ -96,14 +96,14 @@ class ContextCachingEndpoints(VertexBase): client: HTTPHandler, headers: dict, api_key: str, - api_base: Optional[str], + api_base: str | None, logging_obj: Logging, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], - model: Optional[str] = None, - ) -> Optional[str]: + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, + model: str | None = None, + ) -> str | None: """ Checks if content already cached. @@ -125,7 +125,7 @@ class ContextCachingEndpoints(VertexBase): model=model, ) - page_token: Optional[str] = None + page_token: str | None = None # Iterate through all pages for _ in range(MAX_PAGINATION_PAGES): @@ -188,14 +188,14 @@ class ContextCachingEndpoints(VertexBase): client: AsyncHTTPHandler, headers: dict, api_key: str, - api_base: Optional[str], + api_base: str | None, logging_obj: Logging, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], - model: Optional[str] = None, - ) -> Optional[str]: + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, + model: str | None = None, + ) -> str | None: """ Checks if content already cached. @@ -217,7 +217,7 @@ class ContextCachingEndpoints(VertexBase): model=model, ) - page_token: Optional[str] = None + page_token: str | None = None # Iterate through all pages for _ in range(MAX_PAGINATION_PAGES): @@ -276,21 +276,21 @@ class ContextCachingEndpoints(VertexBase): def check_and_create_cache( self, - messages: List[AllMessageValues], # receives openai format messages + messages: list[AllMessageValues], # receives openai format messages optional_params: dict, # cache the tools if present, in case cache content exists in messages api_key: str, - api_base: Optional[str], + api_base: str | None, model: str, - client: Optional[HTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: HTTPHandler | None, + timeout: float | httpx.Timeout | None, logging_obj: Logging, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], - extra_headers: Optional[dict] = None, - cached_content: Optional[str] = None, - ) -> Tuple[List[AllMessageValues], dict, Optional[str]]: + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, + extra_headers: dict | None = None, + cached_content: str | None = None, + ) -> tuple[list[AllMessageValues], dict, str | None]: """ Receives - messages: List of dict - messages in the openai format @@ -435,21 +435,21 @@ class ContextCachingEndpoints(VertexBase): async def async_check_and_create_cache( self, - messages: List[AllMessageValues], # receives openai format messages + messages: list[AllMessageValues], # receives openai format messages optional_params: dict, # cache the tools if present, in case cache content exists in messages api_key: str, - api_base: Optional[str], + api_base: str | None, model: str, - client: Optional[AsyncHTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], + client: AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, logging_obj: Logging, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], - extra_headers: Optional[dict] = None, - cached_content: Optional[str] = None, - ) -> Tuple[List[AllMessageValues], dict, Optional[str]]: + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, + extra_headers: dict | None = None, + cached_content: str | None = None, + ) -> tuple[list[AllMessageValues], dict, str | None]: """ Receives - messages: List of dict - messages in the openai format diff --git a/litellm/llms/vertex_ai/cost_calculator.py b/litellm/llms/vertex_ai/cost_calculator.py index 84c9108847b..1cd4c0a9e97 100644 --- a/litellm/llms/vertex_ai/cost_calculator.py +++ b/litellm/llms/vertex_ai/cost_calculator.py @@ -1,6 +1,6 @@ # What is this? ## Cost calculation for Google AI Studio / Vertex AI models -from typing import Literal, Optional, Tuple, Union +from typing import Literal import litellm from litellm import verbose_logger @@ -30,7 +30,7 @@ models_without_dynamic_pricing = ["gemini-1.0-pro", "gemini-pro", "gemini-2"] def cost_router( model: str, custom_llm_provider: str, - call_type: Union[Literal["embedding", "aembedding"], str], + call_type: Literal["embedding", "aembedding"] | str, ) -> Literal["cost_per_character", "cost_per_token"]: """ Route the cost calc to the right place, based on model/call_type/etc. @@ -38,19 +38,22 @@ def cost_router( Returns - str, the specific google cost calc function it should route to. """ - if custom_llm_provider == "vertex_ai" and ( - "claude" in model - or "llama" in model - or "mistral" in model - or "jamba" in model - or "codestral" in model - or "gemma" in model + if ( + custom_llm_provider == "vertex_ai" + and ( + "claude" in model + or "llama" in model + or "mistral" in model + or "jamba" in model + or "codestral" in model + or "gemma" in model + ) + or custom_llm_provider == "vertex_ai" + and (call_type == "embedding" or call_type == "aembedding") + or custom_llm_provider == "vertex_ai" + and ("gemini-2" in model) ): return "cost_per_token" - elif custom_llm_provider == "vertex_ai" and (call_type == "embedding" or call_type == "aembedding"): - return "cost_per_token" - elif custom_llm_provider == "vertex_ai" and ("gemini-2" in model): - return "cost_per_token" return "cost_per_character" @@ -58,9 +61,9 @@ def cost_per_character( model: str, custom_llm_provider: str, usage: Usage, - prompt_characters: Optional[float] = None, - completion_characters: Optional[float] = None, -) -> Tuple[float, float]: + prompt_characters: float | None = None, + completion_characters: float | None = None, +) -> tuple[float, float]: """ Calculates the cost per character for a given VertexAI model, input messages, and response object. @@ -99,23 +102,19 @@ def cost_per_character( "input_cost_per_character_above_128k_tokens" in model_info and model_info["input_cost_per_character_above_128k_tokens"] is not None ), ( - "model info for model={} does not have 'input_cost_per_character_above_128k_tokens'-pricing for > 128k tokens\nmodel_info={}".format( - model, model_info - ) + f"model info for model={model} does not have 'input_cost_per_character_above_128k_tokens'-pricing for > 128k tokens\nmodel_info={model_info}" ) prompt_cost = prompt_characters * model_info["input_cost_per_character_above_128k_tokens"] else: assert ( "input_cost_per_character" in model_info and model_info["input_cost_per_character"] is not None - ), "model info for model={} does not have 'input_cost_per_character'-pricing\nmodel_info={}".format( - model, model_info + ), ( + f"model info for model={model} does not have 'input_cost_per_character'-pricing\nmodel_info={model_info}" ) prompt_cost = prompt_characters * model_info["input_cost_per_character"] except Exception as e: verbose_logger.debug( - "litellm.litellm_core_utils.llm_cost_calc.google.py::cost_per_character(): Exception occured - {}\nDefaulting to None".format( - str(e) - ) + f"litellm.litellm_core_utils.llm_cost_calc.google.py::cost_per_character(): Exception occured - {e!s}\nDefaulting to None" ) prompt_cost, _ = cost_per_token( model=model, @@ -141,23 +140,19 @@ def cost_per_character( "output_cost_per_character_above_128k_tokens" in model_info and model_info["output_cost_per_character_above_128k_tokens"] is not None ), ( - "model info for model={} does not have 'output_cost_per_character_above_128k_tokens' pricing\nmodel_info={}".format( - model, model_info - ) + f"model info for model={model} does not have 'output_cost_per_character_above_128k_tokens' pricing\nmodel_info={model_info}" ) completion_cost = completion_tokens * model_info["output_cost_per_character_above_128k_tokens"] else: assert ( "output_cost_per_character" in model_info and model_info["output_cost_per_character"] is not None - ), "model info for model={} does not have 'output_cost_per_character'-pricing\nmodel_info={}".format( - model, model_info + ), ( + f"model info for model={model} does not have 'output_cost_per_character'-pricing\nmodel_info={model_info}" ) completion_cost = completion_characters * model_info["output_cost_per_character"] except Exception as e: verbose_logger.debug( - "litellm.litellm_core_utils.llm_cost_calc.google.py::cost_per_character(): Exception occured - {}\nDefaulting to None".format( - str(e) - ) + f"litellm.litellm_core_utils.llm_cost_calc.google.py::cost_per_character(): Exception occured - {e!s}\nDefaulting to None" ) _, completion_cost = cost_per_token( model=model, @@ -171,7 +166,7 @@ def cost_per_character( def _handle_128k_pricing( model_info: ModelInfo, usage: Usage, -) -> Tuple[float, float]: +) -> tuple[float, float]: ## CALCULATE INPUT COST input_cost_per_token_above_128k_tokens = model_info.get("input_cost_per_token_above_128k_tokens") output_cost_per_token_above_128k_tokens = model_info.get("output_cost_per_token_above_128k_tokens") @@ -198,8 +193,8 @@ def cost_per_token( model: str, custom_llm_provider: str, usage: Usage, - service_tier: Optional[str] = None, -) -> Tuple[float, float]: + service_tier: str | None = None, +) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. diff --git a/litellm/llms/vertex_ai/count_tokens/handler.py b/litellm/llms/vertex_ai/count_tokens/handler.py index 9f2826a4bb4..de7f56875df 100644 --- a/litellm/llms/vertex_ai/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/count_tokens/handler.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -7,12 +7,12 @@ from litellm.llms.vertex_ai.vertex_llm_base import VertexBase class VertexAITokenCounter(GoogleAIStudioTokenCounter, VertexBase): async def validate_environment( self, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - headers: Optional[Dict[str, Any]] = None, + api_base: str | None = None, + api_key: str | None = None, + headers: dict[str, Any] | None = None, model: str = "", - litellm_params: Optional[Dict[str, Any]] = None, - ) -> Tuple[Dict[str, Any], str]: + litellm_params: dict[str, Any] | None = None, + ) -> tuple[dict[str, Any], str]: """ Returns a Tuple of headers and url for the Vertex AI countTokens endpoint. """ diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index aa8b446c6ff..21a31d67622 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -3,7 +3,7 @@ import json import os import time from collections.abc import Coroutine, Mapping -from typing import Any, Optional, Tuple, Union +from typing import Any from urllib.parse import unquote import httpx @@ -76,8 +76,8 @@ class VertexAIFilesHandler(GCSBucketBase): self, file_id: str, configured_bucket_name: str, - litellm_params: Optional[dict] = None, - ) -> Tuple[str, str]: + litellm_params: dict | None = None, + ) -> tuple[str, str]: """ Validate and extract bucket name and object path from file_id. @@ -99,12 +99,12 @@ class VertexAIFilesHandler(GCSBucketBase): async def afile_content( self, file_content_request: FileContentRequest, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - litellm_params: Optional[dict] = None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + litellm_params: dict | None = None, ) -> HttpxBinaryResponseContent: """ Download file content from GCS bucket for VertexAI files. @@ -186,14 +186,14 @@ class VertexAIFilesHandler(GCSBucketBase): self, _is_async: bool, file_content_request: FileContentRequest, - api_base: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - vertex_project: Optional[str], - vertex_location: Optional[str], - timeout: Union[float, httpx.Timeout], - max_retries: Optional[int], - litellm_params: Optional[dict] = None, - ) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: + api_base: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + vertex_project: str | None, + vertex_location: str | None, + timeout: float | httpx.Timeout, + max_retries: int | None, + litellm_params: dict | None = None, + ) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: """ Download file content from GCS bucket for VertexAI files. Supports both sync and async operations. diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 17349fdc618..f8eea399bf6 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -8,11 +8,6 @@ import time from collections.abc import Callable, Iterable, Iterator from typing import ( Any, - Dict, - List, - Optional, - Tuple, - Union, ) import httpx @@ -88,7 +83,7 @@ def _sanitize_gcp_label_value(value: str) -> str: return sanitized[:_GCP_LABEL_VALUE_MAX_LEN] -def _encode_gcp_label_value_chunks(value: str) -> List[str]: +def _encode_gcp_label_value_chunks(value: str) -> list[str]: """Encode arbitrary text across one or more GCP-label-safe values.""" max_encoded_len = _GCP_LABEL_VALUE_MAX_LEN - len(_CUSTOM_ID_RAW_LABEL_PREFIX) encoded = base64.b32encode(value.encode("utf-8")).decode("ascii").rstrip("=").lower() @@ -98,7 +93,7 @@ def _encode_gcp_label_value_chunks(value: str) -> List[str]: ] or [_CUSTOM_ID_RAW_LABEL_PREFIX] -def _decode_gcp_label_value_chunks(values: List[str]) -> Optional[str]: +def _decode_gcp_label_value_chunks(values: list[str]) -> str | None: """Decode values produced by _encode_gcp_label_value_chunks.""" encoded_parts = [] for value in values: @@ -113,7 +108,7 @@ def _decode_gcp_label_value_chunks(values: List[str]) -> Optional[str]: return None -def _set_litellm_batch_custom_id_labels(labels: Dict[str, str], custom_id: Any) -> None: +def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: Any) -> None: """ Store OpenAI batch custom_id for Vertex batch correlation. @@ -129,7 +124,7 @@ def _set_litellm_batch_custom_id_labels(labels: Dict[str, str], custom_id: Any) labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk -def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str: +def _get_litellm_batch_custom_id_from_labels(labels: dict[str, Any]) -> str: """Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels).""" raw = labels.get("litellm_custom_id_raw") if raw: @@ -148,9 +143,9 @@ def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str: def _openai_batch_jsonl_entry_to_vertex_wrapped_request( - openai_entry: Dict[str, Any], - map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]], -) -> Dict[str, Any]: + openai_entry: dict[str, Any], + map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]], +) -> dict[str, Any]: """ Transforms a single OpenAI JSONL batch entry into its Vertex wrapped request. @@ -177,7 +172,7 @@ def _openai_batch_jsonl_entry_to_vertex_wrapped_request( return {"request": vertex_request_body} -def _iter_stripped_lines(raw_lines: Iterable[Union[str, bytes]]) -> Iterator[str]: +def _iter_stripped_lines(raw_lines: Iterable[str | bytes]) -> Iterator[str]: """Decode (when needed), strip, and drop blank lines from an iterable of lines.""" for raw in raw_lines: line = raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else raw @@ -248,7 +243,7 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]: def _iter_openai_jsonl_entries( openai_file_content: FileTypes, -) -> Iterator[Dict[str, Any]]: +) -> Iterator[dict[str, Any]]: for line in _iter_openai_jsonl_lines(openai_file_content): yield json.loads(line) @@ -264,7 +259,7 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream): def __init__( self, openai_file_content: FileTypes, - map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]], + map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]], ) -> None: self._openai_file_content = openai_file_content self._map_openai_to_vertex_params = map_openai_to_vertex_params @@ -297,11 +292,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if not api_key: api_key, _ = self.get_access_token( @@ -315,7 +310,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def _get_gcs_object_name_from_batch_jsonl( self, - openai_jsonl_content: List[Dict[str, Any]], + openai_jsonl_content: list[dict[str, Any]], ) -> str: """ Gets a unique GCS object name for the VertexAI batch prediction job @@ -350,7 +345,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): fallback_filename="file", ) - def _get_configured_bucket_name(self, litellm_params: Dict) -> str: + def _get_configured_bucket_name(self, litellm_params: dict) -> str: bucket_name = ( litellm_params.get("gcs_bucket_name") or litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME") ) @@ -360,11 +355,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def get_complete_file_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, - optional_params: Dict, - litellm_params: Dict, + optional_params: dict, + litellm_params: dict, data: CreateFileRequest, ) -> str: """ @@ -389,7 +384,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): return f"{api_base}/{endpoint}" - def get_supported_openai_params(self, model: str) -> List[OpenAICreateFileRequestOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: return [] def map_openai_params( @@ -403,8 +398,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def _map_openai_to_vertex_params( self, - openai_request_body: Dict[str, Any], - ) -> Dict[str, Any]: + openai_request_body: dict[str, Any], + ) -> dict[str, Any]: """ wrapper to call VertexGeminiConfig.map_openai_params """ @@ -428,7 +423,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): create_file_data: CreateFileRequest, optional_params: dict, litellm_params: dict, - ) -> Union[bytes, str, dict]: + ) -> bytes | str | dict: """ 2 Cases: 1. Handle basic file upload @@ -462,7 +457,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def transform_create_file_response( self, - model: Optional[str], + model: str | None, raw_response: Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -497,10 +492,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): object="file", ) - def get_error_class(self, error_message: str, status_code: int, headers: Union[Dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return VertexAIError(status_code=status_code, message=error_message, headers=headers) - def _parse_gcs_uri(self, file_id: str, litellm_params: Optional[Dict] = None) -> Tuple[str, str]: + def _parse_gcs_uri(self, file_id: str, litellm_params: dict | None = None) -> tuple[str, str]: """ Validate a managed GCS file_id and return (bucket, url-encoded-object-path). """ @@ -575,7 +570,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def transform_list_files_request( self, - purpose: Optional[str], + purpose: str | None, optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: @@ -586,7 +581,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): raw_response: Response, logging_obj: LiteLLMLoggingObj, litellm_params: dict, - ) -> List[OpenAIFileObject]: + ) -> list[OpenAIFileObject]: raise NotImplementedError("VertexAIFilesConfig does not support file listing") def transform_file_content_request( @@ -648,7 +643,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): return HttpxBinaryResponseContent(response=raw_response) def _try_transform_vertex_batch_output_to_openai( - self, content: bytes, logging_obj: Optional[LiteLLMLoggingObj] = None + self, content: bytes, logging_obj: LiteLLMLoggingObj | None = None ) -> bytes: """ Try to transform Vertex AI batch output to OpenAI format. @@ -749,11 +744,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def _transform_single_vertex_batch_output_to_openai( self, - vertex_output: Dict[str, Any], + vertex_output: dict[str, Any], vertex_gemini_config: VertexGeminiConfig, logging_obj: Logging, mock_httpx_response: httpx.Response, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Transform a single Vertex AI batch output line to OpenAI format. Uses the existing VertexGeminiConfig transformation for the response. @@ -820,6 +815,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): "response": None, "error": { "code": "transformation_error", - "message": f"Failed to transform response: {str(e)}", + "message": f"Failed to transform response: {e!s}", }, } diff --git a/litellm/llms/vertex_ai/fine_tuning/handler.py b/litellm/llms/vertex_ai/fine_tuning/handler.py index 20f7b3905ab..184287bb688 100644 --- a/litellm/llms/vertex_ai/fine_tuning/handler.py +++ b/litellm/llms/vertex_ai/fine_tuning/handler.py @@ -2,7 +2,7 @@ import json import traceback from collections.abc import Coroutine from datetime import datetime -from typing import Any, Literal, Optional, Union +from typing import Any, Literal import httpx @@ -51,7 +51,7 @@ class VertexFineTuningAPI(VertexLLM): self, create_fine_tuning_job_data: FineTuningJobCreate, original_hyperparameters: dict = {}, - kwargs: Optional[dict] = None, + kwargs: dict | None = None, ) -> FineTuneJobCreate: """ convert request from OpenAI format to Vertex format @@ -87,7 +87,7 @@ class VertexFineTuningAPI(VertexLLM): self, create_fine_tuning_job_data: FineTuningJobCreate, original_hyperparameters: dict = {}, - kwargs: Optional[dict] = None, + kwargs: dict | None = None, ) -> FineTuneHyperparameters: _oai_hyperparameters = create_fine_tuning_job_data.hyperparameters _vertex_hyperparameters = FineTuneHyperparameters() @@ -200,14 +200,14 @@ class VertexFineTuningAPI(VertexLLM): self, _is_async: bool, create_fine_tuning_job_data: FineTuningJobCreate, - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], - kwargs: Optional[dict] = None, - original_hyperparameters: Optional[dict] = {}, - ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]: + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + api_base: str | None, + timeout: float | httpx.Timeout, + kwargs: dict | None = None, + original_hyperparameters: dict | None = {}, + ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: verbose_logger.debug("creating fine tuning job, args= %s", create_fine_tuning_job_data) _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -310,13 +310,12 @@ class VertexFineTuningAPI(VertexLLM): url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs" elif "/tuningJobs/" in request_route and "cancel" in request_route: url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/tuningJobs{request_route}" - elif "generateContent" in request_route: - url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}" - elif "predict" in request_route: - url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}" - elif "/batchPredictionJobs" in request_route: - url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}" - elif "countTokens" in request_route: + elif ( + "generateContent" in request_route + or "predict" in request_route + or "/batchPredictionJobs" in request_route + or "countTokens" in request_route + ): url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}{request_route}" elif "cachedContents" in request_route: _model = request_data.get("model") diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index c0da1536c42..9d4d8a5a02e 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -7,7 +7,7 @@ Why separate file? Make it easy to see how transformation works import json import os import re -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, cast from urllib.parse import quote import httpx @@ -64,11 +64,11 @@ from ..common_utils import ( # Typed as Any to avoid introducing a module-load-time cyclic import to # vertex_llm_base. The instance is lazily constructed by _get_vertex_base() # the first time GCS metadata needs to be fetched. -_GCS_METADATA_VERTEX_BASE: Optional[Any] = None +_GCS_METADATA_VERTEX_BASE: Any | None = None # Shared sync client for GCS JSON API metadata reads so proxy/SSL settings # from litellm's HTTP stack apply (see Greptile review on PR #27278). -_GCS_METADATA_HTTP_HANDLER: Optional[HTTPHandler] = None -_GEMINI_MIME_TYPE_ALIASES: Dict[str, str] = { +_GCS_METADATA_HTTP_HANDLER: HTTPHandler | None = None +_GEMINI_MIME_TYPE_ALIASES: dict[str, str] = { "image/jpg": "image/jpeg", } @@ -109,8 +109,8 @@ else: def _convert_detail_to_media_resolution_enum( - detail: Optional[str], -) -> Optional[Dict[str, str]]: + detail: str | None, +) -> dict[str, str] | None: if detail == "low": return {"level": "MEDIA_RESOLUTION_LOW"} elif detail == "medium": @@ -122,7 +122,7 @@ def _convert_detail_to_media_resolution_enum( return None -def _get_highest_media_resolution(current: Optional[str], new_detail: Optional[str]) -> Optional[str]: +def _get_highest_media_resolution(current: str | None, new_detail: str | None) -> str | None: """ Compare two media resolution values and return the highest one. Resolution hierarchy: ultra_high > high > medium > low > None @@ -137,8 +137,8 @@ def _get_highest_media_resolution(current: Optional[str], new_detail: Optional[s def _extract_max_media_resolution_from_messages( - messages: List[AllMessageValues], -) -> Optional[str]: + messages: list[AllMessageValues], +) -> str | None: """ Extract the highest media resolution (detail) from image content in messages. @@ -151,14 +151,14 @@ def _extract_max_media_resolution_from_messages( Returns: The highest detail level found ("high", "low", or None) """ - max_resolution: Optional[str] = None + max_resolution: str | None = None for msg in messages: content = msg.get("content") if isinstance(content, list): for item in content: if not isinstance(item, dict): continue - detail: Optional[str] = None + detail: str | None = None if item.get("type") == "image_url": image_url = item.get("image_url") if isinstance(image_url, dict): @@ -174,9 +174,9 @@ def _extract_max_media_resolution_from_messages( def _apply_gemini_metadata( part: PartType, - model: Optional[str], - media_resolution_enum: Optional[Dict[str, str]], - video_metadata: Optional[Dict[str, Any]], + model: str | None, + media_resolution_enum: dict[str, str] | None, + video_metadata: dict[str, Any] | None, ) -> PartType: """ Apply media_resolution and video_metadata parameters to a Gemini part. @@ -208,7 +208,7 @@ def _apply_gemini_metadata( return cast(PartType, part_dict) -def _parse_gs_uri(gs_uri: str) -> Tuple[str, str]: +def _parse_gs_uri(gs_uri: str) -> tuple[str, str]: if not gs_uri.startswith("gs://"): raise ValueError(f"Invalid gs URI: {gs_uri}") uri_without_scheme = gs_uri[5:] # drop gs:// @@ -256,8 +256,8 @@ def _image_url_payload_may_need_sync_gcs_metadata_fetch( True when this image_url value (content-part image_url or assistant ``images[]`` entry) can trigger a blocking GCS metadata read for MIME resolution. """ - fmt: Optional[str] = None - url: Optional[str] = None + fmt: str | None = None + url: str | None = None if isinstance(raw_image_url, dict): url = raw_image_url.get("url") # type: ignore[assignment] if not isinstance(url, str): @@ -273,7 +273,7 @@ def _image_url_payload_may_need_sync_gcs_metadata_fetch( def _openai_messages_may_need_sync_gcs_metadata_fetch( - messages: List[AllMessageValues], + messages: list[AllMessageValues], ) -> bool: """ Heuristic: True if any message part can trigger a blocking GCS JSON @@ -325,9 +325,9 @@ def _openai_messages_may_need_sync_gcs_metadata_fetch( def _get_gcs_object_content_type( image_url: str, - vertex_project: Optional[str] = None, - vertex_credentials: Optional[Any] = None, -) -> Optional[str]: + vertex_project: str | None = None, + vertex_credentials: Any | None = None, +) -> str | None: """ Resolve content type from GCS object metadata. @@ -345,7 +345,7 @@ def _get_gcs_object_content_type( if not _is_valid_gcs_bucket_name(bucket): return None - headers: Dict[str, str] = {} + headers: dict[str, str] = {} explicit_vertex_auth_provided = vertex_project is not None or vertex_credentials is not None if explicit_vertex_auth_provided: try: @@ -357,7 +357,7 @@ def _get_gcs_object_content_type( except Exception as e: raise litellm.BadRequestError( message=( - f"Unable to fetch GCS metadata with provided Vertex credentials/project. Original error: {str(e)}" + f"Unable to fetch GCS metadata with provided Vertex credentials/project. Original error: {e!s}" ), model=None, llm_provider="vertex_ai", @@ -448,7 +448,7 @@ def _get_gcs_object_content_type( return None -def _normalize_and_validate_gemini_mime_type(mime_type: str, model: Optional[str]) -> str: +def _normalize_and_validate_gemini_mime_type(mime_type: str, model: str | None) -> str: # Import lazily to avoid a module-level cyclic-import alert with # litellm.types.files. from litellm.types.files import get_file_extension_from_mime_type @@ -476,12 +476,12 @@ def _normalize_and_validate_gemini_mime_type(mime_type: str, model: Optional[str def _process_gemini_media( image_url: str, - format: Optional[str] = None, - media_resolution_enum: Optional[Dict[str, str]] = None, - model: Optional[str] = None, - video_metadata: Optional[Dict[str, Any]] = None, - vertex_project: Optional[str] = None, - vertex_credentials: Optional[Any] = None, + format: str | None = None, + media_resolution_enum: dict[str, str] | None = None, + model: str | None = None, + video_metadata: dict[str, Any] | None = None, + vertex_project: str | None = None, + vertex_credentials: Any | None = None, ) -> PartType: """ Given a media URL (image, audio, or video), return the appropriate PartType for Gemini @@ -503,7 +503,7 @@ def _process_gemini_media( explicit_gcs_format = False if not format: - mime_type: Optional[str] = None + mime_type: str | None = None # For extension-less gs:// URIs, we cannot infer from path. # If callers pass `format`/`mime_type`, this branch is skipped. if extension: @@ -579,7 +579,7 @@ def _process_gemini_media( _blob: BlobType = {"data": image["data"], "mime_type": image["media_type"]} part = {"inline_data": cast(BlobType, _blob)} return _apply_gemini_metadata(part, model, media_resolution_enum, video_metadata) - raise Exception("Invalid image received - {}".format(image_url)) + raise Exception(f"Invalid image received - {image_url}") except Exception as e: raise e @@ -595,7 +595,7 @@ def _camel_to_snake(camel_str: str) -> str: return re.sub(r"(? Optional[str]: +def _get_equivalent_key(key: str, available_keys: set) -> str | None: """ Get the equivalent key from available keys, checking both camelCase and snake_case variants """ @@ -615,7 +615,7 @@ def _get_equivalent_key(key: str, available_keys: set) -> Optional[str]: return None -def check_if_part_exists_in_parts(parts: List[PartType], part: PartType, excluded_keys: List[str] = []) -> bool: +def check_if_part_exists_in_parts(parts: list[PartType], part: PartType, excluded_keys: list[str] = []) -> bool: """ Check if a part exists in a list of parts Handles both camelCase and snake_case key variations (e.g., function_call vs functionCall) @@ -689,11 +689,11 @@ def _collect_tool_call_thought_signatures( def _gemini_convert_messages_with_history( - messages: List[AllMessageValues], - model: Optional[str] = None, - litellm_params: Optional[dict] = None, - custom_llm_provider: Optional[str] = None, -) -> List[ContentType]: + messages: list[AllMessageValues], + model: str | None = None, + litellm_params: dict | None = None, + custom_llm_provider: str | None = None, +) -> list[ContentType]: """ Converts given messages from OpenAI format to Gemini format @@ -702,7 +702,7 @@ def _gemini_convert_messages_with_history( - Please ensure that function response turn comes immediately after a function call turn """ user_message_types = {"user", "system"} - contents: List[ContentType] = [] + contents: list[ContentType] = [] last_message_with_tool_calls = None @@ -720,13 +720,13 @@ def _gemini_convert_messages_with_history( try: while msg_i < len(messages): - user_content: List[PartType] = [] + user_content: list[PartType] = [] init_msg_i = msg_i ## MERGE CONSECUTIVE USER CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types: _message_content = messages[msg_i].get("content") if _message_content is not None and isinstance(_message_content, list): - _parts: List[PartType] = [] + _parts: list[PartType] = [] for element_idx, element in enumerate(_message_content): if element["type"] == "text" and "text" in element and len(element["text"]) > 0: element = cast(ChatCompletionTextObject, element) @@ -735,8 +735,8 @@ def _gemini_convert_messages_with_history( elif element["type"] == "image_url": element = cast(ChatCompletionImageObject, element) img_element = element - format: Optional[str] = None - media_resolution_enum: Optional[Dict[str, str]] = None + format: str | None = None + media_resolution_enum: dict[str, str] | None = None raw_image_url = img_element.get("image_url") if raw_image_url is None: raise litellm.BadRequestError( @@ -754,7 +754,7 @@ def _gemini_convert_messages_with_history( ) # TypedDict does not declare mime_type/content_type; # read via Dict[str, Any] for caller-provided MIME fields. - image_url_dict = cast(Dict[str, Any], raw_image_url) + image_url_dict = cast(dict[str, Any], raw_image_url) format = ( image_url_dict.get("format") or image_url_dict.get("mime_type") @@ -809,7 +809,7 @@ def _gemini_convert_messages_with_history( ) # TypedDict does not declare mime_type/content_type; # read via Dict[str, Any] for caller-provided MIME fields. - file_dict = cast(Dict[str, Any], _file_field) + file_dict = cast(dict[str, Any], _file_field) file_id = file_dict.get("file_id") format = ( file_dict.get("format") or file_dict.get("mime_type") or file_dict.get("content_type") @@ -844,7 +844,7 @@ def _gemini_convert_messages_with_history( f"{file_id or 'provided data'}, set this explicitly " f"using message[{msg_i}].content[{element_idx}].file.format " f"(or file.mime_type/content_type). " - f"Original error: {str(e)}" + f"Original error: {e!s}" ), model=model, llm_provider="vertex_ai", @@ -875,7 +875,7 @@ def _gemini_convert_messages_with_history( ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": if isinstance(messages[msg_i], BaseModel): - msg_dict: Union[ChatCompletionAssistantMessage, dict] = messages[msg_i].model_dump() # type: ignore + msg_dict: ChatCompletionAssistantMessage | dict = messages[msg_i].model_dump() # type: ignore else: msg_dict = messages[msg_i] # type: ignore assistant_msg = ChatCompletionAssistantMessage(**msg_dict) # type: ignore @@ -1005,7 +1005,7 @@ def _gemini_convert_messages_with_history( if isinstance(_ss_invocations, list): for invocation in _ss_invocations: # Re-inject toolCall part - tc_part: Dict[str, Any] = { + tc_part: dict[str, Any] = { "toolCall": { "toolType": invocation.get("tool_type"), "id": invocation.get("id"), @@ -1018,13 +1018,13 @@ def _gemini_convert_messages_with_history( # Re-inject toolResponse part if response is present if "response" in invocation: - tr_dict: Dict[str, Any] = { + tr_dict: dict[str, Any] = { "id": invocation.get("id"), "response": invocation.get("response"), } if invocation.get("tool_type"): tr_dict["toolType"] = invocation["tool_type"] - tr_part: Dict[str, Any] = {"toolResponse": tr_dict} + tr_part: dict[str, Any] = {"toolResponse": tr_dict} if "response_thought_signature" in invocation: tr_part["thoughtSignature"] = invocation["response_thought_signature"] assistant_content.append(tr_part) # type: ignore @@ -1055,9 +1055,7 @@ def _gemini_convert_messages_with_history( if msg_i == init_msg_i: # prevent infinite loops raise Exception( - "Invalid Message passed in - {}. File an issue https://github.com/BerriAI/litellm/issues".format( - messages[msg_i] - ) + f"Invalid Message passed in - {messages[msg_i]}. File an issue https://github.com/BerriAI/litellm/issues" ) if len(tool_call_responses) > 0: contents.append(ContentType(role="user", parts=tool_call_responses)) @@ -1083,7 +1081,7 @@ _LITELLM_INTERNAL_EXTRA_BODY_KEYS: frozenset = frozenset({"cache", "tags"}) def _pop_and_merge_extra_body(data: RequestBody, optional_params: dict) -> None: """Pop extra_body from optional_params and shallow-merge into data, deep-merging dict values.""" - extra_body: Optional[dict] = optional_params.pop("extra_body", None) + extra_body: dict | None = optional_params.pop("extra_body", None) if extra_body is not None: data_dict: dict = data # type: ignore[assignment] for k, v in extra_body.items(): @@ -1095,7 +1093,7 @@ def _pop_and_merge_extra_body(data: RequestBody, optional_params: dict) -> None: data_dict[k] = v -def _has_google_maps_tool(tools: Optional[Any]) -> bool: +def _has_google_maps_tool(tools: Any | None) -> bool: """Return True if any tool object in the list has a 'googleMaps' key.""" if not isinstance(tools, list): return False @@ -1132,14 +1130,14 @@ def _rewrite_mime_type_to_response_format(generation_config: GenerationConfig) - schema = generation_config.pop("response_schema", None) # type: ignore[misc] generation_config.pop("response_mime_type", None) # type: ignore[misc] - response_format: Dict[str, Any] = {"text": {"mimeType": "APPLICATION_JSON"}} + response_format: dict[str, Any] = {"text": {"mimeType": "APPLICATION_JSON"}} if schema is not None: response_format["text"]["schema"] = schema generation_config["responseFormat"] = response_format # type: ignore[typeddict-unknown-key] def _rewrite_google_maps_response_format(data: RequestBody) -> None: - generation_config = cast(Optional[GenerationConfig], data.get("generationConfig")) + generation_config = cast(GenerationConfig | None, data.get("generationConfig")) if ( isinstance(generation_config, dict) and _has_google_maps_tool(data.get("tools")) @@ -1149,12 +1147,12 @@ def _rewrite_google_maps_response_format(data: RequestBody) -> None: def _transform_request_body( - messages: List[AllMessageValues], + messages: list[AllMessageValues], model: str, optional_params: dict, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], litellm_params: dict, - cached_content: Optional[str], + cached_content: str | None, ) -> RequestBody: """ Common transformation logic across sync + async Gemini /generateContent calls. @@ -1194,10 +1192,10 @@ def _transform_request_body( content = litellm.VertexGeminiConfig()._transform_messages( messages=messages, model=model, litellm_params=litellm_params ) - tools: Optional[Tools] = optional_params.pop("tools", None) - tool_choice: Optional[ToolConfig] = optional_params.pop("tool_choice", None) + tools: Tools | None = optional_params.pop("tools", None) + tool_choice: ToolConfig | None = optional_params.pop("tool_choice", None) include_server_side_tool_invocations: bool = optional_params.pop("include_server_side_tool_invocations", False) - safety_settings: Optional[List[SafetSettingsConfig]] = optional_params.pop("safety_settings", None) # type: ignore + safety_settings: list[SafetSettingsConfig] | None = optional_params.pop("safety_settings", None) # type: ignore # Drop output_config as it's not supported by Vertex AI optional_params.pop("output_config", None) config_fields = GenerationConfig.__annotations__.keys() @@ -1207,7 +1205,7 @@ def _transform_request_body( filtered_params = {k: v for k, v in optional_params.items() if _get_equivalent_key(k, set(config_fields))} - generation_config: Optional[GenerationConfig] = GenerationConfig(**filtered_params) + generation_config: GenerationConfig | None = GenerationConfig(**filtered_params) # For Gemini 2.x models, also add media_resolution to generation_config (global) # as a fallback, since some 2.x versions may not support per-part media_resolution. @@ -1262,20 +1260,20 @@ def _transform_request_body( def sync_transform_request_body( - gemini_api_key: Optional[str], - messages: List[AllMessageValues], - api_base: Optional[str], + gemini_api_key: str | None, + messages: list[AllMessageValues], + api_base: str | None, model: str, - client: Optional[HTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[dict], + client: HTTPHandler | None, + timeout: float | httpx.Timeout | None, + extra_headers: dict | None, optional_params: dict, logging_obj: LiteLLMLoggingObj, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], litellm_params: dict, - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, ) -> RequestBody: from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints @@ -1313,20 +1311,20 @@ def sync_transform_request_body( async def async_transform_request_body( - gemini_api_key: Optional[str], - messages: List[AllMessageValues], - api_base: Optional[str], + gemini_api_key: str | None, + messages: list[AllMessageValues], + api_base: str | None, model: str, - client: Optional[AsyncHTTPHandler], - timeout: Optional[Union[float, httpx.Timeout]], - extra_headers: Optional[dict], + client: AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + extra_headers: dict | None, optional_params: dict, logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], litellm_params: dict, - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_auth_header: Optional[str], + vertex_project: str | None, + vertex_location: str | None, + vertex_auth_header: str | None, ) -> RequestBody: from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints @@ -1387,8 +1385,8 @@ def _default_user_message_when_system_message_passed() -> ChatCompletionUserMess def _transform_system_message( - supports_system_message: bool, messages: List[AllMessageValues] -) -> Tuple[Optional[SystemInstructions], List[AllMessageValues]]: + supports_system_message: bool, messages: list[AllMessageValues] +) -> tuple[SystemInstructions | None, list[AllMessageValues]]: """ Extracts the system message from the openai message list. @@ -1400,11 +1398,11 @@ def _transform_system_message( """ # Separate system prompt from rest of message system_prompt_indices = [] - system_content_blocks: List[PartType] = [] + system_content_blocks: list[PartType] = [] if supports_system_message is True: for idx, message in enumerate(messages): if message["role"] == "system": - _system_content_block: Optional[PartType] = None + _system_content_block: PartType | None = None if isinstance(message["content"], str): _system_content_block = PartType(text=message["content"]) elif isinstance(message["content"], list): diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 5ed57bee6c8..19c43d8000c 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -9,12 +9,8 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Type, Union, cast, ) @@ -134,7 +130,7 @@ class VertexAIBaseConfig: optional_params[mapped_params[param]] = value return optional_params - def get_eu_regions(self) -> List[str]: + def get_eu_regions(self) -> list[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -151,7 +147,7 @@ class VertexAIBaseConfig: "europe-west9", ] - def get_us_regions(self) -> List[str]: + def get_us_regions(self) -> list[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -197,29 +193,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Note: Please make sure to modify the default parameters as required for your use case. """ - temperature: Optional[float] = None - max_output_tokens: Optional[int] = None - top_p: Optional[float] = None - top_k: Optional[int] = None - response_mime_type: Optional[str] = None - candidate_count: Optional[int] = None - stop_sequences: Optional[list] = None - frequency_penalty: Optional[float] = None - presence_penalty: Optional[float] = None - seed: Optional[int] = None + temperature: float | None = None + max_output_tokens: int | None = None + top_p: float | None = None + top_k: int | None = None + response_mime_type: str | None = None + candidate_count: int | None = None + stop_sequences: list | None = None + frequency_penalty: float | None = None + presence_penalty: float | None = None + seed: int | None = None def __init__( self, - temperature: Optional[float] = None, - max_output_tokens: Optional[int] = None, - top_p: Optional[float] = None, - top_k: Optional[int] = None, - response_mime_type: Optional[str] = None, - candidate_count: Optional[int] = None, - stop_sequences: Optional[list] = None, - frequency_penalty: Optional[float] = None, - presence_penalty: Optional[float] = None, - seed: Optional[int] = None, + temperature: float | None = None, + max_output_tokens: int | None = None, + top_p: float | None = None, + top_k: int | None = None, + response_mime_type: str | None = None, + candidate_count: int | None = None, + stop_sequences: list | None = None, + frequency_penalty: float | None = None, + presence_penalty: float | None = None, + seed: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -230,9 +226,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def get_config(cls): return super().get_config() - def get_json_schema_from_pydantic_object( - self, response_format: Optional[Union[Type["BaseModel"], dict]] - ) -> Optional[dict]: + def get_json_schema_from_pydantic_object(self, response_format: type["BaseModel"] | dict | None) -> dict | None: """ Override to use Pydantic's model_json_schema() instead of OpenAI's to_strict_json_schema(). @@ -306,7 +300,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False return True - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: supported_params = [ "temperature", "top_p", @@ -341,7 +335,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): supported_params.append("thinking") return supported_params - def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> Optional[ToolConfig]: + def map_tool_choice_values(self, model: str, tool_choice: str | dict) -> ToolConfig | None: if tool_choice == "none": return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE")) elif tool_choice == "required": @@ -354,9 +348,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="ANY", allowed_function_names=[name])) else: raise litellm.utils.UnsupportedParamsError( - message="VertexAI doesn't support tool_choice={}. Supported tool_choice values=['auto', 'required', json object]. To drop it from the call, set `litellm.drop_params = True.".format( - tool_choice - ), + message=f"VertexAI doesn't support tool_choice={tool_choice}. Supported tool_choice values=['auto', 'required', json object]. To drop it from the call, set `litellm.drop_params = True.", status_code=400, ) @@ -467,7 +459,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return transformed_config - def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> Tuple[dict, Optional[dict]]: + def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> tuple[dict, dict | None]: """ Extract location configuration from googleMaps tool for Vertex AI toolConfig. @@ -505,7 +497,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return cleaned_config, retrieval_config - def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]: + def get_tool_value(self, tool: dict, tool_name: str) -> dict | None: """ Helper function to get tool value handling both camelCase and underscore_case variants @@ -530,10 +522,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _resolve_search_tool_conflict( gtool_func_declarations: list, - googleSearch: Optional[dict], - googleSearchRetrieval: Optional[dict], - enterpriseWebSearch: Optional[dict], - urlContext: Optional[dict], + googleSearch: dict | None, + googleSearchRetrieval: dict | None, + enterpriseWebSearch: dict | None, + urlContext: dict | None, optional_params: dict, ) -> tuple: """ @@ -577,7 +569,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext - def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]: + def _map_function(self, value: list[dict], optional_params: dict) -> list[Tools]: """ Map OpenAI-style tools/functions to Vertex AI format. @@ -593,21 +585,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): googleMaps tools contain location data """ gtool_func_declarations = [] - googleSearch: Optional[dict] = None - googleSearchRetrieval: Optional[dict] = None - enterpriseWebSearch: Optional[dict] = None - urlContext: Optional[dict] = None - code_execution: Optional[dict] = None - googleMaps: Optional[dict] = None - google_maps_retrieval_config: Optional[dict] = None - computerUse: Optional[dict] = None + googleSearch: dict | None = None + googleSearchRetrieval: dict | None = None + enterpriseWebSearch: dict | None = None + urlContext: dict | None = None + code_execution: dict | None = None + googleMaps: dict | None = None + google_maps_retrieval_config: dict | None = None + computerUse: dict | None = None # remove 'additionalProperties' from tools value = _remove_additional_properties(value) # remove 'strict' from tools value = _remove_strict_from_schema(value) for tool in value: - openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = None + openai_function_object: ChatCompletionToolParamFunctionChunk | None = None if "function" in tool: # tools list _openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore **tool["function"] @@ -697,7 +689,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Build list of Tool objects - each Tool should contain exactly one type # per Vertex AI API spec: "A Tool object should contain exactly one type of Tool" - _tools_list: List[Tools] = [] + _tools_list: list[Tools] = [] ( googleSearch, @@ -815,7 +807,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: if reasoning_effort == "minimal": # Use model-specific minimum thinking budget or fallback @@ -864,7 +856,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_level( reasoning_effort: str, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: """ Map reasoning_effort to thinking_level for Gemini 3+ models. @@ -910,12 +902,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raise ValueError(f"Invalid reasoning effort: {reasoning_effort}") @staticmethod - def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool: + def _is_thinking_budget_zero(thinking_budget: int | None) -> bool: return thinking_budget is not None and thinking_budget == 0 @staticmethod def _validate_thinking_config_conflicts( - optional_params: Dict, + optional_params: dict, param_name: str, param_description: str = "thinking_budget", ) -> None: @@ -936,7 +928,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _validate_thinking_level_conflicts( - optional_params: Dict, + optional_params: dict, ) -> None: """ Validate that thinking_level and thinking_budget are not both specified. @@ -956,7 +948,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_thinking_param( thinking_param: AnthropicThinkingParam, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: thinking_enabled = thinking_param.get("type") == "enabled" thinking_budget = thinking_param.get("budget_tokens") @@ -1063,8 +1055,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _apply_include_server_side_tool_invocations( - non_default_params: Dict, - optional_params: Dict, + non_default_params: dict, + optional_params: dict, ) -> None: """ Set include_server_side_tool_invocations before tools are mapped. @@ -1083,11 +1075,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def map_openai_params( self, - non_default_params: Dict, - optional_params: Dict, + non_default_params: dict, + optional_params: dict, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: self._apply_include_server_side_tool_invocations(non_default_params, optional_params) gemini_sampling_params_warned: bool = False for param, value in non_default_params.items(): @@ -1179,7 +1171,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif param == "reasoning_effort": # Extract effort value - handle both string and dict formats # Dict format comes from OpenAI Agents SDK: {"effort": "high", "summary": "auto"} - effort_value: Optional[str] = None + effort_value: str | None = None if isinstance(value, str): effort_value = value elif isinstance(value, dict): @@ -1255,7 +1247,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params[mapped_params[param]] = value return optional_params - def get_eu_regions(self) -> List[str]: + def get_eu_regions(self) -> list[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -1292,7 +1284,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model @staticmethod - def _is_model_gemini_spec_model(model: Optional[str]) -> bool: + def _is_model_gemini_spec_model(model: str | None) -> bool: """ Returns true if user is trying to call custom model in `/gemini` request/response format """ @@ -1315,7 +1307,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model.split("/")[-1] return model - def get_flagged_finish_reasons(self) -> Dict[str, str]: + def get_flagged_finish_reasons(self) -> dict[str, str]: """ Return Dictionary of finish reasons which indicate response was flagged @@ -1352,7 +1344,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) @staticmethod - def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]: + def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]: """ Return Dictionary of Gemini/Vertex AI finish reasons and their OpenAI-compatible mappings. @@ -1366,14 +1358,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "GenerateContentRequest.tools[0].function_declarations[0].parameters.properties: should be non-empty for OBJECT type" in exception_string ): - return "'properties' field in tools[0]['function']['parameters'] cannot be empty if 'type' == 'object'. Received error from provider - {}".format( - exception_string - ) + return f"'properties' field in tools[0]['function']['parameters'] cannot be empty if 'type' == 'object'. Received error from provider - {exception_string}" return exception_string - def get_assistant_content_message(self, parts: List[HttpxPartType]) -> Tuple[Optional[str], Optional[str]]: - content_str: Optional[str] = None - reasoning_content_str: Optional[str] = None + def get_assistant_content_message(self, parts: list[HttpxPartType]) -> tuple[str | None, str | None]: + content_str: str | None = None + reasoning_content_str: str | None = None for part in parts: _content_str = "" @@ -1398,7 +1388,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Images and audio are now handled separately in their respective response fields if mime_type.startswith("audio/") or mime_type.startswith("image/"): continue - _content_str += "data:{};base64,{}".format(mime_type, data) + _content_str += f"data:{mime_type};base64,{data}" if len(_content_str) > 0: if part.get("thought") is True: @@ -1412,7 +1402,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return content_str, reasoning_content_str - def _extract_thinking_blocks_from_parts(self, parts: List[HttpxPartType]) -> List[ChatCompletionThinkingBlock]: + def _extract_thinking_blocks_from_parts(self, parts: list[HttpxPartType]) -> list[ChatCompletionThinkingBlock]: """Extract thinking blocks from parts if present. Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking): @@ -1421,7 +1411,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): it does NOT indicate that the content is thinking (a part can have thoughtSignature without thought: true, e.g., function calls) """ - thinking_blocks: List[ChatCompletionThinkingBlock] = [] + thinking_blocks: list[ChatCompletionThinkingBlock] = [] for part in parts: if part.get("thought") is True: thinking_text = part.get("text", "") @@ -1435,7 +1425,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): thinking_blocks.append(block) return thinking_blocks - def _extract_thought_signatures_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[str]]: + def _extract_thought_signatures_from_parts(self, parts: list[HttpxPartType]) -> list[str] | None: """Extract thoughtSignature values from parts. Per Google's docs, thoughtSignature is returned for multi-turn context preservation @@ -1445,7 +1435,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Returns: List of thoughtSignature strings if any are found, None otherwise """ - signatures: List[str] = [] + signatures: list[str] = [] for part in parts: signature = part.get("thoughtSignature") if signature is not None: @@ -1454,8 +1444,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _extract_server_side_tool_invocations( - parts: List[HttpxPartType], - ) -> Optional[List[Dict[str, Any]]]: + parts: list[HttpxPartType], + ) -> list[dict[str, Any]] | None: """Extract server-side tool invocations (toolCall/toolResponse) from parts. These are returned by Gemini when context circulation is enabled @@ -1466,15 +1456,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Returns: List of server-side invocation dicts if any found, None otherwise. """ - invocations: List[Dict[str, Any]] = [] + invocations: list[dict[str, Any]] = [] # Index toolCalls by id so we can pair them with responses - tool_calls_by_id: Dict[str, Dict[str, Any]] = {} - tool_responses_by_id: Dict[str, Dict[str, Any]] = {} + tool_calls_by_id: dict[str, dict[str, Any]] = {} + tool_responses_by_id: dict[str, dict[str, Any]] = {} for part in parts: if "toolCall" in part: tc = part["toolCall"] - entry: Dict[str, Any] = { + entry: dict[str, Any] = { "tool_type": tc.get("toolType"), "id": tc.get("id"), "args": tc.get("args"), @@ -1514,9 +1504,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return invocations if invocations else None - def _extract_image_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[ImageURLListItem]]: + def _extract_image_response_from_parts(self, parts: list[HttpxPartType]) -> list[ImageURLListItem] | None: """Extract image response from parts if present""" - images: List[ImageURLListItem] = [] + images: list[ImageURLListItem] = [] for part in parts: if "inlineData" in part: inline_data = part.get("inlineData", {}) @@ -1534,7 +1524,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) return images - def _extract_audio_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[ChatCompletionAudioResponse]: + def _extract_audio_response_from_parts(self, parts: list[HttpxPartType]) -> ChatCompletionAudioResponse | None: """Extract audio response from parts if present""" for part in parts: if "text" in part: @@ -1572,16 +1562,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _transform_parts( - parts: List[HttpxPartType], + parts: list[HttpxPartType], cumulative_tool_call_idx: int, - is_function_call: Optional[bool], - ) -> Tuple[ - Optional[ChatCompletionToolCallFunctionChunk], - Optional[List[ChatCompletionToolCallChunk]], + is_function_call: bool | None, + ) -> tuple[ + ChatCompletionToolCallFunctionChunk | None, + list[ChatCompletionToolCallChunk] | None, int, ]: - function: Optional[ChatCompletionToolCallFunctionChunk] = None - _tools: List[ChatCompletionToolCallChunk] = [] + function: ChatCompletionToolCallFunctionChunk | None = None + _tools: list[ChatCompletionToolCallChunk] = [] for part in parts: if "functionCall" in part: _function_chunk: ChatCompletionToolCallFunctionChunk = { @@ -1596,7 +1586,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_call_id = part["functionCall"].get("id") if is_function_call is True: - function_dict: Dict[str, Any] = dict(_function_chunk) + function_dict: dict[str, Any] = dict(_function_chunk) if thought_signature: if "provider_specific_fields" not in function_dict: function_dict["provider_specific_fields"] = {} @@ -1625,22 +1615,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tools.append(_tool_response_chunk) cumulative_tool_call_idx += 1 if len(_tools) == 0: - tools: Optional[List[ChatCompletionToolCallChunk]] = None + tools: list[ChatCompletionToolCallChunk] | None = None else: tools = _tools return function, tools, cumulative_tool_call_idx @staticmethod def _transform_logprobs( - logprobs_result: Optional[LogprobsResult], - ) -> Optional[ChoiceLogprobs]: + logprobs_result: LogprobsResult | None, + ) -> ChoiceLogprobs | None: if logprobs_result is None: return None if "chosenCandidates" not in logprobs_result: return None - logprobs_list: List[ChatCompletionTokenLogprob] = [] + logprobs_list: list[ChatCompletionTokenLogprob] = [] for index, candidate in enumerate(logprobs_result["chosenCandidates"]): - top_logprobs: List[TopLogprob] = [] + top_logprobs: list[TopLogprob] = [] if "topCandidates" in logprobs_result and index < len(logprobs_result["topCandidates"]): top_candidates_for_index = logprobs_result["topCandidates"][index]["candidates"] @@ -1743,7 +1733,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _response_has_search_grounding( - completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage], + completion_response: GenerateContentResponseBody | BidiGenerateContentServerMessage, ) -> bool: """ Whether the response used Grounding with Google Search, detected via @@ -1767,20 +1757,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _calculate_usage( - completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage], + completion_response: GenerateContentResponseBody | BidiGenerateContentServerMessage, ) -> Usage: if completion_response is not None and "usageMetadata" not in completion_response: raise ValueError(f"usageMetadata not found in completion_response. Got={completion_response}") - cached_tokens: Optional[int] = None + cached_tokens: int | None = None # Separate variables for prompt tokens by modality - prompt_audio_tokens: Optional[int] = None - prompt_image_tokens: Optional[int] = None - prompt_text_tokens: Optional[int] = None - prompt_video_tokens: Optional[int] = None - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None - reasoning_tokens: Optional[int] = None - response_tokens: Optional[int] = None - response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None + prompt_audio_tokens: int | None = None + prompt_image_tokens: int | None = None + prompt_text_tokens: int | None = None + prompt_video_tokens: int | None = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None + reasoning_tokens: int | None = None + response_tokens: int | None = None + response_tokens_details: CompletionTokensDetailsWrapper | None = None usage_metadata = completion_response["usageMetadata"] def _get_token_count(detail: Mapping[str, Any]) -> int: @@ -1859,10 +1849,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## Parse cacheTokensDetails (breakdown of cached tokens by modality) ## When explicit caching is used, Gemini provides this field to show which modalities were cached - cached_text_tokens: Optional[int] = None - cached_audio_tokens: Optional[int] = None - cached_image_tokens: Optional[int] = None - cached_video_tokens: Optional[int] = None + cached_text_tokens: int | None = None + cached_audio_tokens: int | None = None + cached_image_tokens: int | None = None + cached_video_tokens: int | None = None if "cacheTokensDetails" in usage_metadata: for detail in usage_metadata["cacheTokensDetails"]: @@ -1944,8 +1934,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_finish_reason( - chat_completion_message: Optional[ChatCompletionResponseMessage], - finish_reason: Optional[str], + chat_completion_message: ChatCompletionResponseMessage | None, + finish_reason: str | None, ) -> OpenAIChatCompletionFinishReason: from litellm.litellm_core_utils.core_helpers import map_finish_reason @@ -1961,7 +1951,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_prompt_level_content_filter( processed_chunk: GenerateContentResponseBody, - response_id: Optional[str], + response_id: str | None, ) -> Optional["ModelResponseStream"]: """ Check if prompt is blocked due to content filtering at the prompt level. @@ -2005,8 +1995,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return None @staticmethod - def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]: - web_search_requests: Optional[int] = None + def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: + web_search_requests: int | None = None if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0: for grounding_metadata_item in grounding_metadata: @@ -2022,10 +2012,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): chat_completion_message: ChatCompletionResponseMessage, candidate: Candidates, idx: int, - tools: Optional[List[ChatCompletionToolCallChunk]], - functions: Optional[ChatCompletionToolCallFunctionChunk], - chat_completion_logprobs: Optional[ChoiceLogprobs], - image_response: Optional[List[ImageURLListItem]], + tools: list[ChatCompletionToolCallChunk] | None, + functions: ChatCompletionToolCallFunctionChunk | None, + chat_completion_logprobs: ChoiceLogprobs | None, + image_response: list[ImageURLListItem] | None, ) -> StreamingChoices: """ Helper method to create a streaming choice object for Vertex AI @@ -2057,7 +2047,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _extract_candidate_metadata( candidate: Candidates, - ) -> Tuple[List[dict], List[dict], List, List]: + ) -> tuple[list[dict], list[dict], list, list]: """ Extract metadata from a single candidate response. @@ -2067,10 +2057,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): safety_ratings: List citation_metadata: List """ - grounding_metadata: List[dict] = [] - url_context_metadata: List[dict] = [] - safety_ratings: List = [] - citation_metadata: List = [] + grounding_metadata: list[dict] = [] + url_context_metadata: list[dict] = [] + safety_ratings: list = [] + citation_metadata: list = [] if "groundingMetadata" in candidate: if isinstance(candidate["groundingMetadata"], list): @@ -2115,10 +2105,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _set_stream_metadata_on_response( model_response: Any, - grounding_metadata: List[dict], - url_context_metadata: List[dict], - safety_ratings: List[dict], - citation_metadata: List[dict], + grounding_metadata: list[dict], + url_context_metadata: list[dict], + safety_ratings: list[dict], + citation_metadata: list[dict], ) -> None: setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore if grounding_metadata: @@ -2138,10 +2128,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def apply_assembled_streaming_response_metadata( self, response: ModelResponse, - chunks: List[Any], + chunks: list[Any], ) -> None: for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS: - merged: List[Any] = [] + merged: list[Any] = [] for chunk in chunks: value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name) if not value: @@ -2156,14 +2146,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _convert_grounding_metadata_to_annotations( - grounding_metadata: List[dict], - content_text: Optional[str], - ) -> List[ChatCompletionAnnotation]: + grounding_metadata: list[dict], + content_text: str | None, + ) -> list[ChatCompletionAnnotation]: """ Convert Vertex AI grounding metadata to OpenAI-style annotations. """ - annotations: List[ChatCompletionAnnotation] = [] + annotations: list[ChatCompletionAnnotation] = [] for metadata in grounding_metadata: # Extract groundingSupports - these map text segments to sources @@ -2171,7 +2161,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): grounding_chunks = metadata.get("groundingChunks", []) # Build a map of chunk indices to web URIs - chunk_to_uri_map: Dict[int, Dict[str, str]] = {} + chunk_to_uri_map: dict[int, dict[str, str]] = {} for idx, chunk in enumerate(grounding_chunks): if "web" in chunk: web_data = chunk["web"] @@ -2211,11 +2201,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _process_candidates( - _candidates: List[Candidates], + _candidates: list[Candidates], model_response: Union[ModelResponse, "ModelResponseStream"], standard_optional_params: dict, cumulative_tool_call_index: int = 0, - ) -> Tuple[List[dict], List[dict], List, List, int]: + ) -> tuple[list[dict], list[dict], list, list, int]: """ Helper method to process candidates and extract metadata @@ -2231,19 +2221,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) from litellm.types.utils import ModelResponseStream - grounding_metadata: List[dict] = [] - url_context_metadata: List[dict] = [] - image_response: Optional[List[ImageURLListItem]] = None - safety_ratings: List = [] - citation_metadata: List = [] + grounding_metadata: list[dict] = [] + url_context_metadata: list[dict] = [] + image_response: list[ImageURLListItem] | None = None + safety_ratings: list = [] + citation_metadata: list = [] chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} - chat_completion_logprobs: Optional[ChoiceLogprobs] = None - tools: Optional[List[ChatCompletionToolCallChunk]] = [] - functions: Optional[ChatCompletionToolCallFunctionChunk] = None - thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None - reasoning_content: Optional[str] = None - thought_signatures: Optional[Any] = None - server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None + chat_completion_logprobs: ChoiceLogprobs | None = None + tools: list[ChatCompletionToolCallChunk] | None = [] + functions: ChatCompletionToolCallFunctionChunk | None = None + thinking_blocks: list[ChatCompletionThinkingBlock] | None = None + reasoning_content: str | None = None + thought_signatures: Any | None = None + server_side_tool_invocations: list[dict[str, Any]] | None = None for idx, candidate in enumerate(_candidates): if "content" not in candidate: @@ -2290,11 +2280,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) if audio_response is not None: - cast(Dict[str, Any], chat_completion_message)["audio"] = audio_response + cast(dict[str, Any], chat_completion_message)["audio"] = audio_response chat_completion_message["content"] = None # OpenAI spec if image_response is not None: # Handle image response - combine with text content into structured format - cast(Dict[str, Any], chat_completion_message)["images"] = image_response + cast(dict[str, Any], chat_completion_message)["images"] = image_response if content is not None: chat_completion_message["content"] = content @@ -2394,13 +2384,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2415,9 +2405,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): completion_response = GenerateContentResponseBody(**raw_response.json()) # type: ignore except Exception as e: raise VertexAIError( - message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format( - str(e) - ), + message=f"Error converting to valid response block={e!s}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues", status_code=422, headers=raw_response.headers, ) @@ -2432,7 +2420,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _transform_google_generate_content_to_openai_model_response( self, - completion_response: Union[GenerateContentResponseBody, dict], + completion_response: GenerateContentResponseBody | dict, model_response: ModelResponse, model: str, logging_obj: LoggingClass, @@ -2467,11 +2455,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): response_id = completion_response.get("responseId") if response_id: model_response.id = response_id - url_context_metadata: List[dict] = [] + url_context_metadata: list[dict] = [] try: - grounding_metadata: List[dict] = [] - safety_ratings: List[dict] = [] - citation_metadata: List[dict] = [] + grounding_metadata: list[dict] = [] + safety_ratings: list[dict] = [] + citation_metadata: list[dict] = [] if _candidates: ( grounding_metadata, @@ -2524,9 +2512,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): except Exception as e: raise VertexAIError( - message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format( - str(e) - ), + message=f"Error converting to valid response block={e!s}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues", status_code=422, headers=raw_response.headers, ) @@ -2535,10 +2521,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _transform_messages( self, - messages: List[AllMessageValues], - model: Optional[str] = None, - litellm_params: Optional[dict] = None, - ) -> List[ContentType]: + messages: list[AllMessageValues], + model: str | None = None, + litellm_params: dict | None = None, + ) -> list[ContentType]: return _gemini_convert_messages_with_history( messages=messages, model=model, @@ -2546,31 +2532,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): custom_llm_provider="vertex_ai", ) - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VertexAIError(message=error_message, status_code=status_code, headers=headers) def transform_request( self, model: str, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, - headers: Dict, - ) -> Dict: + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: raise NotImplementedError("Vertex AI has a custom implementation of transform_request. Needs sync + async.") def validate_environment( self, - headers: Optional[Dict], + headers: dict | None, model: str, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, - api_key: Optional[Union[str, Dict]] = None, - api_base: Optional[str] = None, - ) -> Dict: + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: str | dict | None = None, + api_base: str | None = None, + ) -> dict: default_headers = { "Content-Type": "application/json", } @@ -2585,8 +2569,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): async def make_call( - client: Optional[AsyncHTTPHandler], # module-level client - gemini_client: Optional[AsyncHTTPHandler], # if passed by user + client: AsyncHTTPHandler | None, # module-level client + gemini_client: AsyncHTTPHandler | None, # if passed by user api_base: str, headers: dict, data: str, @@ -2637,8 +2621,8 @@ async def make_call( def make_sync_call( - client: Optional[HTTPHandler], # module-level client - gemini_client: Optional[HTTPHandler], # if passed by user + client: HTTPHandler | None, # module-level client + gemini_client: HTTPHandler | None, # if passed by user api_base: str, headers: dict, data: str, @@ -2693,20 +2677,20 @@ class VertexLLM(VertexBase): model_response: ModelResponse, print_verbose: Callable, data: dict, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=None, - api_base: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, - gemini_api_key: Optional[str] = None, - extra_headers: Optional[dict] = None, + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, + gemini_api_key: str | None = None, + extra_headers: dict | None = None, ) -> CustomStreamWrapper: should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) @@ -2789,21 +2773,21 @@ class VertexLLM(VertexBase): custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=None, - api_base: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, - gemini_api_key: Optional[str] = None, - extra_headers: Optional[dict] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, + gemini_api_key: str | None = None, + extra_headers: dict | None = None, + ) -> ModelResponse | CustomStreamWrapper: should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) _auth_header, vertex_project = await self._ensure_access_token_async( @@ -2911,18 +2895,18 @@ class VertexLLM(VertexBase): logging_obj, optional_params: dict, acompletion: bool, - timeout: Optional[Union[float, httpx.Timeout]], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - gemini_api_key: Optional[str], + timeout: float | httpx.Timeout | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + gemini_api_key: str | None, litellm_params: dict, logger_fn=None, - extra_headers: Optional[dict] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - api_base: Optional[str] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore + extra_headers: dict | None = None, + client: AsyncHTTPHandler | HTTPHandler | None = None, + api_base: str | None = None, + ) -> ModelResponse | CustomStreamWrapper: + stream: bool | None = optional_params.pop("stream", None) # type: ignore transform_request_params = { "gemini_api_key": gemini_api_key, @@ -3110,7 +3094,7 @@ class ModelResponseIterator: streaming_response, sync_stream: bool, logging_obj: LoggingClass, - response_headers: Optional[Dict[str, str]] = None, + response_headers: dict[str, str] | None = None, response: httpx.Response | None = None, ): from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -3155,9 +3139,9 @@ class ModelResponseIterator: def _apply_stream_candidates( self, - _candidates: List[Candidates], + _candidates: list[Candidates], model_response: Any, - ) -> Tuple[List[dict], List[dict], List[dict], List[dict]]: + ) -> tuple[list[dict], list[dict], list[dict], list[dict]]: ( grounding_metadata, url_context_metadata, @@ -3236,8 +3220,8 @@ class ModelResponseIterator: self, processed_chunk: Any, model_response: Any, - grounding_metadata: List[dict], - ) -> Optional[Usage]: + grounding_metadata: list[dict], + ) -> Usage | None: if "usageMetadata" not in processed_chunk: return None @@ -3284,12 +3268,12 @@ class ModelResponseIterator: if blocked_response is not None: model_response = blocked_response - grounding_metadata: List[dict] = [] - url_context_metadata: List[dict] = [] - safety_ratings: List[dict] = [] - citation_metadata: List[dict] = [] + grounding_metadata: list[dict] = [] + url_context_metadata: list[dict] = [] + safety_ratings: list[dict] = [] + citation_metadata: list[dict] = [] - _candidates: Optional[List[Candidates]] = processed_chunk.get("candidates") + _candidates: list[Candidates] | None = processed_chunk.get("candidates") if _candidates: ( grounding_metadata, diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index d989750a5f3..858cb116a99 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint """ import json -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Any, Literal import httpx @@ -34,7 +34,7 @@ class GoogleBatchEmbeddings(VertexLLM): @staticmethod def _flatten_and_detect_file_refs( input: GeminiEmbeddingInput, - ) -> Tuple[List[str], bool]: + ) -> tuple[list[str], bool]: """Flatten nested input lists and detect file references.""" input_list = [input] if isinstance(input, str) else input flat_elements = [ @@ -48,7 +48,7 @@ class GoogleBatchEmbeddings(VertexLLM): input: GeminiEmbeddingInput, api_key: str, sync_handler: HTTPHandler, - ) -> Dict[str, Dict[str, str]]: + ) -> dict[str, dict[str, str]]: """ Resolve Gemini file references (files/...) to get mime_type and uri. @@ -61,7 +61,7 @@ class GoogleBatchEmbeddings(VertexLLM): Dict mapping file name to {mime_type, uri} """ input_list = [input] if isinstance(input, str) else input - resolved_files: Dict[str, Dict[str, str]] = {} + resolved_files: dict[str, dict[str, str]] = {} for element in input_list: if isinstance(element, str) and _is_file_reference(element): @@ -85,7 +85,7 @@ class GoogleBatchEmbeddings(VertexLLM): input: GeminiEmbeddingInput, api_key: str, async_handler: AsyncHTTPHandler, - ) -> Dict[str, Dict[str, str]]: + ) -> dict[str, dict[str, str]]: """ Async version of _resolve_file_references. @@ -98,7 +98,7 @@ class GoogleBatchEmbeddings(VertexLLM): Dict mapping file name to {mime_type, uri} """ input_list = [input] if isinstance(input, str) else input - resolved_files: Dict[str, Dict[str, str]] = {} + resolved_files: dict[str, dict[str, str]] = {} for element in input_list: if isinstance(element, str) and _is_file_reference(element): @@ -126,16 +126,16 @@ class GoogleBatchEmbeddings(VertexLLM): custom_llm_provider: Literal["gemini", "vertex_ai"], optional_params: dict, logging_obj: Any, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, encoding=None, vertex_project=None, vertex_location=None, vertex_credentials=None, - aembedding: Optional[bool] = False, + aembedding: bool | None = False, timeout=300, client=None, - extra_headers: Optional[dict] = None, + extra_headers: dict | None = None, ) -> EmbeddingResponse: _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -279,18 +279,18 @@ class GoogleBatchEmbeddings(VertexLLM): async def async_batch_embeddings( self, model: str, - api_base: Optional[str], + api_base: str | None, url: str, - data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]], + data: VertexAIBatchEmbeddingsRequestBody | dict | None, model_response: EmbeddingResponse, input: GeminiEmbeddingInput, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, headers={}, - client: Optional[AsyncHTTPHandler] = None, + client: AsyncHTTPHandler | None = None, use_embed_content: bool = False, - api_key: Optional[str] = None, - optional_params: Optional[dict] = None, - logging_obj: Optional[Any] = None, + api_key: str | None = None, + optional_params: dict | None = None, + logging_obj: Any | None = None, ) -> EmbeddingResponse: if client is None: _params = {} diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index c205c40d707..80b57178e47 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -5,7 +5,6 @@ Why separate file? Make it easy to see how transformation works """ from collections.abc import Mapping, Sequence -from typing import Dict, List, Optional, Tuple from pydantic import TypeAdapter, ValidationError @@ -85,7 +84,7 @@ def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: ) -def _parse_data_url(data_url: str) -> Tuple[str, str]: +def _parse_data_url(data_url: str) -> tuple[str, str]: """ Parse a data URL to extract the media type and base64 data. @@ -161,7 +160,7 @@ def _is_multimodal_element(element: str) -> bool: def _build_part_for_input( element: str, - resolved_files: Optional[Dict[str, Dict[str, str]]] = None, + resolved_files: dict[str, dict[str, str]] | None = None, ) -> PartType: """ Build a single PartType for an input element, handling text, data URIs, @@ -210,7 +209,7 @@ def transform_openai_input_gemini_content( input: GeminiEmbeddingInput, model: str, optional_params: dict, - resolved_files: Optional[Dict[str, Dict[str, str]]] = None, + resolved_files: dict[str, dict[str, str]] | None = None, ) -> VertexAIBatchEmbeddingsRequestBody: """ Transform OpenAI embedding input to Gemini batchEmbedContents format. @@ -227,12 +226,12 @@ def transform_openai_input_gemini_content( input=[["text", "image"]] → 1 combined embedding input=[["text", "image"], "x"] → 2 embeddings (1 combined + 1 separate) """ - gemini_model_name = "models/{}".format(model) + gemini_model_name = f"models/{model}" gemini_params = _filter_embed_params(optional_params) input_list = [input] if isinstance(input, str) else input - requests: List[EmbedContentRequest] = [] + requests: list[EmbedContentRequest] = [] for element in input_list: if isinstance(element, list): @@ -258,7 +257,7 @@ def transform_openai_input_gemini_embed_content( input: GeminiEmbeddingInput, model: str, optional_params: dict, - resolved_files: Optional[Dict[str, Dict[str, str]]] = None, + resolved_files: dict[str, dict[str, str]] | None = None, ) -> dict: """ Transform OpenAI embedding input to Gemini embedContent format (multimodal). @@ -277,7 +276,7 @@ def transform_openai_input_gemini_embed_content( gemini_params = _filter_embed_params(optional_params) input_list = [input] if isinstance(input, str) else input - parts: List[PartType] = [] + parts: list[PartType] = [] for element in input_list: if isinstance(element, list): @@ -303,7 +302,7 @@ _AUDIO_TOKENS_PER_SECOND = 32.0 _usage_metadata_adapter = TypeAdapter(UsageMetadata) -def _parse_usage_metadata(raw_usage_metadata: object) -> Optional[UsageMetadata]: +def _parse_usage_metadata(raw_usage_metadata: object) -> UsageMetadata | None: if not isinstance(raw_usage_metadata, dict): return None try: @@ -450,7 +449,7 @@ def process_response( model: str, _predictions: VertexAIBatchEmbeddingsResponseObject, ) -> EmbeddingResponse: - openai_embeddings: List[Embedding] = [] + openai_embeddings: list[Embedding] = [] for idx, embedding in enumerate(_predictions["embeddings"]): openai_embedding = Embedding( embedding=embedding["values"], @@ -465,7 +464,7 @@ def process_response( has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input) if _is_multimodal_input(input) or has_nested: input_list = input if isinstance(input, list) else [input] - text_elements: List[str] = [] + text_elements: list[str] = [] for e in input_list: if isinstance(e, list): text_elements.extend(sub for sub in e if isinstance(sub, str) and not _is_multimodal_element(sub)) diff --git a/litellm/llms/vertex_ai/google_genai/transformation.py b/litellm/llms/vertex_ai/google_genai/transformation.py index c1120d9ab8b..c22b3e9e811 100644 --- a/litellm/llms/vertex_ai/google_genai/transformation.py +++ b/litellm/llms/vertex_ai/google_genai/transformation.py @@ -2,7 +2,7 @@ Transformation for Calling Google models in their native format. """ -from typing import Any, Dict, Literal, Optional, Union +from typing import Any, Literal from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig from litellm.types.router import GenericLiteLLMParams @@ -22,10 +22,10 @@ class VertexAIGoogleGenAIConfig(GoogleGenAIConfig): def validate_environment( self, - api_key: Optional[str], - headers: Optional[dict], + api_key: str | None, + headers: dict | None, model: str, - litellm_params: Optional[Union[GenericLiteLLMParams, dict]], + litellm_params: GenericLiteLLMParams | dict | None, ) -> dict: default_headers = { "Content-Type": "application/json", @@ -60,7 +60,7 @@ class VertexAIGoogleGenAIConfig(GoogleGenAIConfig): Mapped parameters for the provider """ - _generate_content_config_dict: Dict = {} + _generate_content_config_dict: dict = {} for param, value in generate_content_config_dict.items(): camel_case_key = self._camel_to_snake(param) @@ -71,9 +71,9 @@ class VertexAIGoogleGenAIConfig(GoogleGenAIConfig): self, model: str, contents: Any, - tools: Optional[Any], - generate_content_config_dict: Dict, - system_instruction: Optional[Any] = None, + tools: Any | None, + generate_content_config_dict: dict, + system_instruction: Any | None = None, ) -> dict: """ Transform the generate content request for Vertex AI. diff --git a/litellm/llms/vertex_ai/image_edit/__init__.py b/litellm/llms/vertex_ai/image_edit/__init__.py index 51bb1511653..4ff15a3928d 100644 --- a/litellm/llms/vertex_ai/image_edit/__init__.py +++ b/litellm/llms/vertex_ai/image_edit/__init__.py @@ -11,8 +11,8 @@ from .vertex_imagen_transformation import VertexAIImagenImageEditConfig __all__ = [ "VertexAIGeminiImageEditConfig", "VertexAIImagenImageEditConfig", - "get_vertex_ai_image_edit_config", "cost_calculator", + "get_vertex_ai_image_edit_config", ] diff --git a/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py index a2020149ef2..82507c5cd2f 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_gemini_transformation.py @@ -2,7 +2,7 @@ import base64 import json import os from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -32,13 +32,13 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): Uses generateContent API for Gemini models on Vertex AI """ - SUPPORTED_PARAMS: List[str] = ["size"] + SUPPORTED_PARAMS: list[str] = ["size"] def __init__(self) -> None: BaseImageEditConfig.__init__(self) VertexLLM.__init__(self) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return list(self.SUPPORTED_PARAMS) def map_openai_params( @@ -46,11 +46,11 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: supported_params = self.get_supported_openai_params(model) filtered_params = {key: value for key, value in image_edit_optional_params.items() if key in supported_params} - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} if "size" in filtered_params: mapped_params["aspectRatio"] = self._map_size_to_aspect_ratio( @@ -59,7 +59,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): return mapped_params - def _resolve_vertex_project(self) -> Optional[str]: + def _resolve_vertex_project(self) -> str | None: return ( getattr(self, "_vertex_project", None) or os.environ.get("VERTEXAI_PROJECT") @@ -67,7 +67,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): or get_secret_str("VERTEXAI_PROJECT") ) - def _resolve_vertex_location(self) -> Optional[str]: + def _resolve_vertex_location(self) -> str | None: return ( getattr(self, "_vertex_location", None) or os.environ.get("VERTEXAI_LOCATION") @@ -77,7 +77,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): or get_secret_str("VERTEX_LOCATION") ) - def _resolve_vertex_credentials(self) -> Optional[str]: + def _resolve_vertex_credentials(self) -> str | None: return ( getattr(self, "_vertex_credentials", None) or os.environ.get("VERTEXAI_CREDENTIALS") @@ -90,9 +90,9 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: headers = headers or {} litellm_params = litellm_params or {} @@ -117,7 +117,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -148,12 +148,12 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): def transform_image_edit_request( # type: ignore[override] self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict[str, Any], + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict[str, Any], Optional[RequestFiles]]: + ) -> tuple[dict[str, Any], RequestFiles | None]: inline_parts = self._prepare_inline_image_parts(image) if image else [] if not inline_parts: raise ValueError("Vertex AI Gemini image edit requires at least one image.") @@ -166,13 +166,13 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): # Correct format for Vertex AI Gemini image editing contents = {"role": "USER", "parts": parts} - request_body: Dict[str, Any] = {"contents": contents} + request_body: dict[str, Any] = {"contents": contents} # Generation config with proper structure for image editing - generation_config: Dict[str, Any] = {"response_modalities": ["IMAGE"]} + generation_config: dict[str, Any] = {"response_modalities": ["IMAGE"]} # Add image-specific configuration - image_config: Dict[str, Any] = {} + image_config: dict[str, Any] = {} if "aspectRatio" in image_edit_optional_request_params: image_config["aspect_ratio"] = image_edit_optional_request_params["aspectRatio"] @@ -183,7 +183,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): payload: Any = json.dumps(request_body) empty_files = cast(RequestFiles, []) - return cast(Tuple[Dict[str, Any], Optional[RequestFiles]], (payload, empty_files)) + return cast(tuple[dict[str, Any], RequestFiles | None], (payload, empty_files)) def transform_image_edit_response( self, @@ -202,7 +202,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): ) candidates = response_json.get("candidates", []) - data_list: List[ImageObject] = [] + data_list: list[ImageObject] = [] for candidate in candidates: content = candidate.get("content", {}) @@ -217,7 +217,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): ) ) - model_response.data = cast(List[OpenAIImage], data_list) + model_response.data = cast(list[OpenAIImage], data_list) return model_response def _map_size_to_aspect_ratio(self, size: str) -> str: @@ -231,14 +231,14 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM): } return aspect_ratio_map.get(size, "1:1") - def _prepare_inline_image_parts(self, image: Union[FileTypes, List[FileTypes]]) -> List[Dict[str, Any]]: - images: List[FileTypes] + def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, Any]]: + images: list[FileTypes] if isinstance(image, list): images = image else: images = [image] - inline_parts: List[Dict[str, Any]] = [] + inline_parts: list[dict[str, Any]] = [] for img in images: if img is None: continue diff --git a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py index d9127a1929f..a91af091bab 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py @@ -3,7 +3,7 @@ import json import os from io import BufferedRandom, BufferedReader, BytesIO from pathlib import Path -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -33,13 +33,13 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): Uses predict API for Imagen models on Vertex AI """ - SUPPORTED_PARAMS: List[str] = ["n", "size", "mask"] + SUPPORTED_PARAMS: list[str] = ["n", "size", "mask"] def __init__(self) -> None: BaseImageEditConfig.__init__(self) VertexLLM.__init__(self) - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: return list(self.SUPPORTED_PARAMS) def map_openai_params( @@ -47,11 +47,11 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: supported_params = self.get_supported_openai_params(model) filtered_params = {key: value for key, value in image_edit_optional_params.items() if key in supported_params} - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} # Map OpenAI parameters to Imagen format if "n" in filtered_params: @@ -67,7 +67,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): return mapped_params - def _resolve_vertex_project(self) -> Optional[str]: + def _resolve_vertex_project(self) -> str | None: return ( getattr(self, "_vertex_project", None) or os.environ.get("VERTEXAI_PROJECT") @@ -75,7 +75,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): or get_secret_str("VERTEXAI_PROJECT") ) - def _resolve_vertex_location(self) -> Optional[str]: + def _resolve_vertex_location(self) -> str | None: return ( getattr(self, "_vertex_location", None) or os.environ.get("VERTEXAI_LOCATION") @@ -85,7 +85,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): or get_secret_str("VERTEX_LOCATION") ) - def _resolve_vertex_credentials(self) -> Optional[str]: + def _resolve_vertex_credentials(self) -> str | None: return ( getattr(self, "_vertex_credentials", None) or os.environ.get("VERTEXAI_CREDENTIALS") @@ -98,9 +98,9 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[dict] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + litellm_params: dict | None = None, + api_base: str | None = None, ) -> dict: headers = headers or {} litellm_params = litellm_params or {} @@ -121,7 +121,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -148,12 +148,12 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): def transform_image_edit_request( # type: ignore[override] self, model: str, - prompt: Optional[str], - image: Optional[FileTypes], - image_edit_optional_request_params: Dict[str, Any], + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict[str, Any], litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict[str, Any], Optional[RequestFiles]]: + ) -> tuple[dict[str, Any], RequestFiles | None]: # Prepare reference images in the correct Imagen format if image is None: raise ValueError("Vertex AI Imagen image edit requires at least one reference image.") @@ -184,14 +184,14 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): parameters["guidanceScale"] = 7.5 # Default guidance scale parameters["seed"] = None # Let Vertex AI choose random seed - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "instances": instances, "parameters": parameters, } payload: Any = json.dumps(request_body) empty_files = cast(RequestFiles, []) - return cast(Tuple[Dict[str, Any], Optional[RequestFiles]], (payload, empty_files)) + return cast(tuple[dict[str, Any], RequestFiles | None], (payload, empty_files)) def transform_image_edit_response( self, @@ -210,7 +210,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): ) predictions = response_json.get("predictions", []) - data_list: List[ImageObject] = [] + data_list: list[ImageObject] = [] for prediction in predictions: # Imagen returns images as bytesBase64Encoded @@ -222,7 +222,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): ) ) - model_response.data = cast(List[OpenAIImage], data_list) + model_response.data = cast(list[OpenAIImage], data_list) return model_response def _map_size_to_aspect_ratio(self, size: str) -> str: @@ -238,19 +238,19 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): def _prepare_reference_images( self, - image: Union[FileTypes, List[FileTypes]], - image_edit_optional_request_params: Dict[str, Any], - ) -> List[Dict[str, Any]]: + image: FileTypes | list[FileTypes], + image_edit_optional_request_params: dict[str, Any], + ) -> list[dict[str, Any]]: """ Prepare reference images in the correct Imagen API format """ - images: List[FileTypes] + images: list[FileTypes] if isinstance(image, list): images = image else: images = [image] - reference_images: List[Dict[str, Any]] = [] + reference_images: list[dict[str, Any]] = [] for idx, img in enumerate(images): if img is None: @@ -323,7 +323,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): image.seek(current_pos) return data if isinstance(image, (BufferedReader, BufferedRandom)): - stream_pos: Optional[int] = None + stream_pos: int | None = None try: stream_pos = image.tell() except Exception: 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 d265352ca0a..e7763ce4ce4 100644 --- a/litellm/llms/vertex_ai/image_generation/image_generation_handler.py +++ b/litellm/llms/vertex_ai/image_generation/image_generation_handler.py @@ -1,5 +1,5 @@ import json -from typing import Any, Dict, List, Optional +from typing import Any import httpx from openai.types.image import Image @@ -18,9 +18,9 @@ from litellm.types.utils import ImageResponse class VertexImageGeneration(VertexLLM): def process_image_generation_response( self, - json_response: Dict[str, Any], + json_response: dict[str, Any], model_response: ImageResponse, - model: Optional[str] = None, + model: str | None = None, ) -> ImageResponse: if "predictions" not in json_response: raise litellm.InternalServerError( @@ -30,7 +30,7 @@ class VertexImageGeneration(VertexLLM): ) predictions = json_response["predictions"] - response_data: List[Image] = [] + response_data: list[Image] = [] for prediction in predictions: bytes_base64_encoded = prediction["bytesBase64Encoded"] @@ -40,7 +40,7 @@ class VertexImageGeneration(VertexLLM): model_response.data = response_data return model_response - def transform_optional_params(self, optional_params: Optional[dict]) -> dict: + def transform_optional_params(self, optional_params: dict | None) -> dict: """ Transform the optional params to the format expected by the Vertex AI API. For example, "aspect_ratio" is transformed to "aspectRatio". @@ -69,18 +69,18 @@ class VertexImageGeneration(VertexLLM): def image_generation( self, prompt: str, - api_base: Optional[str], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], + api_base: str | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, model_response: ImageResponse, logging_obj: Any, model: str = "imagegeneration", # vertex ai uses imagegeneration as the default model - client: Optional[Any] = None, - optional_params: Optional[dict] = None, - timeout: Optional[int] = None, + client: Any | None = None, + optional_params: dict | None = None, + timeout: int | None = None, aimg_generation=False, - extra_headers: Optional[dict] = None, + extra_headers: dict | None = None, ) -> ImageResponse: if aimg_generation is True: return self.aimage_generation( # type: ignore @@ -112,7 +112,7 @@ class VertexImageGeneration(VertexLLM): # url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:predict" - auth_header: Optional[str] = None + auth_header: str | None = None auth_header, _ = self._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, @@ -168,17 +168,17 @@ class VertexImageGeneration(VertexLLM): async def aimage_generation( self, prompt: str, - api_base: Optional[str], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], + api_base: str | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, model_response: ImageResponse, logging_obj: Any, model: str = "imagegeneration", # vertex ai uses imagegeneration as the default model - client: Optional[AsyncHTTPHandler] = None, - optional_params: Optional[dict] = None, - timeout: Optional[int] = None, - extra_headers: Optional[dict] = None, + client: AsyncHTTPHandler | None = None, + optional_params: dict | None = None, + timeout: int | None = None, + extra_headers: dict | None = None, ): response = None if client is None: @@ -217,7 +217,7 @@ class VertexImageGeneration(VertexLLM): } \ "https://us-central1-aiplatform.googleapis.com/v1/projects/PROJECT_ID/locations/us-central1/publishers/google/models/imagegeneration:predict" """ - auth_header: Optional[str] = None + auth_header: str | None = None auth_header, _ = self._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, @@ -269,7 +269,7 @@ class VertexImageGeneration(VertexLLM): json_response = response.json() return self.process_image_generation_response(json_response, model_response, model) - def is_image_generation_response(self, json_response: Dict[str, Any]) -> bool: + def is_image_generation_response(self, json_response: dict[str, Any]) -> bool: if "predictions" in json_response: if "bytesBase64Encoded" in json_response["predictions"][0]: return True 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 572725ac789..518b0069893 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py @@ -1,11 +1,10 @@ import os -from typing import TYPE_CHECKING, Any, Optional - -from litellm._logging import verbose_logger +from typing import TYPE_CHECKING, Any import httpx import litellm +from litellm._logging import verbose_logger from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) @@ -74,7 +73,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): mapped_params = {} for k, v in non_default_params.items(): - if k not in optional_params.keys(): + if k not in optional_params: if k in supported_params: # Map OpenAI parameters to Gemini format if k == "n": @@ -110,7 +109,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): } return aspect_ratio_map.get(size, "1:1") - def _resolve_vertex_project(self) -> Optional[str]: + def _resolve_vertex_project(self) -> str | None: return ( getattr(self, "_vertex_project", None) or os.environ.get("VERTEXAI_PROJECT") @@ -118,7 +117,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): or get_secret_str("VERTEXAI_PROJECT") ) - def _resolve_vertex_location(self) -> Optional[str]: + def _resolve_vertex_location(self) -> str | None: return ( getattr(self, "_vertex_location", None) or os.environ.get("VERTEXAI_LOCATION") @@ -128,7 +127,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): or get_secret_str("VERTEX_LOCATION") ) - def _resolve_vertex_credentials(self) -> Optional[str]: + def _resolve_vertex_credentials(self) -> str | None: return ( getattr(self, "_vertex_credentials", None) or os.environ.get("VERTEXAI_CREDENTIALS") @@ -139,12 +138,12 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Vertex AI Gemini generateContent API @@ -178,8 +177,8 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = headers or {} @@ -284,8 +283,8 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Gemini image generation response to litellm ImageResponse format diff --git a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py index 2cd3df010d6..fae99346b46 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py @@ -1,5 +1,5 @@ import os -from typing import TYPE_CHECKING, Any, List, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -39,7 +39,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): BaseImageGenerationConfig.__init__(self) VertexLLM.__init__(self) - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: """ Imagen API supported parameters """ @@ -56,7 +56,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): mapped_params = {} for k, v in non_default_params.items(): - if k not in optional_params.keys(): + if k not in optional_params: if k in supported_params: # Map OpenAI parameters to Imagen format if k == "n": @@ -82,7 +82,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): } return aspect_ratio_map.get(size, "1:1") - def _resolve_vertex_project(self) -> Optional[str]: + def _resolve_vertex_project(self) -> str | None: return ( getattr(self, "_vertex_project", None) or os.environ.get("VERTEXAI_PROJECT") @@ -90,7 +90,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): or get_secret_str("VERTEXAI_PROJECT") ) - def _resolve_vertex_location(self) -> Optional[str]: + def _resolve_vertex_location(self) -> str | None: return ( getattr(self, "_vertex_location", None) or os.environ.get("VERTEXAI_LOCATION") @@ -100,7 +100,7 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): or get_secret_str("VERTEX_LOCATION") ) - def _resolve_vertex_credentials(self) -> Optional[str]: + def _resolve_vertex_credentials(self) -> str | None: return ( getattr(self, "_vertex_credentials", None) or os.environ.get("VERTEXAI_CREDENTIALS") @@ -111,12 +111,12 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for Vertex AI Imagen predict API @@ -147,11 +147,11 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: headers = headers or {} @@ -213,8 +213,8 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ImageResponse: """ Transform Imagen image generation response to litellm ImageResponse format diff --git a/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py index d0ffc7be0a6..085157c2a9d 100644 --- a/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/multimodal_embeddings/embedding_handler.py @@ -1,5 +1,5 @@ import json -from typing import Literal, Optional, Union +from typing import Literal import httpx @@ -32,21 +32,21 @@ class VertexMultimodalEmbedding(VertexLLM): def multimodal_embedding( self, model: str, - input: Union[list, str], + input: list | str, print_verbose, model_response: EmbeddingResponse, custom_llm_provider: Literal["gemini", "vertex_ai"], optional_params: dict, litellm_params: dict, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, headers: dict = {}, encoding=None, vertex_project=None, vertex_location=None, vertex_credentials=None, - aembedding: Optional[bool] = False, + aembedding: bool | None = False, timeout=300, client=None, ) -> EmbeddingResponse: @@ -148,11 +148,11 @@ class VertexMultimodalEmbedding(VertexLLM): litellm_params: dict, data: dict, model_response: EmbeddingResponse, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, logging_obj: LiteLLMLoggingObj, headers={}, - client: Optional[AsyncHTTPHandler] = None, - api_key: Optional[str] = None, + client: AsyncHTTPHandler | None = None, + api_key: str | None = None, ) -> EmbeddingResponse: if client is None: _params = {} diff --git a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py index 4bcfdee2d17..a4815b01ecc 100644 --- a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py @@ -1,4 +1,4 @@ -from typing import List, Optional, Union, cast +from typing import cast from httpx import Headers, Response @@ -45,11 +45,11 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: default_headers = { "Content-Type": "application/json; charset=utf-8", @@ -104,7 +104,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): else: return Instance(text=input_element) - def _try_merge_text_with_media(self, text_str: str, next_elem: Optional[str]) -> tuple[Instance, bool]: + def _try_merge_text_with_media(self, text_str: str, next_elem: str | None) -> tuple[Instance, bool]: """ Try to merge a text element with a following media element into a single instance. @@ -127,7 +127,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): return instance_args, False - def process_openai_embedding_input(self, _input: Union[list, str]) -> List[Instance]: + def process_openai_embedding_input(self, _input: list | str) -> list[Instance]: """ Process the input for multimodal embedding requests. @@ -138,7 +138,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): List[Instance]: List of Instance objects for the embedding request. """ _input_list = [_input] if not isinstance(_input, list) else _input - processed_instances: List[Instance] = [] + processed_instances: list[Instance] = [] i = 0 while i < len(_input_list): @@ -177,7 +177,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): if "instances" in optional_params: request_data["instances"] = optional_params["instances"] elif isinstance(input, list): - vertex_instances: List[Instance] = self.process_openai_embedding_input(_input=input) + vertex_instances: list[Instance] = self.process_openai_embedding_input(_input=input) request_data["instances"] = vertex_instances else: @@ -200,7 +200,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): raw_response: Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -233,8 +233,8 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): vertex_predictions: MultimodalPredictions, ) -> Usage: ## Calculate text embeddings usage - prompt: Optional[str] = None - character_count: Optional[int] = None + prompt: str | None = None + character_count: int | None = None for instance in request_data["instances"]: text = instance.get("text") @@ -275,8 +275,8 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): prompt_tokens_details=prompt_tokens_details, ) - def transform_embedding_response_to_openai(self, predictions: MultimodalPredictions) -> List[Embedding]: - openai_embeddings: List[Embedding] = [] + def transform_embedding_response_to_openai(self, predictions: MultimodalPredictions) -> list[Embedding]: + openai_embeddings: list[Embedding] = [] if "predictions" in predictions: for idx, _prediction in enumerate(predictions["predictions"]): if _prediction: @@ -304,5 +304,5 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): openai_embeddings.append(openai_embedding_object) return openai_embeddings - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException: return VertexAIError(status_code=status_code, message=error_message, headers=headers) diff --git a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py index 68836a64027..b66dead91b1 100644 --- a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py +++ b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py @@ -3,7 +3,7 @@ Vertex AI DeepSeek OCR transformation implementation. """ import json -from typing import TYPE_CHECKING, Any, Dict +from typing import TYPE_CHECKING, Any import httpx @@ -43,13 +43,13 @@ class VertexAIDeepSeekOCRConfig(BaseOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers for Vertex AI OCR. diff --git a/litellm/llms/vertex_ai/ocr/transformation.py b/litellm/llms/vertex_ai/ocr/transformation.py index d67c5f2b089..0fb9523f3eb 100644 --- a/litellm/llms/vertex_ai/ocr/transformation.py +++ b/litellm/llms/vertex_ai/ocr/transformation.py @@ -2,8 +2,6 @@ Vertex AI Mistral OCR transformation implementation. """ -from typing import Dict - from litellm._logging import verbose_logger from litellm.litellm_core_utils.prompt_templates.image_handling import ( async_convert_url_to_base64, @@ -39,13 +37,13 @@ class VertexAIOCRConfig(MistralOCRConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, api_base: str | None = None, litellm_params: dict | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Validate environment and return headers for Vertex AI OCR. diff --git a/litellm/llms/vertex_ai/rag_engine/ingestion.py b/litellm/llms/vertex_ai/rag_engine/ingestion.py index d9e0035aa99..edafa2f8f7a 100644 --- a/litellm/llms/vertex_ai/rag_engine/ingestion.py +++ b/litellm/llms/vertex_ai/rag_engine/ingestion.py @@ -14,7 +14,7 @@ Key differences from OpenAI: from __future__ import annotations import os -from typing import TYPE_CHECKING, Any, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm import get_secret_str from litellm._logging import verbose_logger @@ -26,7 +26,7 @@ if TYPE_CHECKING: from litellm.types.rag import RAGIngestOptions -def _get_str_or_none(value: Any) -> Optional[str]: +def _get_str_or_none(value: Any) -> str | None: """Cast config value to Optional[str].""" return str(value) if value is not None else None @@ -65,8 +65,8 @@ class VertexAIRAGIngestion(BaseRAGIngestion): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): super().__init__(ingest_options=ingest_options, router=router) @@ -227,7 +227,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion): transformation = VertexAIRAGTransformation() chunking_config = transformation.transform_chunking_strategy_to_vertex_format( - cast(Optional[RAGChunkingStrategy], self.chunking_strategy) + cast(RAGChunkingStrategy | None, self.chunking_strategy) ) chunk_size = chunking_config["chunking_config"]["chunk_size"] @@ -242,8 +242,8 @@ class VertexAIRAGIngestion(BaseRAGIngestion): async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Vertex AI handles embedding internally - skip this step. @@ -254,12 +254,12 @@ class VertexAIRAGIngestion(BaseRAGIngestion): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], - ) -> Tuple[Optional[str], Optional[str]]: + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, + ) -> tuple[str | None, str | None]: """ Store content in Vertex AI RAG corpus. diff --git a/litellm/llms/vertex_ai/rag_engine/transformation.py b/litellm/llms/vertex_ai/rag_engine/transformation.py index 4aa2fcb49be..3e0239e1aba 100644 --- a/litellm/llms/vertex_ai/rag_engine/transformation.py +++ b/litellm/llms/vertex_ai/rag_engine/transformation.py @@ -4,7 +4,7 @@ Transformation utilities for Vertex AI RAG Engine. Handles transforming LiteLLM's unified formats to Vertex AI RAG Engine API format. """ -from typing import Any, Dict, Optional +from typing import Any from litellm._logging import verbose_logger from litellm.constants import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE @@ -54,8 +54,8 @@ class VertexAIRAGTransformation(VertexBase): def transform_chunking_strategy_to_vertex_format( self, - chunking_strategy: Optional[RAGChunkingStrategy], - ) -> Dict[str, Any]: + chunking_strategy: RAGChunkingStrategy | None, + ) -> dict[str, Any]: """ Transform LiteLLM's unified chunking_strategy to Vertex AI RAG format. @@ -104,8 +104,8 @@ class VertexAIRAGTransformation(VertexBase): def build_import_rag_files_request( self, gcs_uri: str, - chunking_strategy: Optional[RAGChunkingStrategy] = None, - ) -> Dict[str, Any]: + chunking_strategy: RAGChunkingStrategy | None = None, + ) -> dict[str, Any]: """ Build the request payload for importing RAG files. @@ -127,9 +127,9 @@ class VertexAIRAGTransformation(VertexBase): def get_auth_headers( self, - vertex_credentials: Optional[str] = None, - vertex_project: Optional[str] = None, - ) -> Dict[str, str]: + vertex_credentials: str | None = None, + vertex_project: str | None = None, + ) -> dict[str, str]: """ Get authentication headers for Vertex AI API calls. diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index beb8bc0be6f..cb4c2dc5ed1 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -12,7 +12,6 @@ Auth: OAuth2 Bearer token (not an API key). """ import json -from typing import List, Optional from litellm import verbose_logger from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig @@ -41,9 +40,9 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, model: str, - api_key: Optional[str] = None, # noqa: ARG002 + api_key: str | None = None, # noqa: ARG002 ) -> str: """ Build the Vertex AI Live WSS endpoint URL. @@ -73,7 +72,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): self, headers: dict, model: str, # noqa: ARG002 - api_key: Optional[str] = None, # noqa: ARG002 + api_key: str | None = None, # noqa: ARG002 ) -> dict: """ Return headers with a Bearer token for Vertex AI. @@ -191,8 +190,8 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): self, message: str, model: str, - session_configuration_request: Optional[str] = None, - ) -> List[str]: + session_configuration_request: str | None = None, + ) -> list[str]: """ Translate OpenAI realtime client messages to Vertex AI format. diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index b9680af20cc..ab5464dfb8c 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -4,7 +4,7 @@ Translates from Cohere's `/v1/rerank` input format to Vertex AI Discovery Engine Why separate file? Make it easy to see how transformation works """ -from typing import Any, Dict, List, Union +from typing import Any import httpx @@ -38,7 +38,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): self, api_base: str | None, model: str, - optional_params: Dict | None = None, + optional_params: dict | None = None, ) -> str: """ Get the complete URL for the Vertex AI Discovery Engine ranking API @@ -73,7 +73,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): headers: dict, model: str, api_key: str | None = None, - optional_params: Dict | None = None, + optional_params: dict | None = None, ) -> dict: """ Validate and set up authentication for Vertex AI Discovery Engine API @@ -106,7 +106,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -225,15 +225,15 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map Cohere rerank params to Vertex AI format """ diff --git a/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py b/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py index e27df956c9d..3f4aefbbbc4 100644 --- a/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py +++ b/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - import httpx from typing_extensions import TypedDict @@ -14,8 +12,8 @@ from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES class VertexInput(TypedDict, total=False): - text: Optional[str] - ssml: Optional[str] + text: str | None + ssml: str | None class VertexVoice(TypedDict, total=False): @@ -31,7 +29,7 @@ class VertexAudioConfig(TypedDict, total=False): class VertexTextToSpeechRequest(TypedDict, total=False): input: VertexInput voice: VertexVoice - audioConfig: Optional[VertexAudioConfig] + audioConfig: VertexAudioConfig | None class VertexTextToSpeechAPI(VertexLLM): @@ -45,17 +43,17 @@ class VertexTextToSpeechAPI(VertexLLM): def audio_speech( self, logging_obj, - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - api_base: Optional[str], - timeout: Union[float, httpx.Timeout], + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + api_base: str | None, + timeout: float | httpx.Timeout, model: str, input: str, - voice: Optional[dict] = None, - _is_async: Optional[bool] = False, - optional_params: Optional[dict] = None, - kwargs: Optional[dict] = None, + voice: dict | None = None, + _is_async: bool | None = False, + optional_params: dict | None = None, + kwargs: dict | None = None, ) -> HttpxBinaryResponseContent: import base64 diff --git a/litellm/llms/vertex_ai/text_to_speech/transformation.py b/litellm/llms/vertex_ai/text_to_speech/transformation.py index aeb40b16c28..642ac7b27c1 100644 --- a/litellm/llms/vertex_ai/text_to_speech/transformation.py +++ b/litellm/llms/vertex_ai/text_to_speech/transformation.py @@ -7,7 +7,7 @@ Reference: https://cloud.google.com/text-to-speech/docs/reference/rest/v1/text/s import base64 from collections.abc import Coroutine -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -76,8 +76,8 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): def _map_voice_to_vertex_format( self, - voice: Optional[Union[str, Dict]], - ) -> Tuple[Optional[str], Optional[Dict]]: + voice: str | dict | None, + ) -> tuple[str | None, dict | None]: """ Map voice to Vertex AI format. @@ -126,16 +126,16 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): self, model: str, input: str, - voice: Optional[Union[str, Dict]], - optional_params: Dict, - litellm_params_dict: Dict, + voice: str | dict | None, + optional_params: dict, + litellm_params_dict: dict, logging_obj: "LiteLLMLoggingObj", - timeout: Union[float, httpx.Timeout], - extra_headers: Optional[Dict[str, Any]], + timeout: float | httpx.Timeout, + extra_headers: dict[str, Any] | None, base_llm_http_handler: Any, aspeech: bool, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, **kwargs: Any, ) -> Union[ "HttpxBinaryResponseContent", @@ -157,7 +157,7 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): # Convert voice to string if it's a dict (extract name) # Actual voice mapping happens in map_openai_params - voice_str: Optional[str] = None + voice_str: str | None = None if isinstance(voice, str): voice_str = voice elif isinstance(voice, dict): @@ -204,11 +204,11 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): def map_openai_params( self, model: str, - optional_params: Dict, - voice: Optional[Union[str, Dict]] = None, + optional_params: dict, + voice: str | dict | None = None, drop_params: bool = False, - kwargs: Dict = {}, - ) -> Tuple[Optional[str], Dict]: + kwargs: dict = {}, + ) -> tuple[str | None, dict]: """ Map OpenAI parameters to Vertex AI TTS parameters @@ -223,7 +223,7 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): Returns: Tuple of (mapped_voice_str, mapped_params) """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} ########################################################## # Map voice using helper @@ -263,8 +263,8 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): self, headers: dict, model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """ Validate Vertex AI environment and set up authentication headers @@ -283,7 +283,7 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -299,7 +299,7 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): def _validate_vertex_input( self, input_data: VertexTextToSpeechInput, - optional_params: Dict, + optional_params: dict, ) -> VertexTextToSpeechInput: """ Validate and transform input for Vertex AI TTS @@ -339,9 +339,9 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): self, model: str, input: str, - voice: Optional[str], - optional_params: Dict, - litellm_params: Dict, + voice: str | None, + optional_params: dict, + litellm_params: dict, headers: dict, ) -> TextToSpeechRequestData: """ @@ -356,8 +356,8 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): TextToSpeechRequestData: Contains dict_body and headers """ # Get Vertex AI credentials from litellm_params - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = litellm_params.get("vertex_credentials") - vertex_project: Optional[str] = litellm_params.get("vertex_project") + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = litellm_params.get("vertex_credentials") + vertex_project: str | None = litellm_params.get("vertex_project") ####### Authenticate with Vertex AI ######## _auth_header, vertex_project = self._ensure_access_token( @@ -424,7 +424,7 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): speakingRate=speaking_rate, ) - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "input": dict(vertex_input), "voice": dict(vertex_voice), "audioConfig": dict(vertex_audio_config), diff --git a/litellm/llms/vertex_ai/vector_stores/__init__.py b/litellm/llms/vertex_ai/vector_stores/__init__.py index 98da2c581a8..fb48eec44af 100644 --- a/litellm/llms/vertex_ai/vector_stores/__init__.py +++ b/litellm/llms/vertex_ai/vector_stores/__init__.py @@ -1,4 +1,4 @@ from .rag_api.transformation import VertexVectorStoreConfig from .search_api.transformation import VertexSearchAPIVectorStoreConfig -__all__ = ["VertexVectorStoreConfig", "VertexSearchAPIVectorStoreConfig"] +__all__ = ["VertexSearchAPIVectorStoreConfig", "VertexVectorStoreConfig"] diff --git a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py index 47a81fc07bf..4e1e41331a5 100644 --- a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -60,7 +60,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): "write": [("POST", "/ragCorpora")], } - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate and set up authentication for Vertex AI RAG API """ @@ -72,7 +72,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -91,13 +91,13 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict[str, Any]]: """ Transform search request for Vertex AI RAG API """ @@ -121,7 +121,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}" # Build the request body for Vertex AI RAG API - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "vertex_rag_store": {"rag_resources": [{"rag_corpus": full_rag_corpus}]}, "query": {"text": query}, } @@ -219,14 +219,14 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict[str, Any]]: + ) -> tuple[str, dict[str, Any]]: """ Transform create request for Vertex AI RAG Corpus """ url = f"{api_base}/ragCorpora" # Base URL for creating RAG corpus # Build the request body for Vertex AI RAG Corpus creation - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "display_name": vector_store_create_optional_params.get("name", "litellm-vector-store"), "description": "Vector store created via LiteLLM", } diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 958839d4a48..603964298e3 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import httpx @@ -75,7 +75,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): return VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS @classmethod - def _filter_extra_body(cls, extra_body: Dict[str, Any], is_engine: bool = False) -> Dict[str, Any]: + def _filter_extra_body(cls, extra_body: dict[str, Any], is_engine: bool = False) -> dict[str, Any]: """ Validate ``extra_body`` against the supported-field allowlist for the active serving config (engine/app vs data store). @@ -141,7 +141,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): "write": [], } - def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate and set up authentication for Vertex AI RAG API """ @@ -152,7 +152,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -191,13 +191,13 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def transform_search_vector_store_request( self, vector_store_id: str, - query: Union[str, List[str]], + query: str | list[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict[str, Any]]: """ Transform a search request for the Vertex AI Search (Discovery Engine) API. @@ -222,7 +222,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): is_engine = bool(litellm_params.get("vertex_engine_id")) - request_body: Dict[str, Any] = {"query": query, "pageSize": 10} + request_body: dict[str, Any] = {"query": query, "pageSize": 10} max_num_results = vector_store_search_optional_params.get("max_num_results") if max_num_results is not None: request_body["pageSize"] = max_num_results @@ -262,7 +262,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): results = response_json.get("results", []) # Transform results to standard format - search_results: List[VectorStoreSearchResult] = [] + search_results: list[VectorStoreSearchResult] = [] for result in results: document = result.get("document", {}) derived_data = document.get("derivedStructData", {}) @@ -346,7 +346,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: raise NotImplementedError def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: @@ -355,7 +355,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def calculate_vector_store_cost( self, response: VectorStoreSearchResponse, - ) -> Tuple[float, float]: + ) -> tuple[float, float]: model_info = get_model_info( model="vertex_ai/search_api", ) diff --git a/litellm/llms/vertex_ai/vertex_ai_aws_wif.py b/litellm/llms/vertex_ai/vertex_ai_aws_wif.py index da95ac72c2f..b230a00da3d 100644 --- a/litellm/llms/vertex_ai/vertex_ai_aws_wif.py +++ b/litellm/llms/vertex_ai/vertex_ai_aws_wif.py @@ -9,8 +9,6 @@ uses BaseAWSLLM to obtain AWS credentials and wraps them in a custom AwsSecurityCredentialsSupplier for google-auth. """ -from typing import Dict - GOOGLE_IMPORT_ERROR_MESSAGE = ( "Google Cloud SDK not found. Install it with: pip install 'litellm[google]' or pip install google-cloud-aiplatform" ) @@ -40,7 +38,7 @@ class VertexAIAwsWifAuth: """ @staticmethod - def extract_aws_params(json_obj: dict) -> Dict[str, str]: + def extract_aws_params(json_obj: dict) -> dict[str, str]: """ Extract LiteLLM-specific aws_* keys from a WIF credential JSON dict. diff --git a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py index acf2ef1ccb4..9768229ee07 100644 --- a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py +++ b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py @@ -2,7 +2,7 @@ import json import os import time from collections.abc import Callable -from typing import Any, Optional, cast +from typing import Any, cast import httpx @@ -55,7 +55,7 @@ class TextStreamer: raise StopAsyncIteration # once we run out of data to stream, we raise this error -def _get_client_cache_key(model: str, vertex_project: Optional[str], vertex_location: Optional[str]): +def _get_client_cache_key(model: str, vertex_project: str | None, vertex_location: str | None): _cache_key = f"{model}-{vertex_project}-{vertex_location}" return _cache_key diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/ai21/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/ai21/transformation.py index c8163708574..3805c0693a7 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/ai21/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/ai21/transformation.py @@ -1,5 +1,4 @@ import types -from typing import Optional import litellm from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig @@ -16,7 +15,7 @@ class VertexAIAi21Config(OpenAIGPTConfig): def __init__( self, - max_tokens: Optional[int] = None, + max_tokens: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index 32aaebab768..8a49274d41b 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( @@ -18,7 +18,7 @@ from ..output_params_utils import sanitize_vertex_anthropic_output_params class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "vertex_ai" def should_strip_billing_metadata(self) -> bool: @@ -28,12 +28,12 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert self, headers: dict, model: str, - messages: List[Any], + messages: list[Any], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: """ OPTIONAL @@ -115,12 +115,12 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base is None: raise ValueError("api_base is required. Unable to determine the correct api_base for the request.") @@ -129,11 +129,11 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert def transform_anthropic_messages_request( self, model: str, - messages: List[Dict], - anthropic_messages_optional_request_params: Dict, + messages: list[dict], + anthropic_messages_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Dict: + ) -> dict: anthropic_messages_request = super().transform_anthropic_messages_request( model=model, messages=messages, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py index 8fcefb04b34..80bf0991b62 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py @@ -1,6 +1,6 @@ # What is this? ## Handler file for calling claude-3 on vertex ai -from typing import Any, List, Optional +from typing import Any import httpx @@ -45,7 +45,7 @@ class VertexAIAnthropicConfig(AnthropicConfig): """ @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "vertex_ai" def should_strip_billing_metadata(self) -> bool: @@ -86,7 +86,7 @@ class VertexAIAnthropicConfig(AnthropicConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -180,12 +180,12 @@ class VertexAIAnthropicConfig(AnthropicConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: response = super().transform_response( model, @@ -211,8 +211,6 @@ class VertexAIAnthropicConfig(AnthropicConfig): """ if custom_llm_provider != "vertex_ai" and custom_llm_provider != "vertex_ai_beta": return False - if "claude" in model.lower(): - return True - elif model in litellm.vertex_anthropic_models: + if "claude" in model.lower() or model in litellm.vertex_anthropic_models: return True return False diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py index d3edf2e9848..f32f07762dd 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -6,7 +6,7 @@ Unlike Gemini models which use Google's token counting API, partner models use their respective publisher-specific count-tokens endpoints. """ -from typing import Any, Dict, Optional +from typing import Any from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.vertex_ai.common_utils import get_vertex_base_url @@ -48,7 +48,7 @@ class VertexAIPartnerModelsTokenCounter(VertexBase): model: str, project_id: str, vertex_location: str, - api_base: Optional[str] = None, + api_base: str | None = None, ) -> str: """ Build the count-tokens endpoint URL for a partner model. @@ -95,9 +95,9 @@ class VertexAIPartnerModelsTokenCounter(VertexBase): async def handle_count_tokens_request( self, model: str, - request_data: Dict[str, Any], - litellm_params: Dict[str, Any], - ) -> Dict[str, Any]: + request_data: dict[str, Any], + litellm_params: dict[str, Any], + ) -> dict[str, Any]: """ Handle token counting request for a Vertex AI partner model. diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index 62172db79fd..abad2bb73ea 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -1,6 +1,6 @@ import types from collections.abc import AsyncIterator, Iterator -from typing import Any, List, Optional, Union +from typing import Any import httpx @@ -32,11 +32,11 @@ class VertexAILlama3Config(OpenAIGPTConfig): Note: Please make sure to modify the default parameters as required for your use case. """ - max_tokens: Optional[int] = None + max_tokens: int | None = None def __init__( self, - max_tokens: Optional[int] = None, + max_tokens: int | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -89,9 +89,9 @@ class VertexAILlama3Config(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return VertexAILlama3StreamingHandler( streaming_response=streaming_response, @@ -106,12 +106,12 @@ class VertexAILlama3Config(OpenAIGPTConfig): model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -127,7 +127,7 @@ class VertexAILlama3Config(OpenAIGPTConfig): except Exception as e: response_headers = getattr(raw_response, "headers", None) raise VertexAIError( - message="Unable to get json response - {}, Original Response: {}".format(str(e), raw_response.text), + message=f"Unable to get json response - {e!s}, Original Response: {raw_response.text}", status_code=raw_response.status_code, headers=response_headers, ) @@ -161,7 +161,7 @@ class VertexAILlama3StreamingHandler(OpenAIChatCompletionStreamingHandler): def __init__(self, **kwargs): super().__init__(**kwargs) self.sent_role = False - self._pending_chunk: Optional[ModelResponseStream] = None + self._pending_chunk: ModelResponseStream | None = None def chunk_parser(self, chunk: dict) -> ModelResponseStream: result = super().chunk_parser(chunk) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index b415e24f864..7d63d983b19 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -2,7 +2,6 @@ ## API Handler for calling Vertex AI Partner Models from collections.abc import Callable from enum import Enum -from typing import Optional, Union import httpx # type: ignore @@ -95,11 +94,11 @@ class VertexAIPartnerModels(VertexBase): print_verbose: Callable, encoding, logging_obj, - api_base: Optional[str], + api_base: str | None, optional_params: dict, custom_prompt_dict: dict, - headers: Optional[dict], - timeout: Union[float, httpx.Timeout], + headers: dict | None, + timeout: float | httpx.Timeout, litellm_params: dict, vertex_project=None, vertex_location=None, @@ -190,7 +189,7 @@ class VertexAIPartnerModels(VertexBase): # Build a new dict so we never mutate the shared deployment extra_headers object. headers = { **(headers or {}), - "Authorization": "Bearer {}".format(access_token), + "Authorization": f"Bearer {access_token}", } optional_params.update( diff --git a/litellm/llms/vertex_ai/vertex_embeddings/bge.py b/litellm/llms/vertex_ai/vertex_embeddings/bge.py index 6525d3342f5..98fcc971186 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/bge.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/bge.py @@ -10,8 +10,6 @@ Model name handling: - This module focuses on request/response transformation only """ -from typing import List, Optional, Union - from litellm.types.utils import EmbeddingResponse, Usage from .types import ( @@ -57,7 +55,7 @@ class VertexBGEConfig: return model_lower.startswith("bge/") or "bge" in model_lower @staticmethod - def transform_request(input: Union[list, str], optional_params: dict, model: str) -> VertexEmbeddingRequest: + def transform_request(input: list | str, optional_params: dict, model: str) -> VertexEmbeddingRequest: """ Transforms an OpenAI request to a Vertex BGE embedding request. @@ -72,8 +70,8 @@ class VertexBGEConfig: VertexEmbeddingRequest: The transformed request """ vertex_request: VertexEmbeddingRequest = VertexEmbeddingRequest() - vertex_text_embedding_input_list: List[TextEmbeddingBGEInput] = [] - task_type: Optional[TaskType] = optional_params.get("task_type") + vertex_text_embedding_input_list: list[TextEmbeddingBGEInput] = [] + task_type: TaskType | None = optional_params.get("task_type") title = optional_params.get("title") if isinstance(input, str): @@ -91,8 +89,8 @@ class VertexBGEConfig: @staticmethod def _create_embedding_input( prompt: str, - task_type: Optional[TaskType] = None, - title: Optional[str] = None, + task_type: TaskType | None = None, + title: str | None = None, ) -> TextEmbeddingBGEInput: """ Creates a TextEmbeddingBGEInput object for BGE models. diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 0e7afd5da3f..c0d1e2922bb 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -1,4 +1,4 @@ -from typing import Dict, Literal, Optional, Union +from typing import Literal import httpx @@ -25,7 +25,7 @@ class VertexEmbedding(VertexBase): def embedding( self, model: str, - input: Union[list, str], + input: list | str, print_verbose, model_response: EmbeddingResponse, optional_params: dict, @@ -33,18 +33,18 @@ class VertexEmbedding(VertexBase): custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - timeout: Optional[Union[float, httpx.Timeout]], - api_key: Optional[str] = None, + timeout: float | httpx.Timeout | None, + api_key: str | None = None, encoding=None, - aembedding: Optional[bool] = False, - api_base: Optional[str] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, - gemini_api_key: Optional[str] = None, - extra_headers: Optional[dict] = None, - litellm_params: Optional[Dict] = None, + aembedding: bool | None = False, + api_base: str | None = None, + client: AsyncHTTPHandler | HTTPHandler | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, + gemini_api_key: str | None = None, + extra_headers: dict | None = None, + litellm_params: dict | None = None, ) -> EmbeddingResponse: if aembedding is True: return self.async_embedding( # type: ignore @@ -139,23 +139,23 @@ class VertexEmbedding(VertexBase): async def async_embedding( self, model: str, - input: Union[list, str], + input: list | str, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObject, optional_params: dict, custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - timeout: Optional[Union[float, httpx.Timeout]], - api_base: Optional[str] = None, - client: Optional[AsyncHTTPHandler] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, - gemini_api_key: Optional[str] = None, - extra_headers: Optional[dict] = None, + timeout: float | httpx.Timeout | None, + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, + gemini_api_key: str | None = None, + extra_headers: dict | None = None, encoding=None, - litellm_params: Optional[Dict] = None, + litellm_params: dict | None = None, ) -> EmbeddingResponse: """ Async embedding implementation diff --git a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py index 6b7e6c036c0..a29317dae53 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py @@ -1,5 +1,5 @@ import types -from typing import List, Literal, Optional, Union +from typing import Literal from pydantic import BaseModel @@ -19,8 +19,8 @@ class VertexAITextEmbeddingConfig(BaseModel): title: Optional(str) The title of the document to be embedded. (only valid with task_type=RETRIEVAL_DOCUMENT). """ - auto_truncate: Optional[bool] = None - task_type: Optional[ + auto_truncate: bool | None = None + task_type: ( Literal[ "RETRIEVAL_QUERY", "RETRIEVAL_DOCUMENT", @@ -30,24 +30,24 @@ class VertexAITextEmbeddingConfig(BaseModel): "QUESTION_ANSWERING", "FACT_VERIFICATION", ] - ] = None - title: Optional[str] = None + | None + ) = None + title: str | None = None def __init__( self, - auto_truncate: Optional[bool] = None, - task_type: Optional[ - Literal[ - "RETRIEVAL_QUERY", - "RETRIEVAL_DOCUMENT", - "SEMANTIC_SIMILARITY", - "CLASSIFICATION", - "CLUSTERING", - "QUESTION_ANSWERING", - "FACT_VERIFICATION", - ] - ] = None, - title: Optional[str] = None, + auto_truncate: bool | None = None, + task_type: Literal[ + "RETRIEVAL_QUERY", + "RETRIEVAL_DOCUMENT", + "SEMANTIC_SIMILARITY", + "CLASSIFICATION", + "CLUSTERING", + "QUESTION_ANSWERING", + "FACT_VERIFICATION", + ] + | None = None, + title: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -100,10 +100,10 @@ class VertexAITextEmbeddingConfig(BaseModel): def transform_openai_request_to_vertex_embedding_request( self, - input: Union[list, str], + input: list | str, optional_params: dict, model: str, - litellm_params: Optional[dict] = None, + litellm_params: dict | None = None, ) -> VertexEmbeddingRequest: """ Transforms an openai request to a vertex embedding request. @@ -129,8 +129,8 @@ class VertexAITextEmbeddingConfig(BaseModel): return vertex_request vertex_request = VertexEmbeddingRequest() - vertex_text_embedding_input_list: List[TextEmbeddingInput] = [] - task_type: Optional[TaskType] = optional_params.get("task_type") + vertex_text_embedding_input_list: list[TextEmbeddingInput] = [] + task_type: TaskType | None = optional_params.get("task_type") title = optional_params.get("title") if isinstance(input, str): @@ -148,7 +148,7 @@ class VertexAITextEmbeddingConfig(BaseModel): return vertex_request def _transform_openai_request_to_fine_tuned_embedding_request( - self, input: Union[list, str], optional_params: dict, model: str + self, input: list | str, optional_params: dict, model: str ) -> VertexEmbeddingRequest: """ Transforms an openai request to a vertex fine-tuned embedding request. @@ -173,7 +173,7 @@ class VertexAITextEmbeddingConfig(BaseModel): ``` """ vertex_request: VertexEmbeddingRequest = VertexEmbeddingRequest() - vertex_text_embedding_input_list: List[TextEmbeddingFineTunedInput] = [] + vertex_text_embedding_input_list: list[TextEmbeddingFineTunedInput] = [] if isinstance(input, str): input = [input] # Convert single string to list for uniform processing @@ -192,8 +192,8 @@ class VertexAITextEmbeddingConfig(BaseModel): def create_embedding_input( self, content: str, - task_type: Optional[TaskType] = None, - title: Optional[str] = None, + task_type: TaskType | None = None, + title: str | None = None, ) -> TextEmbeddingInput: """ Creates a TextEmbeddingInput object. diff --git a/litellm/llms/vertex_ai/vertex_embeddings/types.py b/litellm/llms/vertex_ai/vertex_embeddings/types.py index bf73f4d193a..d1c949b0ca7 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/types.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/types.py @@ -3,7 +3,6 @@ Types for Vertex Embeddings Requests """ from enum import Enum -from typing import Dict, List, Optional, Union from typing_extensions import TypedDict @@ -21,14 +20,14 @@ class TaskType(str, Enum): class TextEmbeddingInput(TypedDict, total=False): content: str - task_type: Optional[TaskType] - title: Optional[str] + task_type: TaskType | None + title: str | None class TextEmbeddingBGEInput(TypedDict, total=False): prompt: str - task_type: Optional[TaskType] - title: Optional[str] + task_type: TaskType | None + title: str | None # Fine-tuned models require a different input format @@ -38,25 +37,21 @@ class TextEmbeddingFineTunedInput(TypedDict, total=False): class TextEmbeddingFineTunedParameters(TypedDict, total=False): - max_new_tokens: Optional[int] - temperature: Optional[float] - top_p: Optional[float] - top_k: Optional[int] + max_new_tokens: int | None + temperature: float | None + top_p: float | None + top_k: int | None class EmbeddingParameters(TypedDict, total=False): - auto_truncate: Optional[bool] - output_dimensionality: Optional[int] + auto_truncate: bool | None + output_dimensionality: int | None class VertexEmbeddingRequest(TypedDict, total=False): - instances: Union[ - List[TextEmbeddingInput], - List[TextEmbeddingBGEInput], - List[TextEmbeddingFineTunedInput], - ] - parameters: Optional[Union[EmbeddingParameters, TextEmbeddingFineTunedParameters]] - labels: Optional[Dict[str, str]] + instances: list[TextEmbeddingInput] | list[TextEmbeddingBGEInput] | list[TextEmbeddingFineTunedInput] + parameters: EmbeddingParameters | TextEmbeddingFineTunedParameters | None + labels: dict[str, str] | None # Example usage: diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/main.py b/litellm/llms/vertex_ai/vertex_gemma_models/main.py index b66c8a4c4e5..28ea006cc87 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/main.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/main.py @@ -20,7 +20,6 @@ https://{ENDPOINT_NUMBER}.{location}-{REGION_NUMBER}.prediction.vertexai.goog/v1 """ from collections.abc import Callable -from typing import Optional, Union import httpx # type: ignore @@ -42,11 +41,11 @@ class VertexAIGemmaModels(VertexBase): print_verbose: Callable, encoding, logging_obj, - api_base: Optional[str], + api_base: str | None, optional_params: dict, custom_prompt_dict: dict, - headers: Optional[dict], - timeout: Union[float, httpx.Timeout], + headers: dict | None, + timeout: float | httpx.Timeout, litellm_params: dict, vertex_project=None, vertex_location=None, diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index 5313e420155..d1ae3740cb7 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -9,7 +9,7 @@ The actual message transformation reuses OpenAIGPTConfig since Gemma uses OpenAI """ from collections.abc import Callable -from typing import Any, Dict, List, Optional, Union, cast +from typing import Any, cast import httpx @@ -32,9 +32,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): def should_fake_stream( self, - model: Optional[str], - stream: Optional[bool], - custom_llm_provider: Optional[str] = None, + model: str | None, + stream: bool | None, + custom_llm_provider: str | None = None, ) -> bool: """ Vertex AI Gemma models do not support streaming. @@ -46,7 +46,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> Union[ModelResponse, Any]: + ) -> ModelResponse | Any: """ Helper method to return fake stream iterator if streaming is requested. @@ -66,7 +66,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -107,8 +107,8 @@ class VertexGemmaConfig(OpenAIGPTConfig): def _unwrap_predictions_response( self, - response_json: Dict[str, Any], - ) -> Dict[str, Any]: + response_json: dict[str, Any], + ) -> dict[str, Any]: """ Unwrap the Vertex Gemma predictions format to OpenAI format. @@ -136,9 +136,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): optional_params: dict, acompletion: bool, litellm_params: dict, - logger_fn: Optional[Callable] = None, - client: Optional[httpx.Client] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + logger_fn: Callable | None = None, + client: httpx.Client | None = None, + timeout: float | httpx.Timeout | None = None, encoding=None, custom_llm_provider: str = "vertex_ai", ): @@ -186,7 +186,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj: Any, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding: Any, ): """Synchronous completion request""" @@ -276,7 +276,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): logging_obj: Any, optional_params: dict, litellm_params: dict, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: float | httpx.Timeout | None, encoding: Any, ): """Asynchronous completion request""" diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 788261ac1fe..b3ffa1d40be 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -8,7 +8,7 @@ import asyncio import json import os import threading -from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal from urllib.parse import urlparse import litellm @@ -40,15 +40,15 @@ else: class VertexBase: def __init__(self) -> None: super().__init__() - self.access_token: Optional[str] = None - self.refresh_token: Optional[str] = None - self._credentials: Optional[GoogleCredentialsObject] = None - self._credentials_project_mapping: Dict[ - Tuple[Optional[VERTEX_CREDENTIALS_TYPES], Optional[str]], - Tuple[GoogleCredentialsObject, Optional[str]], + self.access_token: str | None = None + self.refresh_token: str | None = None + self._credentials: GoogleCredentialsObject | None = None + self._credentials_project_mapping: dict[ + tuple[VERTEX_CREDENTIALS_TYPES | None, str | None], + tuple[GoogleCredentialsObject, str | None], ] = {} - self.project_id: Optional[str] = None - self.async_handler: Optional[AsyncHTTPHandler] = None + self.project_id: str | None = None + self.async_handler: AsyncHTTPHandler | None = None # Per-credential-key asyncio.Lock for single-flight async refresh. # Prevents thundering herd when token expires under high concurrency. # Uses a regular dict (not WeakValueDictionary) so the lock identity is @@ -58,17 +58,17 @@ class VertexBase: # each lock; the entry is pruned when the count reaches zero, so the # dict stays bounded even in long-running high-cardinality deployments # without depending on any private asyncio internals. - self._async_refresh_locks: Dict[tuple, asyncio.Lock] = {} - self._async_refresh_lock_refcounts: Dict[tuple, int] = {} + self._async_refresh_locks: dict[tuple, asyncio.Lock] = {} + self._async_refresh_lock_refcounts: dict[tuple, int] = {} # Tracks in-flight background refresh tasks to avoid duplicate refreshes. - self._background_refresh_tasks: Dict[tuple, asyncio.Task] = {} + self._background_refresh_tasks: dict[tuple, asyncio.Task] = {} # Protects the sync get_access_token refresh path. # Use RLock so that the reauthentication retry path (which calls # back into get_access_token while still holding the lock) can # re-acquire it without deadlocking the current thread. self._sync_refresh_lock = threading.RLock() - def get_vertex_region(self, vertex_region: Optional[str], model: str) -> str: + def get_vertex_region(self, vertex_region: str | None, model: str) -> str: import litellm # Try to get supported_regions directly from model_cost @@ -96,9 +96,9 @@ class VertexBase: def load_auth( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], - ) -> Tuple[Any, str]: + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + ) -> tuple[Any, str]: if credentials is not None: if isinstance(credentials, str): _is_path = os.path.exists( @@ -120,12 +120,12 @@ class VertexBase: raise Exception( "Unable to load vertex credentials from environment. " "Ensure the JSON is valid (check for unescaped newlines in private_key). " - "Parse error: {}".format(type(e).__name__) + f"Parse error: {type(e).__name__}" ) elif isinstance(credentials, dict): json_obj = credentials else: - raise ValueError("Invalid credentials type: {}".format(type(credentials))) + raise ValueError(f"Invalid credentials type: {type(credentials)}") # Check if the JSON object contains Workload Identity Federation configuration if "type" in json_obj and json_obj["type"] == "external_account": @@ -258,7 +258,7 @@ class VertexBase: def get_default_vertex_location(self) -> str: return "us-central1" - def get_api_base(self, api_base: Optional[str], vertex_location: Optional[str]) -> str: + def get_api_base(self, api_base: str | None, vertex_location: str | None) -> str: if api_base: return api_base return get_vertex_base_url(vertex_location or self.get_default_vertex_location()) @@ -268,9 +268,9 @@ class VertexBase: vertex_location: str, vertex_project: str, partner: VertexPartnerProvider, - stream: Optional[bool], + stream: bool | None, model: str, - api_base: Optional[str] = None, + api_base: str | None = None, ) -> str: """Return the base url for the vertex partner models""" @@ -296,12 +296,12 @@ class VertexBase: def get_complete_vertex_url( self, - custom_api_base: Optional[str], - vertex_location: Optional[str], - vertex_project: Optional[str], + custom_api_base: str | None, + vertex_location: str | None, + vertex_project: str | None, project_id: str, partner: VertexPartnerProvider, - stream: Optional[bool], + stream: bool | None, model: str, ) -> str: # Use get_vertex_region to handle global-only models @@ -391,8 +391,8 @@ class VertexBase: def _try_get_cached_token( self, credential_cache_key: tuple, - project_id: Optional[str], - ) -> Optional[Tuple[str, str]]: + project_id: str | None, + ) -> tuple[str, str] | None: """ Look up cached credentials and return (token, project_id) if the token is FRESH. Returns None if not cached or not fresh. @@ -414,8 +414,8 @@ class VertexBase: def _try_get_usable_cached_token( self, credential_cache_key: tuple, - project_id: Optional[str], - ) -> Optional[Tuple[str, str, "TokenState", Any, Optional[str]]]: + project_id: str | None, + ) -> tuple[str, str, "TokenState", Any, str | None] | None: """ Look up cached credentials and return usable token info for FRESH or STALE tokens (both are still valid for outbound requests). STALE @@ -438,7 +438,7 @@ class VertexBase: return None return creds.token, resolved_project, token_state, creds, cached_project_id - def _unpack_cached_credentials(self, credential_cache_key: tuple) -> Tuple[Any, Optional[str]]: + def _unpack_cached_credentials(self, credential_cache_key: tuple) -> tuple[Any, str | None]: """ Return (credentials, project_id) from the cache, or (None, None) if not cached. Handles both tuple and legacy cache formats. @@ -471,10 +471,10 @@ class VertexBase: async def _load_and_cache_credentials( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, credential_cache_key: tuple, - ) -> Tuple[Any, Optional[str]]: + ) -> tuple[Any, str | None]: """Load credentials via load_auth (in thread) and cache the result.""" try: _credentials, credential_project_id = await asyncify(self.load_auth)( @@ -496,7 +496,7 @@ class VertexBase: self, credentials: Any, credential_cache_key: tuple, - credential_project_id: Optional[str], + credential_project_id: str | None, ) -> None: """ Refresh credentials in the background without blocking the calling request. @@ -548,7 +548,7 @@ class VertexBase: self, credentials: Any, credential_cache_key: tuple, - credential_project_id: Optional[str], + credential_project_id: str | None, ) -> None: """Kick off a single background refresh for ``credential_cache_key``. @@ -573,12 +573,12 @@ class VertexBase: def _ensure_access_token( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Returns auth token and project id """ @@ -601,19 +601,19 @@ class VertexBase: def _check_custom_proxy( self, - api_base: Optional[str], + api_base: str | None, custom_llm_provider: str, - gemini_api_key: Optional[str], + gemini_api_key: str | None, endpoint: str, - stream: Optional[bool], - auth_header: Optional[str], + stream: bool | None, + auth_header: str | None, url: str, - model: Optional[str] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - vertex_api_version: Optional[Literal["v1", "v1beta1"]] = None, + model: str | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + vertex_api_version: Literal["v1", "v1beta1"] | None = None, use_psc_endpoint_format: bool = False, - ) -> Tuple[Optional[str], str]: + ) -> tuple[str | None, str]: """ for cloudflare ai gateway - https://github.com/BerriAI/litellm/issues/4317 @@ -637,7 +637,7 @@ class VertexBase: # For Gemini (Google AI Studio), construct the full path like other providers if model is None: raise ValueError("Model parameter is required for Gemini custom API base URLs") - url = "{}/models/{}:{}".format(api_base, model, endpoint) + url = f"{api_base}/models/{model}:{endpoint}" if gemini_api_key is None: raise ValueError( "Missing Gemini API key. Set the GEMINI_API_KEY or GOOGLE_API_KEY environment variable." @@ -668,7 +668,7 @@ class VertexBase: elif urlparse(api_base).path in ("", "/"): url = api_base.rstrip("/") + urlparse(url).path else: - url = "{}:{}".format(api_base, endpoint) + url = f"{api_base}:{endpoint}" if stream is True: url = url + "?alt=sse" return auth_header, url @@ -676,18 +676,18 @@ class VertexBase: def _get_token_and_url( self, model: str, - auth_header: Optional[str], - gemini_api_key: Optional[str], - vertex_project: Optional[str], - vertex_location: Optional[str], - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], - stream: Optional[bool], + auth_header: str | None, + gemini_api_key: str | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + stream: bool | None, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], - api_base: Optional[str], - should_use_v1beta1_features: Optional[bool] = False, + api_base: str | None, + should_use_v1beta1_features: bool | None = False, mode: all_gemini_url_modes = "chat", use_psc_endpoint_format: bool = False, - ) -> Tuple[Optional[str], str]: + ) -> tuple[str | None, str]: """ Internal function. Returns the token and url for the call. @@ -696,7 +696,7 @@ class VertexBase: Returns token, url """ - version: Optional[Literal["v1beta1", "v1"]] = None + version: Literal["v1beta1", "v1"] | None = None if custom_llm_provider == "gemini": if not gemini_api_key: raise ValueError( @@ -742,11 +742,11 @@ class VertexBase: def _handle_reauthentication( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], - credential_cache_key: Tuple, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + credential_cache_key: tuple, error: Exception, - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Handle reauthentication when credentials refresh fails. @@ -783,18 +783,18 @@ class VertexBase: except Exception as retry_error: verbose_logger.error( f"Reauthentication retry failed for project_id: {project_id}. " - f"Original error: {str(error)}. Retry error: {str(retry_error)}" + f"Original error: {error!s}. Retry error: {retry_error!s}" ) # Re-raise the original error for better context raise error async def _handle_reauthentication_async( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], - credential_cache_key: Tuple, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + credential_cache_key: tuple, error: Exception, - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Async reauthentication retry that stays within the per-key async lock. """ @@ -828,9 +828,7 @@ class VertexBase: if _credentials.token is None or not isinstance(_credentials.token, str): raise ValueError( - "Could not resolve credentials token. Got None or non-string token (type={})".format( - type(_credentials.token).__name__ - ) + f"Could not resolve credentials token. Got None or non-string token (type={type(_credentials.token).__name__})" ) if project_id is None: raise ValueError("Could not resolve project_id") @@ -839,16 +837,16 @@ class VertexBase: except Exception as retry_error: verbose_logger.error( f"Async reauthentication retry failed for project_id: {project_id}. " - f"Original error: {str(error)}. Retry error: {str(retry_error)}" + f"Original error: {error!s}. Retry error: {retry_error!s}" ) raise error def get_access_token( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, _retry_reauth: bool = False, - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Get access token and project id @@ -870,7 +868,7 @@ class VertexBase: # Convert dict credentials to string for caching cache_credentials = json.dumps(credentials) if isinstance(credentials, dict) else credentials credential_cache_key = (cache_credentials, project_id) - _credentials: Optional[GoogleCredentialsObject] = None + _credentials: GoogleCredentialsObject | None = None verbose_logger.debug(f"Checking cached credentials for project_id: {project_id}") @@ -899,15 +897,13 @@ class VertexBase: _credentials, credential_project_id = self.load_auth(credentials=credentials, project_id=project_id) except Exception as e: verbose_logger.exception( - f"Failed to load vertex credentials. Check to see if credentials containing partial/invalid information. Error: {str(e)}" + f"Failed to load vertex credentials. Check to see if credentials containing partial/invalid information. Error: {e!s}" ) raise e if _credentials is None: raise ValueError( - "Could not resolve credentials - either dynamically or from environment, for project_id: {}".format( - project_id - ) + f"Could not resolve credentials - either dynamically or from environment, for project_id: {project_id}" ) # Cache the project_id and credentials from load_auth result (resolved project_id) self._credentials_project_mapping[credential_cache_key] = ( @@ -957,9 +953,7 @@ class VertexBase: ## VALIDATION STEP if _credentials.token is None or not isinstance(_credentials.token, str): raise ValueError( - "Could not resolve credentials token. Got None or non-string token (type={})".format( - type(_credentials.token).__name__ - ) + f"Could not resolve credentials token. Got None or non-string token (type={type(_credentials.token).__name__})" ) if project_id is None: @@ -969,9 +963,9 @@ class VertexBase: async def get_access_token_async( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], - ) -> Tuple[str, str]: + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + ) -> tuple[str, str]: """ Async version of get_access_token with single-flight refresh coordination. @@ -1084,9 +1078,7 @@ class VertexBase: # Final validation if _credentials.token is None or not isinstance(_credentials.token, str): raise ValueError( - "Could not resolve credentials token. Got None or non-string token (type={})".format( - type(_credentials.token).__name__ - ) + f"Could not resolve credentials token. Got None or non-string token (type={type(_credentials.token).__name__})" ) if project_id is None: raise ValueError("Could not resolve project_id") @@ -1097,12 +1089,12 @@ class VertexBase: async def _ensure_access_token_async( self, - credentials: Optional[VERTEX_CREDENTIALS_TYPES], - project_id: Optional[str], + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - ) -> Tuple[str, str]: + ) -> tuple[str, str]: """ Async version of _ensure_access_token """ @@ -1114,7 +1106,7 @@ class VertexBase: project_id=project_id, ) - def set_headers(self, auth_header: Optional[str], extra_headers: Optional[dict]) -> dict: + def set_headers(self, auth_header: str | None, extra_headers: dict | None) -> dict: headers = { "Content-Type": "application/json", } @@ -1126,7 +1118,7 @@ class VertexBase: return headers @staticmethod - def get_vertex_ai_project(litellm_params: dict) -> Optional[str]: + def get_vertex_ai_project(litellm_params: dict) -> str | None: return ( litellm_params.pop("vertex_project", None) or litellm_params.pop("vertex_ai_project", None) @@ -1135,7 +1127,7 @@ class VertexBase: ) @staticmethod - def get_vertex_ai_credentials(litellm_params: dict) -> Optional[str]: + def get_vertex_ai_credentials(litellm_params: dict) -> str | None: return ( litellm_params.pop("vertex_credentials", None) or litellm_params.pop("vertex_ai_credentials", None) @@ -1143,7 +1135,7 @@ class VertexBase: ) @staticmethod - def get_vertex_ai_location(litellm_params: dict) -> Optional[str]: + def get_vertex_ai_location(litellm_params: dict) -> str | None: return ( litellm_params.pop("vertex_location", None) or litellm_params.pop("vertex_ai_location", None) @@ -1153,7 +1145,7 @@ class VertexBase: ) @staticmethod - def safe_get_vertex_ai_project(litellm_params: dict) -> Optional[str]: + def safe_get_vertex_ai_project(litellm_params: dict) -> str | None: """ Safely get Vertex AI project without mutating the litellm_params dict. @@ -1174,7 +1166,7 @@ class VertexBase: ) @staticmethod - def safe_get_vertex_ai_credentials(litellm_params: dict) -> Optional[str]: + def safe_get_vertex_ai_credentials(litellm_params: dict) -> str | None: """ Safely get Vertex AI credentials without mutating the litellm_params dict. @@ -1194,7 +1186,7 @@ class VertexBase: ) @staticmethod - def safe_get_vertex_ai_location(litellm_params: dict) -> Optional[str]: + def safe_get_vertex_ai_location(litellm_params: dict) -> str | None: """ Safely get Vertex AI location without mutating the litellm_params dict. diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index e09106ea4f6..75cf62f2ffb 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -17,7 +17,6 @@ Vertex Documentation for using the OpenAI /chat/completions endpoint: https://gi """ from collections.abc import Callable -from typing import Optional, Union import httpx # type: ignore @@ -42,9 +41,9 @@ def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: def create_vertex_url( vertex_location: str, vertex_project: str, - stream: Optional[bool], + stream: bool | None, model: str, - api_base: Optional[str] = None, + api_base: str | None = None, ) -> str: """Return the api base for vertex model garden (without /chat/completions).""" base_url = get_vertex_base_url(vertex_location) @@ -65,11 +64,11 @@ class VertexAIModelGardenModels(VertexBase): print_verbose: Callable, encoding, logging_obj, - api_base: Optional[str], + api_base: str | None, optional_params: dict, custom_prompt_dict: dict, - headers: Optional[dict], - timeout: Union[float, httpx.Timeout], + headers: dict | None, + timeout: float | httpx.Timeout, litellm_params: dict, vertex_project=None, vertex_location=None, diff --git a/litellm/llms/vertex_ai/videos/transformation.py b/litellm/llms/vertex_ai/videos/transformation.py index 98af8ca30ea..003155b44ce 100644 --- a/litellm/llms/vertex_ai/videos/transformation.py +++ b/litellm/llms/vertex_ai/videos/transformation.py @@ -7,7 +7,7 @@ Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-refer import base64 import time -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx from httpx._types import RequestFiles @@ -41,10 +41,10 @@ else: def _build_vertex_video_usage_from_request_data( - request_data: Optional[Dict[str, Any]], -) -> Dict[str, Any]: + request_data: dict[str, Any] | None, +) -> dict[str, Any]: """Build usage metadata (duration, resolution) for video cost calculation.""" - usage_data: Dict[str, Any] = {} + usage_data: dict[str, Any] = {} if not request_data: return usage_data @@ -61,7 +61,7 @@ def _build_vertex_video_usage_from_request_data( return usage_data -def _convert_image_to_vertex_format(image_file) -> Dict[str, str]: +def _convert_image_to_vertex_format(image_file) -> dict[str, str]: """ Convert image file to Vertex AI format with base64 encoding and MIME type. @@ -96,7 +96,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): VertexBase.__init__(self) @staticmethod - def extract_model_from_operation_name(operation_name: str) -> Optional[str]: + def extract_model_from_operation_name(operation_name: str) -> str | None: """ Extract the model name from a Vertex AI operation name. @@ -125,7 +125,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): video_create_optional_params: VideoCreateOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map OpenAI-style parameters to Veo format. @@ -135,7 +135,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): - size → aspectRatio (e.g., "1280x720" → "16:9") - seconds → durationSeconds (defaults to 4 seconds if not provided) """ - mapped_params: Dict[str, Any] = {} + mapped_params: dict[str, Any] = {} # Map input_reference to image (will be processed in transform_video_create_request) if "input_reference" in video_create_optional_params: @@ -168,7 +168,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): return mapped_params - def _convert_size_to_aspect_ratio(self, size: str) -> Optional[str]: + def _convert_size_to_aspect_ratio(self, size: str) -> str | None: """ Convert OpenAI size format to Veo aspectRatio format. @@ -190,8 +190,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): self, headers: dict, model: str, - api_key: Optional[str] = None, - litellm_params: Optional[Union[GenericLiteLLMParams, dict]] = None, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | dict | None = None, ) -> dict: """ Validate environment and return headers for Vertex AI OCR. @@ -201,7 +201,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): # Extract Vertex AI parameters using safe helpers from VertexBase # Use safe_get_* methods that don't mutate litellm_params dict # Ensure litellm_params is a dict for type checking - params_dict: Dict[str, Any] = cast(Dict[str, Any], litellm_params) if litellm_params is not None else {} + params_dict: dict[str, Any] = cast(dict[str, Any], litellm_params) if litellm_params is not None else {} vertex_project = VertexBase.safe_get_vertex_ai_project(litellm_params=params_dict) vertex_credentials = VertexBase.safe_get_vertex_ai_credentials(litellm_params=params_dict) @@ -223,7 +223,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): def get_complete_url( self, model: str, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ @@ -264,10 +264,10 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): model: str, prompt: str, api_base: str, - video_create_optional_request_params: Dict, + video_create_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[Dict, RequestFiles, str]: + ) -> tuple[dict, RequestFiles, str]: """ Transform the video creation request for Veo API. @@ -289,7 +289,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): } """ # Build instance with prompt - instance_dict: Dict[str, Any] = {"prompt": prompt} + instance_dict: dict[str, Any] = {"prompt": prompt} params_copy = video_create_optional_request_params.copy() # Check if user wants to provide full instance dict @@ -324,13 +324,13 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): # {"parameters": {"parameters": {...}}} ← wrong # {"parameters": {...}} ← correct nested_params = params_copy.pop("parameters", None) - vertex_params: Dict[str, Any] = {} + vertex_params: dict[str, Any] = {} if isinstance(nested_params, dict): vertex_params.update(nested_params) vertex_params.update(params_copy) # Build request data directly (TypedDict doesn't have model_dump) - request_data: Dict[str, Any] = {"instances": [instance_dict]} + request_data: dict[str, Any] = {"instances": [instance_dict]} # Only add parameters if there are any if vertex_params: @@ -347,8 +347,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): model: str, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: """ Transform the Veo video creation response. @@ -385,7 +385,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Transform the video status retrieve request for Veo API. @@ -414,7 +414,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """ Transform the Veo operation status response. @@ -488,8 +488,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - variant: Optional[str] = None, - ) -> Tuple[str, Dict]: + variant: str | None = None, + ) -> tuple[str, dict]: """ Transform the video content request for Veo API. @@ -548,8 +548,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Video remix is not supported by Veo API. """ @@ -561,7 +561,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> VideoObject: """Video remix is not supported.""" raise NotImplementedError("Video remix is not supported by Vertex AI Veo.") @@ -571,11 +571,11 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_query: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Video list is not supported by Veo API. """ @@ -588,8 +588,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - ) -> Dict[str, str]: + custom_llm_provider: str | None = None, + ) -> dict[str, str]: """Video list is not supported.""" raise NotImplementedError("Video list is not supported by Vertex AI Veo.") @@ -599,7 +599,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """ Video delete is not supported by Veo API. """ @@ -633,7 +633,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> Tuple[str, Dict]: + ) -> tuple[str, dict]: """Return the fetchPredictOperation URL and body needed to retrieve the source video.""" return self.transform_video_status_retrieve_request( video_id=video_id, @@ -649,9 +649,9 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: Optional[Dict[str, Any]] = None, - prefetched_source_data: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict]: + extra_body: dict[str, Any] | None = None, + prefetched_source_data: dict[str, Any] | None = None, + ) -> tuple[str, dict]: """ Build a predictLongRunning edit request from the pre-fetched source video. @@ -672,7 +672,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): raise ValueError("No videos found in the completed operation. Cannot edit.") source_video = videos[0] - video_input: Dict[str, Any] = {} + video_input: dict[str, Any] = {} if "gcsUri" in source_video: video_input["gcsUri"] = source_video["gcsUri"] elif "bytesBase64Encoded" in source_video: @@ -684,13 +684,13 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): operation_name = extract_original_video_id(video_id) model = self.extract_model_from_operation_name(operation_name) or "" - instance_dict: Dict[str, Any] = {"prompt": prompt, "video": video_input} - request_data: Dict[str, Any] = {"instances": [instance_dict]} + instance_dict: dict[str, Any] = {"prompt": prompt, "video": video_input} + request_data: dict[str, Any] = {"instances": [instance_dict]} if extra_body: extra_body_copy = dict(extra_body) nested_params = extra_body_copy.pop("parameters", None) - vertex_params: Dict[str, Any] = {} + vertex_params: dict[str, Any] = {} if isinstance(nested_params, dict): vertex_params.update(nested_params) vertex_params.update(extra_body_copy) @@ -704,8 +704,8 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict] = None, + custom_llm_provider: str | None = None, + request_data: dict | None = None, ) -> VideoObject: """ Transform the Veo video edit response. @@ -753,9 +753,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None): raise NotImplementedError("video extension is not supported for Vertex AI") - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: from litellm.llms.vertex_ai.common_utils import VertexAIError return VertexAIError( diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 1d6b8d7897e..47614c7d00c 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - import httpx import litellm @@ -15,9 +13,9 @@ class VLLMError(BaseLLMException): self, status_code: int, message: str, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[Union[httpx.Headers, dict]] = None, + request: httpx.Request | None = None, + response: httpx.Response | None = None, + headers: httpx.Headers | dict | None = None, ): super().__init__( status_code=status_code, @@ -33,18 +31,18 @@ class VLLMModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is not None: headers["x-api-key"] = api_key return headers @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: api_base = api_base or get_secret_str("VLLM_API_BASE") if api_base is None: raise ValueError( @@ -53,14 +51,14 @@ class VLLMModelInfo(BaseLLMModelInfo): return api_base @staticmethod - def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + def get_api_key(api_key: str | None = None) -> str | None: return None @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: return model - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base = VLLMModelInfo.get_api_base(api_base) api_key = VLLMModelInfo.get_api_key(api_key) endpoint = "/v1/models" @@ -80,7 +78,5 @@ class VLLMModelInfo(BaseLLMModelInfo): return [model["id"] for model in models] - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VLLMError(status_code=status_code, message=error_message, headers=headers) diff --git a/litellm/llms/vllm/completion/transformation.py b/litellm/llms/vllm/completion/transformation.py index e03b07f9897..9c764074c10 100644 --- a/litellm/llms/vllm/completion/transformation.py +++ b/litellm/llms/vllm/completion/transformation.py @@ -11,5 +11,3 @@ class VLLMConfig(HostedVLLMChatConfig): """ VLLM SDK supports the same OpenAI params as hosted_vllm. """ - - pass diff --git a/litellm/llms/vllm/passthrough/transformation.py b/litellm/llms/vllm/passthrough/transformation.py index cc8a78fb50d..d2f445a17d6 100644 --- a/litellm/llms/vllm/passthrough/transformation.py +++ b/litellm/llms/vllm/passthrough/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig @@ -14,13 +14,13 @@ class VLLMPassthroughConfig(VLLMModelInfo, BasePassthroughConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, endpoint: str, - request_query_params: Optional[dict], + request_query_params: dict | None, litellm_params: dict, - ) -> Tuple["URL", str]: + ) -> tuple["URL", str]: base_target_url = self.get_api_base(api_base) if base_target_url is None: diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py index fc0098e84d9..27db76c164f 100644 --- a/litellm/llms/volcengine/__init__.py +++ b/litellm/llms/volcengine/__init__.py @@ -19,8 +19,8 @@ __all__ = [ "VolcEngineChatConfig", "VolcEngineConfig", # backward compatibility "VolcEngineEmbeddingConfig", - "VolcEngineResponsesAPIConfig", "VolcEngineError", + "VolcEngineResponsesAPIConfig", "get_volcengine_base_url", "get_volcengine_headers", ] diff --git a/litellm/llms/volcengine/chat/transformation.py b/litellm/llms/volcengine/chat/transformation.py index c6dbbdbce60..4063d6ffc3f 100644 --- a/litellm/llms/volcengine/chat/transformation.py +++ b/litellm/llms/volcengine/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import Optional, Union - from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig @@ -8,31 +6,31 @@ class VolcEngineChatConfig(OpenAILikeChatConfig): Reference: https://www.volcengine.com/docs/82379/1494384 """ - frequency_penalty: Optional[int] = None - function_call: Optional[Union[str, dict]] = None - functions: Optional[list] = None - logit_bias: Optional[dict] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - temperature: Optional[int] = None - top_p: Optional[int] = None - response_format: Optional[dict] = None + frequency_penalty: int | None = None + function_call: str | dict | None = None + functions: list | None = None + logit_bias: dict | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: int | None = None + stop: str | list | None = None + temperature: int | None = None + top_p: int | None = None + response_format: dict | None = None def __init__( self, - frequency_penalty: Optional[int] = None, - function_call: Optional[Union[str, dict]] = None, - functions: Optional[list] = None, - logit_bias: Optional[dict] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - temperature: Optional[int] = None, - top_p: Optional[int] = None, - response_format: Optional[dict] = None, + frequency_penalty: int | None = None, + function_call: str | dict | None = None, + functions: list | None = None, + logit_bias: dict | None = None, + max_tokens: int | None = None, + n: int | None = None, + presence_penalty: int | None = None, + stop: str | list | None = None, + temperature: int | None = None, + top_p: int | None = None, + response_format: dict | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): diff --git a/litellm/llms/volcengine/common_utils.py b/litellm/llms/volcengine/common_utils.py index be639086437..160449c951d 100644 --- a/litellm/llms/volcengine/common_utils.py +++ b/litellm/llms/volcengine/common_utils.py @@ -2,8 +2,6 @@ Common utilities for Volcengine LLM provider """ -from typing import Optional - import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -14,14 +12,14 @@ class VolcEngineError(BaseLLMException): Custom exception class for Volcengine provider errors. """ - def __init__(self, status_code: int, message: str, headers: Optional[httpx.Headers] = None): + def __init__(self, status_code: int, message: str, headers: httpx.Headers | None = None): self.status_code = status_code self.message = message self.headers = headers or httpx.Headers() super().__init__(status_code=status_code, message=message, headers=dict(self.headers)) -def get_volcengine_base_url(api_base: Optional[str] = None) -> str: +def get_volcengine_base_url(api_base: str | None = None) -> str: """ Get the base URL for Volcengine API calls. @@ -36,7 +34,7 @@ def get_volcengine_base_url(api_base: Optional[str] = None) -> str: return "https://ark.cn-beijing.volces.com" -def get_volcengine_headers(api_key: str, extra_headers: Optional[dict] = None) -> dict: +def get_volcengine_headers(api_key: str, extra_headers: dict | None = None) -> dict: """ Get headers for Volcengine API calls. diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py index cb497c9f155..5a0b59d411c 100644 --- a/litellm/llms/volcengine/embedding/transformation.py +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -3,13 +3,16 @@ Volcengine Embedding Transformation Transforms OpenAI embedding requests to Volcengine format """ -from typing import List, Optional, Union, Dict, Any +from typing import Any + import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.llms.base_llm.chat.transformation import BaseLLMException + from ..common_utils import get_volcengine_base_url, get_volcengine_headers @@ -21,7 +24,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): def __init__( self, - encoding_format: Optional[str] = None, + encoding_format: str | None = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -32,7 +35,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): def get_config(cls): return super().get_config() - def get_supported_openai_params(self, model: str) -> List[str]: + def get_supported_openai_params(self, model: str) -> list[str]: """ Get the list of OpenAI parameters supported by Volcengine embedding models. @@ -50,12 +53,12 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Get the complete URL for volcengine embedding API calls. @@ -80,11 +83,11 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): def map_openai_params( self, - non_default_params: Dict[str, Any], - optional_params: Dict[str, Any], + non_default_params: dict[str, Any], + optional_params: dict[str, Any], model: str, drop_params: bool, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Map OpenAI embedding parameters to Volcengine format. @@ -150,7 +153,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, @@ -159,7 +162,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): try: response_json = raw_response.json() except Exception as e: - raise ValueError(f"Failed to parse Volcengine response as JSON: {str(e)}") + raise ValueError(f"Failed to parse Volcengine response as JSON: {e!s}") # Volcengine response format matches OpenAI format closely # Just need to ensure all required fields are present @@ -181,11 +184,11 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: """Validate environment and return headers""" # Get Volcengine headers @@ -194,9 +197,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): volcengine_headers = get_volcengine_headers(api_key) return {**headers, **volcengine_headers} - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: """Get error class for Volcengine errors""" from ..common_utils import VolcEngineError diff --git a/litellm/llms/voyage/embedding/transformation.py b/litellm/llms/voyage/embedding/transformation.py index 7193fd2f10a..ee6d99951d8 100644 --- a/litellm/llms/voyage/embedding/transformation.py +++ b/litellm/llms/voyage/embedding/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -15,7 +13,7 @@ class VoyageError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] = {}, + headers: dict | httpx.Headers = {}, ): self.status_code = status_code self.message = message @@ -38,12 +36,12 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: if not api_base.endswith("/embeddings"): @@ -79,11 +77,11 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = ( @@ -114,7 +112,7 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -136,7 +134,5 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VoyageError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index d7cca3c87a8..ec37ccaffb0 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -3,8 +3,6 @@ This module is used to transform the request and response for the Voyage context This would be used for all the contextualized embeddings models in Voyage. """ -from typing import List, Optional, Union - import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -20,7 +18,7 @@ class VoyageError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] = {}, + headers: dict | httpx.Headers = {}, ): self.status_code = status_code self.message = message @@ -43,12 +41,12 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: if not api_base.endswith("/contextualizedembeddings"): @@ -81,11 +79,11 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = ( @@ -100,7 +98,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): def transform_embedding_request( self, model: str, - input: Union[AllEmbeddingInputValues, List[List[str]]], + input: AllEmbeddingInputValues | list[list[str]], optional_params: dict, headers: dict, ) -> dict: @@ -116,7 +114,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -138,9 +136,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VoyageError(message=error_message, status_code=status_code, headers=headers) @staticmethod diff --git a/litellm/llms/voyage/embedding/transformation_multimodal.py b/litellm/llms/voyage/embedding/transformation_multimodal.py index 916037054ef..3d3511c920c 100644 --- a/litellm/llms/voyage/embedding/transformation_multimodal.py +++ b/litellm/llms/voyage/embedding/transformation_multimodal.py @@ -6,7 +6,7 @@ containing content blocks, unlike standard Voyage embeddings which use /v1/embeddings and a string/list `input` field. """ -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -23,7 +23,7 @@ class VoyageMultimodalEmbeddingError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] = {}, + headers: dict | httpx.Headers = {}, ): self.status_code = status_code self.message = message @@ -47,12 +47,12 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: if not api_base.endswith("/multimodalembeddings"): @@ -78,11 +78,11 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = ( @@ -98,7 +98,7 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): ) return {"Authorization": f"Bearer {api_key}"} - def _normalize_content_item(self, item: Dict[str, Any]) -> Dict[str, Any]: + def _normalize_content_item(self, item: dict[str, Any]) -> dict[str, Any]: item_type = item.get("type") if item_type == "image_url": image_url = item.get("image_url") @@ -115,7 +115,7 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): return {"type": "image_url", "image_url": image_url} return item - def _normalize_input_item(self, item: Any) -> Dict[str, Any]: + def _normalize_input_item(self, item: Any) -> dict[str, Any]: if isinstance(item, str): return {"content": [{"type": "text", "text": item}]} if isinstance(item, dict) and "content" in item: @@ -146,7 +146,7 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -168,7 +168,5 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): ) return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VoyageMultimodalEmbeddingError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index e426e39962b..2ce335f8e2b 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -4,7 +4,7 @@ Transformation logic for Voyage AI's /v1/rerank endpoint. Docs - https://docs.voyageai.com/docs/reranker """ -from typing import Any, Dict, List, Tuple, Union +from typing import Any import httpx @@ -32,17 +32,17 @@ class VoyageRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: # Voyage AI uses 'top_k' instead of 'top_n' - optional_params: Dict[str, Any] = {"query": query, "documents": documents} + optional_params: dict[str, Any] = {"query": query, "documents": documents} if top_n is not None: optional_params["top_k"] = top_n if return_documents is not None: @@ -70,10 +70,10 @@ class VoyageRerankConfig(BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, - headers: Dict, + optional_rerank_params: dict, + headers: dict, litellm_params: dict | None = None, - ) -> Dict: + ) -> dict: return {"model": model, **optional_rerank_params} def transform_rerank_response( @@ -83,9 +83,9 @@ class VoyageRerankConfig(BaseRerankConfig): model_response: RerankResponse, logging_obj: LiteLLMLoggingObj, api_key: str | None = None, - request_data: Dict = {}, - optional_params: Dict = {}, - litellm_params: Dict = {}, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, ) -> RerankResponse: if raw_response.status_code != 200: raise VoyageError(message=raw_response.text, status_code=raw_response.status_code) @@ -101,14 +101,14 @@ class VoyageRerankConfig(BaseRerankConfig): ) # Voyage AI returns results in "data" key, not "results" - _results: List[dict] | None = _json_response.get("data") + _results: list[dict] | None = _json_response.get("data") if _results is None: raise ValueError(f"No results found in the response={_json_response}") # Transform to LiteLLM format transformed_results = [] for result in _results: - transformed_result: Dict[str, Any] = { + transformed_result: dict[str, Any] = { "index": result["index"], "relevance_score": result["relevance_score"], } @@ -133,11 +133,11 @@ class VoyageRerankConfig(BaseRerankConfig): def validate_environment( self, - headers: Dict, + headers: dict, model: str, api_key: str | None = None, optional_params: dict | None = None, - ) -> Dict: + ) -> dict: if api_key is None: api_key = get_secret_str("VOYAGE_API_KEY") or get_secret_str("VOYAGE_AI_API_KEY") if api_key is None: @@ -153,7 +153,7 @@ class VoyageRerankConfig(BaseRerankConfig): custom_llm_provider: str | None = None, billed_units: RerankBilledUnits | None = None, model_info: ModelInfo | None = None, - ) -> Tuple[float, float]: + ) -> tuple[float, float]: if ( model_info is None or "input_cost_per_token" not in model_info @@ -166,5 +166,5 @@ class VoyageRerankConfig(BaseRerankConfig): return 0.0, 0.0 return model_info["input_cost_per_token"] * total_tokens, 0.0 - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]): + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers): return VoyageError(message=error_message, status_code=status_code, headers=headers) diff --git a/litellm/llms/watsonx/audio_transcription/transformation.py b/litellm/llms/watsonx/audio_transcription/transformation.py index 6d28790b8d1..6019b2e8355 100644 --- a/litellm/llms/watsonx/audio_transcription/transformation.py +++ b/litellm/llms/watsonx/audio_transcription/transformation.py @@ -4,10 +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 typing import Any, Dict, List, Optional +from typing import Any + +from httpx import Response import litellm -from httpx import Response from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.types.llms.openai import ( AllMessageValues, @@ -36,14 +37,14 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran def validate_environment( self, - headers: Dict, + headers: dict, model: str, - messages: List[AllMessageValues], - optional_params: Dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: """ Validate environment for audio transcription. @@ -63,7 +64,7 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran result.pop("Content-Type", None) return result - def get_supported_openai_params(self, model: str) -> List[OpenAIAudioTranscriptionOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: """ Get the supported OpenAI params for WatsonX audio transcription. """ @@ -123,18 +124,18 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran } # Convert TypedDict to regular dict for AudioTranscriptionRequestData - form_data_dict: Dict[str, Any] = dict(form_data) + form_data_dict: dict[str, Any] = dict(form_data) return AudioTranscriptionRequestData(data=form_data_dict, files=files) def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: """ Construct the complete URL for WatsonX audio transcription. @@ -169,7 +170,7 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran try: raw_response_json = raw_response.json() except Exception as e: - raise ValueError(f"Error transforming response to json: {str(e)}\nResponse: {raw_response.text}") + raise ValueError(f"Error transforming response to json: {e!s}\nResponse: {raw_response.text}") # Extract only valid fields for TranscriptionResponse.__init__() # TranscriptionResponse only accepts 'text' and 'usage' in __init__() diff --git a/litellm/llms/watsonx/chat/handler.py b/litellm/llms/watsonx/chat/handler.py index d7db2d12097..1186ce2926d 100644 --- a/litellm/llms/watsonx/chat/handler.py +++ b/litellm/llms/watsonx/chat/handler.py @@ -1,5 +1,4 @@ from collections.abc import Callable -from typing import Optional, Union import httpx @@ -22,23 +21,23 @@ class WatsonXChatHandler(OpenAILikeChatHandler): *, model: str, messages: list, - api_base: Optional[str], + api_base: str | None, custom_llm_provider: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, encoding, - api_key: Optional[str], + api_key: str | None, logging_obj, optional_params: dict, acompletion=None, litellm_params: dict = {}, - headers: Optional[dict] = None, + headers: dict | None = None, logger_fn=None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - custom_endpoint: Optional[bool] = None, - streaming_decoder: Optional[CustomStreamingDecoder] = None, + timeout: float | httpx.Timeout | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, + custom_endpoint: bool | None = None, + streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ): api_params = _get_api_params(params=optional_params, model=model) diff --git a/litellm/llms/watsonx/chat/transformation.py b/litellm/llms/watsonx/chat/transformation.py index 8c938e8dc4d..adc7035d2f7 100644 --- a/litellm/llms/watsonx/chat/transformation.py +++ b/litellm/llms/watsonx/chat/transformation.py @@ -4,8 +4,6 @@ Translation from OpenAI's `/chat/completions` endpoint to IBM WatsonX's `/text/c Docs: https://cloud.ibm.com/apidocs/watsonx-ai#text-chat """ -from typing import Dict, List, Optional, Tuple, Union - from litellm import verbose_logger from litellm.secret_managers.main import get_secret_str from litellm.types.llms.watsonx import ( @@ -19,7 +17,7 @@ from ..common_utils import IBMWatsonXMixin class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): - def get_supported_openai_params(self, model: str) -> List: + def get_supported_openai_params(self, model: str) -> list: return [ "temperature", # equivalent to temperature "max_tokens", # equivalent to max_new_tokens @@ -38,7 +36,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): "reasoning_effort", ] - def is_tool_choice_option(self, tool_choice: Optional[Union[str, dict]]) -> bool: + def is_tool_choice_option(self, tool_choice: str | dict | None) -> bool: if tool_choice is None: return False if isinstance(tool_choice, str): @@ -72,20 +70,20 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): return super().map_openai_params(non_default_params, optional_params, model, drop_params) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") # type: ignore dynamic_api_key = api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "" # vllm does not require an api key return api_base, dynamic_api_key def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: url = self._get_base_url(api_base=api_base) if model.startswith("deployment/"): @@ -103,7 +101,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): return url @staticmethod - def _apply_prompt_template_core(model: str, messages: List[Dict[str, str]], hf_template_fn) -> Optional[str]: + def _apply_prompt_template_core(model: str, messages: list[dict[str, str]], hf_template_fn) -> str | None: """Core logic for applying prompt templates""" from litellm.litellm_core_utils.prompt_templates.factory import ( custom_prompt, @@ -155,7 +153,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): return None @staticmethod - async def aapply_prompt_template(model: str, messages: List[Dict[str, str]]) -> Optional[str]: + async def aapply_prompt_template(model: str, messages: list[dict[str, str]]) -> str | None: """Apply prompt template (async version)""" import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -219,7 +217,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): return None @staticmethod - def apply_prompt_template(model: str, messages: List[Dict[str, str]]) -> Optional[str]: + def apply_prompt_template(model: str, messages: list[dict[str, str]]) -> str | None: """Apply prompt template (sync version)""" from litellm.litellm_core_utils.prompt_templates.factory import ( hf_chat_template, diff --git a/litellm/llms/watsonx/common_utils.py b/litellm/llms/watsonx/common_utils.py index d1b065dbc6d..6c23ae646c0 100644 --- a/litellm/llms/watsonx/common_utils.py +++ b/litellm/llms/watsonx/common_utils.py @@ -1,4 +1,4 @@ -from typing import Dict, List, Optional, Union, cast +from typing import cast import httpx @@ -17,7 +17,7 @@ class WatsonXAIError(BaseLLMException): self, status_code: int, message: str, - headers: Optional[Union[Dict, httpx.Headers]] = None, + headers: dict | httpx.Headers | None = None, ): super().__init__(status_code=status_code, message=message, headers=headers) @@ -30,7 +30,7 @@ def get_watsonx_iam_url(): def generate_iam_token(api_key=None, **params) -> str: - result: Optional[str] = iam_token_cache.get_cache(api_key) # type: ignore + result: str | None = iam_token_cache.get_cache(api_key) # type: ignore if result is None: headers = {} @@ -70,14 +70,14 @@ def generate_iam_token(api_key=None, **params) -> str: return cast(str, result) -def _generate_watsonx_token(api_key: Optional[str], token: Optional[str]) -> str: +def _generate_watsonx_token(api_key: str | None, token: str | None) -> str: if token is not None: return token token = generate_iam_token(api_key) return token -def _get_api_params(params: dict, model: Optional[str] = None) -> WatsonXAPIParams: +def _get_api_params(params: dict, model: str | None = None) -> WatsonXAPIParams: """ Find watsonx.ai credentials in the params or environment variables and return the headers for authentication. """ @@ -122,9 +122,9 @@ def _get_api_params(params: dict, model: Optional[str] = None) -> WatsonXAPIPara async def _aconvert_watsonx_messages_core( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], provider: str, - custom_prompt_dict: Dict, + custom_prompt_dict: dict, apply_template_fn, ) -> str: """Async core logic for converting watsonx messages to prompt""" @@ -154,9 +154,9 @@ async def _aconvert_watsonx_messages_core( def _convert_watsonx_messages_core( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], provider: str, - custom_prompt_dict: Dict, + custom_prompt_dict: dict, apply_template_fn, ) -> str: """Sync core logic for converting watsonx messages to prompt""" @@ -186,9 +186,9 @@ def _convert_watsonx_messages_core( async def aconvert_watsonx_messages_to_prompt( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], provider: str, - custom_prompt_dict: Dict, + custom_prompt_dict: dict, ) -> str: """Async version of convert_watsonx_messages_to_prompt""" from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig @@ -204,9 +204,9 @@ async def aconvert_watsonx_messages_to_prompt( def convert_watsonx_messages_to_prompt( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], provider: str, - custom_prompt_dict: Dict, + custom_prompt_dict: dict, ) -> str: """Sync version of convert_watsonx_messages_to_prompt""" from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig @@ -224,14 +224,14 @@ def convert_watsonx_messages_to_prompt( class IBMWatsonXMixin: def validate_environment( self, - headers: Dict, + headers: dict, model: str, - messages: List[AllMessageValues], - optional_params: Dict, + messages: list[AllMessageValues], + optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> Dict: + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: default_headers = { "Content-Type": "application/json", "Accept": "application/json", @@ -240,11 +240,11 @@ class IBMWatsonXMixin: if "Authorization" in headers: return {**default_headers, **headers} token = cast( - Optional[str], + str | None, optional_params.get("token") or get_secret_str("WATSONX_TOKEN"), ) zen_api_key = cast( - Optional[str], + str | None, optional_params.pop("zen_api_key", None) or get_secret_str("WATSONX_ZENAPIKEY"), ) if token: @@ -257,7 +257,7 @@ class IBMWatsonXMixin: headers["Authorization"] = f"Bearer {token}" return {**default_headers, **headers} - def _get_base_url(self, api_base: Optional[str]) -> str: + def _get_base_url(self, api_base: str | None) -> str: url = ( api_base or get_secret_str("WATSONX_API_BASE") # consistent with 'AZURE_API_BASE' @@ -273,21 +273,17 @@ class IBMWatsonXMixin: ) return url - def _add_api_version_to_url(self, url: str, api_version: Optional[str]) -> str: + def _add_api_version_to_url(self, url: str, api_version: str | None) -> str: api_version = api_version or litellm.WATSONX_DEFAULT_API_VERSION url = url + f"?version={api_version}" return url - def get_error_class( - self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return WatsonXAIError(status_code=status_code, message=error_message, headers=headers) @staticmethod - def get_watsonx_credentials( - optional_params: dict, api_key: Optional[str], api_base: Optional[str] - ) -> WatsonXCredentials: + def get_watsonx_credentials(optional_params: dict, api_key: str | None, api_base: str | None) -> WatsonXCredentials: api_key = ( api_key or optional_params.pop("apikey", None) @@ -314,7 +310,7 @@ class IBMWatsonXMixin: optional_params.pop("watsonx_credentials", None), # follow {provider}_credentials, same as vertex ai ) - token: Optional[str] = None + token: str | None = None if wx_credentials is not None: api_base = wx_credentials.get("url", api_base) @@ -335,7 +331,7 @@ class IBMWatsonXMixin: status_code=401, message="Error: Watsonx API base not set. Set WATSONX_API_BASE in environment variables or pass in as parameter - 'api_base='.", ) - return WatsonXCredentials(api_key=api_key, api_base=api_base, token=cast(Optional[str], token)) + return WatsonXCredentials(api_key=api_key, api_base=api_base, token=cast(str | None, token)) def _prepare_payload(self, model: str, api_params: WatsonXAPIParams) -> dict: payload: dict = {} diff --git a/litellm/llms/watsonx/completion/transformation.py b/litellm/llms/watsonx/completion/transformation.py index 397d6bde1c6..3c46b25d161 100644 --- a/litellm/llms/watsonx/completion/transformation.py +++ b/litellm/llms/watsonx/completion/transformation.py @@ -4,10 +4,6 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Union, ) import httpx @@ -72,39 +68,39 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): - `stream` (bool): If True, the model will return a stream of responses. """ - decoding_method: Optional[str] = "sample" - temperature: Optional[float] = None - max_new_tokens: Optional[int] = None # litellm.max_tokens - min_new_tokens: Optional[int] = None - length_penalty: Optional[dict] = None # e.g {"decay_factor": 2.5, "start_index": 5} - stop_sequences: Optional[List[str]] = None # e.g ["}", ")", "."] - top_k: Optional[int] = None - top_p: Optional[float] = None - repetition_penalty: Optional[float] = None - truncate_input_tokens: Optional[int] = None - include_stop_sequences: Optional[bool] = False - return_options: Optional[Dict[str, bool]] = None - random_seed: Optional[int] = None # e.g 42 - moderations: Optional[dict] = None - stream: Optional[bool] = False + decoding_method: str | None = "sample" + temperature: float | None = None + max_new_tokens: int | None = None # litellm.max_tokens + min_new_tokens: int | None = None + length_penalty: dict | None = None # e.g {"decay_factor": 2.5, "start_index": 5} + stop_sequences: list[str] | None = None # e.g ["}", ")", "."] + top_k: int | None = None + top_p: float | None = None + repetition_penalty: float | None = None + truncate_input_tokens: int | None = None + include_stop_sequences: bool | None = False + return_options: dict[str, bool] | None = None + random_seed: int | None = None # e.g 42 + moderations: dict | None = None + stream: bool | None = False def __init__( self, - decoding_method: Optional[str] = None, - temperature: Optional[float] = None, - max_new_tokens: Optional[int] = None, - min_new_tokens: Optional[int] = None, - length_penalty: Optional[dict] = None, - stop_sequences: Optional[List[str]] = None, - top_k: Optional[int] = None, - top_p: Optional[float] = None, - repetition_penalty: Optional[float] = None, - truncate_input_tokens: Optional[int] = None, - include_stop_sequences: Optional[bool] = None, - return_options: Optional[dict] = None, - random_seed: Optional[int] = None, - moderations: Optional[dict] = None, - stream: Optional[bool] = None, + decoding_method: str | None = None, + temperature: float | None = None, + max_new_tokens: int | None = None, + min_new_tokens: int | None = None, + length_penalty: dict | None = None, + stop_sequences: list[str] | None = None, + top_k: int | None = None, + top_p: float | None = None, + repetition_penalty: float | None = None, + truncate_input_tokens: int | None = None, + include_stop_sequences: bool | None = None, + return_options: dict | None = None, + random_seed: int | None = None, + moderations: dict | None = None, + stream: bool | None = None, **kwargs, ) -> None: locals_ = locals().copy() @@ -152,11 +148,11 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): def map_openai_params( self, - non_default_params: Dict, - optional_params: Dict, + non_default_params: dict, + optional_params: dict, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: extra_body = {} for k, v in non_default_params.items(): if k == "max_tokens": @@ -210,7 +206,7 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): optional_params[mapped_params[param]] = value return optional_params - def get_eu_regions(self) -> List[str]: + def get_eu_regions(self) -> list[str]: """ Source: https://www.ibm.com/docs/en/watsonx/saas?topic=integrations-regional-availability """ @@ -219,7 +215,7 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): "eu-gb", ] - def get_us_regions(self) -> List[str]: + def get_us_regions(self) -> list[str]: """ Source: https://www.ibm.com/docs/en/watsonx/saas?topic=integrations-regional-availability """ @@ -227,7 +223,7 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): "us-south", ] - def _build_request_payload(self, model: str, prompt: str, optional_params: Dict) -> Dict: + def _build_request_payload(self, model: str, prompt: str, optional_params: dict) -> dict: """Shared logic to build request payload""" extra_body_params = optional_params.pop("extra_body", {}) optional_params.update(extra_body_params) @@ -244,11 +240,11 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): async def atransform_request( self, model: str, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, - headers: Dict, - ) -> Dict: + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: """Async version of transform_request""" from litellm.llms.watsonx.common_utils import ( aconvert_watsonx_messages_to_prompt, @@ -263,11 +259,11 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, - headers: Dict, - ) -> Dict: + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: """Sync version of transform_request""" provider = model.split("/")[0] prompt = convert_watsonx_messages_to_prompt( @@ -281,13 +277,13 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, - request_data: Dict, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, + request_data: dict, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -329,12 +325,12 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: url = self._get_base_url(api_base=api_base) if model.startswith("deployment/"): @@ -356,9 +352,9 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ): return WatsonxTextCompletionResponseIterator( streaming_response=streaming_response, diff --git a/litellm/llms/watsonx/embed/transformation.py b/litellm/llms/watsonx/embed/transformation.py index a841ba9d3ad..a5c84b0cb3e 100644 --- a/litellm/llms/watsonx/embed/transformation.py +++ b/litellm/llms/watsonx/embed/transformation.py @@ -2,8 +2,6 @@ Translates from OpenAI's `/v1/embeddings` to IBM's `/text/embeddings` route. """ -from typing import Optional - import httpx from litellm.llms.base_llm.embedding.transformation import ( @@ -60,12 +58,12 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: url = self._get_base_url(api_base=api_base) endpoint = WatsonXAIEndpoint.EMBEDDINGS.value @@ -84,7 +82,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], + api_key: str | None, request_data: dict, optional_params: dict, litellm_params: dict, diff --git a/litellm/llms/watsonx/passthrough/transformation.py b/litellm/llms/watsonx/passthrough/transformation.py index a89c72dbe10..8235e195eab 100644 --- a/litellm/llms/watsonx/passthrough/transformation.py +++ b/litellm/llms/watsonx/passthrough/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.llms.watsonx.common_utils import IBMWatsonXMixin @@ -18,13 +18,13 @@ class WatsonxPassthroughConfig(IBMWatsonXMixin, BasePassthroughConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, endpoint: str, - request_query_params: Optional[dict], + request_query_params: dict | None, litellm_params: dict, - ) -> Tuple["URL", str]: + ) -> tuple["URL", str]: """ Construct complete Watsonx URL with version parameter. @@ -44,14 +44,14 @@ class WatsonxPassthroughConfig(IBMWatsonXMixin, BasePassthroughConfig): @staticmethod def get_api_base( - api_base: Optional[str] = None, - ) -> Optional[str]: + api_base: str | None = None, + ) -> str | None: return api_base or IBMWatsonXMixin()._get_base_url(api_base=api_base) @staticmethod def get_api_key( - api_key: Optional[str] = None, - ) -> Optional[str]: + api_key: str | None = None, + ) -> str | None: return ( api_key or IBMWatsonXMixin.get_watsonx_credentials(optional_params=dict(), api_base=None, api_key=api_key)[ @@ -60,8 +60,8 @@ class WatsonxPassthroughConfig(IBMWatsonXMixin, BasePassthroughConfig): ) @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: return model - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: return super().get_models(api_key, api_base) diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index 25b593f1c0a..6d9de3f481b 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -5,7 +5,7 @@ Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank """ import uuid -from typing import Any, Dict, List, Union, cast +from typing import Any, cast import httpx @@ -60,7 +60,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): model: str, api_key: str | None = None, optional_params: dict | None = None, - ) -> Dict: + ) -> dict: optional_params = optional_params or {} default_headers = { @@ -94,15 +94,15 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, - ) -> Dict: + ) -> dict: """ Map Cohere rerank params to IBM watsonx.ai rerank params """ @@ -131,7 +131,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): def transform_rerank_request( self, model: str, - optional_rerank_params: Dict, + optional_rerank_params: dict, headers: dict, litellm_params: dict | None = None, ) -> dict: @@ -164,19 +164,19 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): raw_response_json = raw_response.json() except Exception as e: raise self.get_error_class( - error_message=f"Failed to parse response: {str(e)}", + error_message=f"Failed to parse response: {e!s}", status_code=raw_response.status_code, headers=raw_response.headers, ) - _results: List[dict] | None = raw_response_json.get("results") + _results: list[dict] | None = raw_response_json.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") transformed_results = [] for result in _results: - transformed_result: Dict[str, Any] = { + transformed_result: dict[str, Any] = { "index": result["index"], "relevance_score": result["score"], } diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 23ab64a286f..98d8fe5fadd 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator, Iterator -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any import httpx @@ -30,12 +30,12 @@ from ...openai.chat.gpt_transformation import ( class XAIChatConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "xai" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("XAI_API_BASE") or XAI_API_BASE # type: ignore dynamic_api_key = XAIModelInfo.get_api_key(api_key) return api_base, dynamic_api_key @@ -44,11 +44,11 @@ class XAIChatConfig(OpenAIGPTConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: from litellm.llms.xai.oauth import ( XAIOAuthAuthenticator, @@ -82,12 +82,12 @@ class XAIChatConfig(OpenAIGPTConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth @@ -149,11 +149,7 @@ class XAIChatConfig(OpenAIGPTConfig): return base_openai_params def _supports_stop_reason(self, model: str) -> bool: - if "grok-3-mini" in model: - return False - elif "grok-4" in model: - return False - elif "grok-code-fast" in model: + if "grok-3-mini" in model or "grok-4" in model or "grok-code-fast" in model: return False return True @@ -195,9 +191,9 @@ class XAIChatConfig(OpenAIGPTConfig): def get_model_response_iterator( self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, - json_mode: Optional[bool] = False, + json_mode: bool | None = False, ) -> Any: return XAIChatCompletionStreamingHandler( streaming_response=streaming_response, @@ -208,7 +204,7 @@ class XAIChatConfig(OpenAIGPTConfig): def transform_request( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -239,12 +235,12 @@ class XAIChatConfig(OpenAIGPTConfig): model_response: ModelResponse, logging_obj, request_data: dict, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, encoding, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: """ Transform the response from the XAI API. @@ -289,7 +285,7 @@ class XAIChatConfig(OpenAIGPTConfig): @staticmethod def _fold_reasoning_tokens_into_completion( - target: Union[ModelResponse, Usage, Dict[str, Any], None], + target: ModelResponse | Usage | dict[str, Any] | None, ) -> None: """Reconcile xAI Usage to the OpenAI invariant. @@ -306,7 +302,7 @@ class XAIChatConfig(OpenAIGPTConfig): return if isinstance(target, ModelResponse): - usage: Union[Usage, Dict[str, Any], None] = getattr(target, "usage", None) + usage: Usage | dict[str, Any] | None = getattr(target, "usage", None) else: usage = target if usage is None: @@ -377,7 +373,7 @@ class XAIChatConfig(OpenAIGPTConfig): @staticmethod def _normalize_openai_compatible_usage_totals( - usage: Union[Usage, Dict[str, Any], None], + usage: Usage | dict[str, Any] | None, ) -> None: if usage is None: return diff --git a/litellm/llms/xai/common_utils.py b/litellm/llms/xai/common_utils.py index 0e499e33ed1..0560ae117e9 100644 --- a/litellm/llms/xai/common_utils.py +++ b/litellm/llms/xai/common_utils.py @@ -1,5 +1,3 @@ -from typing import List, Optional - import httpx import litellm @@ -13,7 +11,7 @@ class XAIModelInfo(BaseLLMModelInfo): def get_provider_info( self, model: str, - ) -> Optional[ProviderSpecificModelInfo]: + ) -> ProviderSpecificModelInfo | None: """ Default values all models of this provider support. """ @@ -25,11 +23,11 @@ class XAIModelInfo(BaseLLMModelInfo): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is not None: headers["Authorization"] = f"Bearer {api_key}" @@ -41,14 +39,14 @@ class XAIModelInfo(BaseLLMModelInfo): return headers @staticmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base(api_base: str | None = None) -> str | None: return api_base or get_secret_str("XAI_API_BASE") or "https://api.x.ai" @staticmethod def get_api_key( - api_key: Optional[str] = None, + api_key: str | None = None, legacy_generic_before_env: bool = False, - ) -> Optional[str]: + ) -> str | None: """ Resolve xAI API keys while preserving endpoint-specific legacy order. @@ -64,10 +62,10 @@ class XAIModelInfo(BaseLLMModelInfo): return api_key or litellm.xai_key or get_secret_str("XAI_API_KEY") @staticmethod - def get_base_model(model: str) -> Optional[str]: + def get_base_model(model: str) -> str | None: return model.replace("xai/", "") - def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base = self.get_api_base(api_base) api_key = self.get_api_key(api_key) if api_base is None or api_key is None: diff --git a/litellm/llms/xai/cost_calculator.py b/litellm/llms/xai/cost_calculator.py index 284400b0824..59aea1b25f3 100644 --- a/litellm/llms/xai/cost_calculator.py +++ b/litellm/llms/xai/cost_calculator.py @@ -4,16 +4,16 @@ Helper util for handling XAI-specific cost calculation - Handles XAI-specific reasoning token billing (billed as part of completion tokens) """ -from typing import TYPE_CHECKING, Tuple +from typing import TYPE_CHECKING -from litellm.types.utils import Usage from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token +from litellm.types.utils import Usage if TYPE_CHECKING: from litellm.types.utils import ModelInfo -def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: +def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: """ Calculates the cost per token for a given XAI model, prompt tokens, and completion tokens. Uses the generic cost calculator for all pricing logic, with XAI-specific reasoning token handling. diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py index 064e0ff77d6..70e343785b8 100644 --- a/litellm/llms/xai/oauth.py +++ b/litellm/llms/xai/oauth.py @@ -9,7 +9,7 @@ import time import uuid import webbrowser from http.server import BaseHTTPRequestHandler, HTTPServer -from typing import Any, Dict, Optional, Tuple, Union +from typing import Any from urllib.parse import parse_qs, urlencode, urlparse import httpx @@ -81,11 +81,11 @@ class _CallbackHandler(BaseHTTPRequestHandler): class _CallbackServer(HTTPServer): expected_state: str - callback_result: Optional[Dict[str, Optional[str]]] + callback_result: dict[str, str | None] | None class XAIOAuthAuthenticator: - def __init__(self, http_client: Optional[Union[httpx.Client, HTTPHandler]] = None) -> None: + def __init__(self, http_client: httpx.Client | HTTPHandler | None = None) -> None: self.token_dir = get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser("~/.config/litellm/xai_oauth") self.auth_file = os.path.join(self.token_dir, get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json") self.http_client = http_client @@ -115,7 +115,7 @@ class XAIOAuthAuthenticator: refreshed = self._refresh_tokens(locked_auth_data) return refreshed["access_token"] - def login(self, force: bool = False, no_browser: bool = False) -> Dict[str, Any]: + def login(self, force: bool = False, no_browser: bool = False) -> dict[str, Any]: existing = self._read_auth_file() if existing and not force and existing.get("access_token"): if not self._is_expired(existing): @@ -167,7 +167,7 @@ class XAIOAuthAuthenticator: self._write_auth_file(auth_data) return auth_data - def _client(self) -> Union[httpx.Client, HTTPHandler]: + def _client(self) -> httpx.Client | HTTPHandler: return self.http_client or _get_httpx_client() def _ensure_token_dir(self) -> None: @@ -177,15 +177,15 @@ class XAIOAuthAuthenticator: except OSError: verbose_logger.debug("Could not chmod xAI OAuth token directory") - def _read_auth_file(self) -> Optional[Dict[str, Any]]: + def _read_auth_file(self) -> dict[str, Any] | None: try: with open(self.auth_file, "r") as f: data = json.load(f) return data if isinstance(data, dict) else None - except (IOError, json.JSONDecodeError): + except (OSError, json.JSONDecodeError): return None - def _write_auth_file(self, data: Dict[str, Any]) -> None: + def _write_auth_file(self, data: dict[str, Any]) -> None: self._ensure_token_dir() tmp_file = os.path.join( self.token_dir, @@ -216,7 +216,7 @@ class XAIOAuthAuthenticator: pass raise - def _is_expired(self, auth_data: Dict[str, Any]) -> bool: + def _is_expired(self, auth_data: dict[str, Any]) -> bool: expires_at = auth_data.get("expires_at") if expires_at is None: return True @@ -225,7 +225,7 @@ class XAIOAuthAuthenticator: except (TypeError, ValueError): return True - def _discover(self) -> Dict[str, str]: + def _discover(self) -> dict[str, str]: try: response = self._client().get(XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"}) response.raise_for_status() @@ -253,13 +253,13 @@ class XAIOAuthAuthenticator: raise XAIOAuthError(f"xAI OAuth discovery returned unexpected endpoint: {url}") return url - def _pkce_pair(self) -> Tuple[str, str]: + def _pkce_pair(self) -> tuple[str, str]: verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode() challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode() return verifier, challenge - def _start_callback_server(self, state: str) -> Tuple[_CallbackServer, str]: - last_error: Optional[OSError] = None + def _start_callback_server(self, state: str) -> tuple[_CallbackServer, str]: + last_error: OSError | None = None for port in (XAI_OAUTH_REDIRECT_PORT, 0): try: server = _CallbackServer((XAI_OAUTH_REDIRECT_HOST, port), _CallbackHandler) @@ -292,7 +292,7 @@ class XAIOAuthAuthenticator: } return f"{authorization_endpoint}?{urlencode(params)}" - def _wait_for_callback(self, server: _CallbackServer) -> Dict[str, Optional[str]]: + def _wait_for_callback(self, server: _CallbackServer) -> dict[str, str | None]: server.timeout = 1 deadline = time.time() + XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS try: @@ -304,7 +304,7 @@ class XAIOAuthAuthenticator: server.server_close() raise XAIOAuthError("Timed out waiting for xAI OAuth callback") - def _exchange_token(self, token_endpoint: str, data: Dict[str, str]) -> Dict[str, Any]: + def _exchange_token(self, token_endpoint: str, data: dict[str, str]) -> dict[str, Any]: try: response = self._client().post( token_endpoint, @@ -329,10 +329,10 @@ class XAIOAuthAuthenticator: def _build_auth_record( self, - token_payload: Dict[str, Any], + token_payload: dict[str, Any], token_endpoint: str, - fallback_refresh_token: Optional[str] = None, - ) -> Dict[str, Any]: + fallback_refresh_token: str | None = None, + ) -> dict[str, Any]: access_token = token_payload.get("access_token") refresh_token = token_payload.get("refresh_token") or fallback_refresh_token if not access_token: @@ -353,7 +353,7 @@ class XAIOAuthAuthenticator: "expires_at": expires_at, } - def _refresh_tokens(self, auth_data: Dict[str, Any]) -> Dict[str, Any]: + def _refresh_tokens(self, auth_data: dict[str, Any]) -> dict[str, Any]: token_endpoint = auth_data.get("token_endpoint") if not token_endpoint: token_endpoint = self._discover()["token_endpoint"] @@ -379,5 +379,5 @@ class XAIOAuthAuthenticator: return refreshed -def should_use_xai_oauth(litellm_params: Optional[Dict[str, Any]]) -> bool: +def should_use_xai_oauth(litellm_params: dict[str, Any] | None) -> bool: return bool((litellm_params or {}).get("use_xai_oauth")) diff --git a/litellm/llms/xai/realtime/transformation.py b/litellm/llms/xai/realtime/transformation.py index 6d8a8948f06..92a0a82058f 100644 --- a/litellm/llms/xai/realtime/transformation.py +++ b/litellm/llms/xai/realtime/transformation.py @@ -16,7 +16,7 @@ construction time (see ``handler.py``) so all normalization is isolated here and ``RealTimeStreaming`` stays provider-agnostic. """ -from typing import Any, Optional +from typing import Any class XAIRealtimeNormalizer: @@ -243,7 +243,7 @@ class XAIRealtimeNormalizer: } @staticmethod - def _normalize_usage(usage: object, *, empty_as_null: bool) -> Optional[dict[str, Any]]: + def _normalize_usage(usage: object, *, empty_as_null: bool) -> dict[str, Any] | None: """Coerce a usage object into the full OpenAI GA shape. ``empty_as_null=True`` for ``response.created`` (usage optional). diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index 2773444bce9..0ad76196398 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -51,7 +51,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return supported_params - def _transform_web_search_tool(self, tool: Dict[str, Any]) -> Union[XAIWebSearchTool, Dict[str, Any]]: + def _transform_web_search_tool(self, tool: dict[str, Any]) -> XAIWebSearchTool | dict[str, Any]: """ Transform web_search tool to XAI format. @@ -62,7 +62,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): XAI does NOT support search_context_size (OpenAI-specific). """ - xai_tool: Dict[str, Any] = {"type": "web_search"} + xai_tool: dict[str, Any] = {"type": "web_search"} # Remove search_context_size if present (not supported by XAI) if "search_context_size" in tool: @@ -90,7 +90,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return xai_tool - def _transform_x_search_tool(self, tool: Dict[str, Any]) -> Union[XAIXSearchTool, Dict[str, Any]]: + def _transform_x_search_tool(self, tool: dict[str, Any]) -> XAIXSearchTool | dict[str, Any]: """ Transform x_search tool to XAI format. @@ -102,7 +102,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): - enable_image_understanding - enable_video_understanding """ - xai_tool: Dict[str, Any] = {"type": "x_search"} + xai_tool: dict[str, Any] = {"type": "x_search"} # Handle allowed_x_handles if "allowed_x_handles" in tool: @@ -135,7 +135,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): response_api_optional_params: ResponsesAPIOptionalRequestParams, model: str, drop_params: bool, - ) -> Dict: + ) -> dict: """ Map parameters for XAI Responses API. @@ -164,7 +164,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): if not isinstance(tools_list, list): tools_list = [tools_list] - transformed_tools: List[Any] = [] + transformed_tools: list[Any] = [] for tool in tools_list: if isinstance(tool, dict): tool_type = tool.get("type") @@ -194,7 +194,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return params - def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict: + def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: """ Validate environment and set up headers for XAI API. @@ -235,7 +235,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, litellm_params: dict, ) -> str: """ diff --git a/litellm/llms/xinference/image_generation/transformation.py b/litellm/llms/xinference/image_generation/transformation.py index 0d2d890ddf4..ce0b0e00e82 100644 --- a/litellm/llms/xinference/image_generation/transformation.py +++ b/litellm/llms/xinference/image_generation/transformation.py @@ -1,5 +1,3 @@ -from typing import List - from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) @@ -13,7 +11,7 @@ class XInferenceImageGenerationConfig(BaseImageGenerationConfig): https://inference.readthedocs.io/en/v1.1.1/reference/generated/xinference.client.handlers.ImageModelHandle.text_to_image.html#xinference.client.handlers.ImageModelHandle.text_to_image """ - def get_supported_openai_params(self, model: str) -> List[OpenAIImageGenerationOptionalParams]: + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: return ["n", "response_format", "size", "response_format"] def map_openai_params( @@ -24,8 +22,8 @@ class XInferenceImageGenerationConfig(BaseImageGenerationConfig): drop_params: bool, ) -> dict: supported_params = self.get_supported_openai_params(model) - for k in non_default_params.keys(): - if k not in optional_params.keys(): + for k in non_default_params: + if k not in optional_params: if k in supported_params: optional_params[k] = non_default_params[k] elif drop_params: diff --git a/litellm/llms/you_com/search/transformation.py b/litellm/llms/you_com/search/transformation.py index 0cd825c3ab8..efc8194a7df 100644 --- a/litellm/llms/you_com/search/transformation.py +++ b/litellm/llms/you_com/search/transformation.py @@ -5,7 +5,7 @@ You.com API Reference: https://you.com/docs/api-reference/search/v1-search OpenAPI spec: https://you.com/specs/openapi_search_v1.yaml """ -from typing import Dict, List, Optional, TypedDict, Union +from typing import TypedDict import httpx @@ -34,8 +34,8 @@ class YouComSearchRequest(_YouComSearchRequestRequired, total=False): country: str language: str freshness: str - include_domains: List[str] - exclude_domains: List[str] + include_domains: list[str] + exclude_domains: list[str] safesearch: str @@ -52,11 +52,11 @@ class YouComSearchConfig(BaseSearchConfig): def validate_environment( self, - headers: Dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + headers: dict, + api_key: str | None = None, + api_base: str | None = None, **kwargs, - ) -> Dict: + ) -> dict: """ Set headers for the You.com Search API. @@ -83,9 +83,9 @@ class YouComSearchConfig(BaseSearchConfig): def get_complete_url( self, - api_base: Optional[str], + api_base: str | None, optional_params: dict, - data: Optional[Union[Dict, List[Dict]]] = None, + data: dict | list[dict] | None = None, **kwargs, ) -> str: """ @@ -115,10 +115,10 @@ class YouComSearchConfig(BaseSearchConfig): def transform_search_request( self, - query: Union[str, List[str]], + query: str | list[str], optional_params: dict, **kwargs, - ) -> Dict: + ) -> dict: """ Transform Search request to You.com API format. @@ -174,7 +174,7 @@ class YouComSearchConfig(BaseSearchConfig): web_results = raw_results.get("web") or [] news_results = raw_results.get("news") or [] - results: List[SearchResult] = [] + results: list[SearchResult] = [] for item in list(web_results) + list(news_results): snippets = item.get("snippets") or [] snippet = snippets[0] if snippets else item.get("description", "") diff --git a/litellm/llms/zai/chat/transformation.py b/litellm/llms/zai/chat/transformation.py index fb1d67df357..714dc95862e 100644 --- a/litellm/llms/zai/chat/transformation.py +++ b/litellm/llms/zai/chat/transformation.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Tuple - from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam @@ -10,12 +8,12 @@ ZAI_API_BASE = "https://api.z.ai/api/paas/v4" class ZAIChatConfig(OpenAIGPTConfig): @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "zai" def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = api_base or get_secret_str("ZAI_API_BASE") or ZAI_API_BASE dynamic_api_key = api_key or get_secret_str("ZAI_API_KEY") return api_base, dynamic_api_key @@ -23,9 +21,9 @@ class ZAIChatConfig(OpenAIGPTConfig): def remove_cache_control_flag_from_messages_and_tools( self, model: str, - messages: List[AllMessageValues], - tools: Optional[List[ChatCompletionToolParam]] = None, - ) -> Tuple[List[AllMessageValues], Optional[List[ChatCompletionToolParam]]]: + messages: list[AllMessageValues], + tools: list[ChatCompletionToolParam] | None = None, + ) -> tuple[list[AllMessageValues], list[ChatCompletionToolParam] | None]: """ Override to preserve cache_control for GLM/ZAI. GLM supports cache_control - don't strip it. diff --git a/litellm/main.py b/litellm/main.py index 605ec2b066a..cea9d44fb1a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -27,12 +27,8 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Type, Union, cast, get_args, @@ -343,30 +339,28 @@ class LiteLLM: self, *, api_key=None, - organization: Optional[str] = None, - base_url: Optional[str] = None, - timeout: Optional[float] = 600, - max_retries: Optional[int] = litellm.num_retries, - default_headers: Optional[Mapping[str, str]] = None, + organization: str | None = None, + base_url: str | None = None, + timeout: float | None = 600, + max_retries: int | None = litellm.num_retries, + default_headers: Mapping[str, str] | None = None, ): self.params = locals() self.chat = Chat(self.params, router_obj=None) class Chat: - def __init__(self, params, router_obj: Optional[Any]): + def __init__(self, params, router_obj: Any | None): self.params = params if self.params.get("acompletion", False) is True: self.params.pop("acompletion") - self.completions: Union[AsyncCompletions, Completions] = AsyncCompletions( - self.params, router_obj=router_obj - ) + self.completions: AsyncCompletions | Completions = AsyncCompletions(self.params, router_obj=router_obj) else: self.completions = Completions(self.params, router_obj=router_obj) class Completions: - def __init__(self, params, router_obj: Optional[Any]): + def __init__(self, params, router_obj: Any | None): self.params = params self.router_obj = router_obj @@ -382,7 +376,7 @@ class Completions: class AsyncCompletions: - def __init__(self, params, router_obj: Optional[Any]): + def __init__(self, params, router_obj: Any | None): self.params = params self.router_obj = router_obj @@ -402,54 +396,54 @@ class AsyncCompletions: async def acompletion( model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create - messages: List = [], - functions: Optional[List] = None, - function_call: Optional[str] = None, - timeout: Optional[Union[float, int]] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - n: Optional[int] = None, - stream: Optional[bool] = None, - stream_options: Optional[dict] = None, + messages: list = [], + functions: list | None = None, + function_call: str | None = None, + timeout: float | None = None, + temperature: float | None = None, + top_p: float | None = None, + n: int | None = None, + stream: bool | None = None, + stream_options: dict | None = None, stop=None, - max_tokens: Optional[int] = None, - max_completion_tokens: Optional[int] = None, - modalities: Optional[List[ChatCompletionModality]] = None, - prediction: Optional[ChatCompletionPredictionContentParam] = None, - audio: Optional[ChatCompletionAudioParam] = None, - presence_penalty: Optional[float] = None, - frequency_penalty: Optional[float] = None, - logit_bias: Optional[dict] = None, - user: Optional[str] = None, + max_tokens: int | None = None, + max_completion_tokens: int | None = None, + modalities: list[ChatCompletionModality] | None = None, + prediction: ChatCompletionPredictionContentParam | None = None, + audio: ChatCompletionAudioParam | None = None, + presence_penalty: float | None = None, + frequency_penalty: float | None = None, + logit_bias: dict | None = None, + user: str | None = None, # openai v1.0+ new params - response_format: Optional[Union[dict, Type[BaseModel]]] = None, - seed: Optional[int] = None, - tools: Optional[List] = None, - tool_choice: Optional[Union[str, dict]] = None, - parallel_tool_calls: Optional[bool] = None, - logprobs: Optional[bool] = None, - top_logprobs: Optional[int] = None, + response_format: dict | type[BaseModel] | None = None, + seed: int | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, + parallel_tool_calls: bool | None = None, + logprobs: bool | None = None, + top_logprobs: int | None = None, deployment_id=None, - reasoning_effort: Optional[Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"]] = None, - verbosity: Optional[Literal["low", "medium", "high"]] = None, - safety_identifier: Optional[str] = None, - service_tier: Optional[str] = None, + reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None, + verbosity: Literal["low", "medium", "high"] | None = None, + safety_identifier: str | None = None, + service_tier: str | None = None, # set api_base, api_version, api_key - base_url: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - model_list: Optional[list] = None, # pass in a list of api_base,keys, etc. - extra_headers: Optional[dict] = None, + base_url: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + model_list: list | None = None, # pass in a list of api_base,keys, etc. + extra_headers: dict | None = None, # Optional liteLLM function params - thinking: Optional[AnthropicThinkingParam] = None, - web_search_options: Optional[OpenAIWebSearchOptions] = None, - include_server_side_tool_invocations: Optional[bool] = None, + thinking: AnthropicThinkingParam | None = None, + web_search_options: OpenAIWebSearchOptions | None = None, + include_server_side_tool_invocations: bool | None = None, # Session management shared_session: Optional["ClientSession"] = None, # Per-request JSON schema validation (overrides litellm.enable_json_schema_validation) - enable_json_schema_validation: Optional[bool] = None, + enable_json_schema_validation: bool | None = None, **kwargs, -) -> Union[ModelResponse, CustomStreamWrapper]: +) -> ModelResponse | CustomStreamWrapper: """ Asynchronously executes a litellm.completion() call for any of litellm supported llms (example gpt-4, gpt-3.5-turbo, claude-2, command-nightly) @@ -516,7 +510,7 @@ async def acompletion( non_default_params=kwargs, messages=cast(list[AllMessageValues], messages), # cast-ok: acompletion types messages as a bare List model=model, - custom_llm_provider=cast(Optional[str], custom_llm_provider), # cast-ok: read from untyped kwargs + custom_llm_provider=cast(str | None, custom_llm_provider), # cast-ok: read from untyped kwargs tools=tools, ) @@ -730,9 +724,9 @@ async def _async_streaming(response, model, custom_llm_provider, args): def _handle_mock_potential_exceptions( - mock_response: Union[str, Exception], + mock_response: str | Exception, model: str, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ): if isinstance(mock_response, Exception): if isinstance(mock_response, openai.APIError): @@ -773,8 +767,8 @@ def _handle_mock_potential_exceptions( def _handle_mock_timeout( - mock_timeout: Optional[bool], - timeout: Optional[Union[float, str, httpx.Timeout]], + mock_timeout: bool | None, + timeout: float | str | httpx.Timeout | None, model: str, ): if mock_timeout is True and timeout is not None: @@ -787,8 +781,8 @@ def _handle_mock_timeout( async def _handle_mock_timeout_async( - mock_timeout: Optional[bool], - timeout: Optional[Union[float, str, httpx.Timeout]], + mock_timeout: bool | None, + timeout: float | str | httpx.Timeout | None, model: str, ): if mock_timeout is True and timeout is not None: @@ -800,7 +794,7 @@ async def _handle_mock_timeout_async( ) -def _sleep_for_timeout(timeout: Union[float, str, httpx.Timeout]): +def _sleep_for_timeout(timeout: float | str | httpx.Timeout): if isinstance(timeout, float): time.sleep(timeout) elif isinstance(timeout, str): @@ -809,7 +803,7 @@ def _sleep_for_timeout(timeout: Union[float, str, httpx.Timeout]): time.sleep(timeout.connect) -async def _sleep_for_timeout_async(timeout: Union[float, str, httpx.Timeout]): +async def _sleep_for_timeout_async(timeout: float | str | httpx.Timeout): if isinstance(timeout, float): await asyncio.sleep(timeout) elif isinstance(timeout, str): @@ -820,15 +814,15 @@ async def _sleep_for_timeout_async(timeout: Union[float, str, httpx.Timeout]): def mock_completion( model: str, - messages: List, - stream: Optional[bool] = False, - n: Optional[int] = None, - mock_response: Optional[MOCK_RESPONSE_TYPE] = "This is a mock request", - mock_tool_calls: Optional[List] = None, - mock_timeout: Optional[bool] = False, + messages: list, + stream: bool | None = False, + n: int | None = None, + mock_response: MOCK_RESPONSE_TYPE | None = "This is a mock request", + mock_tool_calls: list | None = None, + mock_timeout: bool | None = False, logging=None, custom_llm_provider=None, - timeout: Optional[Union[float, str, httpx.Timeout]] = None, + timeout: float | str | httpx.Timeout | None = None, **kwargs, ): """ @@ -876,7 +870,7 @@ def mock_completion( ) mock_response = cast( - Union[str, dict, ModelResponse, ModelResponseStream], mock_response + str | dict | ModelResponse | ModelResponseStream, mock_response ) # after this point, mock_response is a string, dict, ModelResponse, or ModelResponseStream if isinstance(mock_response, str) and mock_response.startswith("Exception: mock_streaming_error"): mock_response = litellm.MockException( @@ -898,7 +892,7 @@ def mock_completion( # convert to ModelResponseStream mock_response = convert_model_response_to_streaming(mock_response) # type: ignore - model_response: Union[ModelResponse, ModelResponseStream] = ModelResponse() + model_response: ModelResponse | ModelResponseStream = ModelResponse() if stream is True: model_response = ModelResponseStream() @@ -973,18 +967,18 @@ def mock_completion( except Exception as e: if isinstance(e, openai.APIError): raise e - raise Exception("Mock completion response failed - {}".format(e)) + raise Exception(f"Mock completion response failed - {e}") def responses_api_bridge_check( model: str, custom_llm_provider: str, - web_search_options: Optional[OpenAIWebSearchOptions] = None, - tools: Optional[List[Any]] = None, - reasoning_effort: Optional[Any] = None, - reasoning_summary: Optional[Any] = None, -) -> Tuple[dict, str]: - model_info: Dict[str, Any] = {} + web_search_options: OpenAIWebSearchOptions | None = None, + tools: list[Any] | None = None, + reasoning_effort: Any | None = None, + reasoning_summary: Any | None = None, +) -> tuple[dict, str]: + model_info: dict[str, Any] = {} # Global flag: route ALL OpenAI chat completions through Responses API. # Returns early with minimal model_info; callers only inspect the "mode" key. @@ -1008,7 +1002,7 @@ def responses_api_bridge_check( model = model.replace("responses/", "") except Exception as e: - verbose_logger.debug("Error getting model info: {}".format(e)) + verbose_logger.debug(f"Error getting model info: {e}") if model.startswith("responses/"): # handle azure models - `azure/responses/` model = model.replace("responses/", "") @@ -1036,7 +1030,7 @@ def responses_api_bridge_check( return model_info, model -def _should_allow_input_examples(custom_llm_provider: Optional[str], model: str) -> bool: +def _should_allow_input_examples(custom_llm_provider: str | None, model: str) -> bool: if custom_llm_provider == "anthropic": return True if custom_llm_provider == "azure_ai" or custom_llm_provider == "bedrock" or custom_llm_provider == "vertex_ai": @@ -1056,11 +1050,11 @@ def _drop_input_examples_from_tool(tool: dict) -> dict: def _drop_input_examples_from_tools( - tools: Optional[List[dict]], -) -> Optional[List[dict]]: + tools: list[dict] | None, +) -> list[dict] | None: if tools is None: return None - cleaned_tools: List[dict] = [] + cleaned_tools: list[dict] = [] for tool in tools: if isinstance(tool, dict): cleaned_tools.append(_drop_input_examples_from_tool(tool)) @@ -1072,7 +1066,7 @@ def _drop_input_examples_from_tools( def _build_custom_pricing_entry( custom_llm_provider: str, kwargs: dict, - model_info: Optional[dict] = None, + model_info: dict | None = None, ) -> dict: """Build a complete model cost entry from kwargs and model_info. @@ -1095,7 +1089,7 @@ def _build_custom_pricing_entry( return entry -def _get_router_deployment_id(kwargs: dict) -> Optional[str]: +def _get_router_deployment_id(kwargs: dict) -> str | None: for metadata_key in ("litellm_metadata", "metadata"): metadata = kwargs.get(metadata_key) or {} if not isinstance(metadata, dict): @@ -1113,7 +1107,7 @@ def _register_custom_pricing_for_request( model: str, custom_llm_provider: str, kwargs: dict, - model_info: Optional[dict], + model_info: dict | None, ) -> None: """Register per-request custom pricing in litellm.model_cost. @@ -2598,7 +2592,7 @@ def _complete_anthropic_text( api_key = api_key or litellm.anthropic_key or litellm.api_key or os.environ.get("ANTHROPIC_API_KEY") custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict api_base = cast( - Optional[str], + str | None, api_base or litellm.api_base or get_secret("ANTHROPIC_API_BASE") @@ -2654,7 +2648,7 @@ def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchR # call /messages # default route for all anthropic models api_base = cast( - Optional[str], + str | None, api_base or litellm.api_base or get_secret("ANTHROPIC_API_BASE") @@ -4044,7 +4038,7 @@ def _complete_watsonx_text( optional_params.pop("watsonx_credentials", None), # follow {provider}_credentials, same as vertex ai ) - token: Optional[str] = None + token: str | None = None if wx_credentials is not None: api_base = wx_credentials.get("url", api_base) api_key = wx_credentials.get("apikey", wx_credentials.get("api_key", api_key)) @@ -4476,8 +4470,6 @@ def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul provider_config=bytez_transformation, ) - pass - return response @@ -4518,8 +4510,6 @@ def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe provider_config=lemonade_transformation, ) - pass - return response @@ -4567,8 +4557,6 @@ def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe provider_config=ovhcloud_transformation, ) - pass - return response @@ -4662,7 +4650,7 @@ def _complete_custom_providers( stream = ctx.stream timeout = ctx.timeout - custom_handler: Optional[CustomLLM] = None + custom_handler: CustomLLM | None = None for item in litellm.custom_provider_map: if item["provider"] == custom_llm_provider: custom_handler = item["custom_handler"] @@ -4808,55 +4796,55 @@ def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe def completion( # type: ignore model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create - messages: List = [], - timeout: Optional[Union[float, str, httpx.Timeout]] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - n: Optional[int] = None, - stream: Optional[bool] = None, - stream_options: Optional[dict] = None, + messages: list = [], + timeout: float | str | httpx.Timeout | None = None, + temperature: float | None = None, + top_p: float | None = None, + n: int | None = None, + stream: bool | None = None, + stream_options: dict | None = None, stop=None, - max_completion_tokens: Optional[int] = None, - max_tokens: Optional[int] = None, - modalities: Optional[List[ChatCompletionModality]] = None, - prediction: Optional[ChatCompletionPredictionContentParam] = None, - audio: Optional[ChatCompletionAudioParam] = None, - presence_penalty: Optional[float] = None, - frequency_penalty: Optional[float] = None, - logit_bias: Optional[dict] = None, - user: Optional[str] = None, + max_completion_tokens: int | None = None, + max_tokens: int | None = None, + modalities: list[ChatCompletionModality] | None = None, + prediction: ChatCompletionPredictionContentParam | None = None, + audio: ChatCompletionAudioParam | None = None, + presence_penalty: float | None = None, + frequency_penalty: float | None = None, + logit_bias: dict | None = None, + user: str | None = None, # openai v1.0+ new params - reasoning_effort: Optional[Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"]] = None, - verbosity: Optional[Literal["low", "medium", "high"]] = None, - response_format: Optional[Union[dict, Type[BaseModel]]] = None, - seed: Optional[int] = None, - tools: Optional[List] = None, - tool_choice: Optional[Union[str, dict]] = None, - logprobs: Optional[bool] = None, - top_logprobs: Optional[int] = None, - parallel_tool_calls: Optional[bool] = None, - web_search_options: Optional[OpenAIWebSearchOptions] = None, - include_server_side_tool_invocations: Optional[bool] = None, + reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None, + verbosity: Literal["low", "medium", "high"] | None = None, + response_format: dict | type[BaseModel] | None = None, + seed: int | None = None, + tools: list | None = None, + tool_choice: str | dict | None = None, + logprobs: bool | None = None, + top_logprobs: int | None = None, + parallel_tool_calls: bool | None = None, + web_search_options: OpenAIWebSearchOptions | None = None, + include_server_side_tool_invocations: bool | None = None, deployment_id=None, - extra_headers: Optional[dict] = None, - safety_identifier: Optional[str] = None, - service_tier: Optional[str] = None, + extra_headers: dict | None = None, + safety_identifier: str | None = None, + service_tier: str | None = None, # soon to be deprecated params by OpenAI - functions: Optional[List] = None, - function_call: Optional[str] = None, + functions: list | None = None, + function_call: str | None = None, # set api_base, api_version, api_key - base_url: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - model_list: Optional[list] = None, # pass in a list of api_base,keys, etc. + base_url: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + model_list: list | None = None, # pass in a list of api_base,keys, etc. # Optional liteLLM function params - thinking: Optional[AnthropicThinkingParam] = None, + thinking: AnthropicThinkingParam | None = None, # Session management shared_session: Optional["ClientSession"] = None, # Per-request JSON schema validation (overrides litellm.enable_json_schema_validation) - enable_json_schema_validation: Optional[bool] = None, + enable_json_schema_validation: bool | None = None, **kwargs, -) -> Union[ModelResponse, CustomStreamWrapper]: +) -> ModelResponse | CustomStreamWrapper: """ Perform a completion() using any of litellm supported llms (example gpt-4, gpt-3.5-turbo, claude-2, command-nightly) Parameters: @@ -4934,7 +4922,7 @@ def completion( # type: ignore # Check if MCP tools are present (following responses pattern) # Cast tools to Optional[Iterable[ToolParam]] for type checking - tools_for_mcp = cast(Optional[Iterable[ToolParam]], tools) + tools_for_mcp = cast(Iterable[ToolParam] | None, tools) if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools_for_mcp): return acompletion_with_mcp( # pyright: ignore[reportReturnType] # MCP path returns a coroutine that acompletion() awaits; completion()'s sync return type omits it model=model, @@ -4981,9 +4969,9 @@ def completion( # type: ignore **kwargs, ) api_base = kwargs.get("api_base", None) - mock_response: Optional[MOCK_RESPONSE_TYPE] = kwargs.get("mock_response", None) + mock_response: MOCK_RESPONSE_TYPE | None = kwargs.get("mock_response", None) mock_tool_calls = kwargs.get("mock_tool_calls", None) - mock_timeout = cast(Optional[bool], kwargs.get("mock_timeout", None)) + mock_timeout = cast(bool | None, kwargs.get("mock_timeout", None)) force_timeout = kwargs.get("force_timeout", 600) ## deprecated logger_fn = kwargs.get("logger_fn", None) verbose = kwargs.get("verbose", False) @@ -4994,14 +4982,12 @@ def completion( # type: ignore model_info = kwargs.get("model_info", None) proxy_server_request = kwargs.get("proxy_server_request", None) fallbacks = kwargs.get("fallbacks", None) - provider_specific_header = cast(Optional[ProviderSpecificHeader], kwargs.get("provider_specific_header", None)) + provider_specific_header = cast(ProviderSpecificHeader | None, kwargs.get("provider_specific_header", None)) headers = kwargs.get("headers", None) or extra_headers - ensure_alternating_roles: Optional[bool] = kwargs.get("ensure_alternating_roles", None) - user_continue_message: Optional[ChatCompletionUserMessage] = kwargs.get("user_continue_message", None) - assistant_continue_message: Optional[ChatCompletionAssistantMessage] = kwargs.get( - "assistant_continue_message", None - ) + ensure_alternating_roles: bool | None = kwargs.get("ensure_alternating_roles", None) + user_continue_message: ChatCompletionUserMessage | None = kwargs.get("user_continue_message", None) + assistant_continue_message: ChatCompletionAssistantMessage | None = kwargs.get("assistant_continue_message", None) if headers is None: headers = {} if extra_headers is not None: @@ -5050,8 +5036,8 @@ def completion( # type: ignore ### Admin Controls ### no_log = kwargs.get("no-log", False) ### PROMPT MANAGEMENT ### - prompt_id = cast(Optional[str], kwargs.get("prompt_id", None)) - prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None)) + prompt_id = cast(str | None, kwargs.get("prompt_id", None)) + prompt_variables = cast(dict | None, kwargs.get("prompt_variables", None)) litellm_system_prompt = kwargs.get("litellm_system_prompt", None) ### COPY MESSAGES ### - related issue https://github.com/BerriAI/litellm/discussions/4489 messages = get_completion_messages( @@ -5074,7 +5060,7 @@ def completion( # type: ignore non_default_params=non_default_params, messages=cast(list[AllMessageValues], messages), # cast-ok: completion types messages as a bare List model=model, - custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs + custom_llm_provider=cast(str | None, kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs tools=tools, ) @@ -5210,12 +5196,12 @@ def completion( # type: ignore messages=messages, model_id=(kwargs.get("model_info") or {}).get("id", None), model_file_id_mapping=cast( - Dict[str, Dict[str, str]], + dict[str, dict[str, str]], kwargs.get("model_file_id_mapping") or {}, ), ) - provider_config: Optional[BaseConfig] = None + provider_config: BaseConfig | None = None if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]: provider_config = ProviderConfigManager.get_provider_chat_config( model=model, @@ -5599,7 +5585,6 @@ def completion( # type: ignore """ Deprecated. We now do together ai calls via the openai client - https://docs.together.ai/docs/openai-api-compatibility """ - pass elif custom_llm_provider == "palm": raise ValueError( "Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en" @@ -5835,7 +5820,7 @@ async def aembedding(*args, **kwargs) -> EmbeddingResponse: # Await normally init_response = await loop.run_in_executor(None, func_with_context) - response: Optional[EmbeddingResponse] = None + response: EmbeddingResponse | None = None if isinstance(init_response, dict): response = EmbeddingResponse(**init_response) elif isinstance(init_response, EmbeddingResponse): ## CACHING SCENARIO @@ -5867,16 +5852,16 @@ def embedding( model, input=[], # Optional params - dimensions: Optional[int] = None, - encoding_format: Optional[str] = None, + dimensions: int | None = None, + encoding_format: str | None = None, timeout=600, # default to 10 minutes # set api_base, api_version, api_key - api_base: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - api_type: Optional[str] = None, + api_base: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + api_type: str | None = None, caching: bool = False, - user: Optional[str] = None, + user: str | None = None, custom_llm_provider=None, litellm_call_id=None, logger_fn=None, @@ -5893,16 +5878,16 @@ def embedding( model, input=[], # Optional params - dimensions: Optional[int] = None, - encoding_format: Optional[str] = None, + dimensions: int | None = None, + encoding_format: str | None = None, timeout=600, # default to 10 minutes # set api_base, api_version, api_key - api_base: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - api_type: Optional[str] = None, + api_base: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + api_type: str | None = None, caching: bool = False, - user: Optional[str] = None, + user: str | None = None, custom_llm_provider=None, litellm_call_id=None, logger_fn=None, @@ -5920,21 +5905,21 @@ def embedding( model, input=[], # Optional params - dimensions: Optional[int] = None, - encoding_format: Optional[str] = None, + dimensions: int | None = None, + encoding_format: str | None = None, timeout=600, # default to 10 minutes # set api_base, api_version, api_key - api_base: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - api_type: Optional[str] = None, + api_base: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + api_type: str | None = None, caching: bool = False, - user: Optional[str] = None, + user: str | None = None, custom_llm_provider=None, litellm_call_id=None, logger_fn=None, **kwargs, -) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]: +) -> EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse]: """ Embedding function that calls an API to generate embeddings for the given input. @@ -5965,9 +5950,9 @@ def embedding( shared_session = kwargs.get("shared_session", None) max_retries = kwargs.get("max_retries", None) litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore + mock_response: list[float] | None = kwargs.get("mock_response", None) # type: ignore azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) - aembedding: Optional[bool] = kwargs.get("aembedding", None) + aembedding: bool | None = kwargs.get("aembedding", None) extra_headers = kwargs.get("extra_headers", None) headers = kwargs.get("headers", None) or extra_headers if headers is None: @@ -6020,7 +6005,7 @@ def embedding( if dynamic_api_key is not None: api_key = dynamic_api_key - allowed_openai_params: Optional[List[str]] = kwargs.get("allowed_openai_params", None) + allowed_openai_params: list[str] | None = kwargs.get("allowed_openai_params", None) optional_params = get_optional_params_embeddings( model=model, user=user, @@ -6054,7 +6039,7 @@ def embedding( if mock_response is not None: return mock_embedding(model=model, mock_response=mock_response) try: - response: Optional[Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]] = None + response: EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse] | None = None if azure is True or custom_llm_provider == "azure": # azure configs @@ -6646,22 +6631,7 @@ def embedding( aembedding=aembedding, litellm_params={}, ) - elif custom_llm_provider == "voyage": - response = base_llm_http_handler.embedding( - model=model, - input=input, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - api_key=api_key, - logging_obj=logging, - timeout=timeout, - model_response=EmbeddingResponse(), - optional_params=optional_params, - client=client, - aembedding=aembedding, - litellm_params={}, - ) - elif custom_llm_provider == "infinity": + elif custom_llm_provider == "voyage" or custom_llm_provider == "infinity": response = base_llm_http_handler.embedding( model=model, input=input, @@ -6875,7 +6845,7 @@ def embedding( litellm_params={}, ) elif custom_llm_provider in litellm._custom_providers: - custom_handler: Optional[CustomLLM] = None + custom_handler: CustomLLM | None = None for item in litellm.custom_provider_map: if item["provider"] == custom_llm_provider: custom_handler = item["custom_handler"] @@ -6976,7 +6946,7 @@ def embedding( ###### Text Completion ################ @client -async def atext_completion(*args, **kwargs) -> Union[TextCompletionResponse, TextCompletionStreamWrapper]: +async def atext_completion(*args, **kwargs) -> TextCompletionResponse | TextCompletionStreamWrapper: """ Implemented to handle async streaming for the text completion endpoint """ @@ -7047,36 +7017,31 @@ async def atext_completion(*args, **kwargs) -> Union[TextCompletionResponse, Tex @client def text_completion( - prompt: Union[ - str, List[Union[str, List[Union[str, List[int]]]]] - ], # Required: The prompt(s) to generate completions for. - model: Optional[str] = None, # Optional: either `model` or `engine` can be set - best_of: Optional[int] = None, # Optional: Generates best_of completions server-side. - echo: Optional[bool] = None, # Optional: Echo back the prompt in addition to the completion. - frequency_penalty: Optional[float] = None, # Optional: Penalize new tokens based on their existing frequency. - logit_bias: Optional[Dict[int, int]] = None, # Optional: Modify the likelihood of specified tokens. - logprobs: Optional[int] = None, # Optional: Include the log probabilities on the most likely tokens. - max_tokens: Optional[int] = None, # Optional: The maximum number of tokens to generate in the completion. - n: Optional[int] = None, # Optional: How many completions to generate for each prompt. - presence_penalty: Optional[ - float - ] = None, # Optional: Penalize new tokens based on whether they appear in the text so far. - stop: Optional[ - Union[str, List[str]] - ] = None, # Optional: Sequences where the API will stop generating further tokens. - stream: Optional[bool] = None, # Optional: Whether to stream back partial progress. - stream_options: Optional[dict] = None, - suffix: Optional[str] = None, # Optional: The suffix that comes after a completion of inserted text. - temperature: Optional[float] = None, # Optional: Sampling temperature to use. - top_p: Optional[float] = None, # Optional: Nucleus sampling parameter. - user: Optional[str] = None, # Optional: A unique identifier representing your end-user. + prompt: str | list[str | list[str | list[int]]], # Required: The prompt(s) to generate completions for. + model: str | None = None, # Optional: either `model` or `engine` can be set + best_of: int | None = None, # Optional: Generates best_of completions server-side. + echo: bool | None = None, # Optional: Echo back the prompt in addition to the completion. + frequency_penalty: float | None = None, # Optional: Penalize new tokens based on their existing frequency. + logit_bias: dict[int, int] | None = None, # Optional: Modify the likelihood of specified tokens. + logprobs: int | None = None, # Optional: Include the log probabilities on the most likely tokens. + max_tokens: int | None = None, # Optional: The maximum number of tokens to generate in the completion. + n: int | None = None, # Optional: How many completions to generate for each prompt. + presence_penalty: float + | None = None, # Optional: Penalize new tokens based on whether they appear in the text so far. + stop: str | list[str] | None = None, # Optional: Sequences where the API will stop generating further tokens. + stream: bool | None = None, # Optional: Whether to stream back partial progress. + stream_options: dict | None = None, + suffix: str | None = None, # Optional: The suffix that comes after a completion of inserted text. + temperature: float | None = None, # Optional: Sampling temperature to use. + top_p: float | None = None, # Optional: Nucleus sampling parameter. + user: str | None = None, # Optional: A unique identifier representing your end-user. # set api_base, api_version, api_key - api_base: Optional[str] = None, - api_version: Optional[str] = None, - api_key: Optional[str] = None, - model_list: Optional[list] = None, # pass in a list of api_base,keys, etc. + api_base: str | None = None, + api_version: str | None = None, + api_key: str | None = None, + model_list: list | None = None, # pass in a list of api_base,keys, etc. # Optional liteLLM function params - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, *args, **kwargs, ): @@ -7117,7 +7082,7 @@ def text_completion( text_completion_response = TextCompletionResponse() - optional_params: Dict[str, Any] = {} + optional_params: dict[str, Any] = {} # default values for all optional params are none, litellm only passes them to the llm when they are set to non None values if best_of is not None: optional_params["best_of"] = best_of @@ -7287,29 +7252,25 @@ def text_completion( ###### Adapter Completion ################ -async def aadapter_completion( - *, adapter_id: str, **kwargs -) -> Optional[Union[BaseModel, AdapterCompletionStreamWrapper]]: +async def aadapter_completion(*, adapter_id: str, **kwargs) -> BaseModel | AdapterCompletionStreamWrapper | None: """ Implemented to handle async calls for adapter_completion() """ try: - translation_obj: Optional[CustomLogger] = None + translation_obj: CustomLogger | None = None for item in litellm.adapters: if item["id"] == adapter_id: translation_obj = item["adapter"] if translation_obj is None: raise ValueError( - "No matching adapter given. Received 'adapter_id'={}, litellm.adapters={}".format( - adapter_id, litellm.adapters - ) + f"No matching adapter given. Received 'adapter_id'={adapter_id}, litellm.adapters={litellm.adapters}" ) new_kwargs = translation_obj.translate_completion_input_params(kwargs=kwargs) - response: Union[ModelResponse, CustomStreamWrapper] = await acompletion(**new_kwargs) # type: ignore - translated_response: Optional[Union[BaseModel, AdapterCompletionStreamWrapper]] = None + response: ModelResponse | CustomStreamWrapper = await acompletion(**new_kwargs) # type: ignore + translated_response: BaseModel | AdapterCompletionStreamWrapper | None = None if isinstance(response, ModelResponse): translated_response = translation_obj.translate_completion_output_params(response=response) if isinstance(response, CustomStreamWrapper): @@ -7324,33 +7285,31 @@ async def aadapter_completion( async def aadapter_generate_content( **kwargs, -) -> Union[Dict[str, Any], AsyncIterator[bytes]]: +) -> dict[str, Any] | AsyncIterator[bytes]: from litellm.google_genai.adapters.handler import GenerateContentToCompletionHandler coro = cast( - Coroutine[Any, Any, Union[Dict[str, Any], AsyncIterator[bytes]]], + Coroutine[Any, Any, dict[str, Any] | AsyncIterator[bytes]], GenerateContentToCompletionHandler.generate_content_handler(**kwargs, _is_async=True), ) return await coro -def adapter_completion(*, adapter_id: str, **kwargs) -> Optional[Union[BaseModel, AdapterCompletionStreamWrapper]]: - translation_obj: Optional[CustomLogger] = None +def adapter_completion(*, adapter_id: str, **kwargs) -> BaseModel | AdapterCompletionStreamWrapper | None: + translation_obj: CustomLogger | None = None for item in litellm.adapters: if item["id"] == adapter_id: translation_obj = item["adapter"] if translation_obj is None: raise ValueError( - "No matching adapter given. Received 'adapter_id'={}, litellm.adapters={}".format( - adapter_id, litellm.adapters - ) + f"No matching adapter given. Received 'adapter_id'={adapter_id}, litellm.adapters={litellm.adapters}" ) new_kwargs = translation_obj.translate_completion_input_params(kwargs=kwargs) - response: Union[ModelResponse, CustomStreamWrapper] = completion(**new_kwargs) # type: ignore - translated_response: Optional[Union[BaseModel, AdapterCompletionStreamWrapper]] = None + response: ModelResponse | CustomStreamWrapper = completion(**new_kwargs) # type: ignore + translated_response: BaseModel | AdapterCompletionStreamWrapper | None = None if isinstance(response, ModelResponse): translated_response = translation_obj.translate_completion_output_params(response=response) elif isinstance(response, CustomStreamWrapper) or inspect.isgenerator(response): @@ -7362,9 +7321,7 @@ def adapter_completion(*, adapter_id: str, **kwargs) -> Optional[Union[BaseModel ##### Moderation ####################### -def moderation( - input: str, model: Optional[str] = None, api_key: Optional[str] = None, **kwargs -) -> OpenAIModerationResponse: +def moderation(input: str, model: str | None = None, api_key: str | None = None, **kwargs) -> OpenAIModerationResponse: # only supports open ai for now api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") @@ -7383,7 +7340,7 @@ def moderation( else: response = openai_client.moderations.create(input=input) - response_dict: Dict = response.model_dump() + response_dict: dict = response.model_dump() return litellm.utils.LiteLLMResponseObjectHandler.convert_to_moderation_response( response_object=response_dict, ) @@ -7392,9 +7349,9 @@ def moderation( @client async def amoderation( input: str, - model: Optional[str] = None, - api_key: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + model: str | None = None, + api_key: str | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> OpenAIModerationResponse: from openai import AsyncOpenAI @@ -7402,7 +7359,7 @@ async def amoderation( # only supports open ai for now api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") optional_params = GenericLiteLLMParams(**kwargs) - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) + litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj", None) _dynamic_api_base = None try: ( @@ -7449,7 +7406,7 @@ async def amoderation( response = await _openai_client.moderations.create(input=input, model=model) else: response = await _openai_client.moderations.create(input=input) - response_dict: Dict = response.model_dump() + response_dict: dict = response.model_dump() return litellm.utils.LiteLLMResponseObjectHandler.convert_to_moderation_response( response_object=response_dict, ) @@ -7525,21 +7482,21 @@ def transcription( model: str, file: FileTypes, ## OPTIONAL OPENAI PARAMS ## - language: Optional[str] = None, - prompt: Optional[str] = None, - response_format: Optional[Literal["json", "text", "srt", "verbose_json", "vtt"]] = None, - timestamp_granularities: Optional[List[Literal["word", "segment"]]] = None, - temperature: Optional[int] = None, # openai defaults this to 0 + language: str | None = None, + prompt: str | None = None, + response_format: Literal["json", "text", "srt", "verbose_json", "vtt"] | None = None, + timestamp_granularities: list[Literal["word", "segment"]] | None = None, + temperature: int | None = None, # openai defaults this to 0 ## LITELLM PARAMS ## - user: Optional[str] = None, + user: str | None = None, timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - max_retries: Optional[int] = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + max_retries: int | None = None, custom_llm_provider=None, **kwargs, -) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: +) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]: """ Calls openai + azure whisper endpoints. @@ -7556,14 +7513,9 @@ def transcription( kwargs.pop("tags", []) non_default_params = get_non_default_transcription_params(kwargs) - client: Optional[ - Union[ - openai.AsyncOpenAI, - openai.OpenAI, - openai.AzureOpenAI, - openai.AsyncAzureOpenAI, - ] - ] = kwargs.pop("client", None) + client: openai.AsyncOpenAI | openai.OpenAI | openai.AzureOpenAI | openai.AsyncAzureOpenAI | None = kwargs.pop( + "client", None + ) if litellm_logging_obj: litellm_logging_obj.model_call_details["client"] = str(client) @@ -7611,7 +7563,7 @@ def transcription( custom_llm_provider=custom_llm_provider, ) - response: Optional[Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]] = None + response: TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse] | None = None provider_config = ProviderConfigManager.get_provider_audio_transcription_config( model=model, @@ -7828,26 +7780,26 @@ async def aspeech(*args, **kwargs) -> HttpxBinaryResponseContent: def speech( model: str, input: str, - voice: Optional[Union[str, dict]] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, - project: Optional[str] = None, - max_retries: Optional[int] = None, - metadata: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - response_format: Optional[str] = None, - speed: Optional[int] = None, - instructions: Optional[str] = None, + voice: str | dict | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + organization: str | None = None, + project: str | None = None, + max_retries: int | None = None, + metadata: dict | None = None, + timeout: float | httpx.Timeout | None = None, + response_format: str | None = None, + speed: int | None = None, + instructions: str | None = None, client=None, - headers: Optional[dict] = None, - custom_llm_provider: Optional[str] = None, - aspeech: Optional[bool] = None, + headers: dict | None = None, + custom_llm_provider: str | None = None, + aspeech: bool | None = None, **kwargs, -) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]: +) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]: user = kwargs.get("user", None) - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) proxy_server_request = kwargs.get("proxy_server_request", None) extra_headers = kwargs.get("extra_headers", None) model_info = kwargs.get("model_info", None) @@ -7904,11 +7856,7 @@ def speech( }, custom_llm_provider=custom_llm_provider, ) - response: Union[ - HttpxBinaryResponseContent, - Coroutine[Any, Any, HttpxBinaryResponseContent], - None, - ] = None + response: HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent] | None = None if custom_llm_provider == "openai" or custom_llm_provider in litellm.openai_compatible_providers: if voice is None or not (isinstance(voice, str)): raise litellm.BadRequestError( @@ -8015,7 +7963,7 @@ def speech( or get_secret("AZURE_API_KEY") ) # type: ignore - azure_ad_token: Optional[str] = optional_params.get("extra_body", {}).pop( # type: ignore + azure_ad_token: str | None = optional_params.get("extra_body", {}).pop( # type: ignore "azure_ad_token", None ) or get_secret("AZURE_AD_TOKEN") azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) @@ -8202,7 +8150,7 @@ def speech( litellm_params_dict["api_key"] = api_key # Convert voice to string if it's a dict (minimax handler expects Optional[str]) - voice_str: Optional[str] = None + voice_str: str | None = None if isinstance(voice, str): voice_str = voice elif isinstance(voice, dict): @@ -8253,9 +8201,7 @@ def speech( if response is None: raise Exception( - "Unable to map the custom llm provider={} to a known provider={}.".format( - custom_llm_provider, litellm.provider_list - ) + f"Unable to map the custom llm provider={custom_llm_provider} to a known provider={litellm.provider_list}." ) return response @@ -8266,8 +8212,8 @@ def speech( async def ahealth_check( model_params: dict, mode: str | None = "chat", - prompt: Optional[str] = None, - input: Optional[List] = None, + prompt: str | None = None, + input: list | None = None, ): """ Support health checks for different providers. Return remaining rate limit, etc. @@ -8305,7 +8251,7 @@ async def ahealth_check( ) ######################################################### try: - model: Optional[str] = model_params.get("model", None) + model: str | None = model_params.get("model", None) if model is None: raise Exception("model not set") @@ -8357,7 +8303,7 @@ async def ahealth_check( if mode is None: return { - "error": f"error:{str(e)}. Missing `mode`. Set the `mode` for the model - https://docs.litellm.ai/docs/proxy/health#embedding-models \nstacktrace: {stack_trace}", + "error": f"error:{e!s}. Missing `mode`. Set the `mode` for the model - https://docs.litellm.ai/docs/proxy/health#embedding-models \nstacktrace: {stack_trace}", "exception": e, } @@ -8394,7 +8340,7 @@ def config_completion(**kwargs): ) -def stream_chunk_builder_text_completion(chunks: list, messages: Optional[List] = None) -> TextCompletionResponse: +def stream_chunk_builder_text_completion(chunks: list, messages: list | None = None) -> TextCompletionResponse: id = chunks[0]["id"] object = chunks[0]["object"] created = chunks[0]["created"] @@ -8450,11 +8396,11 @@ def stream_chunk_builder_text_completion(chunks: list, messages: Optional[List] def stream_chunk_builder( chunks: list, - messages: Optional[list] = None, + messages: list | None = None, start_time=None, end_time=None, logging_obj: Optional["Logging"] = None, -) -> Optional[Union[ModelResponse, TextCompletionResponse]]: +) -> ModelResponse | TextCompletionResponse | None: try: if chunks is None: raise litellm.APIError( @@ -8485,7 +8431,7 @@ def stream_chunk_builder( # Fast path for the common text-only streaming case: # avoid repeated multi-pass list scans over chunks. - simple_content_parts: List[str] = [] + simple_content_parts: list[str] = [] is_simple_text_stream = True for chunk in chunks: if len(chunk["choices"]) == 0: @@ -8496,7 +8442,7 @@ def stream_chunk_builder( if isinstance(delta_obj, dict): delta = delta_obj elif hasattr(delta_obj, "model_dump"): - delta = cast(Dict[str, Any], delta_obj.model_dump()) + delta = cast(dict[str, Any], delta_obj.model_dump()) else: delta = {} @@ -8672,7 +8618,7 @@ def stream_chunk_builder( ] if len(provider_specific_chunks) > 0: - combined_provider_fields: Dict[str, Any] = {} + combined_provider_fields: dict[str, Any] = {} for chunk in provider_specific_chunks: fields = chunk["choices"][0]["delta"]["provider_specific_fields"] if isinstance(fields, dict): @@ -8723,7 +8669,7 @@ def stream_chunk_builder( processor.apply_provider_assembled_streaming_metadata(response, chunks, logging_obj) return response except Exception as e: - verbose_logger.exception("litellm.main.py::stream_chunk_builder() - Exception occurred - {}".format(str(e))) + verbose_logger.exception(f"litellm.main.py::stream_chunk_builder() - Exception occurred - {e!s}") raise litellm.APIError( status_code=500, message="Error building chunks for logging/streaming usage calculation", @@ -8737,11 +8683,11 @@ def stream_chunk_builder( async def acount_tokens( model: str, - messages: Optional[List[Dict[str, Any]]] = None, - tools: Optional[List[Dict[str, Any]]] = None, - system: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + messages: list[dict[str, Any]] | None = None, + tools: list[dict[str, Any]] | None = None, + system: str | None = None, + api_key: str | None = None, + api_base: str | None = None, ) -> "TokenCountResponse": """ Count tokens for a given model and messages using provider-specific APIs. @@ -8783,7 +8729,7 @@ async def acount_tokens( api_base = dynamic_api_base # Build deployment dict for the token counter - deployment: Dict[str, Any] = { + deployment: dict[str, Any] = { "litellm_params": { "model": model, "api_key": api_key, @@ -8834,7 +8780,7 @@ async def acount_tokens( # Cache for encoding to avoid repeated __getattr__ calls -_encoding_cache: Optional[Any] = None +_encoding_cache: Any | None = None def _get_encoding(): diff --git a/litellm/models/__init__.py b/litellm/models/__init__.py index 7e2d2c0ed9d..07d1ffa743d 100644 --- a/litellm/models/__init__.py +++ b/litellm/models/__init__.py @@ -36,31 +36,31 @@ from litellm.models.user import LiteLLM_UserTable from litellm.models.verification_token import LiteLLM_VerificationToken __all__ = [ + "CreateCredentialItem", + "CredentialBase", + "CredentialItem", "LiteLLM_AccessGroupTable", "LiteLLM_BudgetTable", "LiteLLM_BudgetTableFull", - "LiteLLM_TeamMemberTable", "LiteLLM_Config", - "CredentialBase", - "CredentialItem", - "CreateCredentialItem", "LiteLLM_EndUserTable", + "LiteLLM_ErrorLogs", + "LiteLLM_MCPServerTable", "LiteLLM_ManagedFileTable", "LiteLLM_ManagedObjectTable", "LiteLLM_ManagedVectorStoreTable", "LiteLLM_ManagedVectorStoresTable", - "LiteLLM_MCPServerTable", - "LiteLLM_ProxyModelTable", "LiteLLM_ObjectPermissionTable", - "LiteLLM_OrganizationTable", "LiteLLM_OrganizationMembershipTable", + "LiteLLM_OrganizationTable", "LiteLLM_ProjectTable", + "LiteLLM_ProxyModelTable", "LiteLLM_SkillsTable", - "LiteLLM_ErrorLogs", "LiteLLM_SpendLogs", "LiteLLM_TagTable", - "LiteLLM_TeamTable", + "LiteLLM_TeamMemberTable", "LiteLLM_TeamMembership", + "LiteLLM_TeamTable", "LiteLLM_UserTable", "LiteLLM_VerificationToken", ] diff --git a/litellm/models/access_group.py b/litellm/models/access_group.py index 682e779e531..513a60aa9ba 100644 --- a/litellm/models/access_group.py +++ b/litellm/models/access_group.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_accessgrouptable``. Re-exported from """ from datetime import datetime -from typing import List, Optional from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -14,13 +13,13 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_AccessGroupTable(LiteLLMPydanticObjectBase): access_group_id: str access_group_name: str - description: Optional[str] = None - access_model_names: List[str] = [] - access_mcp_server_ids: List[str] = [] - access_agent_ids: List[str] = [] - assigned_team_ids: List[str] = [] - assigned_key_ids: List[str] = [] - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None - updated_by: Optional[str] = None + description: str | None = None + access_model_names: list[str] = [] + access_mcp_server_ids: list[str] = [] + access_agent_ids: list[str] = [] + assigned_team_ids: list[str] = [] + assigned_key_ids: list[str] = [] + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None + updated_by: str | None = None diff --git a/litellm/models/base.py b/litellm/models/base.py index 01981297bd5..7eedf10212e 100644 --- a/litellm/models/base.py +++ b/litellm/models/base.py @@ -3,7 +3,7 @@ Base model class for domain models. """ from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any from pydantic import BaseModel, ConfigDict @@ -17,8 +17,8 @@ class DomainModel(BaseModel): extra="ignore", ) - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None + created_at: datetime | None = None + updated_at: datetime | None = None @classmethod def from_db_record(cls, record: Any) -> "DomainModel": @@ -33,6 +33,6 @@ class DomainModel(BaseModel): return cls(**record.dict()) return cls(**dict(record)) - def to_db_dict(self, exclude_unset: bool = False) -> Dict[str, Any]: + def to_db_dict(self, exclude_unset: bool = False) -> dict[str, Any]: """Convert domain model to a dictionary for database operations.""" return self.model_dump(exclude_none=True, exclude_unset=exclude_unset) diff --git a/litellm/models/budget.py b/litellm/models/budget.py index 8c35aebd208..335800a49a8 100644 --- a/litellm/models/budget.py +++ b/litellm/models/budget.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_budgettable``. Re-exported from """ from datetime import datetime -from typing import List, Optional from pydantic import ConfigDict @@ -21,15 +20,15 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase): `LiteLLM_BudgetTableFull` so they aren't user-settable. """ - budget_id: Optional[str] = None - soft_budget: Optional[float] = None - max_budget: Optional[float] = None - max_parallel_requests: Optional[int] = None - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - model_max_budget: Optional[dict] = None - budget_duration: Optional[str] = None - allowed_models: Optional[List[str]] = None # per-member model scope; empty = inherit team models + budget_id: str | None = None + soft_budget: float | None = None + max_budget: float | None = None + max_parallel_requests: int | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + model_max_budget: dict | None = None + budget_duration: str | None = None + allowed_models: list[str] | None = None # per-member model scope; empty = inherit team models model_config = ConfigDict(protected_namespaces=()) @@ -37,7 +36,7 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase): class LiteLLM_BudgetTableFull(LiteLLM_BudgetTable): """LiteLLM_BudgetTable + server-managed fields returned on API responses.""" - budget_reset_at: Optional[datetime] = None + budget_reset_at: datetime | None = None created_at: datetime @@ -46,9 +45,9 @@ class LiteLLM_TeamMemberTable(LiteLLM_BudgetTable): Used to track spend of a user_id within a team_id """ - spend: Optional[float] = None - user_id: Optional[str] = None - team_id: Optional[str] = None - budget_id: Optional[str] = None + spend: float | None = None + user_id: str | None = None + team_id: str | None = None + budget_id: str | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/models/config.py b/litellm/models/config.py index 99b5c5692fd..d2c0b23cf4b 100644 --- a/litellm/models/config.py +++ b/litellm/models/config.py @@ -5,11 +5,9 @@ Canonical definition for ``litellm_config``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from typing import Dict - from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_Config(LiteLLMPydanticObjectBase): param_name: str - param_value: Dict + param_value: dict diff --git a/litellm/models/credentials.py b/litellm/models/credentials.py index b74ea055d21..56836234898 100644 --- a/litellm/models/credentials.py +++ b/litellm/models/credentials.py @@ -5,8 +5,6 @@ These are the canonical credential types for the proxy. They live in the model layer; ``litellm.types.utils`` re-exports them for backwards compatibility. """ -from typing import Optional - from pydantic import BaseModel, model_validator @@ -20,8 +18,8 @@ class CredentialItem(CredentialBase): class CreateCredentialItem(CredentialBase): - credential_values: Optional[dict] = None - model_id: Optional[str] = None + credential_values: dict | None = None + model_id: str | None = None @model_validator(mode="before") @classmethod diff --git a/litellm/models/end_user.py b/litellm/models/end_user.py index 9bf895b9447..8dccf1eb5e7 100644 --- a/litellm/models/end_user.py +++ b/litellm/models/end_user.py @@ -5,7 +5,7 @@ Canonical definition for ``litellm_endusertable``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from typing import Literal, Optional +from typing import Literal from pydantic import ConfigDict, model_validator @@ -17,14 +17,14 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_EndUserTable(LiteLLMPydanticObjectBase): user_id: str blocked: bool - alias: Optional[str] = None + alias: str | None = None spend: float = 0.0 - allowed_model_region: Optional[Literal["eu", "us"]] = None - default_model: Optional[str] = None - budget_id: Optional[str] = None - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - object_permission_id: Optional[str] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None + allowed_model_region: Literal["eu", "us"] | None = None + default_model: str | None = None + budget_id: str | None = None + litellm_budget_table: LiteLLM_BudgetTable | None = None + object_permission_id: str | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None @model_validator(mode="before") @classmethod diff --git a/litellm/models/managed_files.py b/litellm/models/managed_files.py index 99ba764dd98..23d70ef5c48 100644 --- a/litellm/models/managed_files.py +++ b/litellm/models/managed_files.py @@ -6,7 +6,7 @@ Canonical definitions for the ``litellm_managed*`` tables. Re-exported from """ from datetime import datetime -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal from litellm.types.llms.base import LiteLLMPydanticObjectBase from litellm.types.llms.openai import OpenAIFileObject, ResponsesAPIResponse @@ -15,48 +15,48 @@ from litellm.types.utils import LiteLLMBatch, LiteLLMFineTuningJob class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): unified_file_id: str - file_object: Optional[OpenAIFileObject] = None - model_mappings: Dict[str, str] - flat_model_file_ids: List[str] - created_by: Optional[str] = None - team_id: Optional[str] = None - updated_by: Optional[str] = None - storage_backend: Optional[str] = None - storage_url: Optional[str] = None + file_object: OpenAIFileObject | None = None + model_mappings: dict[str, str] + flat_model_file_ids: list[str] + created_by: str | None = None + team_id: str | None = None + updated_by: str | None = None + storage_backend: str | None = None + storage_url: str | None = None class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase): unified_object_id: str model_object_id: str file_purpose: Literal["batch", "fine-tune", "response", "container"] - file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, ResponsesAPIResponse] - created_by: Optional[str] = None - team_id: Optional[str] = None + file_object: LiteLLMBatch | LiteLLMFineTuningJob | ResponsesAPIResponse + created_by: str | None = None + team_id: str | None = None class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase): """Table for managing vector stores with target_model_names support.""" unified_resource_id: str - resource_object: Optional[Any] = None - model_mappings: Dict[str, str] - flat_model_resource_ids: List[str] - created_by: Optional[str] = None - team_id: Optional[str] = None - updated_by: Optional[str] = None - storage_backend: Optional[str] = None - storage_url: Optional[str] = None + resource_object: Any | None = None + model_mappings: dict[str, str] + flat_model_resource_ids: list[str] + created_by: str | None = None + team_id: str | None = None + updated_by: str | None = None + storage_backend: str | None = None + storage_url: str | None = None class LiteLLM_ManagedVectorStoresTable(LiteLLMPydanticObjectBase): vector_store_id: str custom_llm_provider: str - vector_store_name: Optional[str] = None - vector_store_description: Optional[str] = None - vector_store_metadata: Optional[Dict[str, Any]] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None - litellm_credential_name: Optional[str] = None - litellm_params: Optional[Dict[str, Any]] = None - team_id: Optional[str] = None - user_id: Optional[str] = None + vector_store_name: str | None = None + vector_store_description: str | None = None + vector_store_metadata: dict[str, Any] | None = None + created_at: datetime | None = None + updated_at: datetime | None = None + litellm_credential_name: str | None = None + litellm_params: dict[str, Any] | None = None + team_id: str | None = None + user_id: str | None = None diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py index e428d20f99d..6cc4a765e46 100644 --- a/litellm/models/mcp_server.py +++ b/litellm/models/mcp_server.py @@ -7,7 +7,7 @@ Canonical definition for ``litellm_mcpservertable``. Re-exported from import enum from datetime import datetime -from typing import Dict, List, Literal, Optional +from typing import Literal from pydantic import Field @@ -41,76 +41,76 @@ class MCPEnvVar(LiteLLMPydanticObjectBase): name: str value: str = "" scope: MCPEnvVarScope = MCPEnvVarScope.global_ - description: Optional[str] = None + description: str | None = None class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): """Represents a LiteLLM_MCPServerTable record""" server_id: str - server_name: Optional[str] = None - alias: Optional[str] = None - description: Optional[str] = None - url: Optional[str] = None - spec_path: Optional[str] = None + server_name: str | None = None + alias: str | None = None + description: str | None = None + url: str | None = None + spec_path: str | None = None transport: MCPTransportType - auth_type: Optional[MCPAuthType] = None - credentials: Optional[MCPCredentials] = None - instructions: Optional[str] = None - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None - updated_by: Optional[str] = None - teams: List[Dict[str, Optional[str]]] = Field(default_factory=list) - mcp_access_groups: List[str] = Field(default_factory=list) - allowed_tools: List[str] = Field(default_factory=list) - tool_name_to_display_name: Optional[Dict[str, str]] = None - tool_name_to_description: Optional[Dict[str, str]] = None - extra_headers: List[str] = Field(default_factory=list) - mcp_info: Optional[MCPInfo] = None - static_headers: Optional[Dict[str, str]] = None - env_vars: Optional[List[MCPEnvVar]] = None - status: Optional[Literal["healthy", "unhealthy", "unknown"]] = Field( + auth_type: MCPAuthType | None = None + credentials: MCPCredentials | None = None + instructions: str | None = None + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None + updated_by: str | None = None + teams: list[dict[str, str | None]] = Field(default_factory=list) + mcp_access_groups: list[str] = Field(default_factory=list) + allowed_tools: list[str] = Field(default_factory=list) + tool_name_to_display_name: dict[str, str] | None = None + tool_name_to_description: dict[str, str] | None = None + extra_headers: list[str] = Field(default_factory=list) + mcp_info: MCPInfo | None = None + static_headers: dict[str, str] | None = None + env_vars: list[MCPEnvVar] | None = None + status: Literal["healthy", "unhealthy", "unknown"] | None = Field( default="unknown", description="Health status: 'healthy', 'unhealthy', 'unknown'", ) - last_health_check: Optional[datetime] = None - health_check_error: Optional[str] = None - command: Optional[str] = None - args: List[str] = Field(default_factory=list) - env: Dict[str, str] = Field(default_factory=dict) - issuer: Optional[str] = None - authorization_url: Optional[str] = None - token_url: Optional[str] = None - registration_url: Optional[str] = None - oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None + last_health_check: datetime | None = None + health_check_error: str | None = None + command: str | None = None + args: list[str] = Field(default_factory=list) + env: dict[str, str] = Field(default_factory=dict) + issuer: str | None = None + authorization_url: str | None = None + token_url: str | None = None + registration_url: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None # Token Exchange (OBO) fields — RFC 8693. ``audience`` is named for the RFC's # request parameter (token-exchange only); RFC 8707 resource indicators are a # separate concept named ``resource`` in the v2 egress types. A null # ``subject_token_type`` means DEFAULT_SUBJECT_TOKEN_TYPE (litellm.types.mcp), # applied at the egress build sites. - token_exchange_endpoint: Optional[str] = None - audience: Optional[str] = None - subject_token_type: Optional[str] = None - token_exchange_profile: Optional[str] = None + token_exchange_endpoint: str | None = None + audience: str | None = None + subject_token_type: str | None = None + token_exchange_profile: str | None = None allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False oauth_passthrough: bool = False - dcr_bridge: Optional[bool] = None + dcr_bridge: bool | None = None is_byok: bool = False - byok_description: List[str] = Field(default_factory=list) - byok_api_key_help_url: Optional[str] = None - has_user_credential: Optional[bool] = None + byok_description: list[str] = Field(default_factory=list) + byok_api_key_help_url: str | None = None + has_user_credential: bool | None = None connected_app_reachable: bool | None = None - source_url: Optional[str] = None - timeout: Optional[float] = None - max_concurrent_requests: Optional[int] = None - approval_status: Optional[str] = Field( + source_url: str | None = None + timeout: float | None = None + max_concurrent_requests: int | None = None + approval_status: str | None = Field( default="active", description="Approval status: 'pending_review', 'active', 'rejected'", ) - submitted_by: Optional[str] = None - submitted_at: Optional[datetime] = None - reviewed_at: Optional[datetime] = None - review_notes: Optional[str] = None + submitted_by: str | None = None + submitted_at: datetime | None = None + reviewed_at: datetime | None = None + review_notes: str | None = None diff --git a/litellm/models/model.py b/litellm/models/model.py index 7657e4d30f8..209f26d4837 100644 --- a/litellm/models/model.py +++ b/litellm/models/model.py @@ -7,7 +7,6 @@ Canonical definition for ``litellm_proxymodeltable``. Re-exported from import json from datetime import datetime -from typing import Optional from pydantic import ConfigDict, model_validator @@ -18,12 +17,12 @@ class LiteLLM_ProxyModelTable(LiteLLMPydanticObjectBase): model_id: str model_name: str litellm_params: dict - model_info: Optional[dict] = None + model_info: dict | None = None blocked: bool = False - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None - updated_by: Optional[str] = None + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None + updated_by: str | None = None model_config = ConfigDict(protected_namespaces=()) @@ -47,13 +46,13 @@ class LiteLLM_ProxyModelTable(LiteLLMPydanticObjectBase): return self.blocked @property - def team_id(self) -> Optional[str]: + def team_id(self) -> str | None: if self.model_info: return self.model_info.get("team_id") return None @property - def team_public_model_name(self) -> Optional[str]: + def team_public_model_name(self) -> str | None: if self.model_info: return self.model_info.get("team_public_model_name") return None diff --git a/litellm/models/object_permission.py b/litellm/models/object_permission.py index 3052a2af459..a09d50ddc33 100644 --- a/litellm/models/object_permission.py +++ b/litellm/models/object_permission.py @@ -5,8 +5,6 @@ Canonical definition for ``litellm_objectpermissiontable``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from typing import Dict, List, Optional - from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -14,14 +12,14 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): """Represents a LiteLLM_ObjectPermissionTable record""" object_permission_id: str - mcp_servers: Optional[List[str]] = [] - mcp_access_groups: Optional[List[str]] = [] - mcp_tool_permissions: Optional[Dict[str, List[str]]] = None - vector_stores: Optional[List[str]] = [] - agents: Optional[List[str]] = [] - agent_access_groups: Optional[List[str]] = [] - models: Optional[List[str]] = [] - mcp_toolsets: Optional[List[str]] = None - blocked_tools: Optional[List[str]] = [] - search_tools: Optional[List[str]] = [] - mcp_tool_search_enabled: Optional[bool] = None + mcp_servers: list[str] | None = [] + mcp_access_groups: list[str] | None = [] + mcp_tool_permissions: dict[str, list[str]] | None = None + vector_stores: list[str] | None = [] + agents: list[str] | None = [] + agent_access_groups: list[str] | None = [] + models: list[str] | None = [] + mcp_toolsets: list[str] | None = None + blocked_tools: list[str] | None = [] + search_tools: list[str] | None = [] + mcp_tool_search_enabled: bool | None = None diff --git a/litellm/models/organization.py b/litellm/models/organization.py index 8b2d95c3e09..894c178af0d 100644 --- a/litellm/models/organization.py +++ b/litellm/models/organization.py @@ -5,8 +5,6 @@ Canonical definition for ``litellm_organizationtable``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from typing import List, Optional - from litellm.models.budget import LiteLLM_BudgetTable from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.models.user import LiteLLM_UserTable @@ -16,16 +14,16 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_OrganizationTable(LiteLLMPydanticObjectBase): """Represents user-controllable params for a LiteLLM_OrganizationTable record""" - organization_id: Optional[str] = None - organization_alias: Optional[str] = None + organization_id: str | None = None + organization_alias: str | None = None budget_id: str spend: float = 0.0 - metadata: Optional[dict] = None - models: List[str] = [] - model_spend: Optional[dict] = {} + metadata: dict | None = None + models: list[str] = [] + model_spend: dict | None = {} created_by: str updated_by: str - users: Optional[List[LiteLLM_UserTable]] = None - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None - object_permission_id: Optional[str] = None + users: list[LiteLLM_UserTable] | None = None + litellm_budget_table: LiteLLM_BudgetTable | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None + object_permission_id: str | None = None diff --git a/litellm/models/organization_membership.py b/litellm/models/organization_membership.py index 9957c0c21af..e9697cf57b6 100644 --- a/litellm/models/organization_membership.py +++ b/litellm/models/organization_membership.py @@ -6,7 +6,7 @@ Canonical definition for ``litellm_organizationmembership``. Re-exported from """ from datetime import datetime -from typing import Any, Optional +from typing import Any from pydantic import ConfigDict, model_validator @@ -19,14 +19,14 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase): user_id: str organization_id: str - user_role: Optional[str] = None + user_role: str | None = None spend: float = 0.0 - budget_id: Optional[str] = None + budget_id: str | None = None created_at: datetime updated_at: datetime - user: Optional[Any] = None - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - user_email: Optional[str] = None + user: Any | None = None + litellm_budget_table: LiteLLM_BudgetTable | None = None + user_email: str | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/models/project.py b/litellm/models/project.py index 083c7ee3cc5..a785b6db502 100644 --- a/litellm/models/project.py +++ b/litellm/models/project.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_projecttable``. Re-exported from """ from datetime import datetime -from typing import List, Optional from litellm.models.budget import LiteLLM_BudgetTable from litellm.models.object_permission import LiteLLM_ObjectPermissionTable @@ -17,24 +16,24 @@ class LiteLLM_ProjectTable(LiteLLMPydanticObjectBase): """Database model representation for project""" project_id: str - project_alias: Optional[str] = None - description: Optional[str] = None - team_id: Optional[str] = None - budget_id: Optional[str] = None - metadata: Optional[dict] = None - models: List[str] = [] + project_alias: str | None = None + description: str | None = None + team_id: str | None = None + budget_id: str | None = None + metadata: dict | None = None + models: list[str] = [] spend: float = 0.0 - model_spend: Optional[dict] = None - model_rpm_limit: Optional[dict] = None - model_tpm_limit: Optional[dict] = None + model_spend: dict | None = None + model_rpm_limit: dict | None = None + model_tpm_limit: dict | None = None blocked: bool = False - object_permission_id: Optional[str] = None - created_by: Optional[str] = None - updated_by: Optional[str] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None + object_permission_id: str | None = None + created_by: str | None = None + updated_by: str | None = None + created_at: datetime | None = None + updated_at: datetime | None = None + litellm_budget_table: LiteLLM_BudgetTable | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None @property def is_blocked(self) -> bool: diff --git a/litellm/models/skills.py b/litellm/models/skills.py index 62091c0ca01..f56dcfff49a 100644 --- a/litellm/models/skills.py +++ b/litellm/models/skills.py @@ -6,7 +6,7 @@ Canonical definition for ``litellm_skillstable``. Re-exported from """ from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -15,16 +15,16 @@ class LiteLLM_SkillsTable(LiteLLMPydanticObjectBase): """Represents a LiteLLM_SkillsTable record""" skill_id: str - display_title: Optional[str] = None - description: Optional[str] = None - instructions: Optional[str] = None + display_title: str | None = None + description: str | None = None + instructions: str | None = None source: str = "custom" - latest_version: Optional[str] = None - file_content: Optional[bytes] = None - file_name: Optional[str] = None - file_type: Optional[str] = None - metadata: Optional[Dict[str, Any]] = None - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None - updated_by: Optional[str] = None + latest_version: str | None = None + file_content: bytes | None = None + file_name: str | None = None + file_type: str | None = None + metadata: dict[str, Any] | None = None + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None + updated_by: str | None = None diff --git a/litellm/models/spend_logs.py b/litellm/models/spend_logs.py index 96bd328c3ca..c5a0522864a 100644 --- a/litellm/models/spend_logs.py +++ b/litellm/models/spend_logs.py @@ -6,7 +6,6 @@ Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ from datetime import datetime -from typing import Optional, Union from pydantic import Json @@ -17,34 +16,34 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_SpendLogs(LiteLLMPydanticObjectBase): request_id: str api_key: str - model: Optional[str] = "" - api_base: Optional[str] = "" + model: str | None = "" + api_base: str | None = "" call_type: str - spend: Optional[float] = 0.0 - total_tokens: Optional[int] = 0 - prompt_tokens: Optional[int] = 0 - completion_tokens: Optional[int] = 0 - startTime: Union[str, datetime, None] - endTime: Union[str, datetime, None] - user: Optional[str] = "" - metadata: Optional[Json] = {} - cache_hit: Optional[str] = "False" - cache_key: Optional[str] = None - request_tags: Optional[Json] = None - requester_ip_address: Optional[str] = None - messages: Optional[Union[str, list, dict]] - response: Optional[Union[str, list, dict]] + spend: float | None = 0.0 + total_tokens: int | None = 0 + prompt_tokens: int | None = 0 + completion_tokens: int | None = 0 + startTime: str | datetime | None + endTime: str | datetime | None + user: str | None = "" + metadata: Json | None = {} + cache_hit: str | None = "False" + cache_key: str | None = None + request_tags: Json | None = None + requester_ip_address: str | None = None + messages: str | list | dict | None + response: str | list | dict | None class LiteLLM_ErrorLogs(LiteLLMPydanticObjectBase): - request_id: Optional[str] = str(uuid.uuid4()) - api_base: Optional[str] = "" - model_group: Optional[str] = "" - litellm_model_name: Optional[str] = "" - model_id: Optional[str] = "" - request_kwargs: Optional[dict] = {} - exception_type: Optional[str] = "" - status_code: Optional[str] = "" - exception_string: Optional[str] = "" - startTime: Union[str, datetime, None] - endTime: Union[str, datetime, None] + request_id: str | None = str(uuid.uuid4()) + api_base: str | None = "" + model_group: str | None = "" + litellm_model_name: str | None = "" + model_id: str | None = "" + request_kwargs: dict | None = {} + exception_type: str | None = "" + status_code: str | None = "" + exception_string: str | None = "" + startTime: str | datetime | None + endTime: str | datetime | None diff --git a/litellm/models/tag.py b/litellm/models/tag.py index 02d8f58916d..3b8cd37c003 100644 --- a/litellm/models/tag.py +++ b/litellm/models/tag.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_tagtable``. Re-exported from """ from datetime import datetime -from typing import List, Optional from pydantic import model_validator @@ -16,15 +15,15 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_TagTable(LiteLLMPydanticObjectBase): tag_name: str - description: Optional[str] = None - models: List[str] = [] - model_info: Optional[dict] = None + description: str | None = None + models: list[str] = [] + model_info: dict | None = None spend: float = 0.0 - budget_id: Optional[str] = None - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None + budget_id: str | None = None + litellm_budget_table: LiteLLM_BudgetTable | None = None + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None @model_validator(mode="before") @classmethod diff --git a/litellm/models/team.py b/litellm/models/team.py index f11c21a078e..9cb6f81ab17 100644 --- a/litellm/models/team.py +++ b/litellm/models/team.py @@ -8,7 +8,7 @@ budget-window value types and the team-model alias table). Re-exported from import json from datetime import datetime -from typing import List, Literal, Optional, Union +from typing import Literal, Optional from pydantic import BaseModel, ConfigDict, Field, model_validator @@ -17,11 +17,11 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class MemberBase(LiteLLMPydanticObjectBase): - user_id: Optional[str] = Field( + user_id: str | None = Field( default=None, description="The unique ID of the user to add. Either user_id or user_email must be provided", ) - user_email: Optional[str] = Field( + user_email: str | None = Field( default=None, description="The email address of the user to add. Either user_id or user_email must be provided", ) @@ -47,12 +47,12 @@ class BudgetLimitEntry(LiteLLMPydanticObjectBase): budget_duration: str max_budget: float - reset_at: Optional[datetime] = None + reset_at: datetime | None = None class LiteLLM_ModelTable(LiteLLMPydanticObjectBase): - id: Optional[int] = None - model_aliases: Optional[Union[str, dict]] = None + id: int | None = None + model_aliases: str | dict | None = None created_by: str updated_by: str team: Optional["LiteLLM_TeamTable"] = None @@ -61,43 +61,43 @@ class LiteLLM_ModelTable(LiteLLMPydanticObjectBase): class TeamBase(LiteLLMPydanticObjectBase): - team_alias: Optional[str] = None - team_id: Optional[str] = None - organization_id: Optional[str] = None + team_alias: str | None = None + team_id: str | None = None + organization_id: str | None = None admins: list = [] members: list = [] - members_with_roles: List[Member] = [] - team_member_permissions: Optional[List[str]] = None - metadata: Optional[dict] = None - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - max_budget: Optional[float] = None - soft_budget: Optional[float] = None - budget_duration: Optional[str] = None - budget_limits: Optional[List[BudgetLimitEntry]] = None + members_with_roles: list[Member] = [] + team_member_permissions: list[str] | None = None + metadata: dict | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + max_budget: float | None = None + soft_budget: float | None = None + budget_duration: str | None = None + budget_limits: list[BudgetLimitEntry] | None = None models: list = [] blocked: bool = False - router_settings: Optional[dict] = None - access_group_ids: Optional[List[str]] = None - default_team_member_models: Optional[List[str]] = None + router_settings: dict | None = None + access_group_ids: list[str] | None = None + default_team_member_models: list[str] | None = None class LiteLLM_TeamTable(TeamBase): team_id: str # type: ignore - spend: Optional[float] = None - max_parallel_requests: Optional[int] = None - budget_duration: Optional[str] = None - budget_reset_at: Optional[datetime] = None - model_id: Optional[int] = None - model_spend: Optional[dict] = {} - model_max_budget: Optional[dict] = {} - policies: Optional[List[str]] = None - allow_team_guardrail_config: Optional[bool] = False - litellm_model_table: Optional[LiteLLM_ModelTable] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None - object_permission_id: Optional[str] = None - updated_at: Optional[datetime] = None - created_at: Optional[datetime] = None + spend: float | None = None + max_parallel_requests: int | None = None + budget_duration: str | None = None + budget_reset_at: datetime | None = None + model_id: int | None = None + model_spend: dict | None = {} + model_max_budget: dict | None = {} + policies: list[str] | None = None + allow_team_guardrail_config: bool | None = False + litellm_model_table: LiteLLM_ModelTable | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None + object_permission_id: str | None = None + updated_at: datetime | None = None + created_at: datetime | None = None model_config = ConfigDict(protected_namespaces=()) @@ -133,17 +133,17 @@ class LiteLLM_TeamTable(TeamBase): class LiteLLM_TeamTableCachedObj(LiteLLM_TeamTable): - last_refreshed_at: Optional[float] = None + last_refreshed_at: float | None = None class LiteLLM_DeletedTeamTable(LiteLLM_TeamTable): """Audit record for deleted teams; mirrors the team plus deletion metadata.""" - id: Optional[str] = None - deleted_at: Optional[datetime] = None - deleted_by: Optional[str] = None - deleted_by_api_key: Optional[str] = None - litellm_changed_by: Optional[str] = None + id: str | None = None + deleted_at: datetime | None = None + deleted_by: str | None = None + deleted_by_api_key: str | None = None + litellm_changed_by: str | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/models/team_membership.py b/litellm/models/team_membership.py index e79b64977d4..0ffe8f8dbcc 100644 --- a/litellm/models/team_membership.py +++ b/litellm/models/team_membership.py @@ -5,8 +5,6 @@ Canonical definition for ``litellm_teammembership``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from typing import Optional, Union - from litellm.models.budget import LiteLLM_BudgetTable, LiteLLM_BudgetTableFull from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -14,17 +12,17 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_TeamMembership(LiteLLMPydanticObjectBase): user_id: str team_id: str - budget_id: Optional[str] = None - spend: Optional[float] = 0.0 - total_spend: Optional[float] = 0.0 - litellm_budget_table: Optional[Union[LiteLLM_BudgetTableFull, LiteLLM_BudgetTable]] = None + budget_id: str | None = None + spend: float | None = 0.0 + total_spend: float | None = 0.0 + litellm_budget_table: LiteLLM_BudgetTableFull | LiteLLM_BudgetTable | None = None - def safe_get_team_member_rpm_limit(self) -> Optional[int]: + def safe_get_team_member_rpm_limit(self) -> int | None: if self.litellm_budget_table is not None: return self.litellm_budget_table.rpm_limit return None - def safe_get_team_member_tpm_limit(self) -> Optional[int]: + def safe_get_team_member_tpm_limit(self) -> int | None: if self.litellm_budget_table is not None: return self.litellm_budget_table.tpm_limit return None diff --git a/litellm/models/user.py b/litellm/models/user.py index cd7e9db4aec..259c3440d87 100644 --- a/litellm/models/user.py +++ b/litellm/models/user.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_usertable``. Re-exported from """ from datetime import datetime -from typing import Dict, List, Optional from pydantic import ConfigDict, Field, model_validator @@ -19,32 +18,32 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_UserTable(LiteLLMPydanticObjectBase): user_id: str - user_alias: Optional[str] = None - team_id: Optional[str] = None - sso_user_id: Optional[str] = None - organization_id: Optional[str] = None - object_permission_id: Optional[str] = None - password: Optional[str] = Field(default=None, exclude=True) - teams: List[str] = [] - user_role: Optional[str] = None - max_budget: Optional[float] = None + user_alias: str | None = None + team_id: str | None = None + sso_user_id: str | None = None + organization_id: str | None = None + object_permission_id: str | None = None + password: str | None = Field(default=None, exclude=True) + teams: list[str] = [] + user_role: str | None = None + max_budget: float | None = None spend: float = 0.0 - user_email: Optional[str] = None + user_email: str | None = None models: list = [] - metadata: Optional[dict] = None - max_parallel_requests: Optional[int] = None - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - budget_duration: Optional[str] = None - budget_reset_at: Optional[datetime] = None - allowed_cache_controls: List[str] = [] - policies: List[str] = [] - model_spend: Optional[Dict] = {} - model_max_budget: Optional[Dict] = {} - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None - organization_memberships: Optional[List[LiteLLM_OrganizationMembershipTable]] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None + metadata: dict | None = None + max_parallel_requests: int | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + budget_duration: str | None = None + budget_reset_at: datetime | None = None + allowed_cache_controls: list[str] = [] + policies: list[str] = [] + model_spend: dict | None = {} + model_max_budget: dict | None = {} + created_at: datetime | None = None + updated_at: datetime | None = None + organization_memberships: list[LiteLLM_OrganizationMembershipTable] | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/models/verification_token.py b/litellm/models/verification_token.py index 519066b8266..ea822c2dab0 100644 --- a/litellm/models/verification_token.py +++ b/litellm/models/verification_token.py @@ -6,7 +6,6 @@ Canonical definition for ``litellm_verificationtoken``. Re-exported from """ from datetime import datetime -from typing import Dict, List, Optional, Union from pydantic import ConfigDict @@ -15,62 +14,62 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase): - token: Optional[str] = None - key_name: Optional[str] = None - key_alias: Optional[str] = None + token: str | None = None + key_name: str | None = None + key_alias: str | None = None spend: float = 0.0 - max_budget: Optional[float] = None - expires: Optional[Union[str, datetime]] = None - models: List = [] - aliases: Dict = {} - config: Dict = {} - user_id: Optional[str] = None - team_id: Optional[str] = None - agent_id: Optional[str] = None - project_id: Optional[str] = None - max_parallel_requests: Optional[int] = None - metadata: Dict = {} - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - budget_duration: Optional[str] = None - budget_reset_at: Optional[datetime] = None - allowed_cache_controls: Optional[list] = [] - allowed_routes: Optional[list] = [] + max_budget: float | None = None + expires: str | datetime | None = None + models: list = [] + aliases: dict = {} + config: dict = {} + user_id: str | None = None + team_id: str | None = None + agent_id: str | None = None + project_id: str | None = None + max_parallel_requests: int | None = None + metadata: dict = {} + tpm_limit: int | None = None + rpm_limit: int | None = None + budget_duration: str | None = None + budget_reset_at: datetime | None = None + allowed_cache_controls: list | None = [] + allowed_routes: list | None = [] key_type: str | None = None - permissions: Dict = {} - model_spend: Dict = {} - model_max_budget: Dict = {} + permissions: dict = {} + model_spend: dict = {} + model_max_budget: dict = {} budget_fallbacks: dict[str, list[str]] = {} soft_budget_cooldown: bool = False - blocked: Optional[bool] = None - litellm_budget_table: Optional[dict] = None - budget_id: Optional[str] = None - org_id: Optional[str] = None # org id for a given key - created_at: Optional[datetime] = None - created_by: Optional[str] = None - updated_at: Optional[datetime] = None - updated_by: Optional[str] = None - last_active: Optional[datetime] = None - object_permission_id: Optional[str] = None - object_permission: Optional[LiteLLM_ObjectPermissionTable] = None - access_group_ids: Optional[List[str]] = None - rotation_count: Optional[int] = 0 - auto_rotate: Optional[bool] = False - rotation_interval: Optional[str] = None - last_rotation_at: Optional[datetime] = None - key_rotation_at: Optional[datetime] = None - router_settings: Optional[dict] = None - budget_limits: Optional[List[dict]] = None + blocked: bool | None = None + litellm_budget_table: dict | None = None + budget_id: str | None = None + org_id: str | None = None # org id for a given key + created_at: datetime | None = None + created_by: str | None = None + updated_at: datetime | None = None + updated_by: str | None = None + last_active: datetime | None = None + object_permission_id: str | None = None + object_permission: LiteLLM_ObjectPermissionTable | None = None + access_group_ids: list[str] | None = None + rotation_count: int | None = 0 + auto_rotate: bool | None = False + rotation_interval: str | None = None + last_rotation_at: datetime | None = None + key_rotation_at: datetime | None = None + router_settings: dict | None = None + budget_limits: list[dict] | None = None model_config = ConfigDict(protected_namespaces=()) class LiteLLM_DeletedVerificationToken(LiteLLM_VerificationToken): """Audit record for deleted keys; mirrors the token plus deletion metadata.""" - id: Optional[str] = None - deleted_at: Optional[datetime] = None - deleted_by: Optional[str] = None - deleted_by_api_key: Optional[str] = None - litellm_changed_by: Optional[str] = None + id: str | None = None + deleted_at: datetime | None = None + deleted_by: str | None = None + deleted_by_api_key: str | None = None + litellm_changed_by: str | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/ocr/__init__.py b/litellm/ocr/__init__.py index e97497b2db7..a39141c0b5a 100644 --- a/litellm/ocr/__init__.py +++ b/litellm/ocr/__init__.py @@ -2,4 +2,4 @@ from .main import aocr, ocr -__all__ = ["ocr", "aocr"] +__all__ = ["aocr", "ocr"] diff --git a/litellm/passthrough/__init__.py b/litellm/passthrough/__init__.py index bfd13e7a74e..bd89d352e37 100644 --- a/litellm/passthrough/__init__.py +++ b/litellm/passthrough/__init__.py @@ -2,7 +2,7 @@ from .main import allm_passthrough_route, llm_passthrough_route from .utils import BasePassthroughUtils __all__ = [ + "BasePassthroughUtils", "allm_passthrough_route", "llm_passthrough_route", - "BasePassthroughUtils", ] diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index b3fdb31997a..7667e74256f 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -9,9 +9,7 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - List, Optional, - Union, cast, ) @@ -39,20 +37,20 @@ async def allm_passthrough_route( method: str, endpoint: str, model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - request_query_params: Optional[dict] = None, - request_headers: Optional[dict] = None, - content: Optional[Any] = None, - data: Optional[dict] = None, - files: Optional[RequestFiles] = None, - json: Optional[Any] = None, - params: Optional[QueryParamTypes] = None, - cookies: Optional[CookieTypes] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, + request_query_params: dict | None = None, + request_headers: dict | None = None, + content: Any | None = None, + data: dict | None = None, + files: RequestFiles | None = None, + json: Any | None = None, + params: QueryParamTypes | None = None, + cookies: CookieTypes | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, **kwargs, -) -> Union[httpx.Response, AsyncGenerator[Any, Any]]: +) -> httpx.Response | AsyncGenerator[Any, Any]: """ Async: Reranks a list of documents based on their relevance to the query """ @@ -164,26 +162,26 @@ def llm_passthrough_route( method: str, endpoint: str, model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - request_query_params: Optional[dict] = None, - request_headers: Optional[dict] = None, - content: Optional[Any] = None, - data: Optional[dict] = None, - files: Optional[RequestFiles] = None, - json: Optional[Any] = None, - params: Optional[QueryParamTypes] = None, - cookies: Optional[CookieTypes] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, + request_query_params: dict | None = None, + request_headers: dict | None = None, + content: Any | None = None, + data: dict | None = None, + files: RequestFiles | None = None, + json: Any | None = None, + params: QueryParamTypes | None = None, + cookies: CookieTypes | None = None, + client: HTTPHandler | AsyncHTTPHandler | None = None, **kwargs, -) -> Union[ - httpx.Response, - Coroutine[Any, Any, httpx.Response], - Coroutine[Any, Any, Union[httpx.Response, AsyncGenerator[Any, Any]]], - Generator[Any, Any, Any], - AsyncGenerator[Any, Any], -]: +) -> ( + httpx.Response + | Coroutine[Any, Any, httpx.Response] + | Coroutine[Any, Any, httpx.Response | AsyncGenerator[Any, Any]] + | Generator[Any, Any, Any] + | AsyncGenerator[Any, Any] +): """ Pass through requests to the LLM APIs. @@ -360,12 +358,12 @@ def llm_passthrough_route( async def _async_passthrough_request( - client: Union[HTTPHandler, AsyncHTTPHandler], + client: HTTPHandler | AsyncHTTPHandler, request: httpx.Request, is_streaming_request: bool, litellm_logging_obj: "LiteLLMLoggingObj", provider_config: "BasePassthroughConfig", -) -> Union[httpx.Response, AsyncGenerator[Any, Any]]: +) -> httpx.Response | AsyncGenerator[Any, Any]: """ Handle async passthrough requests. Uses async client to send request and properly handles streaming. @@ -399,7 +397,7 @@ def _sync_streaming( ): from litellm.utils import executor - raw_bytes: List[bytes] = [] + raw_bytes: list[bytes] = [] flush_scheduled = False try: for chunk in response.iter_bytes(): # type: ignore @@ -439,7 +437,7 @@ async def _async_streaming( pass raise - raw_bytes: List[bytes] = [] + raw_bytes: list[bytes] = [] flush_scheduled = False try: async for chunk in iter_response.aiter_bytes(): # type: ignore diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index 84ec89b7e2a..cc3bcb691ed 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -1,11 +1,10 @@ import sys -from typing import Optional DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS = 600.0 def resolve_pass_through_request_timeout( - endpoint_timeout: Optional[float] = None, + endpoint_timeout: float | None = None, ) -> float: """ Resolve the upstream httpx timeout for pass_through_request. @@ -31,9 +30,9 @@ def resolve_pass_through_request_timeout( def resolve_llm_passthrough_timeout( - kwargs: Optional[dict] = None, - litellm_params: Optional[dict] = None, - router_timeout: Optional[float] = None, + kwargs: dict | None = None, + litellm_params: dict | None = None, + router_timeout: float | None = None, ) -> float: """ Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse). diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index a5c03acb4f5..b7c64e2014d 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -1,5 +1,4 @@ from collections.abc import Mapping -from typing import Dict, List, Optional, Union from urllib.parse import parse_qs import httpx @@ -29,9 +28,9 @@ class BasePassthroughUtils: @staticmethod def get_merged_query_parameters( existing_url: httpx.URL, - request_query_params: Mapping[str, Union[str, list]], - default_query_params: Optional[Dict[str, Union[str, list]]] = None, - ) -> Dict[str, Union[str, List[str]]]: + request_query_params: Mapping[str, str | list], + default_query_params: dict[str, str | list] | None = None, + ) -> dict[str, str | list[str]]: # Get the existing query params from the target URL existing_query_string = existing_url.query.decode("utf-8") existing_query_params = parse_qs(existing_query_string) @@ -56,7 +55,7 @@ class BasePassthroughUtils: def forward_headers_from_request( request_headers: dict, headers: dict, - forward_headers: Optional[bool] = False, + forward_headers: bool | None = False, ): """ Helper to forward headers from original request. diff --git a/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py b/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py index 7122c64ec64..c6bc4c93009 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py +++ b/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py @@ -1,5 +1,3 @@ -from typing import Dict, List, Optional - from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser from litellm.proxy._types import UserAPIKeyAuth @@ -20,14 +18,14 @@ class MCPAuthenticatedUser(AuthenticatedUser): def __init__( self, - user_api_key_auth: Optional[UserAPIKeyAuth], - mcp_auth_header: Optional[str] = None, - mcp_servers: Optional[List[str]] = None, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, - oauth2_headers: Optional[Dict[str, str]] = None, - mcp_protocol_version: Optional[str] = None, - raw_headers: Optional[Dict[str, str]] = None, - client_ip: Optional[str] = None, + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + mcp_protocol_version: str | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ): self.user_api_key_auth = user_api_key_auth self.mcp_auth_header = mcp_auth_header diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 973544366de..e8f39daa758 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1,7 +1,7 @@ import re from collections.abc import Sequence from datetime import datetime, timezone -from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple, cast +from typing import TYPE_CHECKING, cast from fastapi import HTTPException from starlette.datastructures import Headers @@ -80,7 +80,7 @@ class UnloadableEntitlementError(Exception): it places no ceiling; denying there would refuse MCP to every caller during a cold-cache fault.""" -def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[List[str]] = None) -> Optional[List[str]]: +def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | None = None) -> list[str] | None: """Resolve the single MCP server name a cold-start passthrough bypass may target. Delegates parsing to :meth:`MCPRequestHandler._extract_target_server_names_from_path` so the @@ -114,7 +114,7 @@ def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[Li return servers -def _is_mcp_passthrough_cold_start(mcp_servers: Optional[List[str]], client_ip: Optional[str]) -> bool: +def _is_mcp_passthrough_cold_start(mcp_servers: list[str] | None, client_ip: str | None) -> bool: """True only when EVERY targeted server is a pass-through server with no auth headers — the cold-start OAuth discovery case per RFC 9728 / MCP Authorization spec. Lets the route handler's 401 emitter produce the @@ -150,8 +150,8 @@ def _is_litellm_auth_admission_error(exc: Exception) -> bool: def _has_client_supplied_mcp_auth( - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, ) -> bool: return bool(mcp_auth_header) or bool(mcp_server_auth_headers) @@ -336,13 +336,13 @@ class MCPRequestHandler: @staticmethod async def process_mcp_request( scope: Scope, - ) -> Tuple[ + ) -> tuple[ UserAPIKeyAuth, - Optional[str], - Optional[List[str]], - Optional[Dict[str, Dict[str, str]]], - Optional[Dict[str, str]], - Optional[Dict[str, str]], + str | None, + list[str] | None, + dict[str, dict[str, str]] | None, + dict[str, str] | None, + dict[str, str] | None, ]: """ Process and validate MCP request headers from the ASGI scope. @@ -579,7 +579,7 @@ class MCPRequestHandler: return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers @staticmethod - def _extract_target_server_names_from_path(path: str) -> List[str]: + def _extract_target_server_names_from_path(path: str) -> list[str]: """ Extract the target MCP server name(s) from the standard MCP transport URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and @@ -638,7 +638,7 @@ class MCPRequestHandler: @staticmethod def _target_servers_delegate_auth_to_upstream( - path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] + path: str, mcp_servers: list[str] | None, client_ip: str | None ) -> bool: """ True only when EVERY MCP server the request targets is configured for @@ -695,9 +695,7 @@ class MCPRequestHandler: return True @staticmethod - def _target_servers_are_true_passthrough( - path: str, mcp_servers: Optional[list[str]], client_ip: Optional[str] - ) -> bool: + def _target_servers_are_true_passthrough(path: str, mcp_servers: list[str] | None, client_ip: str | None) -> bool: """ True only when EVERY MCP server the request targets is ``auth_type == true_passthrough``. Fails closed when any target does not opt in or cannot be resolved. @@ -723,8 +721,8 @@ class MCPRequestHandler: @staticmethod def _single_dcr_bridge_delegate_target( - path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] - ) -> Optional[MCPServer]: + path: str, mcp_servers: list[str] | None, client_ip: str | None + ) -> MCPServer | None: """The one DCR-bridge ``oauth_delegate`` server this request targets, or ``None``. Returns the server only when EXACTLY ONE target resolves and it is both @@ -753,10 +751,10 @@ class MCPRequestHandler: async def _admit_dcr_bridge_delegate( server: MCPServer, authorization_value: str, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + mcp_server_auth_headers: dict[str, dict[str, str]] | None, request: Request, route: str, - ) -> Tuple[UserAPIKeyAuth, Optional[Dict[str, Dict[str, str]]]]: + ) -> tuple[UserAPIKeyAuth, dict[str, dict[str, str]] | None]: """Open the bridge envelope and admit the caller under the live key it references. The envelope's signature proves the user authenticated when it was minted, but @@ -1010,7 +1008,7 @@ class MCPRequestHandler: limits[source.team_id] = applicable return limits or None except Exception as e: # noqa: BLE001 # throttling metadata must never fail an allowed request - verbose_logger.warning(f"Failed to resolve per-team MCP rpm limits for admitted subject: {str(e)}") + verbose_logger.warning(f"Failed to resolve per-team MCP rpm limits for admitted subject: {e!s}") return None @staticmethod @@ -1158,7 +1156,7 @@ class MCPRequestHandler: return expiry >= datetime.now(timezone.utc) @staticmethod - def _resolve_target_server_names(path: str, mcp_servers_header: Optional[List[str]]) -> List[str]: + def _resolve_target_server_names(path: str, mcp_servers_header: list[str] | None) -> list[str]: """ Resolve the target MCP server names exactly as downstream routing does (``server.py::extract_mcp_auth_context``). @@ -1178,7 +1176,7 @@ class MCPRequestHandler: return mcp_servers_header if mcp_servers_header is not None else [] @staticmethod - def _get_mcp_auth_header_from_headers(headers: Headers) -> Optional[str]: + def _get_mcp_auth_header_from_headers(headers: Headers) -> str | None: """ Get the header passed to LiteLLM to pass to downstream MCP servers @@ -1204,7 +1202,7 @@ class MCPRequestHandler: @staticmethod def _get_mcp_server_auth_headers_from_headers( headers: Headers, - ) -> Dict[str, Dict[str, str]]: + ) -> dict[str, dict[str, str]]: """ Parse server-specific MCP auth headers from the request headers. @@ -1217,7 +1215,7 @@ class MCPRequestHandler: Returns: Dict[str, Dict[str, str]]: Mapping of server alias to header dict """ - server_auth_headers: Dict[str, Dict[str, str]] = {} + server_auth_headers: dict[str, dict[str, str]] = {} prefix = "x-mcp-" for header_name, header_value in headers.items(): @@ -1253,7 +1251,7 @@ class MCPRequestHandler: return server_auth_headers @staticmethod - def _get_oauth2_headers_from_headers(headers: Headers) -> Dict[str, str]: + def _get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]: """ Get the oauth2 headers from the request headers. """ @@ -1287,7 +1285,7 @@ class MCPRequestHandler: return MCP_CLIENT_SIDE_AUTH_HEADER_NAME @staticmethod - def get_litellm_api_key_from_headers(headers: Headers) -> Optional[str]: + def get_litellm_api_key_from_headers(headers: Headers) -> str | None: """ Get the Litellm API key from the headers using case-insensitive lookup @@ -1347,9 +1345,12 @@ class MCPRequestHandler: if not isinstance(entry, (list, tuple)) or len(entry) < 1: continue name = entry[0] - if isinstance(name, (bytes, bytearray)) and bytes(name).lower() == b"authorization": - count += 1 - elif isinstance(name, str) and name.lower() == "authorization": + if ( + isinstance(name, (bytes, bytearray)) + and bytes(name).lower() == b"authorization" + or isinstance(name, str) + and name.lower() == "authorization" + ): count += 1 if count > 1: raise HTTPException( @@ -1359,10 +1360,10 @@ class MCPRequestHandler: @staticmethod async def get_allowed_mcp_servers( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, *, keyless_source: bool = False, - ) -> List[str]: + ) -> list[str]: """ Get list of allowed MCP servers for the given user/key based on permissions. @@ -1440,7 +1441,7 @@ class MCPRequestHandler: # 2. Add the key's access-group grants on top. These are additive: # attaching a group to the key grants its servers regardless of the # team ceiling. - allowed_mcp_servers: List[str] = list(base | grants_set) + allowed_mcp_servers: list[str] = list(base | grants_set) ######################################################### # Check end_user permissions if end_user_id is set @@ -1513,9 +1514,9 @@ class MCPRequestHandler: if isinstance(e, UnloadableEntitlementError): # A ceiling we KNOW exists and cannot read. Denying is the only answer that does not # widen this caller past what an operator configured, for both caller shapes. - verbose_logger.warning(f"Denying MCP access, entitlement unreadable: {str(e)}") + verbose_logger.warning(f"Denying MCP access, entitlement unreadable: {e!s}") else: - verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers: {e!s}") return [] @staticmethod @@ -1648,7 +1649,7 @@ class MCPRequestHandler: # Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for # it alone, access only narrows) while every other source stands. Raising would collapse the # whole union to deny-all over one momentarily-unreadable row. - verbose_logger.warning(f"MCP admitted-subject source team {team_id!r} unresolvable, skipping: {str(e)}") + verbose_logger.warning(f"MCP admitted-subject source team {team_id!r} unresolvable, skipping: {e!s}") return None if team_obj is None: return None @@ -1681,10 +1682,10 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) except BudgetExceededError as e: - verbose_logger.info(f"MCP admitted-subject source team {team_id!r} over budget, not a grantor: {str(e)}") + verbose_logger.info(f"MCP admitted-subject source team {team_id!r} over budget, not a grantor: {e!s}") return None except Exception as e: # noqa: BLE001 # per-source isolation: a budget-check fault narrows, never raises - verbose_logger.warning(f"MCP budget check failed for source team {team_id!r}, skipping source: {str(e)}") + verbose_logger.warning(f"MCP budget check failed for source team {team_id!r}, skipping source: {e!s}") return None return team_obj @@ -1737,7 +1738,7 @@ class MCPRequestHandler: billed.org_id = source.org_id return billed except Exception as e: # noqa: BLE001 # attribution must never fail an authorized call - verbose_logger.warning(f"MCP billing attribution failed for {tool_name!r}, billing the user: {str(e)}") + verbose_logger.warning(f"MCP billing attribution failed for {tool_name!r}, billing the user: {e!s}") return auth @staticmethod @@ -1796,7 +1797,7 @@ class MCPRequestHandler: @staticmethod def _get_key_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ): """ Get key object_permission - already loaded by get_key_object() in main auth flow. @@ -1811,7 +1812,7 @@ class MCPRequestHandler: @staticmethod async def _get_team_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ): """ Get team object_permission - automatically loaded by get_team_object() in main auth flow. @@ -1836,7 +1837,7 @@ class MCPRequestHandler: return None # Get the team object (which has object_permission already loaded) - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( + team_obj: LiteLLM_TeamTable | None = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -1852,10 +1853,10 @@ class MCPRequestHandler: @staticmethod async def get_allowed_tools_for_server( server_id: str, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, *, keyless_source: bool = False, - ) -> Optional[List[str]]: + ) -> list[str] | None: """ Get list of allowed tool names for a specific server based on key/team permissions. Follows same inheritance logic as get_allowed_mcp_servers. @@ -1928,7 +1929,7 @@ class MCPRequestHandler: allowed_tools = team_tools else: # No team restrictions → use key restrictions - allowed_tools = cast(List[str], key_tools) + allowed_tools = cast(list[str], key_tools) allowed_tools = _as_list( await MCPRequestHandler._apply_user_tool_ceiling( @@ -1945,9 +1946,9 @@ class MCPRequestHandler: # than the None (allow-all) key auth gets for an indeterminate fault. unreadable_entitlement = isinstance(e, UnloadableEntitlementError) if unreadable_entitlement: - verbose_logger.warning(f"Denying MCP tools, entitlement unreadable: {str(e)}") + verbose_logger.warning(f"Denying MCP tools, entitlement unreadable: {e!s}") else: - verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}") + verbose_logger.warning(f"Failed to get allowed tools for server: {e!s}") # Fail CLOSED for a keyless admitted subject: ANY error must deny the server's tools ([]), # not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both # keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so @@ -1998,7 +1999,7 @@ class MCPRequestHandler: raise verbose_logger.warning( f"MCP org tool ceiling unresolvable for org_id={user_api_key_auth.org_id!r}; " - f"skipping org intersect, key/team/agent restrictions stand: {str(e)}" + f"skipping org intersect, key/team/agent restrictions stand: {e!s}" ) return allowed_tools org_tools = ( @@ -2029,7 +2030,7 @@ class MCPRequestHandler: async def is_tool_allowed_for_server( tool_name: str, server_id: str, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> bool: """ Check if a specific tool is allowed for a server based on key/team permissions. @@ -2050,7 +2051,7 @@ class MCPRequestHandler: @staticmethod def is_tool_allowed( - allowed_mcp_servers: List[str], + allowed_mcp_servers: list[str], server_name: str, ) -> bool: """ @@ -2064,8 +2065,8 @@ class MCPRequestHandler: @staticmethod async def _get_key_access_group_mcp_server_extras( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Resolve the key's unified `access_group_ids` (LiteLLM_AccessGroupTable) to MCP server IDs as additive grants: a group attached to the key extends the @@ -2101,13 +2102,13 @@ class MCPRequestHandler: # Permission entries may be server_ids OR names/aliases — expand to ids. return global_mcp_server_manager.expand_permission_list(raw_server_ids) except Exception as e: - verbose_logger.warning(f"Failed to get key access group MCP server grants: {str(e)}") + verbose_logger.warning(f"Failed to get key access group MCP server grants: {e!s}") return [] @staticmethod async def _get_allowed_mcp_servers_for_key( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get the key's own MCP ceiling from its object_permission (mcp_servers, tag-style mcp_access_groups, mcp_tool_permissions). @@ -2179,7 +2180,7 @@ class MCPRequestHandler: all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers return list(set(all_servers)) except Exception as e: - verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for key: {e!s}") return [] @staticmethod @@ -2237,7 +2238,7 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises - verbose_logger.warning(f"Failed to resolve user teams for MCP grant: {str(e)}") + verbose_logger.warning(f"Failed to resolve user teams for MCP grant: {e!s}") return [] if user_object is None or not user_object.teams: return [] @@ -2322,7 +2323,7 @@ class MCPRequestHandler: servers = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers) return list(servers) except Exception as e: - verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for team: {e!s}") return [] @staticmethod @@ -2360,7 +2361,7 @@ class MCPRequestHandler: @staticmethod async def _get_org_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """ Get org object_permission via the established ``get_org_object`` / @@ -2419,7 +2420,7 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_org( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> list[str] | None: """ Get allowed MCP servers for an organization. @@ -2461,7 +2462,7 @@ class MCPRequestHandler: # A NAMED-but-unreadable ceiling is a stronger fact than "unresolved" and denies everywhere. if isinstance(e, UnloadableEntitlementError): raise - verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for org: {e!s}") return None @staticmethod @@ -2489,7 +2490,7 @@ class MCPRequestHandler: route="/mcp", ) except Exception as e: # noqa: BLE001 # entitlement unknown, not known-absent: no ceiling, as before this level - verbose_logger.warning(f"Failed to resolve end_user for MCP permissions: {str(e)}") + verbose_logger.warning(f"Failed to resolve end_user for MCP permissions: {e!s}") return None if end_user_obj is None: @@ -2509,8 +2510,8 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_end_user( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get allowed MCP servers for an end user. @@ -2553,7 +2554,7 @@ class MCPRequestHandler: all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers return list(set(all_servers)) except Exception as e: - verbose_logger.warning(f"Failed to get allowed MCP servers for end_user: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for end_user: {e!s}") return [] @staticmethod @@ -2636,7 +2637,7 @@ class MCPRequestHandler: ) return object_permission_id except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before - verbose_logger.warning(f"MCP user entitlement: link for {user_id!r} unresolved, no ceiling: {str(e)}") + verbose_logger.warning(f"MCP user entitlement: link for {user_id!r} unresolved, no ceiling: {e!s}") return None @staticmethod @@ -2668,7 +2669,7 @@ class MCPRequestHandler: ) return list(set(direct_mcp_servers + access_group_servers + tool_perm_servers)) except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling" - verbose_logger.warning(f"Failed to get allowed MCP servers for user: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for user: {e!s}") return None @staticmethod @@ -2738,7 +2739,7 @@ class MCPRequestHandler: try: object_permissions = await MCPRequestHandler._get_user_object_permission(user_api_key_auth) except Exception as e: # noqa: BLE001 # an unresolved human entitlement must deny, not widen - verbose_logger.warning(f"MCP user tool ceiling unresolvable, denying tools on {server_id!r}: {str(e)}") + verbose_logger.warning(f"MCP user tool ceiling unresolvable, denying tools on {server_id!r}: {e!s}") return [] if object_permissions is None or not object_permissions.mcp_tool_permissions: @@ -2784,12 +2785,12 @@ class MCPRequestHandler: ) return object_permission_id except Exception as e: # noqa: BLE001 # entitlement unknown, not known-absent: no ceiling, as before this level - verbose_logger.warning(f"Failed to resolve object_permission_id for agent {agent_id!r}: {str(e)}") + verbose_logger.warning(f"Failed to resolve object_permission_id for agent {agent_id!r}: {e!s}") return None @staticmethod async def _get_agent_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """ Get agent object_permission via the established ``get_object_permission`` @@ -2824,9 +2825,9 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_agent( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, agent_object_permission=None, - ) -> List[str]: + ) -> list[str]: """ Get allowed MCP servers for an agent (from the agent's object_permission). @@ -2868,15 +2869,15 @@ class MCPRequestHandler: all_servers = expanded_direct_servers + access_group_servers return list(set(all_servers)) except Exception as e: - verbose_logger.warning(f"Failed to get allowed MCP servers for agent: {str(e)}") + verbose_logger.warning(f"Failed to get allowed MCP servers for agent: {e!s}") return [] @staticmethod async def _get_agent_tool_permissions_for_server( server_id: str, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, agent_object_permission=None, - ) -> Optional[List[str]]: + ) -> list[str] | None: """ Get allowed tool names for a server from the agent's object_permission. Returns None if agent has no tool restrictions for this server. An entitlement the agent @@ -2910,15 +2911,15 @@ class MCPRequestHandler: tools = global_mcp_server_manager.expand_tool_permissions(mcp_tool_permissions).get(server_id) return list(tools) if tools else None except Exception as e: - verbose_logger.warning(f"Failed to get agent tool permissions for server: {str(e)}") + verbose_logger.warning(f"Failed to get agent tool permissions for server: {e!s}") return None @staticmethod - def _get_config_server_ids_for_access_groups(config_mcp_servers, access_groups: List[str]) -> Set[str]: + def _get_config_server_ids_for_access_groups(config_mcp_servers, access_groups: list[str]) -> set[str]: """ Helper to get server_ids from config-loaded servers that match any of the given access groups. """ - server_ids: Set[str] = set() + server_ids: set[str] = set() for server_id, server in config_mcp_servers.items(): if server.access_groups: if any(group in server.access_groups for group in access_groups): @@ -2926,11 +2927,11 @@ class MCPRequestHandler: return server_ids @staticmethod - async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: List[str]) -> Set[str]: + async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]: """ Helper to get server_ids from DB servers that match any of the given access groups. """ - server_ids: Set[str] = set() + server_ids: set[str] = set() if access_groups and prisma_client is not None: try: mcp_servers = await MCPServerRepository(prisma_client).table.find_many( @@ -2944,8 +2945,8 @@ class MCPRequestHandler: @staticmethod async def _get_mcp_servers_from_access_groups( - access_groups: List[str], - ) -> List[str]: + access_groups: list[str], + ) -> list[str]: """ Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers """ @@ -2968,17 +2969,17 @@ class MCPRequestHandler: return list(server_ids) except Exception as e: - verbose_logger.warning(f"Failed to get MCP servers from access groups: {str(e)}") + verbose_logger.warning(f"Failed to get MCP servers from access groups: {e!s}") return [] @staticmethod async def get_mcp_access_groups( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get list of MCP access groups for the given user/key based on permissions """ - access_groups: List[str] = [] + access_groups: list[str] = [] access_groups_for_key = await MCPRequestHandler._get_mcp_access_groups_for_key(user_api_key_auth) access_groups_for_team = await MCPRequestHandler._get_mcp_access_groups_for_team(user_api_key_auth) @@ -2996,8 +2997,8 @@ class MCPRequestHandler: @staticmethod async def _get_mcp_access_groups_for_key( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: from litellm.proxy.auth.auth_checks import get_object_permission from litellm.proxy.proxy_server import ( prisma_client, @@ -3028,13 +3029,13 @@ class MCPRequestHandler: return key_object_permission.mcp_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get MCP access groups for key: {str(e)}") + verbose_logger.warning(f"Failed to get MCP access groups for key: {e!s}") return [] @staticmethod async def _get_mcp_access_groups_for_team( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get MCP access groups for the team """ @@ -3059,7 +3060,7 @@ class MCPRequestHandler: return [] try: - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( + team_obj: LiteLLM_TeamTable | None = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -3076,11 +3077,11 @@ class MCPRequestHandler: return object_permissions.mcp_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get MCP access groups for team: {str(e)}") + verbose_logger.warning(f"Failed to get MCP access groups for team: {e!s}") return [] @staticmethod - def get_mcp_access_groups_from_headers(headers: Headers) -> Optional[List[str]]: + def get_mcp_access_groups_from_headers(headers: Headers) -> list[str] | None: """ Extract and parse the x-mcp-access-groups header as a list of strings. """ @@ -3093,7 +3094,7 @@ class MCPRequestHandler: return None @staticmethod - def get_mcp_access_groups_from_scope(scope: Scope) -> Optional[List[str]]: + def get_mcp_access_groups_from_scope(scope: Scope) -> list[str] | None: """ Extract and parse the x-mcp-access-groups header from an ASGI scope. """ diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 19048e2eb7c..6a2f42658a3 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -3,7 +3,7 @@ import math from dataclasses import dataclass from datetime import datetime, timezone -from typing import TYPE_CHECKING, Literal, Optional +from typing import TYPE_CHECKING, Literal from fastapi import HTTPException, Request from fastapi.responses import JSONResponse @@ -25,7 +25,7 @@ if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth -def _litellm_key_from_request(request: Request) -> Optional[str]: +def _litellm_key_from_request(request: Request) -> str | None: """Return the LiteLLM API key presented on the request, or ``None``. Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index 4f58f4bdbb3..1414387d6d6 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -18,7 +18,7 @@ import hashlib import html as _html_module import time import uuid -from typing import Dict, Optional, cast +from typing import cast from urllib.parse import urlencode import jwt @@ -40,7 +40,7 @@ from litellm.proxy._types import UserAPIKeyAuth # In-memory store for pending authorization codes. # Each entry: {code: {api_key, server_id, code_challenge, redirect_uri, user_id, expires_at}} # --------------------------------------------------------------------------- -_byok_auth_codes: Dict[str, dict] = {} +_byok_auth_codes: dict[str, dict] = {} # Authorization codes expire after 5 minutes. _AUTH_CODE_TTL_SECONDS = 300 @@ -82,7 +82,7 @@ def _oauth_token_error(code: str, status: int = 400) -> JSONResponse: return JSONResponse(status_code=status, content={"error": code}, headers=TOKEN_NO_CACHE_HEADERS) -def _user_id_from_session_cookie(request: Request) -> Optional[str]: +def _user_id_from_session_cookie(request: Request) -> str | None: """Return user_id from the UI ``token`` cookie (HS256-signed with ``master_key``), or None if missing/invalid. @@ -632,13 +632,13 @@ async def oauth_protected_resource_metadata(request: Request) -> JSONResponse: @router.get("/v1/mcp/oauth/authorize", include_in_schema=False) async def byok_authorize_get( request: Request, - client_id: Optional[str] = None, - redirect_uri: Optional[str] = None, - response_type: Optional[str] = None, - code_challenge: Optional[str] = None, - code_challenge_method: Optional[str] = None, - state: Optional[str] = None, - server_id: Optional[str] = None, + client_id: str | None = None, + redirect_uri: str | None = None, + response_type: str | None = None, + code_challenge: str | None = None, + code_challenge_method: str | None = None, + state: str | None = None, + server_id: str | None = None, ) -> HTMLResponse: """ Show the BYOK API-key entry form. diff --git a/litellm/proxy/_experimental/mcp_server/cost_calculator.py b/litellm/proxy/_experimental/mcp_server/cost_calculator.py index 9b6f89bc7bd..43d12756e73 100644 --- a/litellm/proxy/_experimental/mcp_server/cost_calculator.py +++ b/litellm/proxy/_experimental/mcp_server/cost_calculator.py @@ -2,7 +2,7 @@ Cost calculator for MCP tools. """ -from typing import TYPE_CHECKING, Any, Optional, cast +from typing import TYPE_CHECKING, Any, cast from litellm.types.mcp import MCPServerCostInfo from litellm.types.utils import StandardLoggingMCPToolCall @@ -18,7 +18,7 @@ else: class MCPCostCalculator: @staticmethod def calculate_mcp_tool_call_cost( - litellm_logging_obj: Optional[LitellmLoggingObject], + litellm_logging_obj: LitellmLoggingObject | None, ) -> float: """ Calculate the cost of an MCP tool call. diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 3221f3b8dd4..3c8e7d9f1ef 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -570,9 +570,7 @@ async def get_all_mcp_servers( decrypt_global_env_var_values(table.env_vars) return tables except Exception as e: - verbose_proxy_logger.debug( - "litellm.proxy._experimental.mcp_server.db.py::get_all_mcp_servers - {}".format(str(e)) - ) + verbose_proxy_logger.debug(f"litellm.proxy._experimental.mcp_server.db.py::get_all_mcp_servers - {e!s}") return [] @@ -721,14 +719,12 @@ async def delete_mcp_server_from_team(prisma_client: PrismaClient, server_id: st """ Remove the mcp server from the team """ - pass async def delete_mcp_server_from_virtualkey(): """ Remove the mcp server from the virtual key """ - pass async def delete_mcp_server( diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 865787d5a07..e16fb0d0e00 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -5,7 +5,7 @@ import secrets import time from collections.abc import Mapping from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal, Optional from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx @@ -23,7 +23,6 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, normalize_token_endpoint_auth_method, ) -from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _bridge_mint_error_response, _BridgeMintReady, @@ -66,7 +65,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.mcp import MCPAuth, MCPCredentials -from litellm.types.mcp_server.mcp_server_manager import MCPServer +from litellm.types.mcp_server.mcp_server_manager import MCPServer, MCPTokenEndpointAuthMethod if TYPE_CHECKING: from litellm.proxy._types import LiteLLM_MCPServerTable @@ -76,20 +75,20 @@ if TYPE_CHECKING: # Keyed by (server_id, resource_url) → (expires_at_epoch, payload). # A payload of ``None`` is a negative-result entry that prevents repeated # upstream fetches when the IdP consistently has no metadata to serve. -_OAUTH_METADATA_CACHE: Dict[Tuple[str, str], Tuple[float, Optional[dict]]] = {} +_OAUTH_METADATA_CACHE: dict[tuple[str, str], tuple[float, dict | None]] = {} _OAUTH_METADATA_CACHE_TTL_SECONDS = 300 _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS = 60 _OAUTH_METADATA_CACHE_MAX_SIZE = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. -_OAUTH_METADATA_FETCH_LOCKS: Dict[Tuple[str, str], asyncio.Lock] = {} +_OAUTH_METADATA_FETCH_LOCKS: dict[tuple[str, str], asyncio.Lock] = {} router = APIRouter( tags=["mcp"], ) -def _prune_oauth_metadata_cache(now: Optional[float] = None) -> None: +def _prune_oauth_metadata_cache(now: float | None = None) -> None: now = now if now is not None else time.time() expired_cache_keys = [ cache_key for cache_key, (expires_at, _payload) in _OAUTH_METADATA_CACHE.items() if expires_at <= now @@ -120,9 +119,9 @@ def _prune_oauth_metadata_cache(now: Optional[float] = None) -> None: def encode_state_with_base_url( base_url: str, original_state: str, - code_challenge: Optional[str] = None, - code_challenge_method: Optional[str] = None, - client_redirect_uri: Optional[str] = None, + code_challenge: str | None = None, + code_challenge_method: str | None = None, + client_redirect_uri: str | None = None, litellm_user_id: str | None = None, mcp_server_id: str | None = None, dcr_client_id: str | None = None, @@ -421,7 +420,7 @@ def _clear_oauth_state_cookie(response: Response, request: Request, state: str) ) -def _get_validated_client_redirect_uri(request: Request, state_data: Dict[str, Any]) -> str: +def _get_validated_client_redirect_uri(request: Request, state_data: dict[str, Any]) -> str: """Return a trusted (same-origin, loopback, or ops-allowlisted) client redirect URI from OAuth state. """ @@ -432,7 +431,7 @@ def _get_validated_client_redirect_uri(request: Request, state_data: Dict[str, A return redirect_uri -def _append_query_params(url: str, params: Dict[str, str]) -> str: +def _append_query_params(url: str, params: dict[str, str]) -> str: parsed = urlparse(url) query_params = parse_qsl(parsed.query, keep_blank_values=True) query_params.extend(params.items()) @@ -440,8 +439,8 @@ def _append_query_params(url: str, params: Dict[str, str]) -> str: def _resolve_oauth2_server_for_root_endpoints( - client_ip: Optional[str] = None, -) -> Optional[MCPServer]: + client_ip: str | None = None, +) -> MCPServer | None: """ Resolve the MCP server for root-level OAuth endpoints (no server name in path). @@ -472,8 +471,8 @@ def _normalize_for_token_comparison(value: Any) -> str: def _validate_token_response( - token_response: Dict[str, Any], - validation_rules: Dict[str, Any], + token_response: dict[str, Any], + validation_rules: dict[str, Any], server_id: str, ) -> None: """Raise HTTPException 403 if any validation rule doesn't match the token response. @@ -524,7 +523,7 @@ def _validate_token_response( async def _store_per_user_token_server_side( server: MCPServer, user_id: str, - token_response: Dict[str, Any], + token_response: dict[str, Any], ) -> None: """Persist the OAuth token server-side and warm the Redis cache. @@ -538,19 +537,19 @@ async def _store_per_user_token_server_side( ) from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 - access_token: Optional[str] = token_response.get("access_token") + access_token: str | None = token_response.get("access_token") if not access_token: return raw_expires = token_response.get("expires_in") try: - expires_in: Optional[int] = int(raw_expires) if raw_expires is not None else None + expires_in: int | None = int(raw_expires) if raw_expires is not None else None except (TypeError, ValueError): expires_in = None - refresh_token: Optional[str] = token_response.get("refresh_token") or None + refresh_token: str | None = token_response.get("refresh_token") or None raw_scope = token_response.get("scope") - scopes: Optional[list] = raw_scope.split() if isinstance(raw_scope, str) and raw_scope else None + scopes: list | None = raw_scope.split() if isinstance(raw_scope, str) and raw_scope else None try: prisma_client = get_prisma_client_or_throw("Database not connected. Cannot store per-user OAuth token.") @@ -657,8 +656,8 @@ def _endpoint_not_configured_detail( def _raise_unless_oauth2_discovery_server( - mcp_server: Optional[MCPServer], - mcp_server_name: Optional[str], + mcp_server: MCPServer | None, + mcp_server_name: str | None, description: str, ) -> None: """404 a NAMED discovery request unless it resolves to an oauth2 or DCR-bridge server. @@ -694,9 +693,9 @@ def _dcr_bridge_relays_client_registration(mcp_server: MCPServer) -> bool: def _require_s256_pkce( - code_challenge: Optional[str], - code_challenge_method: Optional[str], -) -> Tuple[str, str]: + code_challenge: str | None, + code_challenge_method: str | None, +) -> tuple[str, str]: """DCR-bridge servers serve unauthenticated public OAuth clients, so the PKCE downgrade paths (no challenge, or a non-S256 method; RFC 7636 defaults a missing method to ``plain``) are rejected at the gateway instead of relying on upstream enforcement. Returns the @@ -720,8 +719,8 @@ def _redirect_to_upstream_authorize( state: str, code_challenge: str, code_challenge_method: str, - response_type: Optional[str], - scope: Optional[str], + response_type: str | None, + scope: str | None, ) -> RedirectResponse: """The bridge relay arm's authorize redirect: every client-supplied parameter passes through to the upstream authorize endpoint verbatim, no relay state cookie is set, and the upstream @@ -749,10 +748,10 @@ async def authorize_with_server( client_id: str, redirect_uri: str, state: str = "", - code_challenge: Optional[str] = None, - code_challenge_method: Optional[str] = None, - response_type: Optional[str] = None, - scope: Optional[str] = None, + code_challenge: str | None = None, + code_challenge_method: str | None = None, + response_type: str | None = None, + scope: str | None = None, ephemeral_dcr_client: "EphemeralDcrClient | None" = None, ): _raise_if_not_oauth2(mcp_server) @@ -869,13 +868,13 @@ async def exchange_token_with_server( request: Request, mcp_server: MCPServer, grant_type: str, - code: Optional[str], - redirect_uri: Optional[str], + code: str | None, + redirect_uri: str | None, client_id: str, - client_secret: Optional[str], - code_verifier: Optional[str], - refresh_token: Optional[str] = None, - scope: Optional[str] = None, + client_secret: str | None, + code_verifier: str | None, + refresh_token: str | None = None, + scope: str | None = None, client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, ): _raise_if_not_oauth2(mcp_server) @@ -1111,15 +1110,15 @@ class _DcrClientRegistration(BaseModel): must persist to authenticate later token-endpoint calls. Extra members are ignored.""" client_id: str - client_secret: Optional[str] = None - token_endpoint_auth_method: Optional[str] = None + client_secret: str | None = None + token_endpoint_auth_method: str | None = None class _PersistedDcrCredentials(BaseModel): - client_id: Optional[str] = None - client_secret: Optional[str] = None - token_endpoint_auth_method: Optional[str] = None - redirect_uris: Optional[list[str]] = None + client_id: str | None = None + client_secret: str | None = None + token_endpoint_auth_method: str | None = None + redirect_uris: list[str] | None = None def _redirect_uri_not_registered(credentials: _PersistedDcrCredentials, current_redirect_uri: str) -> bool: @@ -1137,7 +1136,7 @@ def _redirect_uri_not_registered(credentials: _PersistedDcrCredentials, current_ return current_redirect_uri not in recorded -def _get_persisted_dcr_credentials(credentials: object) -> Optional[_PersistedDcrCredentials]: +def _get_persisted_dcr_credentials(credentials: object) -> _PersistedDcrCredentials | None: if not credentials: return None try: @@ -1150,7 +1149,7 @@ def _get_persisted_dcr_credentials(credentials: object) -> Optional[_PersistedDc return None -def _decrypt_persisted_dcr_credential(value: Optional[str], key: str) -> Optional[str]: +def _decrypt_persisted_dcr_credential(value: str | None, key: str) -> str | None: if value is None: return None return decrypt_value_helper( @@ -1253,7 +1252,7 @@ async def _resolve_persisted_dcr_client( async def _reuse_persisted_dcr_client_if_available( - mcp_server: MCPServer, current_redirect_uri: Optional[str] = None + mcp_server: MCPServer, current_redirect_uri: str | None = None ) -> bool: persisted_mcp_server, credentials = await _resolve_persisted_dcr_client(mcp_server) if credentials is None: @@ -1567,10 +1566,10 @@ async def register_client_with_server( request: Request, mcp_server: MCPServer, client_name: str, - grant_types: Optional[list], - response_types: Optional[list], - token_endpoint_auth_method: Optional[str], - fallback_client_id: Optional[str] = None, + grant_types: list | None, + response_types: list | None, + token_endpoint_auth_method: str | None, + fallback_client_id: str | None = None, persist_credentials: bool = False, client_redirect_uris: list[str] | None = None, ): @@ -1649,13 +1648,13 @@ async def register_client_with_server( async def authorize( request: Request, redirect_uri: str, - client_id: Optional[str] = None, + client_id: str | None = None, state: str = "", - mcp_server_name: Optional[str] = None, - code_challenge: Optional[str] = None, - code_challenge_method: Optional[str] = None, - response_type: Optional[str] = None, - scope: Optional[str] = None, + mcp_server_name: str | None = None, + code_challenge: str | None = None, + code_challenge_method: str | None = None, + response_type: str | None = None, + scope: str | None = None, ): # Redirect to real OAuth provider with PKCE support from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1674,7 +1673,7 @@ async def authorize( session_user_id=_session_cookie_user_id(request), ) - lookup_name: Optional[str] = mcp_server_name or client_id + lookup_name: str | None = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = ( global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) if lookup_name else None @@ -1718,11 +1717,11 @@ async def token_endpoint( code: str = Form(None), redirect_uri: str = Form(None), client_id: str = Form(...), - client_secret: Optional[str] = Form(None), + client_secret: str | None = Form(None), code_verifier: str = Form(None), - refresh_token: Optional[str] = Form(None), - scope: Optional[str] = Form(None), - mcp_server_name: Optional[str] = None, + refresh_token: str | None = Form(None), + scope: str | None = Form(None), + mcp_server_name: str | None = None, ): """ Accept the authorization code from client and exchange it for OAuth token. @@ -1805,7 +1804,7 @@ async def authorize_complete(request: Request, flow: str = Form(...), delivery: # which strands the MCP client waiting on the loopback (see LIT-2750). -def _render_oauth_error_html(error: str, description: Optional[str]) -> HTMLResponse: +def _render_oauth_error_html(error: str, description: str | None) -> HTMLResponse: """Render an actionable HTML page for an IdP-reported OAuth error. Used when we cannot propagate the error back to the registered @@ -1830,11 +1829,11 @@ def _render_oauth_error_html(error: str, description: Optional[str]) -> HTMLResp @router.get("/callback") async def callback( request: Request, - code: Optional[str] = None, - state: Optional[str] = None, - error: Optional[str] = None, - error_description: Optional[str] = None, - error_uri: Optional[str] = None, + code: str | None = None, + state: str | None = None, + error: str | None = None, + error_description: str | None = None, + error_uri: str | None = None, ): """OAuth 2.0 authorization response handler for MCP loopback clients. @@ -1871,7 +1870,7 @@ async def callback( _clear_oauth_state_cookie(response, request, state) return response - params: Dict[str, str] = {"error": error} + params: dict[str, str] = {"error": error} if error_description: params["error_description"] = error_description if error_uri: @@ -1974,7 +1973,7 @@ async def callback( async def fetch_upstream_oauth_protected_resource( mcp_server: MCPServer, -) -> Optional[dict]: +) -> dict | None: """Fetch the upstream MCP server's ``.well-known/oauth-protected-resource`` metadata for a pass-through server. @@ -2081,7 +2080,7 @@ def is_network_error(exc: Exception) -> bool: async def _build_oauth_protected_resource_response( request: Request, - mcp_server_name: Optional[str], + mcp_server_name: str | None, use_standard_pattern: bool, ) -> dict: """ @@ -2129,7 +2128,7 @@ async def _build_oauth_protected_resource_response( if resolved: mcp_server_name = resolved.server_name or resolved.name - mcp_server: Optional[MCPServer] = None + mcp_server: MCPServer | None = None if mcp_server_name: mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip) @@ -2212,7 +2211,7 @@ async def _build_oauth_protected_resource_response( } -def _obo_protected_resource_response(mcp_server: Optional[MCPServer], resource_url: str) -> Optional[dict]: +def _obo_protected_resource_response(mcp_server: MCPServer | None, resource_url: str) -> dict | None: """The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured. The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery @@ -2360,7 +2359,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam # Kept for backward compatibility with existing deployments @router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/{{mcp_server_name}}/mcp") @router.get("/.well-known/oauth-protected-resource") -async def oauth_protected_resource_mcp(request: Request, mcp_server_name: Optional[str] = None): +async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None): """ OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern. @@ -2379,7 +2378,7 @@ async def oauth_protected_resource_mcp(request: Request, mcp_server_name: Option def _build_oauth_authorization_server_response( request: Request, - mcp_server_name: Optional[str], + mcp_server_name: str | None, ) -> dict: """Build OAuth authorization server metadata response (gateway-as-AS shape). @@ -2405,7 +2404,7 @@ def _build_oauth_authorization_server_response( ) token_endpoint = f"{request_base_url}/{mcp_server_name}/token" if mcp_server_name else f"{request_base_url}/token" - mcp_server: Optional[MCPServer] = None + mcp_server: MCPServer | None = None if mcp_server_name: mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip) @@ -2445,7 +2444,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n # LiteLLM legacy pattern and root endpoint @router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}") @router.get("/.well-known/oauth-authorization-server") -async def oauth_authorization_server_mcp(request: Request, mcp_server_name: Optional[str] = None): +async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None): """ OAuth authorization server discovery endpoint. @@ -2530,7 +2529,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s @router.post("/{mcp_server_name}/register") @router.post("/register") -async def register_client(request: Request, mcp_server_name: Optional[str] = None): +async def register_client(request: Request, mcp_server_name: str | None = None): from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) diff --git a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py index 030f4dfeca6..9927afa20d0 100644 --- a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py +++ b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py @@ -9,7 +9,8 @@ MCP Spec Reference: https://modelcontextprotocol.io/specification/2025-11-25/client/elicitation """ -from typing import Any, Optional, Union +from typing import Any, Union + from litellm._logging import verbose_logger # Guard imports that require the mcp package @@ -30,8 +31,8 @@ except ImportError: async def handle_elicitation_request( context: Any, params: "ElicitRequestParams", - downstream_session: Optional[Any] = None, - downstream_capabilities: Optional[Any] = None, + downstream_session: Any | None = None, + downstream_capabilities: Any | None = None, ) -> Union["ElicitResult", "ErrorData"]: """ Handle an MCP elicitation/create request from an upstream MCP server. @@ -78,14 +79,14 @@ async def handle_elicitation_request( verbose_logger.exception("MCP elicitation handler failed: %s", e) return ErrorData( code=-1, - message=f"Elicitation failed: {str(e)}", + message=f"Elicitation failed: {e!s}", ) async def _relay_elicitation_to_downstream( params: "ElicitRequestParams", downstream_session: Any, - downstream_capabilities: Optional[Any] = None, + downstream_capabilities: Any | None = None, ) -> Union["ElicitResult", "ErrorData"]: """ Relay an elicitation request to the downstream MCP client. diff --git a/litellm/proxy/_experimental/mcp_server/exceptions.py b/litellm/proxy/_experimental/mcp_server/exceptions.py index ca2261139c9..eb2f3bdde74 100644 --- a/litellm/proxy/_experimental/mcp_server/exceptions.py +++ b/litellm/proxy/_experimental/mcp_server/exceptions.py @@ -1,7 +1,5 @@ """Exceptions raised by the LiteLLM MCP proxy.""" -from typing import Optional - from fastapi import HTTPException @@ -21,7 +19,7 @@ class MCPUpstreamAuthError(Exception): def __init__( self, status_code: int, - www_authenticate: Optional[str], + www_authenticate: str | None, server_name: str, ) -> None: self.status_code = status_code @@ -31,8 +29,8 @@ class MCPUpstreamAuthError(Exception): def to_http_exception( self, - base_url: Optional[str] = None, - request_path: Optional[str] = None, + base_url: str | None = None, + request_path: str | None = None, ) -> HTTPException: """Convert this upstream-auth error into an ``HTTPException`` that preserves the upstream status code and any ``WWW-Authenticate`` @@ -59,7 +57,7 @@ class MCPUpstreamAuthError(Exception): the client originally targeted, matching the path-aware behaviour of ``get_passthrough_resource_metadata_url`` in ``oauth_utils.py``. """ - challenge: Optional[str] = self.www_authenticate + challenge: str | None = self.www_authenticate if challenge is None and self.status_code == 401 and base_url: prefix = base_url.rstrip("/") if request_path and request_path.startswith(f"/{self.server_name}/mcp"): diff --git a/litellm/proxy/_experimental/mcp_server/guardrail_translation/__init__.py b/litellm/proxy/_experimental/mcp_server/guardrail_translation/__init__.py index e0fd610e678..12f0df6545b 100644 --- a/litellm/proxy/_experimental/mcp_server/guardrail_translation/__init__.py +++ b/litellm/proxy/_experimental/mcp_server/guardrail_translation/__init__.py @@ -13,4 +13,4 @@ guardrail_translation_mappings = { CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler, } -__all__ = ["guardrail_translation_mappings", "MCPGuardrailTranslationHandler"] +__all__ = ["MCPGuardrailTranslationHandler", "guardrail_translation_mappings"] diff --git a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py index 909925da00a..3b6de6dfc3e 100644 --- a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py +++ b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py @@ -11,7 +11,7 @@ when you have a full MCP Tool from list_tools. Here we only have the call payload (name + arguments) so we just build the tool_call. """ -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any from fastapi import HTTPException from mcp.types import Tool as MCPTool @@ -46,10 +46,10 @@ class MCPGuardrailTranslationHandler(BaseTranslation): async def process_input_messages( self, - data: Dict[str, Any], + data: dict[str, Any], guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - ) -> Dict[str, Any]: + litellm_logging_obj: Any | None = None, + ) -> dict[str, Any]: mcp_tool_name = data.get("mcp_tool_name") or data.get("name") mcp_arguments = data.get("mcp_arguments") or data.get("arguments") mcp_tool_description = data.get("mcp_tool_description") or data.get("description") @@ -99,9 +99,9 @@ class MCPGuardrailTranslationHandler(BaseTranslation): self, response: "CallToolResult", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + litellm_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, ) -> Any: """Scan the text content of an MCP tool result and write masked text back. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 8a85c0c516b..42f90e5a530 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -6,18 +6,17 @@ mcp_server_manager.py and server.py. """ from contextvars import ContextVar -from typing import Optional # Set server-side in proxy_server.py route handlers when a request arrives via # /toolset/{name}/mcp or the toolset fallback in dynamic_mcp_route. # Never populated from client-supplied headers. -_mcp_active_toolset_id: ContextVar[Optional[str]] = ContextVar("_mcp_active_toolset_id", default=None) +_mcp_active_toolset_id: ContextVar[str | None] = ContextVar("_mcp_active_toolset_id", default=None) # Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers. -_mcp_gateway_initialize_instructions: ContextVar[Optional[str]] = ContextVar( +_mcp_gateway_initialize_instructions: ContextVar[str | None] = ContextVar( "_mcp_gateway_initialize_instructions", default=None ) # Per-request scoped server name; set in MCP HTTP/SSE handlers when the path # identifies exactly one upstream server. Never populated from client-supplied headers. -_mcp_gateway_server_name: ContextVar[Optional[str]] = ContextVar("_mcp_gateway_server_name", default=None) +_mcp_gateway_server_name: ContextVar[str | None] = ContextVar("_mcp_gateway_server_name", default=None) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index 42e2b17d697..fe093760ab4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -85,7 +85,7 @@ Usage with curl:: http://localhost:4000/mcp/atlassian_mcp """ -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING from starlette.types import Message, Send @@ -125,14 +125,14 @@ class MCPDebug: ) @staticmethod - def _mask(value: Optional[str]) -> str: + def _mask(value: str | None) -> str: """Mask a single value for safe display in headers.""" if not value: return "(none)" return MCPDebug._masker._mask_value(value) @staticmethod - def is_debug_enabled(headers: Dict[str, str]) -> bool: + def is_debug_enabled(headers: dict[str, str]) -> bool: """ Check if the client opted into MCP debug mode. @@ -147,9 +147,9 @@ class MCPDebug: @staticmethod def resolve_auth_resolution( server: "MCPServer", - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], - oauth2_headers: Optional[Dict[str, str]], + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + oauth2_headers: dict[str, str] | None, ) -> str: """ Determine which auth priority will be used for the outbound MCP call. @@ -178,13 +178,13 @@ class MCPDebug: @staticmethod def build_debug_headers( *, - inbound_headers: Dict[str, str], - oauth2_headers: Optional[Dict[str, str]], - litellm_api_key: Optional[str], + inbound_headers: dict[str, str], + oauth2_headers: dict[str, str] | None, + litellm_api_key: str | None, auth_resolution: str, - server_url: Optional[str], - server_auth_type: Optional[str], - ) -> Dict[str, str]: + server_url: str | None, + server_auth_type: str | None, + ) -> dict[str, str]: """ Build masked debug response headers. @@ -209,7 +209,7 @@ class MCPDebug: dict Headers to include in the response (all values masked). """ - debug: Dict[str, str] = {} + debug: dict[str, str] = {} # --- Inbound auth summary --- inbound_parts = [] @@ -244,7 +244,7 @@ class MCPDebug: return debug @staticmethod - def wrap_send_with_debug_headers(send: Send, debug_headers: Dict[str, str]) -> Send: + def wrap_send_with_debug_headers(send: Send, debug_headers: dict[str, str]) -> Send: """ Return a new ASGI ``send`` callable that injects *debug_headers* into the ``http.response.start`` message. @@ -263,14 +263,14 @@ class MCPDebug: @staticmethod def maybe_build_debug_headers( *, - raw_headers: Optional[Dict[str, str]], - scope: Dict, - mcp_servers: Optional[List[str]], - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], - oauth2_headers: Optional[Dict[str, str]], - client_ip: Optional[str], - ) -> Dict[str, str]: + raw_headers: dict[str, str] | None, + scope: dict, + mcp_servers: list[str] | None, + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + oauth2_headers: dict[str, str] | None, + client_ip: str | None, + ) -> dict[str, str]: """ Build debug headers if debug mode is enabled, otherwise return empty dict. @@ -286,8 +286,8 @@ class MCPDebug: global_mcp_server_manager, ) - server_url: Optional[str] = None - server_auth_type: Optional[str] = None + server_url: str | None = None + server_auth_type: str | None = None auth_resolution = "no-auth" for server_name in mcp_servers or []: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e609de71e7a..d8ab34a7ddb 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -19,10 +19,8 @@ from typing import ( TYPE_CHECKING, Any, Literal, - Optional, TypeAlias, TypedDict, - Union, cast, ) from urllib.parse import ParseResult, urlparse @@ -707,8 +705,8 @@ def _write_user_env_vars_cache(user_id: str, server_id: str, values: dict[str, s def _should_strip_caller_authorization( mcp_server: MCPServer, - raw_headers: Optional[dict[str, str]], - user_api_key_auth: Optional[UserAPIKeyAuth], + raw_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, ) -> bool: """Decide whether the caller's ``Authorization`` header must NOT be forwarded upstream when populating ``extra_headers`` for an MCP server. @@ -771,8 +769,8 @@ def _should_strip_caller_authorization( def _without_authorization( - headers: Optional[dict[str, str]], -) -> Optional[dict[str, str]]: + headers: dict[str, str] | None, +) -> dict[str, str] | None: """A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or None if nothing remains. Drops only the credential, keeping other forwarded headers. """ @@ -793,9 +791,9 @@ def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str def _openapi_forwarded_extra_headers( mcp_server: MCPServer, - raw_headers: Optional[dict[str, str]], - user_api_key_auth: Optional[UserAPIKeyAuth], -) -> Optional[dict[str, str]]: + raw_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, +) -> dict[str, str] | None: if not mcp_server.extra_headers or not raw_headers: return None normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} @@ -818,9 +816,9 @@ def _openapi_forwarded_extra_headers( async def _resolve_byok_mcp_auth_header( mcp_server: MCPServer, - user_api_key_auth: Optional[UserAPIKeyAuth], - mcp_auth_header: Optional[str], -) -> Optional[str]: + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, +) -> str | None: """Resolve BYOK credential for tool calls that bypass ``execute_mcp_tool``.""" if not mcp_server.is_byok: return mcp_auth_header @@ -854,10 +852,10 @@ async def _resolve_byok_mcp_auth_header( def _client_forwarded_authorization_headers( mcp_server: MCPServer, - oauth2_headers: Optional[dict[str, str]], - raw_headers: Optional[dict[str, str]], - user_api_key_auth: Optional[UserAPIKeyAuth], -) -> Optional[dict[str, str]]: + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, +) -> dict[str, str] | None: """Egress headers for the client-forwarded-token modes (``true_passthrough`` / ``oauth_delegate``). Forwards the caller's ``Authorization`` to the upstream, stripped when @@ -876,8 +874,8 @@ def _client_forwarded_authorization_headers( def _take_forwarded_authorization( - headers: Optional[dict[str, str]], -) -> tuple[Optional[str], Optional[dict[str, str]]]: + headers: dict[str, str] | None, +) -> tuple[str | None, dict[str, str] | None]: """Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the remaining headers, so the passthrough resolver arm is the single Authorization source rather than the header also riding in ``extra_headers`` (which the resolved auth would then defer to).""" @@ -888,8 +886,8 @@ def _take_forwarded_authorization( def _passthrough_token_from_mcp_auth_header( - mcp_auth_header: Optional[Union[str, dict[str, str]]], -) -> Optional[str]: + mcp_auth_header: str | dict[str, str] | None, +) -> str | None: """The caller's per-server upstream credential for a passthrough-mode server, or None. Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated @@ -968,7 +966,7 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str } -def _to_server_spec_fail_closed(server: MCPServer) -> Optional[ServerSpec]: +def _to_server_spec_fail_closed(server: MCPServer) -> ServerSpec | None: """`to_server_spec`, except a half-configured `oauth2_id_jag` server refuses instead of deferring. ID-JAG has no v1 arm, so deferring to v1 would let `resolve_mcp_auth` honor a caller x-mcp-* @@ -989,7 +987,7 @@ def _to_server_spec_fail_closed(server: MCPServer) -> Optional[ServerSpec]: def _caller_authorization_fans_out( server: MCPServer, - scope_servers: Optional[list[MCPServer]], + scope_servers: list[MCPServer] | None, ) -> bool: """True when forwarding the caller's request-wide ``Authorization`` to ``server`` inside a listing fan-out would replay one credential against multiple upstreams: another server in the @@ -1006,7 +1004,7 @@ def _caller_authorization_fans_out( def _extract_upstream_auth_failure( exc: BaseException, -) -> Optional[tuple[int, Optional[str]]]: +) -> tuple[int, str | None] | None: """The upstream 401/403 and its ``WWW-Authenticate`` header from the exception tree, or ``None``. Delegates to the shared traversal in ``faults`` so every consumer (tool listing, @@ -1034,10 +1032,10 @@ def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool: def _warn_on_server_name_fields( *, server_id: str, - alias: Optional[str], - server_name: Optional[str], + alias: str | None, + server_name: str | None, ): - def _warn(field_name: str, value: Optional[str]) -> None: + def _warn(field_name: str, value: str | None) -> None: if not value: return result = validate_tool_name(value) @@ -1079,7 +1077,7 @@ def _warn_internal_delegate_pkce_if_applicable(server: MCPServer, *, source: str ) -def _deserialize_json_dict(data: str | _StringMap | None) -> Optional[dict[str, str]]: +def _deserialize_json_dict(data: str | _StringMap | None) -> dict[str, str] | None: """ Deserialize optional JSON mappings stored in the database. @@ -1100,7 +1098,7 @@ def _deserialize_json_dict(data: str | _StringMap | None) -> Optional[dict[str, return data -def _deserialize_json_list(data: Any) -> Optional[list[dict[str, Any]]]: +def _deserialize_json_list(data: Any) -> list[dict[str, Any]] | None: """Deserialize a JSON array stored in the DB (``env_vars`` and friends). Returns ``None`` for empty / null / unparseable input. Accepts strings @@ -1168,7 +1166,7 @@ def _normalize_mcp_server_cost_info(mcp_info: MCPInfo) -> None: mcp_info["mcp_server_cost_info"] = normalized -def _create_sampling_callback(user_api_key_auth: Optional[UserAPIKeyAuth] = None): +def _create_sampling_callback(user_api_key_auth: UserAPIKeyAuth | None = None): """ Create a sampling callback for MCP ClientSession. Returns a callable that handles sampling/createMessage requests from @@ -1246,8 +1244,8 @@ class MCPServerManager: @staticmethod def _explicit_oauth2_flow( - oauth2_flow: Optional[str], - ) -> Optional[Literal["client_credentials", "authorization_code"]]: + oauth2_flow: str | None, + ) -> Literal["client_credentials", "authorization_code"] | None: """DB rows persist their flow (write-time stamps plus the startup backfill) and config servers must declare it (validated at load), so both builds read the value verbatim: unknown or null resolves to None, which @@ -1262,13 +1260,13 @@ class MCPServerManager: @staticmethod def _resolve_oauth2_flow( *, - auth_type: Optional[MCPAuthType], - oauth2_flow: Optional[str], - token_url: Optional[str], - authorization_url: Optional[str], - client_id: Optional[str], - client_secret: Optional[str], - ) -> Optional[Literal["client_credentials", "authorization_code"]]: + auth_type: MCPAuthType | None, + oauth2_flow: str | None, + token_url: str | None, + authorization_url: str | None, + client_id: str | None, + client_secret: str | None, + ) -> Literal["client_credentials", "authorization_code"] | None: """Infer oauth2_flow from field shape when the value is omitted. SECURITY-SENSITIVE: this is the shape-inference engine both request-time security @@ -1296,7 +1294,7 @@ class MCPServerManager: return None @staticmethod - def effective_oauth2_flow(server: "MCPServer") -> Optional[Literal["client_credentials", "authorization_code"]]: + def effective_oauth2_flow(server: "MCPServer") -> Literal["client_credentials", "authorization_code"] | None: """The oauth2_flow a security decision must use for ``server`` this request. Column-first, shape-fallback: a stamped row returns its explicit value; an @@ -1342,9 +1340,9 @@ class MCPServerManager: @staticmethod def _obo_needs_endpoint_discovery( - auth_type: Optional[MCPAuthType], - token_exchange_endpoint: Optional[str], - token_url: Optional[str], + auth_type: MCPAuthType | None, + token_exchange_endpoint: str | None, + token_url: str | None, ) -> bool: """An ``oauth2_token_exchange`` server with no configured token endpoint can have it discovered (RFC 9728 -> RFC 8414) the same way the ``oauth2`` flow already does; an explicitly @@ -1354,9 +1352,9 @@ class MCPServerManager: def __init__( self, - cred_provider: Optional[UpstreamCredentialProvider] = None, - per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None, - per_user_token_cache: Optional[MCPPerUserTokenCache] = None, + cred_provider: UpstreamCredentialProvider | None = None, + per_user_oauth_token_store: InvalidatableOAuthTokenStore | None = None, + per_user_token_cache: MCPPerUserTokenCache | None = None, ): self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore( self.get_mcp_server_by_id @@ -1492,7 +1490,7 @@ class MCPServerManager: user_api_key_auth=None, raise_on_missing=False, ) - extra_headers: Optional[dict[str, str]] = dict(resolved_static_headers) if resolved_static_headers else None + extra_headers: dict[str, str] | None = dict(resolved_static_headers) if resolved_static_headers else None client = await self._create_mcp_client( server=server, mcp_auth_header=None, @@ -1529,7 +1527,7 @@ class MCPServerManager: async def load_servers_from_config( self, mcp_servers_config: dict[str, Any], - mcp_aliases: Optional[dict[str, str]] = None, + mcp_aliases: dict[str, str] | None = None, ): """ Load the MCP Servers from the config @@ -1953,7 +1951,7 @@ class MCPServerManager: verbose_logger.info(f"Successfully registered {registered_count} OpenAPI tools for server {server.name}") except Exception as e: - verbose_logger.error(f"Failed to register OpenAPI tools for server {server.name}: {str(e)}") + verbose_logger.error(f"Failed to register OpenAPI tools for server {server.name}: {e!s}") raise e def _cleanup_server_tool_routing_artifacts(self, server: MCPServer) -> None: @@ -1985,9 +1983,7 @@ class MCPServerManager: stale_mapping_keys: list[str] = [] for tool_name, mapped_server in list(self.tool_name_to_mcp_server_name_mapping.items()): - if mapped_server in owned_raw: - stale_mapping_keys.append(tool_name) - elif normalize_server_name(str(mapped_server)) in owned_normalized: + if mapped_server in owned_raw or normalize_server_name(str(mapped_server)) in owned_normalized: stale_mapping_keys.append(tool_name) for key in stale_mapping_keys: @@ -1997,7 +1993,7 @@ class MCPServerManager: """ Remove a server from the registry """ - evicted: Optional[MCPServer] = self.registry.pop(mcp_server.server_id, None) + evicted: MCPServer | None = self.registry.pop(mcp_server.server_id, None) if evicted is None and mcp_server.server_name: evicted = self.registry.pop(mcp_server.server_name, None) if evicted is not None: @@ -2011,7 +2007,7 @@ class MCPServerManager: mcp_server: LiteLLM_MCPServerTable, *, env_vars_are_encrypted: bool, - ) -> Optional[_EnvVarList]: + ) -> _EnvVarList | None: env_vars_list = _deserialize_json_list(getattr(mcp_server, "env_vars", None)) if env_vars_are_encrypted: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 @@ -2026,15 +2022,15 @@ class MCPServerManager: *, mcp_server: LiteLLM_MCPServerTable, auth_type: MCPAuthType, - server_url: Optional[str], - manual_issuer: Optional[str], - manual_authorization_url: Optional[str], - manual_token_url: Optional[str], + server_url: str | None, + manual_issuer: str | None, + manual_authorization_url: str | None, + manual_token_url: str | None, is_discovery_auth_type: bool, use_issuer_anchor: bool, - scopes: Optional[list[str]], - token_exchange_endpoint: Optional[str], - ) -> Optional[MCPOAuthMetadata]: + scopes: list[str] | None, + token_exchange_endpoint: str | None, + ) -> MCPOAuthMetadata | None: obo_needs_discovery = self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) needs_authorization_url = ( is_discovery_auth_type and getattr(mcp_server, "oauth2_flow", None) != "client_credentials" @@ -2051,7 +2047,7 @@ class MCPServerManager: (is_discovery_auth_type and not has_all_upstream_oauth_fields) or obo_needs_discovery ) if not needs_discovery: - mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None + mcp_oauth_metadata: MCPOAuthMetadata | None = None elif use_issuer_anchor and manual_issuer is not None: mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url) else: @@ -2090,7 +2086,7 @@ class MCPServerManager: mcp_server: LiteLLM_MCPServerTable, *, credentials_are_encrypted: bool = True, - env_vars_are_encrypted: Optional[bool] = None, + env_vars_are_encrypted: bool | None = None, ) -> MCPServer: _mcp_info: MCPInfo = mcp_server.mcp_info or {} env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None)) @@ -2103,15 +2099,15 @@ class MCPServerManager: ) credentials_dict = _deserialize_json_dict(getattr(mcp_server, "credentials", None)) - encrypted_auth_value: Optional[str] = None - encrypted_client_id: Optional[str] = None - encrypted_client_secret: Optional[str] = None + encrypted_auth_value: str | None = None + encrypted_client_id: str | None = None + encrypted_client_secret: str | None = None if credentials_dict: encrypted_auth_value = credentials_dict.get("auth_value") encrypted_client_id = credentials_dict.get("client_id") encrypted_client_secret = credentials_dict.get("client_secret") - auth_value: Optional[str] = None + auth_value: str | None = None if encrypted_auth_value: if credentials_are_encrypted: auth_value = decrypt_value_helper( @@ -2123,7 +2119,7 @@ class MCPServerManager: else: auth_value = encrypted_auth_value - client_id_value: Optional[str] = None + client_id_value: str | None = None if encrypted_client_id: if credentials_are_encrypted: client_id_value = decrypt_value_helper( @@ -2135,7 +2131,7 @@ class MCPServerManager: else: client_id_value = encrypted_client_id - client_secret_value: Optional[str] = None + client_secret_value: str | None = None if encrypted_client_secret: if credentials_are_encrypted: client_secret_value = decrypt_value_helper( @@ -2150,7 +2146,7 @@ class MCPServerManager: # AWS SigV4 credential fields aws_creds = self._extract_aws_credentials(credentials_dict, credentials_are_encrypted) - scopes: Optional[list[str]] = None + scopes: list[str] | None = None if credentials_dict: scopes_value = credentials_dict.get("scopes") if scopes_value is not None: @@ -2330,7 +2326,7 @@ class MCPServerManager: verbose_logger.debug(f"Added MCP Server: {new_server.name}") except Exception as e: - verbose_logger.debug(f"Failed to add MCP server: {str(e)}") + verbose_logger.debug(f"Failed to add MCP server: {e!s}") raise e async def update_server(self, mcp_server: LiteLLM_MCPServerTable): @@ -2364,7 +2360,7 @@ class MCPServerManager: verbose_logger.debug(f"Updated MCP Server: {new_server.name}") except Exception as e: - verbose_logger.debug(f"Failed to udpate MCP server: {str(e)}") + verbose_logger.debug(f"Failed to udpate MCP server: {e!s}") raise e def get_all_mcp_server_ids(self) -> set[str]: @@ -2390,7 +2386,7 @@ class MCPServerManager: await user_api_key_cache.async_delete_cache(key=self.get_byom_submitted_servers_cache_key(user_id)) except Exception as e: # noqa: BLE001 - verbose_logger.warning(f"Failed to invalidate BYOM submitted MCP server cache: {str(e)}") + verbose_logger.warning(f"Failed to invalidate BYOM submitted MCP server cache: {e!s}") async def _get_active_submitted_mcp_server_ids_for_user( self, user_api_key_auth: UserAPIKeyAuth | None @@ -2405,7 +2401,7 @@ class MCPServerManager: ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache except Exception as e: # noqa: BLE001 - verbose_logger.warning(f"Failed to load BYOM submitted MCP server cache dependencies: {str(e)}") + verbose_logger.warning(f"Failed to load BYOM submitted MCP server cache dependencies: {e!s}") return [] byom_cache_key = self.get_byom_submitted_servers_cache_key(submitter_user_id) @@ -2415,7 +2411,7 @@ class MCPServerManager: if cached_submitted_server_ids is not None: submitted_server_ids = cast(list[str], cached_submitted_server_ids) except Exception as e: # noqa: BLE001 - verbose_logger.warning(f"Failed to read BYOM submitted MCP server cache: {str(e)}") + verbose_logger.warning(f"Failed to read BYOM submitted MCP server cache: {e!s}") if submitted_server_ids is None: if prisma_client is None: @@ -2426,7 +2422,7 @@ class MCPServerManager: prisma_client, submitter_user_id ) except Exception as e: # noqa: BLE001 - verbose_logger.warning(f"Failed to read BYOM submitted MCP servers from database: {str(e)}") + verbose_logger.warning(f"Failed to read BYOM submitted MCP servers from database: {e!s}") submitted_server_ids = [] try: await user_api_key_cache.async_set_cache( @@ -2435,7 +2431,7 @@ class MCPServerManager: ttl=60, ) except Exception as e: # noqa: BLE001 - verbose_logger.warning(f"Failed to write BYOM submitted MCP server cache: {str(e)}") + verbose_logger.warning(f"Failed to write BYOM submitted MCP server cache: {e!s}") return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None] @@ -2489,7 +2485,7 @@ class MCPServerManager: open_ids.update(submitted_server_ids) return open_ids - async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> list[str]: + async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]: """ Get the allowed MCP Servers for the user. @@ -2651,10 +2647,10 @@ class MCPServerManager: ) return tool_permissions except Exception as e: - verbose_logger.warning(f"Failed to resolve toolset permissions: {str(e)}") + verbose_logger.warning(f"Failed to resolve toolset permissions: {e!s}") return {} - def invalidate_toolset_cache(self, toolset_id: Optional[str] = None) -> None: + def invalidate_toolset_cache(self, toolset_id: str | None = None) -> None: """Evict cached toolset permission entries. Called after create/update/delete of a toolset so stale data is not served. @@ -2692,7 +2688,7 @@ class MCPServerManager: self, prisma_client: PrismaClient, toolset_name: str, - ) -> "Optional[MCPToolset]": + ) -> "MCPToolset | None": """Return a toolset by name, cached in ``user_api_key_cache`` (Redis-backed ``DualCache`` in production) to avoid a DB hit on every routed request. @@ -2728,7 +2724,7 @@ class MCPServerManager: ) return toolset - def filter_server_ids_by_ip(self, server_ids: list[str], client_ip: Optional[str]) -> list[str]: + def filter_server_ids_by_ip(self, server_ids: list[str], client_ip: str | None) -> list[str]: """ Filter server IDs by client IP — external callers only see public servers. @@ -2737,9 +2733,7 @@ class MCPServerManager: filtered, _ = self.filter_server_ids_by_ip_with_info(server_ids, client_ip) return filtered - def filter_server_ids_by_ip_with_info( - self, server_ids: list[str], client_ip: Optional[str] - ) -> tuple[list[str], int]: + def filter_server_ids_by_ip_with_info(self, server_ids: list[str], client_ip: str | None) -> tuple[list[str], int]: """ Filter server IDs by client IP — external callers only see public servers. @@ -2770,14 +2764,14 @@ class MCPServerManager: return [] return await self._get_tools_from_server(server) except Exception as e: - verbose_logger.warning(f"Failed to get tools from server {server_id}: {str(e)}") + verbose_logger.warning(f"Failed to get tools from server {server_id}: {e!s}") return [] async def list_tools( self, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[dict[str, Union[str, dict[str, str]]]] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, str | dict[str, str]] | None = None, ) -> list[MCPTool]: """ List all tools available across all MCP Servers. @@ -2803,7 +2797,7 @@ class MCPServerManager: return [] # Get server-specific auth header if available - server_auth_header: Optional[Union[str, dict[str, str]]] = None + server_auth_header: str | dict[str, str] | None = None if mcp_server_auth_headers: from litellm.proxy._experimental.mcp_server.utils import ( lookup_mcp_server_auth_in_headers, @@ -2828,7 +2822,7 @@ class MCPServerManager: return tools except Exception as e: verbose_logger.warning( - f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers." + f"Failed to list tools from server {server.name}: {e!s}. Continuing with other servers." ) return [] @@ -2847,15 +2841,15 @@ class MCPServerManager: ######################################################### @staticmethod def _extract_bearer_token( - oauth2_headers: Optional[dict[str, str]], - raw_headers: Optional[dict[str, str]], - ) -> Optional[str]: + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + ) -> str | None: """Extract the bare Bearer token from oauth2_headers or raw_headers. Returns the token string without the ``Bearer `` prefix, or ``None`` if no Authorization header is found. """ - auth_value: Optional[str] = None + auth_value: str | None = None if oauth2_headers and "Authorization" in oauth2_headers: auth_value = oauth2_headers["Authorization"] elif raw_headers: @@ -2871,8 +2865,8 @@ class MCPServerManager: def _obo_subject_token( self, server: MCPServer, - raw_headers: Optional[dict[str, str]], - ) -> Optional[str]: + raw_headers: dict[str, str] | None, + ) -> str | None: """The caller's bearer as the token_exchange (OBO) subject token, for that mode only. Prompts/resources discovery and reads on a token_exchange server must exchange the caller's @@ -2886,8 +2880,8 @@ class MCPServerManager: def _build_stdio_env( self, server: MCPServer, - raw_headers: Optional[dict[str, str]] = None, - ) -> Optional[dict[str, str]]: + raw_headers: dict[str, str] | None = None, + ) -> dict[str, str] | None: """Resolve stdio env values, supporting header-driven placeholders.""" if server.transport != MCPTransport.stdio or not server.env: @@ -2932,10 +2926,10 @@ class MCPServerManager: async def _resolve_static_headers_with_env_vars( self, server: MCPServer, - user_api_key_auth: Optional[UserAPIKeyAuth], + user_api_key_auth: UserAPIKeyAuth | None, *, raise_on_missing: bool = True, - ) -> Optional[dict[str, str]]: + ) -> dict[str, str] | None: """Return server.static_headers with ``${NAME}`` interpolated. Globals come from ``server.env_vars`` entries with ``scope=="global"``. @@ -3025,7 +3019,7 @@ class MCPServerManager: async def _load_user_env_vars( self, server: MCPServer, - user_api_key_auth: Optional[UserAPIKeyAuth], + user_api_key_auth: UserAPIKeyAuth | None, *, force_refresh: bool = False, ) -> dict[str, str]: @@ -3079,10 +3073,10 @@ class MCPServerManager: server: MCPServer, spec: ServerSpec, provider: UpstreamCredentialProvider, - subject_token: Optional[str], - user_api_key_auth: Optional[UserAPIKeyAuth], - extra_headers: Optional[dict[str, str]], - ) -> tuple[Optional[httpx.Auth], Optional[dict[str, str]]]: + subject_token: str | None, + user_api_key_auth: UserAPIKeyAuth | None, + extra_headers: dict[str, str] | None, + ) -> tuple[httpx.Auth | None, dict[str, str] | None]: """Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``. On a missing/rejected per-user credential this raises the mode's discovery challenge @@ -3134,8 +3128,8 @@ class MCPServerManager: async def preflight_token_exchange( self, server: MCPServer, - oauth2_headers: Optional[dict[str, str]], - user_api_key_auth: Optional[UserAPIKeyAuth], + oauth2_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, ) -> None: """Run the OBO exchange for a caller-supplied subject at the transport edge. @@ -3168,12 +3162,12 @@ class MCPServerManager: async def _create_mcp_client( self, server: MCPServer, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, - stdio_env: Optional[dict[str, str]] = None, - subject_token: Optional[str] = None, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - cred_provider: Optional[UpstreamCredentialProvider] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + stdio_env: dict[str, str] | None = None, + subject_token: str | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + cred_provider: UpstreamCredentialProvider | None = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -3253,7 +3247,7 @@ class MCPServerManager: f"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.", ) - stdio_config: Optional[MCPStdioConfig] = None + stdio_config: MCPStdioConfig | None = None if server.command and server.args is not None: stdio_config = MCPStdioConfig( command=server.command, @@ -3330,12 +3324,12 @@ class MCPServerManager: async def _get_tools_from_server( self, server: MCPServer, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, add_prefix: bool = True, - raw_headers: Optional[dict[str, str]] = None, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - oauth2_headers: Optional[dict[str, str]] = None, + raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + oauth2_headers: dict[str, str] | None = None, ) -> list[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -3482,21 +3476,21 @@ class MCPServerManager: www_authenticate=None if server.is_dcr_bridge else challenge_header, server_name=server.name, ) from e - verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}") + verbose_logger.warning(f"Failed to get tools from server {server.name}: {e!s}") raise MCPServerListError(ServerListFault(tag="internal", status_code=e.status_code), server.name) from e except MCPServerListError: raise except Exception as e: - verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}") + verbose_logger.warning(f"Failed to get tools from server {server.name}: {e!s}") raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge) async def get_prompts_from_server( self, server: MCPServer, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, add_prefix: bool = True, - raw_headers: Optional[dict[str, str]] = None, + raw_headers: dict[str, str] | None = None, ) -> list[Prompt]: """ Helper method to get prompts from a single MCP server with prefixed names. @@ -3538,16 +3532,16 @@ class MCPServerManager: return prefixed_or_original_prompts except Exception as e: - verbose_logger.warning(f"Failed to get prompts from server {server.name}: {str(e)}") + verbose_logger.warning(f"Failed to get prompts from server {server.name}: {e!s}") return [] async def get_resources_from_server( self, server: MCPServer, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, add_prefix: bool = True, - raw_headers: Optional[dict[str, str]] = None, + raw_headers: dict[str, str] | None = None, ) -> list[Resource]: """Fetch available resources from a single MCP server.""" @@ -3580,16 +3574,16 @@ class MCPServerManager: return prefixed_resources except Exception as e: - verbose_logger.warning(f"Failed to get resources from server {server.name}: {str(e)}") + verbose_logger.warning(f"Failed to get resources from server {server.name}: {e!s}") return [] async def get_resource_templates_from_server( self, server: MCPServer, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, add_prefix: bool = True, - raw_headers: Optional[dict[str, str]] = None, + raw_headers: dict[str, str] | None = None, ) -> list[ResourceTemplate]: """Fetch available resource templates from a single MCP server.""" @@ -3624,16 +3618,16 @@ class MCPServerManager: return prefixed_templates except Exception as e: - verbose_logger.warning(f"Failed to get resource templates from server {server.name}: {str(e)}") + verbose_logger.warning(f"Failed to get resource templates from server {server.name}: {e!s}") return [] async def read_resource_from_server( self, server: MCPServer, url: AnyUrl, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, - raw_headers: Optional[dict[str, str]] = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" @@ -3662,10 +3656,10 @@ class MCPServerManager: self, server: MCPServer, prompt_name: str, - arguments: Optional[dict[str, str]] = None, - mcp_auth_header: Optional[Union[str, dict[str, str]]] = None, - extra_headers: Optional[dict[str, str]] = None, - raw_headers: Optional[dict[str, str]] = None, + arguments: dict[str, str] | None = None, + mcp_auth_header: str | dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" @@ -3741,7 +3735,7 @@ class MCPServerManager: *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False, - ) -> Optional[MCPOAuthMetadata]: + ) -> MCPOAuthMetadata | None: """Discover OAuth metadata by following RFC 9728 (protected resource metadata discovery). ``allow_origin_fallback`` controls the last-resort guess that treats the resource server's own @@ -3827,7 +3821,7 @@ class MCPServerManager: exc, ) - header_value: Optional[str] = None + header_value: str | None = None if exc.response is not None: header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get("www-authenticate") status_attempt = ( @@ -3894,7 +3888,7 @@ class MCPServerManager: return metadata, attempts - def _parse_www_authenticate_header(self, header_value: Optional[str]) -> tuple[Optional[str], Optional[list[str]]]: + def _parse_www_authenticate_header(self, header_value: str | None) -> tuple[str | None, list[str] | None]: if not header_value: return None, None @@ -3916,7 +3910,7 @@ class MCPServerManager: async def _fetch_oauth_metadata_from_resource( self, resource_metadata_url: str, server_url: str - ) -> tuple[list[str], Optional[list[str]]]: + ) -> tuple[list[str], list[str] | None]: if not resource_metadata_url: return [], None @@ -3951,7 +3945,7 @@ class MCPServerManager: return authorization_servers, scopes - async def _attempt_well_known_discovery(self, server_url: str) -> tuple[list[str], Optional[list[str]]]: + async def _attempt_well_known_discovery(self, server_url: str) -> tuple[list[str], list[str] | None]: try: parsed = urlparse(server_url) except Exception: @@ -3981,7 +3975,7 @@ class MCPServerManager: async def _fetch_authorization_server_metadata( self, authorization_servers: list[str], server_url: str - ) -> Optional[MCPOAuthMetadata]: + ) -> MCPOAuthMetadata | None: for issuer in authorization_servers: metadata = await self._fetch_single_authorization_server_metadata(issuer, server_url) if metadata is not None: @@ -3989,8 +3983,8 @@ class MCPServerManager: return None async def _fetch_issuer_anchored_oauth_metadata( - self, issuer: str, server_url: Optional[str] - ) -> Optional[MCPOAuthMetadata]: + self, issuer: str, server_url: str | None + ) -> MCPOAuthMetadata | None: """RFC 8414 issuer-anchored discovery for the OAuth endpoints, with resource-driven scopes. Fetch authorization-server metadata from the admin-configured issuer's own origin and adopt @@ -4022,8 +4016,8 @@ class MCPServerManager: return metadata.model_copy(update={"scopes": resource_scopes}) async def _fetch_single_authorization_server_metadata( - self, issuer_url: str, server_url: str, require_issuer: Optional[str] = None - ) -> Optional[MCPOAuthMetadata]: + self, issuer_url: str, server_url: str, require_issuer: str | None = None + ) -> MCPOAuthMetadata | None: try: parsed = urlparse(issuer_url) except Exception: @@ -4110,7 +4104,7 @@ class MCPServerManager: @staticmethod def _build_azure_authorization_server_metadata( parsed_issuer_url: ParseResult, - ) -> Optional[MCPOAuthMetadata]: + ) -> MCPOAuthMetadata | None: path_parts = [part for part in (parsed_issuer_url.path or "").split("/") if part] if parsed_issuer_url.netloc not in _AZURE_ENTRA_HOSTS or len(path_parts) != 2 or path_parts[1] != "v2.0": return None @@ -4124,10 +4118,10 @@ class MCPServerManager: @staticmethod def _decrypt_credential_field( - encrypted_value: Optional[str], + encrypted_value: str | None, key: str, credentials_are_encrypted: bool, - ) -> Optional[str]: + ) -> str | None: """Decrypt a single credential field, or return as-is if not encrypted.""" if not encrypted_value: return None @@ -4142,9 +4136,9 @@ class MCPServerManager: def _extract_aws_credentials( self, - credentials_dict: Optional[dict[str, str]], + credentials_dict: dict[str, str] | None, credentials_are_encrypted: bool, - ) -> dict[str, Optional[str]]: + ) -> dict[str, str | None]: """Extract and decrypt AWS SigV4 credential fields from credentials dict.""" if not credentials_dict: return {} @@ -4170,7 +4164,7 @@ class MCPServerManager: "aws_session_name": credentials_dict.get("aws_session_name"), } - def _extract_scopes(self, scopes_value: str | Sequence[object] | None) -> Optional[list[str]]: + def _extract_scopes(self, scopes_value: str | Sequence[object] | None) -> list[str] | None: if isinstance(scopes_value, str): scopes = [s.strip() for s in scopes_value.split() if s.strip()] return scopes or None @@ -4221,10 +4215,10 @@ class MCPServerManager: verbose_logger.warning(f"Task cancelled while listing tools from {server_name}") raise MCPServerListError(ServerListFault(tag="internal"), server_name) from e except ConnectionError as e: - verbose_logger.warning(f"Connection error while listing tools from {server_name}: {str(e)}") + verbose_logger.warning(f"Connection error while listing tools from {server_name}: {e!s}") raise MCPServerListError(ServerListFault(tag="unreachable"), server_name) from e except Exception as e: - verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}") + verbose_logger.warning(f"Error listing tools from {server_name}: {e!s}") raise_classified_list_failure(e, server_name) _SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024 @@ -4232,7 +4226,7 @@ class MCPServerManager: def _assign_unique_short_prefix( self, server: MCPServer, - registry: Optional[dict[str, MCPServer]] = None, + registry: dict[str, MCPServer] | None = None, ) -> None: """Resolve and cache a collision-free short tool prefix on ``server``. @@ -4448,7 +4442,7 @@ class MCPServerManager: self, tool_name: str, server: MCPServer, - user_api_key_auth: Optional[UserAPIKeyAuth], + user_api_key_auth: UserAPIKeyAuth | None, ) -> None: """ Check if a tool is allowed based on key/team object_permission.mcp_tool_permissions. @@ -4539,7 +4533,7 @@ class MCPServerManager: return result except Exception as e: - error_msg = f"Error calling OpenAPI tool {tool_name}: {str(e)}" + error_msg = f"Error calling OpenAPI tool {tool_name}: {e!s}" verbose_logger.error(error_msg) return CallToolResult( content=[TextContent(type="text", text=error_msg)], @@ -4551,10 +4545,10 @@ class MCPServerManager: name: str, arguments: dict[str, Any], server_name: str, - user_api_key_auth: Optional[UserAPIKeyAuth], + user_api_key_auth: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging | None, server: MCPServer, - raw_headers: Optional[dict[str, str]] = None, + raw_headers: dict[str, str] | None = None, ) -> dict[str, Any]: """ Run pre-call checks and guardrail hooks for an MCP tool call. @@ -4598,7 +4592,7 @@ class MCPServerManager: # Extract incoming Bearer token from raw request headers so # guardrails like MCPJWTSigner can verify + re-sign it (FR-5). normalized_raw = {k.lower(): v for k, v in (raw_headers or {}).items()} - incoming_bearer_token: Optional[str] = None + incoming_bearer_token: str | None = None auth_hdr = normalized_raw.get("authorization", "") if auth_hdr.lower().startswith("bearer "): incoming_bearer_token = auth_hdr[len("bearer ") :] @@ -4645,7 +4639,7 @@ class MCPServerManager: HTTPException, ) as e: # Re-raise guardrail exceptions to properly fail the MCP call - verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}") + verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {e!s}") raise e return hook_result @@ -4654,8 +4648,8 @@ class MCPServerManager: self, name: str, arguments: _ToolArguments, - server_name_from_prefix: Optional[str], - user_api_key_auth: Optional[UserAPIKeyAuth], + server_name_from_prefix: str | None, + user_api_key_auth: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, start_time: datetime.datetime, ): @@ -4688,7 +4682,7 @@ class MCPServerManager: ) ) - def _get_call_semaphore(self, mcp_server: MCPServer) -> Optional[asyncio.Semaphore]: + def _get_call_semaphore(self, mcp_server: MCPServer) -> asyncio.Semaphore | None: limit = mcp_server.max_concurrent_requests if limit is None or limit <= 0: return None @@ -4713,13 +4707,13 @@ class MCPServerManager: *, client: MCPClient, call_tool_params: MCPCallToolRequestParams, - host_progress_callback: Optional[Callable], + host_progress_callback: Callable | None, mcp_server: MCPServer, server_auth_header: str | dict[str, str] | None, - extra_headers: Optional[dict[str, str]], - stdio_env: Optional[dict[str, str]], - subject_token: Optional[str], - user_api_key_auth: Optional[UserAPIKeyAuth], + extra_headers: dict[str, str] | None, + stdio_env: dict[str, str] | None, + subject_token: str | None, + user_api_key_auth: UserAPIKeyAuth | None, ) -> CallToolResult: """Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry. @@ -4754,14 +4748,14 @@ class MCPServerManager: original_tool_name: str, arguments: _ToolArguments, tasks: list, - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[dict[str, dict[str, str]]], - oauth2_headers: Optional[dict[str, str]], - raw_headers: Optional[dict[str, str]], - proxy_logging_obj: Optional[ProxyLogging], - host_progress_callback: Optional[Callable] = None, - hook_extra_headers: Optional[dict[str, str]] = None, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + proxy_logging_obj: ProxyLogging | None, + host_progress_callback: Callable | None = None, + hook_extra_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -4791,7 +4785,7 @@ class MCPServerManager: # Get server-specific auth header if available (case-insensitive) # FIX: Added case-insensitive matching to handle auth header keys that may not match # the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway') - server_auth_header: Optional[Union[dict[str, str], str]] = None + server_auth_header: dict[str, str] | str | None = None if mcp_server_auth_headers: # Normalize keys for case-insensitive lookup from litellm.proxy._experimental.mcp_server.utils import ( @@ -4809,8 +4803,8 @@ class MCPServerManager: server_auth_header = mcp_auth_header # Extract subject token for OAuth2 Token Exchange (OBO) and ID-JAG flows - subject_token: Optional[str] = None - extra_headers: Optional[dict[str, str]] = None + subject_token: str | None = None + extra_headers: dict[str, str] | None = None if mcp_server.auth_type in ( MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, @@ -5001,7 +4995,7 @@ class MCPServerManager: GuardrailRaisedException, HTTPException, ) as e: - verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}") + verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {e!s}") raise e # If proxy_logging_obj is None, the tool call result is at index 0 @@ -5056,7 +5050,7 @@ class MCPServerManager: return mcp_server - async def has_user_oauth_token(self, server: MCPServer, user_api_key_auth: Optional[UserAPIKeyAuth]) -> bool: + async def has_user_oauth_token(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> bool: """Whether the v2 resolver can produce a per-user token for this server right now. This is the preemptive 401's existence check, routed through the same resolver that drives @@ -5093,9 +5087,9 @@ class MCPServerManager: async def _resolve_oauth2_headers_for_tool_call( self, mcp_server: MCPServer, - oauth2_headers: Optional[dict[str, str]], - user_api_key_auth: Optional[UserAPIKeyAuth], - ) -> Optional[dict[str, str]]: + oauth2_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, + ) -> dict[str, str] | None: """Look up per-user OAuth headers when the client did not supply a token.""" if not mcp_server.needs_user_oauth_token or oauth2_headers or user_api_key_auth is None: return oauth2_headers @@ -5188,7 +5182,7 @@ class MCPServerManager: async def _gather_openapi_tool_tasks( self, tasks: list[Any], - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, ) -> CallToolResult: """Await OpenAPI tool tasks and return the tool call result.""" try: @@ -5200,7 +5194,7 @@ class MCPServerManager: GuardrailRaisedException, HTTPException, ) as e: - verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}") + verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {e!s}") raise e async def call_tool( @@ -5208,13 +5202,13 @@ class MCPServerManager: server_name: str, name: str, arguments: _ToolArguments, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, - oauth2_headers: Optional[dict[str, str]] = None, - raw_headers: Optional[dict[str, str]] = None, - host_progress_callback: Optional[Callable] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + proxy_logging_obj: ProxyLogging | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + host_progress_callback: Callable | None = None, ) -> CallToolResult: """ Call a tool with the given name and arguments @@ -5351,7 +5345,7 @@ class MCPServerManager: asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping()) except RuntimeError as e: # no running event loop verbose_logger.exception( - f"No running event loop - skipping tool name to MCP server name mapping initialization: {str(e)}" + f"No running event loop - skipping tool name to MCP server name mapping initialization: {e!s}" ) async def _initialize_tool_name_to_mcp_server_name_mapping(self): @@ -5370,12 +5364,12 @@ class MCPServerManager: # at startup we have none, so an upstream 401 is normal. # Swallow it so we keep mapping the remaining servers. verbose_logger.debug( - f"Skipping tool name mapping for server {server.name} due to upstream auth error: {str(e)}" + f"Skipping tool name mapping for server {server.name} due to upstream auth error: {e!s}" ) continue except Exception as e: verbose_logger.warning( - f"Failed to get tools from server {server.name} during tool name mapping initialization: {str(e)}" + f"Failed to get tools from server {server.name} during tool name mapping initialization: {e!s}" ) continue for tool in tools: @@ -5385,7 +5379,7 @@ class MCPServerManager: self.tool_name_to_mcp_server_name_mapping[original_name] = server.name self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name - def _get_mcp_server_from_tool_name(self, tool_name: str) -> Optional[MCPServer]: + def _get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None: """ Get the MCP Server from the tool name (handles both prefixed and non-prefixed names) @@ -5558,7 +5552,7 @@ class MCPServerManager: # Fallback if proxy_server not available return {} - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: Optional[str]) -> bool: + def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: """ Check if a server is accessible from the given client IP. @@ -5579,7 +5573,7 @@ class MCPServerManager: internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges")) return IPAddressUtils.is_internal_ip(client_ip, internal_networks) - def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]: + def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None: """ Get the MCP Server from the server id """ @@ -5661,7 +5655,7 @@ class MCPServerManager: def expand_tool_permissions( self, - tool_permissions: Optional[dict[str, list[str]]], + tool_permissions: dict[str, list[str]] | None, ) -> dict[str, list[str]]: """ Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so @@ -5683,7 +5677,7 @@ class MCPServerManager: result.setdefault(server_id, []).extend(tools or []) return result - def get_mcp_server_by_name(self, server_name: str, client_ip: Optional[str] = None) -> Optional[MCPServer]: + def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: """ Get the MCP Server from the server name. @@ -5718,7 +5712,7 @@ class MCPServerManager: return server return None - def get_filtered_registry(self, client_ip: Optional[str] = None) -> dict[str, MCPServer]: + def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ Get registry filtered by client IP access control. @@ -5736,8 +5730,8 @@ class MCPServerManager: server_name: str, url: str, transport: str, - auth_type: Optional[str] = None, - alias: Optional[str] = None, + auth_type: str | None = None, + alias: str | None = None, ) -> str: """ Generate a stable server ID based on server parameters using a hash function. @@ -5767,9 +5761,7 @@ class MCPServerManager: # Take first 32 characters and format as UUID-like string return hash_hex[:32] - async def health_check_server( - self, server_id: str, mcp_auth_header: Optional[str] = None - ) -> LiteLLM_MCPServerTable: + async def health_check_server(self, server_id: str, mcp_auth_header: str | None = None) -> LiteLLM_MCPServerTable: """ Perform a health check on a specific MCP server. @@ -5801,22 +5793,17 @@ class MCPServerManager: should_skip_health_check = False # Skip if server requires per-user authentication (OAuth2 or passthrough auth) - if server.requires_per_user_auth: - should_skip_health_check = True - # Skip if auth_type is not none and authentication_token is missing - # (except aws_sigv4 which uses its own credential fields) - elif ( - server.auth_type - and server.auth_type != MCPAuth.none - and server.auth_type != MCPAuth.aws_sigv4 - and not server.authentication_token + if ( + server.requires_per_user_auth + or ( + server.auth_type + and server.auth_type != MCPAuth.none + and server.auth_type != MCPAuth.aws_sigv4 + and not server.authentication_token + ) + or self._references_per_user_env_var(server) ): should_skip_health_check = True - # Skip if static_headers reference a per-user env var: a userless probe - # can't fill ${NAME} and would forward the literal placeholder upstream, - # flipping the server to unhealthy even though real user calls succeed. - elif self._references_per_user_env_var(server): - should_skip_health_check = True if not should_skip_health_check: resolved_static_headers = await self._resolve_static_headers_with_env_vars( @@ -5893,8 +5880,8 @@ class MCPServerManager: async def get_all_mcp_servers_with_health_and_teams( self, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - server_ids: Optional[list[str]] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + server_ids: list[str] | None = None, ) -> list[LiteLLM_MCPServerTable]: """ Get all MCP servers that the user has access to, with health status and team information. @@ -5923,7 +5910,7 @@ class MCPServerManager: async def get_all_allowed_mcp_servers( self, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> list[LiteLLM_MCPServerTable]: """ Get all MCP servers that the user has access to. @@ -5952,8 +5939,8 @@ class MCPServerManager: @staticmethod def _env_vars_to_models( - env_vars: Optional[_EnvVarList], - ) -> Optional[list[MCPEnvVar]]: + env_vars: _EnvVarList | None, + ) -> list[MCPEnvVar] | None: if env_vars is None: return None return [MCPEnvVar.model_validate(env_var) for env_var in env_vars] @@ -6021,7 +6008,7 @@ class MCPServerManager: return servers async def get_all_mcp_servers_with_health_unfiltered( - self, server_ids: Optional[list[str]] = None + self, server_ids: list[str] | None = None ) -> list[LiteLLM_MCPServerTable]: """Return health info for all servers in registry regardless of user access.""" diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py index 02cec2475e2..cf85b4b4ab6 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py @@ -36,7 +36,7 @@ a healed fleet has no null rows and the backfill exits after one query. import json from collections import Counter -from typing import Any, Literal, Optional +from typing import Any, Literal from litellm._logging import verbose_proxy_logger from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials @@ -55,7 +55,7 @@ BackfillRule = Literal[ _BACKFILL_AUDIT_ACTOR = "oauth2_flow_backfill" -def _decrypted_credentials(raw_credentials: Any) -> Optional[MCPCredentials]: +def _decrypted_credentials(raw_credentials: Any) -> MCPCredentials | None: if raw_credentials is None: return None if isinstance(raw_credentials, str): @@ -73,11 +73,11 @@ def _decrypted_credentials(raw_credentials: Any) -> Optional[MCPCredentials]: def classify_null_flow_row( *, has_per_user_tokens: bool, - authorization_url: Optional[str], - registration_url: Optional[str], - token_url: Optional[str], - credentials: Optional[MCPCredentials], -) -> tuple[Optional[OAuth2Flow], BackfillRule]: + authorization_url: str | None, + registration_url: str | None, + token_url: str | None, + credentials: MCPCredentials | None, +) -> tuple[OAuth2Flow | None, BackfillRule]: if has_per_user_tokens: return "authorization_code", "per_user_tokens" if authorization_url: diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index b2b3f70d200..d74f1d31655 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -7,7 +7,7 @@ with ``client_id``, ``client_secret``, and ``token_url``. import asyncio import hashlib -from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING import httpx @@ -23,14 +23,14 @@ from litellm.constants import ( MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX, ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - decrypt_value_helper, - encrypt_value_helper, -) from litellm.proxy._experimental.mcp_server.oauth_utils import ( build_upstream_oauth2_token_request, resolve_upstream_resource, ) +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) from litellm.types.llms.custom_http import httpxSpecialProvider if TYPE_CHECKING: @@ -58,7 +58,7 @@ class MCPOAuth2TokenCache(InMemoryCache): max_size_in_memory=MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, ) - self._locks: Dict[str, asyncio.Lock] = {} + self._locks: dict[str, asyncio.Lock] = {} @staticmethod def _token_identity(server: "MCPServer") -> str: @@ -84,7 +84,7 @@ class MCPOAuth2TokenCache(InMemoryCache): def _has_client_credentials_config(server: "MCPServer") -> bool: return bool(server.client_id and server.client_secret and server.token_url) - async def async_get_token(self, server: "MCPServer") -> Optional[str]: + async def async_get_token(self, server: "MCPServer") -> str | None: """Return a valid access token, fetching or refreshing as needed. Returns ``None`` when the server lacks client credentials config. @@ -111,7 +111,7 @@ class MCPOAuth2TokenCache(InMemoryCache): self.set_cache(identity, token, ttl=ttl) return token - async def _fetch_token(self, server: "MCPServer") -> Tuple[str, int]: + async def _fetch_token(self, server: "MCPServer") -> tuple[str, int]: """POST to ``token_url`` with ``grant_type=client_credentials``. Returns ``(access_token, ttl_seconds)`` where ttl accounts for the @@ -133,7 +133,7 @@ class MCPOAuth2TokenCache(InMemoryCache): client_id=server.client_id, client_secret=server.client_secret, ) - data: Dict[str, str] = { + data: dict[str, str] = { "grant_type": "client_credentials", **token_request.body, } @@ -200,7 +200,7 @@ class MCPOAuth2TokenCache(InMemoryCache): mcp_oauth2_token_cache = MCPOAuth2TokenCache() -def _compute_per_user_token_ttl(server: "MCPServer", expires_in: Optional[int]) -> int: +def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: """Compute Redis TTL for a per-user token. Uses server.token_storage_ttl_seconds when configured, capped at the token's @@ -232,7 +232,7 @@ class MCPPerUserTokenCache: def _cache_key(self, user_id: str, server_id: str) -> str: return f"{MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX}:{user_id}:{server_id}" - async def get(self, user_id: str, server_id: str) -> Optional[str]: + async def get(self, user_id: str, server_id: str) -> str | None: """Return the plaintext access_token, or None on miss/error.""" try: from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 @@ -305,8 +305,8 @@ mcp_per_user_token_cache = MCPPerUserTokenCache() async def resolve_mcp_auth( server: "MCPServer", - mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, -) -> Optional[Union[str, Dict[str, str]]]: + mcp_auth_header: str | dict[str, str] | None = None, +) -> str | dict[str, str] | None: """Resolve the auth value for an MCP server. Priority: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 8f47aa7344d..6cea56dd5d3 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -3,7 +3,7 @@ import os from ipaddress import ip_address -from typing import TYPE_CHECKING, Any, Dict, List, NoReturn, Optional +from typing import TYPE_CHECKING, Any, NoReturn from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit from fastapi import HTTPException, Request @@ -49,17 +49,17 @@ _TRUSTED_REDIRECT_ORIGINS_ENV = "MCP_TRUSTED_REDIRECT_ORIGINS" _TRUSTED_NATIVE_REDIRECT_URIS_ENV = "MCP_TRUSTED_NATIVE_REDIRECT_URIS" # Default allowlist for trusted native redirect URIs. -_DEFAULT_NATIVE_REDIRECT_URIS: List[str] = [ +_DEFAULT_NATIVE_REDIRECT_URIS: list[str] = [ "cursor://anysphere.cursor-mcp/oauth/callback", ] -_warned_invalid_proxy_base_url: Optional[str] = None +_warned_invalid_proxy_base_url: str | None = None def _oauth_invalid_request( error_description: str, *, - hint: Optional[str] = None, + hint: str | None = None, **extra: Any, ) -> NoReturn: """Raise ``invalid_request`` (RFC 6749) with a debuggable description. @@ -68,7 +68,7 @@ def _oauth_invalid_request( ``invalid_request``; ``error_description`` and ``hint`` explain what failed and how to fix it (e.g. reverse-proxy / PROXY_BASE_URL issues). """ - detail: Dict[str, Any] = { + detail: dict[str, Any] = { "error": "invalid_request", "error_description": error_description, } @@ -83,7 +83,7 @@ def _origin_label(scheme: str, netloc: str) -> str: return f"{scheme}://{netloc}" if netloc else f"{scheme}://" -def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: +def _redact_mcp_resource_url(url: str | None) -> str | None: """Reduce an MCP server URL to its origin (scheme + host + port) for logging. Everything else is dropped: userinfo (``user:pass@``), the query string, the @@ -106,7 +106,7 @@ def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: return urlunsplit((parts.scheme, netloc, "", "", "")) or None -def _resolve_proxy_base_url_env() -> Optional[str]: +def _resolve_proxy_base_url_env() -> str | None: global _warned_invalid_proxy_base_url configured = os.environ.get("PROXY_BASE_URL", "").strip() if not configured: @@ -281,7 +281,7 @@ def _strip_default_port(scheme: str, netloc: str) -> str: return lowered -def _parse_trusted_redirect_origins() -> List[str]: +def _parse_trusted_redirect_origins() -> list[str]: """Parse ``MCP_TRUSTED_REDIRECT_ORIGINS`` into normalized entries. Empty / unset env var → empty list. Entries are lowercased and any scheme / path component the operator included is stripped. Default @@ -293,7 +293,7 @@ def _parse_trusted_redirect_origins() -> List[str]: raw = os.environ.get(_TRUSTED_REDIRECT_ORIGINS_ENV, "").strip() if not raw: return [] - entries: List[str] = [] + entries: list[str] = [] for token in raw.split(","): entry = token.strip().lower() if not entry: @@ -345,9 +345,9 @@ def _normalize_native_redirect_uri( ) -def _parse_trusted_native_redirect_uris() -> List[str]: +def _parse_trusted_native_redirect_uris() -> list[str]: """Built-in native MCP callbacks plus ``MCP_TRUSTED_NATIVE_REDIRECT_URIS``.""" - entries: List[str] = [uri.lower() for uri in _DEFAULT_NATIVE_REDIRECT_URIS] + entries: list[str] = [uri.lower() for uri in _DEFAULT_NATIVE_REDIRECT_URIS] raw = os.environ.get(_TRUSTED_NATIVE_REDIRECT_URIS_ENV, "").strip() if not raw: return entries @@ -467,7 +467,7 @@ def validate_redirect_uri_shape(parsed: ParseResult) -> bool: return False -def _resolve_proxy_base_for_redirect(request: Request) -> Optional[str]: +def _resolve_proxy_base_for_redirect(request: Request) -> str | None: try: return get_request_base_url(request) except Exception as exc: @@ -482,7 +482,7 @@ def _resolve_proxy_base_for_redirect(request: Request) -> Optional[str]: def _trusted_redirect_uri_is_allowed( parsed: ParseResult, redirect_netloc: str, - proxy_base: Optional[str], + proxy_base: str | None, ) -> bool: if proxy_base: proxy_parsed = urlparse(proxy_base) @@ -505,7 +505,7 @@ def _build_trusted_redirect_rejection_message( redirect_uri: str, parsed: ParseResult, redirect_netloc: str, - proxy_base: Optional[str], + proxy_base: str | None, ) -> str: """Build a client-facing rejection message. @@ -520,7 +520,7 @@ def _build_trusted_redirect_rejection_message( _strip_default_port(proxy_parsed.scheme, proxy_parsed.netloc) if proxy_parsed and proxy_parsed.netloc else "" ) - mismatch_parts: List[str] = [] + mismatch_parts: list[str] = [] if proxy_parsed and proxy_parsed.netloc: if parsed.scheme != proxy_parsed.scheme: mismatch_parts.append( @@ -546,7 +546,7 @@ def _raise_trusted_redirect_uri_rejected( redirect_uri: str, parsed: ParseResult, redirect_netloc: str, - proxy_base: Optional[str], + proxy_base: str | None, ) -> NoReturn: description = _build_trusted_redirect_rejection_message(redirect_uri, parsed, redirect_netloc, proxy_base) diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 0b795057837..db2851c60aa 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -8,7 +8,7 @@ import json import os import re from pathlib import PurePosixPath -from typing import Any, Dict, List, Optional +from typing import Any from urllib.parse import quote # Tool names emitted from OpenAPI specs must work across all major LLM providers. @@ -35,30 +35,28 @@ def sanitize_openapi_tool_name(raw_name: str) -> str: from litellm._logging import verbose_logger +from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) # Store the base URL and headers globally BASE_URL = "" -HEADERS: Dict[str, str] = {} +HEADERS: dict[str, str] = {} # Per-request auth header override for BYOK servers. # Set this ContextVar before calling a local tool handler to inject the user's # stored credential into the HTTP request made by the tool function closure. -_request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar( - "_request_auth_header", default=None -) +_request_auth_header: contextvars.ContextVar[str | None] = contextvars.ContextVar("_request_auth_header", default=None) # Per-request extra headers forwarded from the client request. # Populated from MCPServer.extra_headers names matched against raw request # headers in server.py before dispatching to a local/OpenAPI tool handler. -_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = contextvars.ContextVar( +_request_extra_headers: contextvars.ContextVar[dict[str, str] | None] = contextvars.ContextVar( "_request_extra_headers", default=None ) @@ -90,7 +88,7 @@ def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str: return quote(value_str, safe="") -def load_openapi_spec(filepath: str) -> Dict[str, Any]: +def load_openapi_spec(filepath: str) -> dict[str, Any]: """ Sync wrapper. For URL specs, use the shared/custom MCP httpx client. """ @@ -108,7 +106,7 @@ def load_openapi_spec(filepath: str) -> Dict[str, Any]: return asyncio.run(load_openapi_spec_async(filepath)) -async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]: +async def load_openapi_spec_async(filepath: str) -> dict[str, Any]: if filepath.startswith("http://") or filepath.startswith("https://"): client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) r = await async_safe_get(client, filepath) @@ -123,7 +121,7 @@ async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]: return json.load(f) -def get_base_url(spec: Dict[str, Any], spec_path: Optional[str] = None) -> str: +def get_base_url(spec: dict[str, Any], spec_path: str | None = None) -> str: """Extract base URL from OpenAPI spec.""" # OpenAPI 3.x if "servers" in spec and spec["servers"]: @@ -173,7 +171,7 @@ def get_base_url(spec: Dict[str, Any], spec_path: Optional[str] = None) -> str: return "" -def _resolve_ref(param: Dict[str, Any], component_params: Dict[str, Any]) -> Optional[Dict[str, Any]]: +def _resolve_ref(param: dict[str, Any], component_params: dict[str, Any]) -> dict[str, Any] | None: """Resolve a single parameter, following a $ref if present. Returns the resolved param dict, or None if the $ref target is absent from @@ -186,7 +184,7 @@ def _resolve_ref(param: Dict[str, Any], component_params: Dict[str, Any]) -> Opt return component_params.get(ref.split("/")[-1]) -def _resolve_param_list(raw: List[Dict[str, Any]], component_params: Dict[str, Any]) -> List[Dict[str, Any]]: +def _resolve_param_list(raw: list[dict[str, Any]], component_params: dict[str, Any]) -> list[dict[str, Any]]: """Resolve $refs in a parameter list, dropping any unresolvable entries.""" result = [] for p in raw: @@ -197,10 +195,10 @@ def _resolve_param_list(raw: List[Dict[str, Any]], component_params: Dict[str, A def resolve_operation_params( - operation: Dict[str, Any], - path_item: Dict[str, Any], - components: Dict[str, Any], -) -> Dict[str, Any]: + operation: dict[str, Any], + path_item: dict[str, Any], + components: dict[str, Any], +) -> dict[str, Any]: """Return a copy of *operation* with fully-resolved, merged parameters. Handles two common patterns in real-world OpenAPI specs: @@ -225,7 +223,7 @@ def resolve_operation_params( return result -def extract_parameters(operation: Dict[str, Any]) -> tuple: +def extract_parameters(operation: dict[str, Any]) -> tuple: """Extract parameter names from OpenAPI operation.""" path_params = [] query_params = [] @@ -251,7 +249,7 @@ def extract_parameters(operation: Dict[str, Any]) -> tuple: return path_params, query_params, body_params -def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]: +def build_input_schema(operation: dict[str, Any]) -> dict[str, Any]: """Build MCP input schema from OpenAPI operation.""" properties = {} required = [] @@ -297,8 +295,8 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]: def _merge_openapi_tool_request_headers( - static_headers: Dict[str, str], -) -> Dict[str, str]: + static_headers: dict[str, str], +) -> dict[str, str]: """Merge static closure headers with per-request ContextVar overrides. Precedence (highest to lowest): @@ -327,7 +325,7 @@ def _merge_openapi_tool_request_headers( static = static_headers or {} static_lower_names = {k.lower() for k in static} - effective_headers: Dict[str, str] = {k: v for k, v in request_extra.items() if k.lower() not in static_lower_names} + effective_headers: dict[str, str] = {k: v for k, v in request_extra.items() if k.lower() not in static_lower_names} effective_headers.update(static) override_auth = _request_auth_header.get() @@ -348,9 +346,9 @@ def _merge_openapi_tool_request_headers( def create_tool_function( path: str, method: str, - operation: Dict[str, Any], + operation: dict[str, Any], base_url: str, - headers: Optional[Dict[str, str]] = None, + headers: dict[str, str] | None = None, ): """Create a tool function for an OpenAPI operation. @@ -402,7 +400,7 @@ def create_tool_function( url = url.replace("{{" + param_name + "}}", safe_value) # Build query params using original parameter names - params: Dict[str, Any] = {} + params: dict[str, Any] = {} for param_name in query_params: param_value = kwargs.get(param_name, "") if param_value: @@ -410,7 +408,7 @@ def create_tool_function( params[param_name] = param_value # Build request body - json_body: Optional[Dict[str, Any]] = None + json_body: dict[str, Any] | None = None if body_params: # Try "body" first (most common), then check all body param names body_value = kwargs.get("body", {}) @@ -449,7 +447,7 @@ def create_tool_function( return tool_function -def register_tools_from_openapi(spec: Dict[str, Any], base_url: str): +def register_tools_from_openapi(spec: dict[str, Any], base_url: str): """Register MCP tools from OpenAPI specification.""" paths = spec.get("paths", {}) used_names: set = set() diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py index 2bdb8770e4e..a5dc75e3829 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py @@ -48,34 +48,34 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ) __all__ = [ - "Ok", - "Error", - "Result", - "NoOpAuth", - "StaticHeaderAuth", - "UpstreamCredentialProvider", - "AuthSpecKind", - "CredError", - "Subject", - "ServerSpec", - "AuthConfig", - "parse_auth_spec_kind", - "AuthorizationCodeConfig", - "ClientCredentialsConfig", - "TokenExchangeConfig", - "IdJagConfig", - "ClientAuth", - "PrivateKeyJwtAuth", - "ClientSecretAuth", + "Ambient", "ApiKeyConfig", "ApiKeySource", - "SharedKey", - "Byok", - "PassthroughConfig", - "NoneConfig", - "AwsSigV4Config", - "AwsCredentialSource", - "StaticKeys", "AssumeRole", - "Ambient", + "AuthConfig", + "AuthSpecKind", + "AuthorizationCodeConfig", + "AwsCredentialSource", + "AwsSigV4Config", + "Byok", + "ClientAuth", + "ClientCredentialsConfig", + "ClientSecretAuth", + "CredError", + "Error", + "IdJagConfig", + "NoOpAuth", + "NoneConfig", + "Ok", + "PassthroughConfig", + "PrivateKeyJwtAuth", + "Result", + "ServerSpec", + "SharedKey", + "StaticHeaderAuth", + "StaticKeys", + "Subject", + "TokenExchangeConfig", + "UpstreamCredentialProvider", + "parse_auth_spec_kind", ] diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index efaa7b742c2..42131a0c207 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -12,7 +12,7 @@ every other mode so the caller defers to v1 (parity-safe); it grows one branch p from __future__ import annotations import base64 -from typing import TYPE_CHECKING, Literal, NoReturn, Optional +from typing import TYPE_CHECKING, Literal, NoReturn from fastapi import HTTPException from pydantic import SecretStr @@ -45,7 +45,7 @@ _TOKEN_EXCHANGE_SUBJECT_TOKEN_DEFAULT = "urn:ietf:params:oauth:token-type:access _ID_JAG_SUBJECT_TOKEN_DEFAULT = "urn:ietf:params:oauth:token-type:id_token" -def to_subject(user_api_key_auth: Optional[UserAPIKeyAuth], subject_token: Optional[str]) -> Subject: +def to_subject(user_api_key_auth: UserAPIKeyAuth | None, subject_token: str | None) -> Subject: """Map v1's authenticated principal onto the resolver's Subject. tenant_id / subject_id are empty for an unauthenticated caller; the per-user arms must reject @@ -61,7 +61,7 @@ def to_subject(user_api_key_auth: Optional[UserAPIKeyAuth], subject_token: Optio ) -def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: +def to_server_spec(server: MCPServer) -> ServerSpec | None: """Map a v1 server onto a ServerSpec for a migrated mode, or None to defer to v1. BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just @@ -151,7 +151,7 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec: ) -def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]: +def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None: """Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured. An OBO server with ``client_id``/``client_secret`` is owned by the v2 arm even if the @@ -192,7 +192,7 @@ def _shared_key_spec( value_prefix: str, *, encode: bool = False, -) -> Optional[ServerSpec]: +) -> ServerSpec | None: """Build an api_key spec from the server's static token, or defer (None) if it is absent. Covers the whole shared-key static-header family: ``api_key`` on ``X-API-Key`` and the @@ -213,7 +213,7 @@ def _shared_key_spec( ) -def _id_jag_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]: +def _id_jag_spec(server: MCPServer, resource: str) -> ServerSpec | None: """Build an ID-JAG spec from the v1 server's raw fields, or defer (None) if half-configured. The enum already routes here, but a server missing an endpoint, ``client_id``, or any client-auth @@ -243,7 +243,7 @@ def _id_jag_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]: ) -def _id_jag_client_auth(server: MCPServer) -> Optional[ClientAuth]: +def _id_jag_client_auth(server: MCPServer) -> ClientAuth | None: """Private-key JWT when a key is configured, else client_secret, else None (defer to v1).""" if server.client_private_key: return PrivateKeyJwtAuth( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index b70db64ba94..99bb9cc7553 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -27,7 +27,6 @@ import httpx from typing_extensions import assert_never from litellm._logging import verbose_proxy_logger - from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( ClientCredentialsBearerAuth, ClientCredentialsTokenSource, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchanger.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchanger.py index 02b6d4eafb1..2cc4735e877 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchanger.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchanger.py @@ -27,6 +27,9 @@ from typing import Literal, Protocol from typing_extensions import assert_never from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( + build_token_endpoint_client_auth, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( InMemoryTokenCacheBackend, InProcessRefreshCoordinator, @@ -39,9 +42,6 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Ok, Result, ) -from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( - build_token_endpoint_client_auth, -) from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, ServerSpec, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 24d37f81787..6f4fde6fbfe 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -654,7 +654,7 @@ if MCP_AVAILABLE: return { "tools": [], "error": "server_error", - "message": f"Failed to get tools from server {server.name}: {str(e)}", + "message": f"Failed to get tools from server {server.name}: {e!s}", } return { "tools": list_tools_result, @@ -866,7 +866,7 @@ if MCP_AVAILABLE: errors.append( f"{get_server_prefix(server)}: {classify_list_exception(e).tag}" if isinstance(e, (MCPServerListError, MCPUpstreamAuthError)) - else f"{get_server_prefix(server)}: {str(e)}" + else f"{get_server_prefix(server)}: {e!s}" ) continue @@ -905,7 +905,7 @@ if MCP_AVAILABLE: return { "tools": [], "error": "unexpected_error", - "message": f"An unexpected error occurred: {str(e)}", + "message": f"An unexpected error occurred: {e!s}", } @router.post("/tools/call", dependencies=[Depends(user_api_key_auth)]) @@ -1052,7 +1052,7 @@ if MCP_AVAILABLE: }, ) except BlockedPiiEntityError as e: - verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {e!s}") raise HTTPException( status_code=400, detail={ @@ -1063,7 +1063,7 @@ if MCP_AVAILABLE: }, ) except GuardrailRaisedException as e: - verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {e!s}") raise HTTPException( status_code=400, detail={ @@ -1082,15 +1082,15 @@ if MCP_AVAILABLE: # Locally generated denials (tool/server permission, IP filtering, BYOK) stay at error level # so restriction probing keeps full monitoring visibility; the relayed upstream 401 above is # the only status demoted to info. - verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + verbose_logger.error(f"HTTPException in MCP tool call: {e!s}") raise e except Exception as e: - verbose_logger.exception(f"Unexpected error in MCP tool call: {str(e)}") + verbose_logger.exception(f"Unexpected error in MCP tool call: {e!s}") raise HTTPException( status_code=500, detail={ "error": "internal_server_error", - "message": f"An unexpected error occurred: {str(e)}", + "message": f"An unexpected error occurred: {e!s}", }, ) diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 5bad530f37b..779cc5861d4 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -12,7 +12,7 @@ MCP Spec Reference: import typing from collections.abc import Mapping, Sequence -from typing import Any, Dict, List, NamedTuple, Optional, Protocol, Union +from typing import Any, NamedTuple, Optional, Protocol, Union if typing.TYPE_CHECKING: from fastapi import Request @@ -55,7 +55,7 @@ except ImportError as _sampling_import_err: def _resolve_model_from_preferences( model_preferences: Optional["ModelPreferences"], - default_model: Optional[str] = None, + default_model: str | None = None, ) -> str: """ Resolve an LLM model name from MCP ModelPreferences. @@ -168,9 +168,9 @@ class _ScoredModel(NamedTuple): def _select_model_by_priority( - model_names: List[str], + model_names: list[str], model_preferences: "ModelPreferences", -) -> Optional[str]: +) -> str | None: """Score available models by MCP priority weights and return the best. Scoring strategy (per the MCP spec, priorities are 0-1 floats): @@ -226,7 +226,7 @@ def _select_model_by_priority( return None # Min-max normalisation helpers - def _normalise(values: List[float], invert: bool = False) -> List[float]: + def _normalise(values: list[float], invert: bool = False) -> list[float]: """Normalise to [0, 1]. If *invert*, lower raw → higher score.""" lo, hi = min(values), max(values) if hi == lo: @@ -370,8 +370,8 @@ def _convert_single_content( def _convert_mcp_messages_to_openai( - messages: List["SamplingMessage"], - system_prompt: Optional[str] = None, + messages: list["SamplingMessage"], + system_prompt: str | None = None, ) -> "Sequence[Mapping[str, object]]": """ Convert MCP SamplingMessage list to OpenAI messages format. @@ -497,7 +497,7 @@ def _extract_tool_calls( def _extract_text_parts( content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]", -) -> Optional[str]: +) -> str | None: """Extract text parts from mixed content.""" items = content if isinstance(content, list) else [content] texts = [] @@ -534,7 +534,7 @@ def _extract_tool_results( def _convert_mcp_tools_to_openai( - tools: Optional[List["Tool"]], + tools: list["Tool"] | None, ) -> "Sequence[Mapping[str, object]] | None": """ Convert MCP Tool definitions to OpenAI function calling format. @@ -850,8 +850,8 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N async def _run_budget_checks( model: str, user_api_key_auth: "UserAPIKeyAuth", - raw_headers: Optional[Dict[str, str]] = None, - client_ip: Optional[str] = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> Optional["ErrorData"]: """Enforce key/team/user/org/global budget checks for sampling requests. @@ -991,8 +991,8 @@ async def _run_budget_checks( def _build_sampling_request( - raw_headers: Optional[Dict[str, str]] = None, - client_ip: Optional[str] = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> "Request": """Build a synthetic FastAPI Request for sampling sub-calls. @@ -1057,7 +1057,7 @@ def _build_sampling_request( _server_host = "127.0.0.1" _server_port = 4000 # LiteLLM default try: - import litellm.proxy.proxy_server as proxy_server + from litellm.proxy import proxy_server _proxy_host: str | None = getattr(proxy_server, "server_host", None) _proxy_port: str | int | None = getattr(proxy_server, "server_port", None) @@ -1094,14 +1094,14 @@ async def _build_completion_kwargs( params: "CreateMessageRequestParams", model: str, user_api_key_auth: "UserAPIKeyAuth", - raw_headers: Optional[Dict[str, str]], - client_ip: Optional[str], -) -> Dict[str, Any]: + raw_headers: dict[str, str] | None, + client_ip: str | None, +) -> dict[str, Any]: openai_messages = _convert_mcp_messages_to_openai( messages=params.messages, system_prompt=params.systemPrompt, ) - completion_kwargs: Dict[str, Any] = { + completion_kwargs: dict[str, Any] = { "model": model, "messages": openai_messages, "max_tokens": params.maxTokens, @@ -1135,7 +1135,7 @@ async def _build_completion_kwargs( async def _run_guardrails_and_call_llm( - completion_kwargs: Dict[str, Any], + completion_kwargs: dict[str, Any], user_api_key_auth: "UserAPIKeyAuth", ) -> Any: try: @@ -1171,10 +1171,10 @@ async def _run_guardrails_and_call_llm( async def handle_sampling_create_message( context: "RequestContext[ClientSession, object]", params: "CreateMessageRequestParams", - default_model: Optional[str] = None, + default_model: str | None = None, user_api_key_auth: "UserAPIKeyAuth | None" = None, - raw_headers: Optional[Dict[str, str]] = None, - client_ip: Optional[str] = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> Union["CreateMessageResult", "CreateMessageResultWithTools", "ErrorData"]: """ Handle an MCP sampling/createMessage request by routing through LiteLLM. @@ -1292,5 +1292,5 @@ async def handle_sampling_create_message( verbose_logger.exception("MCP sampling handler failed: %s", e) return ErrorData( code=-1, - message=f"Sampling failed: {str(e)}", + message=f"Sampling failed: {e!s}", ) diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index ed78f7c6fb8..a7cc9fe3ed0 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -5,7 +5,7 @@ Filters MCP tools semantically for /chat/completions and /responses endpoints. """ import asyncio -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.exceptions import ContextWindowExceededError @@ -35,7 +35,7 @@ class SemanticToolFilterContextWindowError(Exception): ) -def _is_context_window_error(error: Optional[BaseException]) -> bool: +def _is_context_window_error(error: BaseException | None) -> bool: """Detect a context-window overflow anywhere in an exception's tree.""" if error is None: return False @@ -72,9 +72,9 @@ class SemanticMCPToolFilter: self.similarity_threshold = similarity_threshold self.embedding_model = embedding_model self.router_instance = litellm_router_instance - self.tool_router: Optional["SemanticRouter"] = None - self.context_window_error: Optional[str] = None - self._tool_map: Dict[str, Any] = {} # MCPTool objects or OpenAI function dicts + self.tool_router: SemanticRouter | None = None + self.context_window_error: str | None = None + self._tool_map: dict[str, Any] = {} # MCPTool objects or OpenAI function dicts self._index_sync_lock = asyncio.Lock() async def build_router_from_mcp_registry(self) -> None: @@ -130,7 +130,7 @@ class SemanticMCPToolFilter: return name, description - def _build_router(self, tools: List) -> None: + def _build_router(self, tools: list) -> None: """Build semantic router with tools (MCPTool objects or OpenAI function dicts).""" from semantic_router.routers import SemanticRouter from semantic_router.routers.base import Route @@ -260,9 +260,9 @@ class SemanticMCPToolFilter: async def filter_tools( self, query: str, - available_tools: List[Any], - top_k: Optional[int] = None, - ) -> List[Any]: + available_tools: list[Any], + top_k: int | None = None, + ) -> list[Any]: """ Filter tools semantically based on query. @@ -304,8 +304,7 @@ class SemanticMCPToolFilter: return available_tools limit = top_k or self.top_k - if self.tool_router.top_k < limit: - self.tool_router.top_k = limit + self.tool_router.top_k = max(self.tool_router.top_k, limit) matches = self.tool_router(text=query, limit=limit, route_filter=available_names) matched_tool_names = self._extract_tool_names_from_matches(matches) @@ -333,7 +332,7 @@ class SemanticMCPToolFilter: verbose_logger.error(f"Semantic tool filter failed: {e}", exc_info=True) return available_tools - def _extract_tool_names_from_matches(self, matches) -> List[str]: + def _extract_tool_names_from_matches(self, matches) -> list[str]: """Extract tool names from semantic router match results.""" if not matches: return [] @@ -385,7 +384,7 @@ class SemanticMCPToolFilter: separator = client_name[-len(canonical) - 1] return separator in ("_", "-") - def _get_tools_by_names(self, tool_names: List[str], available_tools: List[Any]) -> List[Any]: + def _get_tools_by_names(self, tool_names: list[str], available_tools: list[Any]) -> list[Any]: """ Get tools from available_tools by their names, preserving the semantic router's ordering. @@ -401,13 +400,13 @@ class SemanticMCPToolFilter: # Exact matches win over suffix matches when both are present, and # each incoming tool is returned at most once even if two canonical # names happen to be tail-compatible with the same incoming name. - available_by_name: Dict[str, Any] = {} + available_by_name: dict[str, Any] = {} for tool in available_tools: client_name, _ = self._extract_tool_info(tool) if client_name and client_name not in available_by_name: available_by_name[client_name] = tool - matched: List[Any] = [] + matched: list[Any] = [] used_ids: set = set() for canonical in tool_names: tool = available_by_name.get(canonical) @@ -417,7 +416,7 @@ class SemanticMCPToolFilter: # "my_search" and "my_tag_search" both end in "search"), # the one closest in length to the canonical is the # least-wrapped and most likely the intended target. - best_name: Optional[str] = None + best_name: str | None = None for client_name in available_by_name: if not self._name_matches_canonical(client_name, canonical): continue @@ -430,7 +429,7 @@ class SemanticMCPToolFilter: used_ids.add(id(tool)) return matched - def extract_user_query(self, messages: List[Dict[str, Any]]) -> str: + def extract_user_query(self, messages: list[dict[str, Any]]) -> str: """ Extract user query from messages for /chat/completions or /responses. diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d8431b9b3bb..fe47c264dfa 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -805,7 +805,7 @@ if MCP_AVAILABLE: } return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta}) except Exception as e: - verbose_logger.exception(f"Error in list_tools endpoint: {str(e)}") + verbose_logger.exception(f"Error in list_tools endpoint: {e!s}") # Return empty list instead of failing completely # This prevents the HTTP stream from failing and allows the client to get a response return [] @@ -1080,26 +1080,26 @@ if MCP_AVAILABLE: isError=True, ) except BlockedPiiEntityError as e: - verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {e!s}") return CallToolResult( content=[ TextContent( - text=f"Error: Blocked PII entity detected - {str(e)}", + text=f"Error: Blocked PII entity detected - {e!s}", type="text", ) ], isError=True, ) except GuardrailRaisedException as e: - verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {e!s}") return CallToolResult( - content=[TextContent(text=f"Error: Guardrail violation - {str(e)}", type="text")], + content=[TextContent(text=f"Error: Guardrail violation - {e!s}", type="text")], isError=True, ) except HTTPException as e: - verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + verbose_logger.error(f"HTTPException in MCP tool call: {e!s}") return CallToolResult( - content=[TextContent(text=f"Error: {str(e.detail)}", type="text")], + content=[TextContent(text=f"Error: {e.detail!s}", type="text")], isError=True, ) except MCPUpstreamAuthError as e: @@ -1121,7 +1121,7 @@ if MCP_AVAILABLE: except Exception as e: verbose_logger.exception(f"MCP mcp_server_tool_call - error: {e}") return CallToolResult( - content=[TextContent(text=f"Error: {str(e)}", type="text")], + content=[TextContent(text=f"Error: {e!s}", type="text")], isError=True, ) @@ -1173,7 +1173,7 @@ if MCP_AVAILABLE: verbose_logger.info(f"MCP list_prompts - Successfully returned {len(prompts)} prompts") return prompts except Exception as e: - verbose_logger.exception(f"Error in list_prompts endpoint: {str(e)}") + verbose_logger.exception(f"Error in list_prompts endpoint: {e!s}") # Return empty list instead of failing completely # This prevents the HTTP stream from failing and allows the client to get a response return [] @@ -1265,7 +1265,7 @@ if MCP_AVAILABLE: verbose_logger.info(f"MCP list_resources - Successfully returned {len(resources)} resources") return resources except Exception as e: - verbose_logger.exception(f"Error in list_resources endpoint: {str(e)}") + verbose_logger.exception(f"Error in list_resources endpoint: {e!s}") return [] finally: if _session_reset_token is not None: @@ -1310,7 +1310,7 @@ if MCP_AVAILABLE: ) return resource_templates except Exception as e: - verbose_logger.exception(f"Error in list_resource_templates endpoint: {str(e)}") + verbose_logger.exception(f"Error in list_resource_templates endpoint: {e!s}") return [] finally: if _session_reset_token is not None: @@ -2036,7 +2036,7 @@ if MCP_AVAILABLE: verbose_logger.debug(f"MCP list_tools: omitting {server.name}; it needs upstream auth") return [], classify_list_exception(e) except Exception as e: - verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}") + verbose_logger.exception(f"Error getting tools from server {server.name}: {e!s}") return [], classify_list_exception(e) # Fetch tools from all servers in parallel @@ -2169,7 +2169,7 @@ if MCP_AVAILABLE: verbose_logger.debug(f"Successfully fetched {len(prompts)} prompts from server {server.name}") except Exception as e: - verbose_logger.exception(f"Error getting prompts from server {server.name}: {str(e)}") + verbose_logger.exception(f"Error getting prompts from server {server.name}: {e!s}") # Continue with other servers instead of failing completely verbose_logger.info(f"Successfully fetched {len(all_prompts)} prompts total from all MCP servers") @@ -2221,7 +2221,7 @@ if MCP_AVAILABLE: verbose_logger.debug(f"Successfully fetched {len(resources)} resources from server {server.name}") except Exception as e: - verbose_logger.exception(f"Error getting resources from server {server.name}: {str(e)}") + verbose_logger.exception(f"Error getting resources from server {server.name}: {e!s}") verbose_logger.info(f"Successfully fetched {len(all_resources)} resources total from all MCP servers") @@ -2359,7 +2359,7 @@ if MCP_AVAILABLE: verbose_logger.debug(f"Successfully fetched {len(listing.tools)} tools from managed MCP servers") return listing except Exception as e: - verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}") + verbose_logger.exception(f"Error getting tools from managed MCP servers: {e!s}") # Continue with an empty listing instead of failing completely return AggregateToolListing(tools=[], outcomes={}) @@ -2398,7 +2398,7 @@ if MCP_AVAILABLE: ) verbose_logger.debug(f"Successfully fetched {len(managed_prompts)} prompts from managed MCP servers") except Exception as e: - verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}") + verbose_logger.exception(f"Error getting tools from managed MCP servers: {e!s}") # Continue with empty managed tools list instead of failing completely return managed_prompts @@ -2428,7 +2428,7 @@ if MCP_AVAILABLE: ) verbose_logger.debug(f"Successfully fetched {len(managed_resources)} resources from managed MCP servers") except Exception as e: - verbose_logger.exception(f"Error getting resources from managed MCP servers: {str(e)}") + verbose_logger.exception(f"Error getting resources from managed MCP servers: {e!s}") return managed_resources @@ -3335,8 +3335,8 @@ if MCP_AVAILABLE: result = tool.handler(**arguments) return [TextContent(text=str(result), type="text")] except Exception as e: - verbose_logger.exception(f"Error executing local tool {name}: {str(e)}") - return [TextContent(text=f"Error: {str(e)}", type="text")] + verbose_logger.exception(f"Error executing local tool {name}: {e!s}") + return [TextContent(text=f"Error: {e!s}", type="text")] def _get_mcp_servers_in_path(path: str) -> list[str] | None: """ diff --git a/litellm/proxy/_experimental/mcp_server/sse_transport.py b/litellm/proxy/_experimental/mcp_server/sse_transport.py index 0a896328dde..09863a7d391 100644 --- a/litellm/proxy/_experimental/mcp_server/sse_transport.py +++ b/litellm/proxy/_experimental/mcp_server/sse_transport.py @@ -11,10 +11,10 @@ from urllib.parse import quote from uuid import UUID, uuid4 import anyio -import mcp.types as types from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from fastapi.requests import Request from fastapi.responses import Response +from mcp import types from pydantic import ValidationError from sse_starlette import EventSourceResponse from starlette.types import Receive, Scope, Send diff --git a/litellm/proxy/_experimental/mcp_server/tool_registry.py b/litellm/proxy/_experimental/mcp_server/tool_registry.py index 27d313e33a8..1ebeac9993a 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_registry.py +++ b/litellm/proxy/_experimental/mcp_server/tool_registry.py @@ -1,6 +1,6 @@ import json from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.proxy.types_utils.utils import get_instance_fn @@ -22,13 +22,13 @@ class MCPToolRegistry: def __init__(self): # Registry to store all registered tools - self.tools: Dict[str, MCPTool] = {} + self.tools: dict[str, MCPTool] = {} def register_tool( self, name: str, description: str, - input_schema: Dict[str, Any], + input_schema: dict[str, Any], handler: Callable, ) -> None: """ @@ -42,13 +42,13 @@ class MCPToolRegistry: ) verbose_logger.debug(f"Registered tool: {name}") - def get_tool(self, name: str) -> Optional[MCPTool]: + def get_tool(self, name: str) -> MCPTool | None: """ Get a tool from the registry by name """ return self.tools.get(name) - def list_tools(self, tool_prefix: Optional[str] = None) -> List[MCPTool]: + def list_tools(self, tool_prefix: str | None = None) -> list[MCPTool]: """ List all registered tools """ @@ -72,7 +72,7 @@ class MCPToolRegistry: verbose_logger.debug("Unregistered MCP tool %s", name) return removed - def convert_tools_to_mcp_sdk_tool_type(self, tools: List[MCPTool]) -> List["MCPToolSDKTool"]: + def convert_tools_to_mcp_sdk_tool_type(self, tools: list[MCPTool]) -> list["MCPToolSDKTool"]: if MCPToolSDKTool is None: raise ImportError("MCP SDK is not installed. Please install it with: pip install 'litellm[proxy]'") return [ @@ -86,8 +86,8 @@ class MCPToolRegistry: def load_tools_from_config( self, - mcp_tools_config: Optional[Dict[str, Any]] = None, - config_file_path: Optional[str] = None, + mcp_tools_config: dict[str, Any] | None = None, + config_file_path: str | None = None, ) -> None: """ Load and register tools from the proxy config diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 2f6b54a264a..34accfb1485 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from datetime import datetime -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from mcp.types import CallToolResult @@ -80,12 +80,12 @@ async def handle_mcp_tool_search( query: str, top_k: int, user_api_key_dict: UserAPIKeyAuth, - client_ip: Optional[str] = None, - mcp_servers: Optional[list[str]] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None, - oauth2_headers: Optional[dict[str, str]] = None, - raw_headers: Optional[dict[str, str]] = None, + client_ip: str | None = None, + mcp_servers: list[str] | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, ) -> CallToolResult: from mcp.types import CallToolResult, TextContent @@ -117,13 +117,13 @@ async def handle_mcp_tool_call( tool_name: str, arguments: dict[str, Any], user_api_key_dict: UserAPIKeyAuth, - client_ip: Optional[str] = None, - mcp_servers: Optional[list[str]] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None, - oauth2_headers: Optional[dict[str, str]] = None, - raw_headers: Optional[dict[str, str]] = None, - litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, + client_ip: str | None = None, + mcp_servers: list[str] | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + litellm_logging_obj: LiteLLMLoggingObj | None = None, ) -> CallToolResult: from litellm.proxy._experimental.mcp_server.server import ( _get_allowed_mcp_servers, diff --git a/litellm/proxy/_experimental/mcp_server/toolset_db.py b/litellm/proxy/_experimental/mcp_server/toolset_db.py index 9652a3a2888..62733edf378 100644 --- a/litellm/proxy/_experimental/mcp_server/toolset_db.py +++ b/litellm/proxy/_experimental/mcp_server/toolset_db.py @@ -1,5 +1,4 @@ import json -from typing import List, Optional from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -38,7 +37,7 @@ async def create_mcp_toolset( async def get_mcp_toolset( prisma_client: PrismaClient, toolset_id: str, -) -> Optional[MCPToolset]: +) -> MCPToolset | None: row = await MCPToolsetRepository(prisma_client).table.find_unique(where={"toolset_id": toolset_id}) if row is None: return None @@ -47,8 +46,8 @@ async def get_mcp_toolset( async def list_mcp_toolsets( prisma_client: PrismaClient, - toolset_ids: Optional[List[str]] = None, -) -> List[MCPToolset]: + toolset_ids: list[str] | None = None, +) -> list[MCPToolset]: try: where = {} if toolset_ids is not None: @@ -56,16 +55,14 @@ async def list_mcp_toolsets( rows = await MCPToolsetRepository(prisma_client).table.find_many(where=where) return [_toolset_from_row(r) for r in rows] except Exception as e: - verbose_proxy_logger.warning( - "litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - {}".format(str(e)) - ) + verbose_proxy_logger.warning(f"litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - {e!s}") return [] async def get_mcp_toolset_by_name( prisma_client: PrismaClient, toolset_name: str, -) -> Optional[MCPToolset]: +) -> MCPToolset | None: row = await MCPToolsetRepository(prisma_client).table.find_first(where={"toolset_name": toolset_name}) if row is None: return None @@ -76,7 +73,7 @@ async def update_mcp_toolset( prisma_client: PrismaClient, data: UpdateMCPToolsetRequest, touched_by: str, -) -> Optional[MCPToolset]: +) -> MCPToolset | None: data_dict = data.model_dump(exclude_none=True, exclude={"toolset_id"}) if "tools" in data_dict: data_dict["tools"] = json.dumps(data_dict["tools"]) @@ -98,7 +95,7 @@ async def update_mcp_toolset( async def delete_mcp_toolset( prisma_client: PrismaClient, toolset_id: str, -) -> Optional[MCPToolset]: +) -> MCPToolset | None: try: row = await MCPToolsetRepository(prisma_client).table.delete(where={"toolset_id": toolset_id}) except Exception as e: diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index d1d28574988..14726cba3a7 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -2,8 +2,6 @@ from __future__ import annotations -from typing import List - from fastapi import HTTPException from litellm._logging import verbose_logger @@ -34,7 +32,7 @@ def is_ui_session_credential(user_api_key_auth: UserAPIKeyAuth) -> bool: async def resolve_ui_session_team_ids( user_api_key_auth: UserAPIKeyAuth, -) -> List[str]: +) -> list[str]: """Resolve the real team ids backing a UI session token.""" if not is_ui_session_credential(user_api_key_auth): @@ -70,7 +68,7 @@ async def resolve_ui_session_team_ids( if user_obj is None or not user_obj.teams: return [] - resolved_team_ids: List[str] = [] + resolved_team_ids: list[str] = [] for team_id in user_obj.teams: if team_id and team_id not in resolved_team_ids: resolved_team_ids.append(team_id) @@ -122,7 +120,7 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth: async def build_effective_auth_contexts( user_api_key_auth: UserAPIKeyAuth, -) -> List[UserAPIKeyAuth]: +) -> list[UserAPIKeyAuth]: """Every auth context a management or listing surface must resolve a UI session token through: one per real team backing the session, plus the session user's own admitted identity, so a grant made directly to the user row is as visible to the dashboard as it is to a gateway session.""" diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 06ff9034a44..85effdd0f58 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -10,12 +10,6 @@ import re from collections.abc import Iterable, Iterator, Mapping, MutableMapping, MutableSequence from typing import ( Any, - Dict, - List, - Optional, - Set, - Tuple, - Union, ) from urllib.parse import quote @@ -152,11 +146,11 @@ def sanitize_mcp_alias_for_header(alias: str) -> str: def lookup_mcp_server_auth_in_headers( - mcp_server_auth_headers: Mapping[str, Union[str, Dict[str, str]]], + mcp_server_auth_headers: Mapping[str, str | dict[str, str]], *, - alias: Optional[str] = None, - server_name: Optional[str] = None, -) -> Optional[Union[str, Dict[str, str]]]: + alias: str | None = None, + server_name: str | None = None, +) -> str | dict[str, str] | None: """ Resolve server-specific auth headers with case-insensitive matching. @@ -184,7 +178,7 @@ def lookup_mcp_server_auth_in_headers( MCP_TOOL_ALLOWLIST_ENFORCED_KEY = "tool_allowlist_enforced" -def _parse_mcp_info_dict(mcp_info: Any) -> Optional[Dict[str, Any]]: +def _parse_mcp_info_dict(mcp_info: Any) -> dict[str, Any] | None: if mcp_info is None: return None if isinstance(mcp_info, dict): @@ -304,7 +298,7 @@ def iter_known_server_prefixes(server: Any) -> Iterator[str]: """ seen = set() - def _emit(value: Optional[str]) -> Iterator[str]: + def _emit(value: str | None) -> Iterator[str]: if value and value not in seen: seen.add(value) yield value @@ -367,7 +361,7 @@ def match_known_tool_name(tool_name: str, server: MCPServer, names: Iterable[str return next((name for name in names if normalize(name) in spellings), None) -def split_server_prefix_from_name(prefixed_name: str) -> Tuple[str, str]: +def split_server_prefix_from_name(prefixed_name: str) -> tuple[str, str]: """Return the unprefixed name plus the server name used as prefix. Cuts at the FIRST separator, so the two halves are only trustworthy as a @@ -405,7 +399,7 @@ def match_known_server_prefix(name: str, known_prefixes: Iterable[str]) -> tuple return None -def strip_known_server_prefix(name: str, server: Optional[Any]) -> str: +def strip_known_server_prefix(name: str, server: Any | None) -> str: """Strip ``server``'s registered prefix from a prefixed tool/resource name. Unlike :func:`split_server_prefix_from_name`, which guesses the boundary at @@ -428,7 +422,7 @@ def strip_known_server_prefix(name: str, server: Optional[Any]) -> str: def is_tool_name_prefixed( tool_name: str, - known_server_prefixes: Optional[set] = None, + known_server_prefixes: set | None = None, ) -> bool: """ Check if tool name has a known MCP server prefix. @@ -483,7 +477,7 @@ def validate_mcp_server_name(server_name: str, raise_http_exception: bool = Fals raise Exception(error_message) -def extract_mcp_tool_result_error_message(result: object) -> Optional[str]: +def extract_mcp_tool_result_error_message(result: object) -> str | None: """The first text content of an ``isError=True`` tool result, or ``None`` when the result is not an error. @@ -555,7 +549,7 @@ def with_mcp_content_item_text(item: object, text: str) -> object: TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$") -def validate_tool_display_names(tool_name_to_display_name: Optional[Mapping[str, str]]) -> None: +def validate_tool_display_names(tool_name_to_display_name: Mapping[str, str] | None) -> None: """ Validate tool display name overrides against Bedrock's tool-name constraint. @@ -600,8 +594,8 @@ class MCPMissingUserEnvVarsError(Exception): self, *, server_id: str, - server_name: Optional[str], - missing: List[str], + server_name: str | None, + missing: list[str], setup_url: str, ) -> None: self.server_id = server_id @@ -627,8 +621,8 @@ _ENV_VAR_PATTERN = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}") def parse_admin_env_vars( - env_vars: Optional[Iterable[Any]], -) -> Tuple[Dict[str, str], List[Dict[str, Any]]]: + env_vars: Iterable[Any] | None, +) -> tuple[dict[str, str], list[dict[str, Any]]]: """Split admin-configured env var entries into globals and per-user specs. Accepts the raw value of ``MCPServer.env_vars`` (list of dicts or Pydantic @@ -640,8 +634,8 @@ def parse_admin_env_vars( Unknown / malformed entries are skipped silently. """ - global_values: Dict[str, str] = {} - user_specs: List[Dict[str, Any]] = [] + global_values: dict[str, str] = {} + user_specs: list[dict[str, Any]] = [] if not env_vars: return global_values, user_specs for raw in env_vars: @@ -665,16 +659,16 @@ def parse_admin_env_vars( return global_values, user_specs -def find_env_var_references(value: str) -> Set[str]: +def find_env_var_references(value: str) -> set[str]: """Return the set of ``${NAME}`` identifiers referenced inside ``value``.""" if not value: return set() return set(_ENV_VAR_PATTERN.findall(value)) -def collect_env_var_references(*, strings: Iterable[str]) -> Set[str]: +def collect_env_var_references(*, strings: Iterable[str]) -> set[str]: """Union of every ``${NAME}`` reference across a collection of strings.""" - refs: Set[str] = set() + refs: set[str] = set() for s in strings: if isinstance(s, str): refs |= find_env_var_references(s) @@ -698,7 +692,7 @@ def interpolate_env_vars(value: str, variables: Mapping[str, str]) -> str: return _ENV_VAR_PATTERN.sub(_sub, value) -def interpolate_headers(headers: Mapping[str, str], variables: Mapping[str, str]) -> Dict[str, str]: +def interpolate_headers(headers: Mapping[str, str], variables: Mapping[str, str]) -> dict[str, str]: """Return a copy of ``headers`` with every value passed through ``interpolate_env_vars``.""" return {k: interpolate_env_vars(v, variables) for k, v in headers.items()} @@ -712,9 +706,9 @@ def build_env_var_setup_url(server_id: str) -> str: def merge_mcp_headers( *, - extra_headers: Optional[Mapping[str, str]] = None, - static_headers: Optional[Mapping[str, str]] = None, -) -> Optional[Dict[str, str]]: + extra_headers: Mapping[str, str] | None = None, + static_headers: Mapping[str, str] | None = None, +) -> dict[str, str] | None: """Merge outbound HTTP headers for MCP calls. This is used when calling out to external MCP servers (or OpenAPI-based MCP tools). @@ -727,7 +721,7 @@ def merge_mcp_headers( behavior in `MCPServerManager` where `server.static_headers` is applied after any caller-provided headers. """ - merged: Dict[str, str] = {} + merged: dict[str, str] = {} if extra_headers: merged.update({str(k): str(v) for k, v in extra_headers.items()}) diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 6571c3ec4b5..97ac9b609f9 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -11,7 +11,7 @@ import importlib import sys from collections.abc import Callable from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Dict, Tuple +from typing import TYPE_CHECKING from starlette.types import Receive, Scope, Send @@ -39,11 +39,11 @@ def _mount_app(prefix: str, attr_name: str = "app") -> Callable[["FastAPI", obje class LazyFeature: name: str module_path: str - path_prefixes: Tuple[str, ...] + path_prefixes: tuple[str, ...] register_fn: Callable[["FastAPI", object], None] = field(default_factory=lambda: _include_router("router")) # For routes whose path has a leading parameter (e.g. /{server}/authorize) # — startswith can't match those, so the matcher also checks endswith. - path_suffixes: Tuple[str, ...] = () + path_suffixes: tuple[str, ...] = () # Keep the stub injected even after load — for mounted ASGI sub-apps # whose routes don't appear in the parent app's openapi spec. persistent_swagger_stub: bool = False @@ -52,7 +52,7 @@ class LazyFeature: return any(path.startswith(p) for p in self.path_prefixes) or any(path.endswith(s) for s in self.path_suffixes) -LAZY_FEATURES: Tuple[LazyFeature, ...] = ( +LAZY_FEATURES: tuple[LazyFeature, ...] = ( LazyFeature( name="guardrails", module_path="litellm.proxy.guardrails.guardrail_endpoints", @@ -265,7 +265,7 @@ class LazyFeatureMiddleware: self, app, fastapi_app: "FastAPI", - features: Tuple[LazyFeature, ...] = LAZY_FEATURES, + features: tuple[LazyFeature, ...] = LAZY_FEATURES, ): self.app = app self._fastapi_app = fastapi_app @@ -395,7 +395,7 @@ def _make_warmup_router(app: "FastAPI") -> "APIRouter": return router -def inject_lazy_stubs(schema: Dict) -> Dict: +def inject_lazy_stubs(schema: dict) -> dict: """Inject openapi entries for unloaded features. Uses the snapshot file when available (full route info), otherwise falls back to a single placeholder per feature. Any failure logs and returns the schema unchanged @@ -434,7 +434,7 @@ def inject_lazy_stubs(schema: Dict) -> Dict: return schema -def lazy_tag_to_prefix() -> Dict[str, str]: +def lazy_tag_to_prefix() -> dict[str, str]: """feature.name -> first prefix, used by the Swagger warmup JS plugin. Returns empty when the snapshot is loaded — the plugin is unnecessary because /openapi.json already has full route info.""" diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 818232d650e..56c08d4176b 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -11,7 +11,6 @@ import json import re import sys from pathlib import Path -from typing import Dict, Optional, Set SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" HTTP_METHOD_SUFFIXES = { @@ -39,7 +38,7 @@ def _stabilize_multi_method_route_ids(routes) -> None: route.unique_id = f"{operation_id}_{methods[0].lower()}" -def load_snapshot() -> Optional[Dict[str, Dict]]: +def load_snapshot() -> dict[str, dict] | None: if not SNAPSHOT_FILE.exists(): return None try: @@ -49,7 +48,7 @@ def load_snapshot() -> Optional[Dict[str, Dict]]: return None -def _normalize_operation_ids(paths: Dict[str, Dict]) -> None: +def _normalize_operation_ids(paths: dict[str, dict]) -> None: """Make FastAPI-generated operation IDs stable for multi-method routes. FastAPI derives the default operation ID suffix from the first item in the @@ -80,7 +79,7 @@ def _normalize_operation_ids(paths: Dict[str, Dict]) -> None: break -def generate_snapshot() -> Dict[str, Dict]: +def generate_snapshot() -> dict[str, dict]: import importlib from fastapi.openapi.utils import get_openapi @@ -97,8 +96,8 @@ def generate_snapshot() -> Dict[str, Dict]: except Exception as exc: sys.stderr.write(f"warning: skip {feat.name}: {exc}\n") - fragments: Dict[str, Dict] = {} - used_operation_ids: Set[str] = set() + fragments: dict[str, dict] = {} + used_operation_ids: set[str] = set() for feat in LAZY_FEATURES: feat_routes = [r for r in app.routes if any(getattr(r, "path", "").startswith(p) for p in feat.path_prefixes)] if not feat_routes: diff --git a/litellm/proxy/_logging.py b/litellm/proxy/_logging.py index dc6b34fd360..8f3cb08eb5d 100644 --- a/litellm/proxy/_logging.py +++ b/litellm/proxy/_logging.py @@ -14,7 +14,7 @@ numeric_level: str = getattr(logging, log_level.upper()) class JsonFormatter(Formatter): def __init__(self): - super(JsonFormatter, self).__init__() + super().__init__() def format(self, record): json_record = { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 077c4a624f5..3ccf3ea9952 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3,7 +3,7 @@ import json import os from collections.abc import Callable from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Literal, Union import httpx from pydantic import ( @@ -141,7 +141,7 @@ class LitellmUserRoles(str, enum.Enum): def __str__(self): return str(self.value) - def values(self) -> List[str]: + def values(self) -> list[str]: return list(self.__annotations__.keys()) @property @@ -881,10 +881,10 @@ class LiteLLMPromptInjectionParams(LiteLLMPydanticObjectBase): heuristics_check: bool = False vector_db_check: bool = False llm_api_check: bool = False - llm_api_name: Optional[str] = None - llm_api_system_prompt: Optional[str] = None - llm_api_fail_call_string: Optional[str] = None - reject_as_response: Optional[bool] = Field( + llm_api_name: str | None = None + llm_api_system_prompt: str | None = None + llm_api_fail_call_string: str | None = None + reject_as_response: bool | None = Field( default=False, description="Return rejected request error message as a string to the user. Default behaviour is to raise an exception.", ) @@ -912,40 +912,40 @@ class ProxyChatCompletionRequest(LiteLLMPydanticObjectBase): # Required fields (from ChatCompletionRequest) model: str - messages: List[AllMessageValues] + messages: list[AllMessageValues] # Standard OpenAI completion parameters (all optional) - frequency_penalty: Optional[float] = None - logit_bias: Optional[Dict[str, float]] = None - logprobs: Optional[bool] = None - top_logprobs: Optional[int] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[float] = None - response_format: Optional[Dict[str, Any]] = None - seed: Optional[int] = None - service_tier: Optional[str] = None - stop: Optional[Union[str, List[str]]] = None - stream_options: Optional[Dict[str, Any]] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - tools: Optional[List[Dict[str, Any]]] = None - tool_choice: Optional[Union[str, Dict[str, Any]]] = None - parallel_tool_calls: Optional[bool] = None - function_call: Optional[Union[str, Dict[str, Any]]] = None - functions: Optional[List[Dict[str, Any]]] = None - user: Optional[str] = None - stream: Optional[bool] = None + frequency_penalty: float | None = None + logit_bias: dict[str, float] | None = None + logprobs: bool | None = None + top_logprobs: int | None = None + max_tokens: int | None = None + n: int | None = None + presence_penalty: float | None = None + response_format: dict[str, Any] | None = None + seed: int | None = None + service_tier: str | None = None + stop: str | list[str] | None = None + stream_options: dict[str, Any] | None = None + temperature: float | None = None + top_p: float | None = None + tools: list[dict[str, Any]] | None = None + tool_choice: str | dict[str, Any] | None = None + parallel_tool_calls: bool | None = None + function_call: str | dict[str, Any] | None = None + functions: list[dict[str, Any]] | None = None + user: str | None = None + stream: bool | None = None # LiteLLM-specific metadata param (from original ChatCompletionRequest) - metadata: Optional[Dict[str, Any]] = None + metadata: dict[str, Any] | None = None # Optional LiteLLM params - guardrails: Optional[List[str]] = None - caching: Optional[bool] = None - num_retries: Optional[int] = None - context_window_fallback_dict: Optional[Dict[str, str]] = None - fallbacks: Optional[List[str]] = None + guardrails: list[str] | None = None + caching: bool | None = None + num_retries: int | None = None + context_window_fallback_dict: dict[str, str] | None = None + fallbacks: list[str] | None = None class ModelInfoDelete(LiteLLMPydanticObjectBase): @@ -953,24 +953,20 @@ class ModelInfoDelete(LiteLLMPydanticObjectBase): class ModelInfo(LiteLLMPydanticObjectBase): - id: Optional[str] - mode: Optional[Literal["embedding", "chat", "completion"]] - input_cost_per_token: Optional[float] = 0.0 - output_cost_per_token: Optional[float] = 0.0 - max_tokens: Optional[int] = 2048 # assume 2048 if not set + id: str | None + mode: Literal["embedding", "chat", "completion"] | None + input_cost_per_token: float | None = 0.0 + output_cost_per_token: float | None = 0.0 + max_tokens: int | None = 2048 # assume 2048 if not set # for azure models we need users to specify the base model, one azure you can call deployments - azure/my-random-model # we look up the base model in model_prices_and_context_window.json - base_model: Optional[ + base_model: ( Literal[ - "gpt-4-1106-preview", - "gpt-4-32k", - "gpt-4", - "gpt-3.5-turbo-16k", - "gpt-3.5-turbo", - "text-embedding-ada-002", + "gpt-4-1106-preview", "gpt-4-32k", "gpt-4", "gpt-3.5-turbo-16k", "gpt-3.5-turbo", "text-embedding-ada-002" ] - ] + | None + ) model_config = ConfigDict(protected_namespaces=(), extra="allow") @@ -994,11 +990,11 @@ class ModelInfo(LiteLLMPydanticObjectBase): class ProviderInfo(LiteLLMPydanticObjectBase): name: str - fields: List[ProviderField] + fields: list[ProviderField] class BlockUsers(LiteLLMPydanticObjectBase): - user_ids: List[str] # required + user_ids: list[str] # required class ModelParams(LiteLLMPydanticObjectBase): @@ -1017,17 +1013,17 @@ class ModelParams(LiteLLMPydanticObjectBase): class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): - mcp_servers: Optional[List[str]] = None - mcp_access_groups: Optional[List[str]] = None - mcp_tool_permissions: Optional[Dict[str, List[str]]] = None - mcp_toolsets: Optional[List[str]] = None - blocked_tools: Optional[List[str]] = None - vector_stores: Optional[List[str]] = None - agents: Optional[List[str]] = None - agent_access_groups: Optional[List[str]] = None - models: Optional[List[str]] = None - search_tools: Optional[List[str]] = None - mcp_tool_search_enabled: Optional[bool] = None + mcp_servers: list[str] | None = None + mcp_access_groups: list[str] | None = None + mcp_tool_permissions: dict[str, list[str]] | None = None + mcp_toolsets: list[str] | None = None + blocked_tools: list[str] | None = None + vector_stores: list[str] | None = None + agents: list[str] | None = None + agent_access_groups: list[str] | None = None + models: list[str] | None = None + search_tools: list[str] | None = None + mcp_tool_search_enabled: bool | None = None from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 @@ -1041,38 +1037,38 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): Overlapping schema between key and user generate/update requests """ - key_alias: Optional[str] = None - duration: Optional[str] = None - models: Optional[list] = [] - spend: Optional[float] = 0 - max_budget: Optional[float] = None - user_id: Optional[str] = None - team_id: Optional[str] = None - agent_id: Optional[str] = None - max_parallel_requests: Optional[int] = None - metadata: Optional[dict] = {} - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None + key_alias: str | None = None + duration: str | None = None + models: list | None = [] + spend: float | None = 0 + max_budget: float | None = None + user_id: str | None = None + team_id: str | None = None + agent_id: str | None = None + max_parallel_requests: int | None = None + metadata: dict | None = {} + tpm_limit: int | None = None + rpm_limit: int | None = None - budget_duration: Optional[str] = None - budget_limits: Optional[List[BudgetLimitEntry]] = None # multiple concurrent budget windows - allowed_cache_controls: Optional[list] = [] - config: Optional[dict] = {} - permissions: Optional[dict] = {} - model_max_budget: Optional[dict] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} - budget_fallbacks: Optional[dict[str, list[str]]] = None + budget_duration: str | None = None + budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows + allowed_cache_controls: list | None = [] + config: dict | None = {} + permissions: dict | None = {} + model_max_budget: dict | None = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} + budget_fallbacks: dict[str, list[str]] | None = None model_config = ConfigDict(protected_namespaces=()) - model_rpm_limit: Optional[dict] = None - model_tpm_limit: Optional[dict] = None - mcp_rpm_limit: Optional[Dict[str, int]] = None - tag_rpm_limit: Optional[dict[str, int]] = None - guardrails: Optional[List[str]] = None - policies: Optional[List[str]] = None - prompts: Optional[List[str]] = None - blocked: Optional[bool] = None - aliases: Optional[dict] = {} - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + model_rpm_limit: dict | None = None + model_tpm_limit: dict | None = None + mcp_rpm_limit: dict[str, int] | None = None + tag_rpm_limit: dict[str, int] | None = None + guardrails: list[str] | None = None + policies: list[str] | None = None + prompts: list[str] | None = None + blocked: bool | None = None + aliases: dict | None = {} + object_permission: LiteLLM_ObjectPermissionBase | None = None @field_validator("max_budget", mode="before") @classmethod @@ -1084,27 +1080,27 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase): index_name: str - index_permissions: List[Literal["read", "write"]] + index_permissions: list[Literal["read", "write"]] class KeyRequestBase(GenerateRequestBase): - key: Optional[str] = None - budget_id: Optional[str] = None - tags: Optional[List[str]] = None - disable_global_guardrails: Optional[bool] = None - throttle_on_budget_exceeded: Optional[bool] = None - enforced_params: Optional[List[str]] = None - allowed_routes: Optional[list] = [] - allowed_passthrough_routes: Optional[list] = None - allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None - rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]] = ( + key: str | None = None + budget_id: str | None = None + tags: list[str] | None = None + disable_global_guardrails: bool | None = None + throttle_on_budget_exceeded: bool | None = None + enforced_params: list[str] | None = None + allowed_routes: list | None = [] + allowed_passthrough_routes: list | None = None + allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None + rpm_limit_type: Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"] | None = ( None # raise an error if 'guaranteed_throughput' is set and we're overallocating rpm ) - tpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]] = ( + tpm_limit_type: Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"] | None = ( None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm ) - router_settings: Optional[UpdateRouterConfig] = None - access_group_ids: Optional[List[str]] = None + router_settings: UpdateRouterConfig | None = None + access_group_ids: list[str] | None = None class LiteLLMKeyType(str, enum.Enum): @@ -1119,36 +1115,36 @@ class LiteLLMKeyType(str, enum.Enum): class GenerateKeyRequest(KeyRequestBase): - soft_budget: Optional[float] = None - send_invite_email: Optional[bool] = None - key_type: Optional[LiteLLMKeyType] = Field( + soft_budget: float | None = None + send_invite_email: bool | None = None + key_type: LiteLLMKeyType | None = Field( default=LiteLLMKeyType.DEFAULT, description="Type of key that determines default allowed routes.", ) - auto_rotate: Optional[bool] = Field(default=False, description="Whether this key should be automatically rotated") - rotation_interval: Optional[str] = Field( + auto_rotate: bool | None = Field(default=False, description="Whether this key should be automatically rotated") + rotation_interval: str | None = Field( default=None, description="How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True", ) - organization_id: Optional[str] = None - project_id: Optional[str] = None + organization_id: str | None = None + project_id: str | None = None class GenerateKeyResponse(KeyRequestBase): key: str # type: ignore - key_name: Optional[str] = None + key_name: str | None = None key_type: str | None = None - expires: Optional[datetime] = None - user_id: Optional[str] = None - token_id: Optional[str] = None - organization_id: Optional[str] = None - project_id: Optional[str] = None - litellm_budget_table: Optional[Any] = None - token: Optional[str] = None - created_by: Optional[str] = None - updated_by: Optional[str] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None + expires: datetime | None = None + user_id: str | None = None + token_id: str | None = None + organization_id: str | None = None + project_id: str | None = None + litellm_budget_table: Any | None = None + token: str | None = None + created_by: str | None = None + updated_by: str | None = None + created_at: datetime | None = None + updated_at: datetime | None = None @model_validator(mode="before") @classmethod @@ -1179,14 +1175,14 @@ class GenerateKeyResponse(KeyRequestBase): class UpdateKeyRequest(KeyRequestBase): # Note: the defaults of all Params here MUST BE NONE # else they will get overwritten - duration: Optional[str] = None - spend: Optional[float] = None - metadata: Optional[dict] = None - temp_budget_increase: Optional[float] = None - temp_budget_expiry: Optional[datetime] = None - auto_rotate: Optional[bool] = None - rotation_interval: Optional[str] = None - organization_id: Optional[str] = None + duration: str | None = None + spend: float | None = None + metadata: dict | None = None + temp_budget_increase: float | None = None + temp_budget_expiry: datetime | None = None + auto_rotate: bool | None = None + rotation_interval: str | None = None + organization_id: str | None = None @model_validator(mode="after") def validate_temp_budget(self) -> "UpdateKeyRequest": @@ -1204,13 +1200,13 @@ class UpdateKeyRequest(KeyRequestBase): class RegenerateKeyRequest(GenerateKeyRequest): # This needs to be different from UpdateKeyRequest, because "key" is optional for this - key: Optional[str] = None - new_key: Optional[str] = None - duration: Optional[str] = None - spend: Optional[float] = None - metadata: Optional[dict] = None - new_master_key: Optional[str] = None - grace_period: Optional[str] = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke + key: str | None = None + new_key: str | None = None + duration: str | None = None + spend: float | None = None + metadata: dict | None = None + new_master_key: str | None = None + grace_period: str | None = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke class ResetSpendRequest(LiteLLMPydanticObjectBase): @@ -1218,8 +1214,8 @@ class ResetSpendRequest(LiteLLMPydanticObjectBase): class KeyRequest(LiteLLMPydanticObjectBase): - keys: Optional[List[str]] = None - key_aliases: Optional[List[str]] = None + keys: list[str] | None = None + key_aliases: list[str] | None = None @model_validator(mode="before") @classmethod @@ -1265,63 +1261,63 @@ def _dcr_bridge_auth_type_error(auth_type: object) -> ValueError: class NewMCPServerRequest(LiteLLMPydanticObjectBase): - server_id: Optional[str] = None - server_name: Optional[str] = None - alias: Optional[str] = None - description: Optional[str] = None + server_id: str | None = None + server_name: str | None = None + alias: str | None = None + description: str | None = None transport: MCPTransportType = MCPTransport.sse - auth_type: Optional[MCPAuthType] = None - credentials: Optional[MCPCredentials] = None - url: Optional[str] = None - spec_path: Optional[str] = None - mcp_info: Optional[MCPInfo] = None - mcp_access_groups: List[str] = Field(default_factory=list) - allowed_tools: Optional[List[str]] = None - tool_name_to_display_name: Optional[Dict[str, str]] = None - tool_name_to_description: Optional[Dict[str, str]] = None - extra_headers: Optional[List[str]] = None - static_headers: Optional[Dict[str, str]] = None - env_vars: Optional[List[MCPEnvVar]] = None - instructions: Optional[str] = None + auth_type: MCPAuthType | None = None + credentials: MCPCredentials | None = None + url: str | None = None + spec_path: str | None = None + mcp_info: MCPInfo | None = None + mcp_access_groups: list[str] = Field(default_factory=list) + allowed_tools: list[str] | None = None + tool_name_to_display_name: dict[str, str] | None = None + tool_name_to_description: dict[str, str] | None = None + extra_headers: list[str] | None = None + static_headers: dict[str, str] | None = None + env_vars: list[MCPEnvVar] | None = None + instructions: str | None = None # Stdio-specific fields - command: Optional[str] = None - args: List[str] = Field(default_factory=list) - env: Dict[str, str] = Field(default_factory=dict) - issuer: Optional[str] = None - authorization_url: Optional[str] = None - token_url: Optional[str] = None - registration_url: Optional[str] = None - oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None + command: str | None = None + args: list[str] = Field(default_factory=list) + env: dict[str, str] = Field(default_factory=dict) + issuer: str | None = None + authorization_url: str | None = None + token_url: str | None = None + registration_url: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None # Token Exchange (OBO) fields — RFC 8693. These top-level fields are the # canonical shape; the same keys inside ``credentials`` are the legacy # pre-column REST shape and are lifted into these columns on write (an # explicit top-level value wins) and stripped from the stored blob. - token_exchange_endpoint: Optional[str] = None - audience: Optional[str] = None - subject_token_type: Optional[str] = None - token_exchange_profile: Optional[str] = None + token_exchange_endpoint: str | None = None + audience: str | None = None + subject_token_type: str | None = None + token_exchange_profile: str | None = None allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False oauth_passthrough: bool = False - dcr_bridge: Optional[bool] = None + dcr_bridge: bool | None = None is_byok: bool = False - byok_description: List[str] = Field(default_factory=list) - byok_api_key_help_url: Optional[str] = None - source_url: Optional[str] = None - timeout: Optional[float] = None - max_concurrent_requests: Optional[int] = None + byok_description: list[str] = Field(default_factory=list) + byok_api_key_help_url: str | None = None + source_url: str | None = None + timeout: float | None = None + max_concurrent_requests: int | None = None # BYOM submission fields — set by the endpoint, not by the caller. # Any caller-provided values are silently overridden before persistence. - approval_status: Optional[str] = Field( + approval_status: str | None = Field( None, description="Server-managed: set by the endpoint; caller values are overridden.", ) - submitted_by: Optional[str] = Field( + submitted_by: str | None = Field( None, description="Server-managed: set by the endpoint; caller values are overridden.", ) - submitted_at: Optional[datetime] = Field( + submitted_at: datetime | None = Field( None, description="Server-managed: set by the endpoint; caller values are overridden.", ) @@ -1372,51 +1368,51 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): server_id: str - server_name: Optional[str] = None - alias: Optional[str] = None - description: Optional[str] = None + server_name: str | None = None + alias: str | None = None + description: str | None = None transport: MCPTransportType = MCPTransport.sse - auth_type: Optional[MCPAuthType] = None - credentials: Optional[MCPCredentials] = None - url: Optional[str] = None - spec_path: Optional[str] = None - mcp_info: Optional[MCPInfo] = None - mcp_access_groups: List[str] = Field(default_factory=list) - allowed_tools: Optional[List[str]] = None - tool_name_to_display_name: Optional[Dict[str, str]] = None - tool_name_to_description: Optional[Dict[str, str]] = None - extra_headers: Optional[List[str]] = None - static_headers: Optional[Dict[str, str]] = None - env_vars: Optional[List[MCPEnvVar]] = None - instructions: Optional[str] = None + auth_type: MCPAuthType | None = None + credentials: MCPCredentials | None = None + url: str | None = None + spec_path: str | None = None + mcp_info: MCPInfo | None = None + mcp_access_groups: list[str] = Field(default_factory=list) + allowed_tools: list[str] | None = None + tool_name_to_display_name: dict[str, str] | None = None + tool_name_to_description: dict[str, str] | None = None + extra_headers: list[str] | None = None + static_headers: dict[str, str] | None = None + env_vars: list[MCPEnvVar] | None = None + instructions: str | None = None # Stdio-specific fields - command: Optional[str] = None - args: List[str] = Field(default_factory=list) - env: Dict[str, str] = Field(default_factory=dict) - issuer: Optional[str] = None - authorization_url: Optional[str] = None - token_url: Optional[str] = None - registration_url: Optional[str] = None - oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None + command: str | None = None + args: list[str] = Field(default_factory=list) + env: dict[str, str] = Field(default_factory=dict) + issuer: str | None = None + authorization_url: str | None = None + token_url: str | None = None + registration_url: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None # Token Exchange (OBO) fields — RFC 8693. These top-level fields are the # canonical shape; the same keys inside ``credentials`` are the legacy # pre-column REST shape and are lifted into these columns on write (an # explicit top-level value wins) and stripped from the stored blob. - token_exchange_endpoint: Optional[str] = None - audience: Optional[str] = None - subject_token_type: Optional[str] = None - token_exchange_profile: Optional[str] = None + token_exchange_endpoint: str | None = None + audience: str | None = None + subject_token_type: str | None = None + token_exchange_profile: str | None = None allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False oauth_passthrough: bool = False - dcr_bridge: Optional[bool] = None + dcr_bridge: bool | None = None is_byok: bool = False - byok_description: List[str] = Field(default_factory=list) - byok_api_key_help_url: Optional[str] = None - source_url: Optional[str] = None - timeout: Optional[float] = None - max_concurrent_requests: Optional[int] = None + byok_description: list[str] = Field(default_factory=list) + byok_api_key_help_url: str | None = None + source_url: str | None = None + timeout: float | None = None + max_concurrent_requests: int | None = None @model_validator(mode="before") @classmethod @@ -1462,7 +1458,7 @@ from litellm.models.mcp_server import ( # noqa: E402 class MakeMCPServersPublicRequest(LiteLLMPydanticObjectBase): - mcp_server_ids: List[str] + mcp_server_ids: list[str] class MCPUserCredentialRequest(LiteLLMPydanticObjectBase): @@ -1479,9 +1475,9 @@ class MCPOAuthUserCredentialRequest(LiteLLMPydanticObjectBase): """Stores a user's OAuth2 token for an OpenAPI MCP server.""" access_token: str - refresh_token: Optional[str] = None - expires_in: Optional[int] = None # seconds until expiry - scopes: Optional[List[str]] = None + refresh_token: str | None = None + expires_in: int | None = None # seconds until expiry + scopes: list[str] | None = None class MCPOAuthUserCredentialStatus(LiteLLMPydanticObjectBase): @@ -1489,27 +1485,27 @@ class MCPOAuthUserCredentialStatus(LiteLLMPydanticObjectBase): server_id: str has_credential: bool - expires_at: Optional[str] = None # ISO-8601 + expires_at: str | None = None # ISO-8601 is_expired: bool = False - connected_at: Optional[str] = None # ISO-8601 + connected_at: str | None = None # ISO-8601 class MCPUserCredentialListItem(LiteLLMPydanticObjectBase): """One entry in the /user-credentials list.""" server_id: str - server_name: Optional[str] = None - alias: Optional[str] = None + server_name: str | None = None + alias: str | None = None credential_type: str # "oauth2" or "byok" has_credential: bool - expires_at: Optional[str] = None # ISO-8601; None means non-expiring - connected_at: Optional[str] = None # ISO-8601 + expires_at: str | None = None # ISO-8601; None means non-expiring + connected_at: str | None = None # ISO-8601 class MCPUserEnvVarsRequest(LiteLLMPydanticObjectBase): """Payload for storing the calling user's per-user env var values.""" - values: Dict[str, str] + values: dict[str, str] class MCPUserEnvVarSpec(LiteLLMPydanticObjectBase): @@ -1520,7 +1516,7 @@ class MCPUserEnvVarSpec(LiteLLMPydanticObjectBase): """ name: str - description: Optional[str] = None + description: str | None = None is_set: bool = False @@ -1528,15 +1524,15 @@ class MCPUserEnvVarsStatus(LiteLLMPydanticObjectBase): """Per-user env var status for a single MCP server.""" server_id: str - server_name: Optional[str] = None - alias: Optional[str] = None - required: List[MCPUserEnvVarSpec] = Field(default_factory=list) + server_name: str | None = None + alias: str | None = None + required: list[MCPUserEnvVarSpec] = Field(default_factory=list) missing_count: int = 0 - setup_url: Optional[str] = None # frontend URL where the user can fill these in + setup_url: str | None = None # frontend URL where the user can fill these in class RejectMCPServerRequest(LiteLLMPydanticObjectBase): - review_notes: Optional[str] = None + review_notes: str | None = None class MCPSubmissionsSummary(LiteLLMPydanticObjectBase): @@ -1544,7 +1540,7 @@ class MCPSubmissionsSummary(LiteLLMPydanticObjectBase): pending_review: int active: int rejected: int - items: List["LiteLLM_MCPServerTable"] + items: list["LiteLLM_MCPServerTable"] ######## Skills API Types ######## @@ -1553,29 +1549,29 @@ class MCPSubmissionsSummary(LiteLLMPydanticObjectBase): class NewSkillRequest(LiteLLMPydanticObjectBase): """Request to create a new skill in LiteLLM database""" - display_title: Optional[str] = None - description: Optional[str] = None - instructions: Optional[str] = None - file_content: Optional[bytes] = None # Binary content of skill files (zip) - file_name: Optional[str] = None # Original filename - file_type: Optional[str] = None # MIME type (e.g., "application/zip") - metadata: Optional[Dict[str, Any]] = None - authorization_url: Optional[str] = None - token_url: Optional[str] = None - registration_url: Optional[str] = None + display_title: str | None = None + description: str | None = None + instructions: str | None = None + file_content: bytes | None = None # Binary content of skill files (zip) + file_name: str | None = None # Original filename + file_type: str | None = None # MIME type (e.g., "application/zip") + metadata: dict[str, Any] | None = None + authorization_url: str | None = None + token_url: str | None = None + registration_url: str | None = None class UpdateSkillRequest(LiteLLMPydanticObjectBase): """Request to update an existing skill""" skill_id: str - display_title: Optional[str] = None - description: Optional[str] = None - instructions: Optional[str] = None - file_content: Optional[bytes] = None # Binary content of skill files (zip) - file_name: Optional[str] = None # Original filename - file_type: Optional[str] = None # MIME type - metadata: Optional[Dict[str, Any]] = None + display_title: str | None = None + description: str | None = None + instructions: str | None = None + file_content: bytes | None = None # Binary content of skill files (zip) + file_name: str | None = None # Original filename + file_type: str | None = None # MIME type + metadata: dict[str, Any] | None = None from litellm.models.skills import ( # noqa: E402 @@ -1586,74 +1582,77 @@ from litellm.models.skills import ( # noqa: E402 class ListSkillsRequest(LiteLLMPydanticObjectBase): """Request to list skills from LiteLLM database""" - limit: Optional[int] = 20 - offset: Optional[int] = 0 + limit: int | None = 20 + offset: int | None = 0 class NewUserRequestTeam(LiteLLMPydanticObjectBase): team_id: str - max_budget_in_team: Optional[float] = None + max_budget_in_team: float | None = None user_role: Literal["user", "admin"] = "user" class NewUserRequest(GenerateRequestBase): - max_budget: Optional[float] = None - user_email: Optional[str] = None - user_alias: Optional[str] = None - user_role: Optional[ + max_budget: float | None = None + user_email: str | None = None + user_alias: str | None = None + user_role: ( Literal[ LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - ] = None - teams: Optional[Union[List[str], List[NewUserRequestTeam]]] = None + | None + ) = None + teams: list[str] | list[NewUserRequestTeam] | None = None auto_create_key: bool = True # flag used for returning a key as part of the /user/new response - send_invite_email: Optional[bool] = None - sso_user_id: Optional[str] = None - organizations: Optional[List[str]] = None + send_invite_email: bool | None = None + sso_user_id: str | None = None + organizations: list[str] | None = None class NewUserResponse(GenerateKeyResponse): - max_budget: Optional[float] = None - user_email: Optional[str] = None - user_role: Optional[ + max_budget: float | None = None + user_email: str | None = None + user_role: ( Literal[ LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - ] = None - teams: Optional[list] = None - user_alias: Optional[str] = None - model_max_budget: Optional[dict] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None + | None + ) = None + teams: list | None = None + user_alias: str | None = None + model_max_budget: dict | None = None + created_at: datetime | None = None + updated_at: datetime | None = None class UpdateUserRequestNoUserIDorEmail(GenerateRequestBase): # shared with BulkUpdateUserRequest - password: Optional[str] = None - spend: Optional[float] = None - metadata: Optional[dict] = None - user_alias: Optional[str] = None - user_role: Optional[ + password: str | None = None + spend: float | None = None + metadata: dict | None = None + user_alias: str | None = None + user_role: ( Literal[ LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - ] = None - max_budget: Optional[float] = None + | None + ) = None + max_budget: float | None = None class UpdateUserRequest(UpdateUserRequestNoUserIDorEmail): # Note: the defaults of all Params here MUST BE NONE # else they will get overwritten - user_id: Optional[str] = None - user_email: Optional[str] = None + user_id: str | None = None + user_email: str | None = None @model_validator(mode="before") @classmethod @@ -1664,43 +1663,43 @@ class UpdateUserRequest(UpdateUserRequestNoUserIDorEmail): class DeleteUserRequest(LiteLLMPydanticObjectBase): - user_ids: List[str] # required + user_ids: list[str] # required AllowedModelRegion = Literal["eu", "us"] class BudgetNewRequest(LiteLLMPydanticObjectBase): - budget_id: Optional[str] = Field(default=None, description="The unique budget id.") - max_budget: Optional[float] = Field( + budget_id: str | None = Field(default=None, description="The unique budget id.") + max_budget: float | None = Field( default=None, description="Requests will fail if this budget (in USD) is exceeded.", ) - soft_budget: Optional[float] = Field( + soft_budget: float | None = Field( default=None, description="Requests will NOT fail if this is exceeded. Will fire alerting though.", ) - max_parallel_requests: Optional[int] = Field( + max_parallel_requests: int | None = Field( default=None, description="Max concurrent requests allowed for this budget id." ) - tpm_limit: Optional[int] = Field(default=None, description="Max tokens per minute, allowed for this budget id.") - rpm_limit: Optional[int] = Field(default=None, description="Max requests per minute, allowed for this budget id.") - budget_duration: Optional[str] = Field( + tpm_limit: int | None = Field(default=None, description="Max tokens per minute, allowed for this budget id.") + rpm_limit: int | None = Field(default=None, description="Max requests per minute, allowed for this budget id.") + budget_duration: str | None = Field( default=None, description="Max duration budget should be set for (e.g. '1hr', '1d', '28d')", ) - model_max_budget: Optional[GenericBudgetConfigType] = Field( + model_max_budget: GenericBudgetConfigType | None = Field( default=None, description="Max budget for each model (e.g. {'gpt-4o': {'max_budget': '0.0000001', 'budget_duration': '1d', 'tpm_limit': 1000, 'rpm_limit': 1000}})", ) - budget_reset_at: Optional[datetime] = Field( + budget_reset_at: datetime | None = Field( default=None, description="Datetime when the budget is reset", ) class BudgetRequest(LiteLLMPydanticObjectBase): - budgets: List[str] + budgets: list[str] class BudgetDeleteRequest(LiteLLMPydanticObjectBase): @@ -1709,12 +1708,12 @@ class BudgetDeleteRequest(LiteLLMPydanticObjectBase): class CustomerBase(LiteLLMPydanticObjectBase): user_id: str - alias: Optional[str] = None + alias: str | None = None spend: float = 0.0 - allowed_model_region: Optional[AllowedModelRegion] = None - default_model: Optional[str] = None - budget_id: Optional[str] = None - litellm_budget_table: Optional[BudgetNewRequest] = None + allowed_model_region: AllowedModelRegion | None = None + default_model: str | None = None + budget_id: str | None = None + litellm_budget_table: BudgetNewRequest | None = None blocked: bool = False @@ -1724,15 +1723,15 @@ class NewCustomerRequest(BudgetNewRequest): """ user_id: str - alias: Optional[str] = None # human-friendly alias + alias: str | None = None # human-friendly alias blocked: bool = False # allow/disallow requests for this end-user - budget_id: Optional[str] = None # give either a budget_id or max_budget - spend: Optional[float] = None - allowed_model_region: Optional[AllowedModelRegion] = ( + budget_id: str | None = None # give either a budget_id or max_budget + spend: float | None = None + allowed_model_region: AllowedModelRegion | None = ( None # require all user requests to use models in this specific region ) - default_model: Optional[str] = None # if no equivalent model in allowed region - default all requests to this model - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + default_model: str | None = None # if no equivalent model in allowed region - default all requests to this model + object_permission: LiteLLM_ObjectPermissionBase | None = None @model_validator(mode="before") @classmethod @@ -1750,15 +1749,15 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase): """ user_id: str - alias: Optional[str] = None # human-friendly alias + alias: str | None = None # human-friendly alias blocked: bool = False # allow/disallow requests for this end-user - max_budget: Optional[float] = None - budget_id: Optional[str] = None # give either a budget_id or max_budget - allowed_model_region: Optional[AllowedModelRegion] = ( + max_budget: float | None = None + budget_id: str | None = None # give either a budget_id or max_budget + allowed_model_region: AllowedModelRegion | None = ( None # require all user requests to use models in this specific region ) - default_model: Optional[str] = None # if no equivalent model in allowed region - default all requests to this model - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + default_model: str | None = None # if no equivalent model in allowed region - default all requests to this model + object_permission: LiteLLM_ObjectPermissionBase | None = None class DeleteCustomerRequest(LiteLLMPydanticObjectBase): @@ -1766,7 +1765,7 @@ class DeleteCustomerRequest(LiteLLMPydanticObjectBase): Delete multiple Customers """ - user_ids: List[str] + user_ids: list[str] from litellm.models.team import Member as Member # noqa: E402 @@ -1785,41 +1784,41 @@ from litellm.models.team import TeamBase as TeamBase # noqa: E402 class NewTeamRequest(TeamBase): - model_aliases: Optional[dict] = None - tags: Optional[list] = None - guardrails: Optional[List[str]] = None - policies: Optional[List[str]] = None - prompts: Optional[List[str]] = None - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None - allowed_passthrough_routes: Optional[list] = None - disable_global_guardrails: Optional[bool] = None - secret_manager_settings: Optional[dict] = None - model_rpm_limit: Optional[Dict[str, int]] = None - rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] = ( + model_aliases: dict | None = None + tags: list | None = None + guardrails: list[str] | None = None + policies: list[str] | None = None + prompts: list[str] | None = None + object_permission: LiteLLM_ObjectPermissionBase | None = None + allowed_passthrough_routes: list | None = None + disable_global_guardrails: bool | None = None + secret_manager_settings: dict | None = None + model_rpm_limit: dict[str, int] | None = None + rpm_limit_type: Literal["guaranteed_throughput", "best_effort_throughput"] | None = ( None # raise an error if 'guaranteed_throughput' is set and we're overallocating rpm ) - tpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] = ( + tpm_limit_type: Literal["guaranteed_throughput", "best_effort_throughput"] | None = ( None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm ) - model_tpm_limit: Optional[Dict[str, int]] = None - mcp_rpm_limit: Optional[Dict[str, int]] = None - team_member_budget: Optional[float] = None # allow user to set a budget for all team members - team_member_rpm_limit: Optional[int] = None # allow user to set RPM limit for all team members - team_member_tpm_limit: Optional[int] = None # allow user to set TPM limit for all team members - team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m" - team_member_budget_duration: Optional[str] = None # e.g. "30d", "1mo" - allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None - enforced_batch_output_expires_after: Optional[dict] = None - enforced_file_expires_after: Optional[dict] = None + model_tpm_limit: dict[str, int] | None = None + mcp_rpm_limit: dict[str, int] | None = None + team_member_budget: float | None = None # allow user to set a budget for all team members + team_member_rpm_limit: int | None = None # allow user to set RPM limit for all team members + team_member_tpm_limit: int | None = None # allow user to set TPM limit for all team members + team_member_key_duration: str | None = None # e.g. "1d", "1w", "1m" + team_member_budget_duration: str | None = None # e.g. "30d", "1mo" + allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None + enforced_batch_output_expires_after: dict | None = None + enforced_file_expires_after: dict | None = None model_config = ConfigDict(protected_namespaces=()) class GlobalEndUsersSpend(LiteLLMPydanticObjectBase): - api_key: Optional[str] = None - startTime: Optional[datetime] = None - endTime: Optional[datetime] = None + api_key: str | None = None + startTime: datetime | None = None + endTime: datetime | None = None class UpdateTeamRequest(LiteLLMPydanticObjectBase): @@ -1841,40 +1840,40 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): """ team_id: str # required - team_alias: Optional[str] = None - organization_id: Optional[str] = None - metadata: Optional[dict] = None - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - max_budget: Optional[float] = None - soft_budget: Optional[float] = None - models: Optional[list] = None - blocked: Optional[bool] = None - budget_duration: Optional[str] = None - tags: Optional[list] = None - model_aliases: Optional[dict] = None - guardrails: Optional[List[str]] = None - policies: Optional[List[str]] = None - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None - disable_global_guardrails: Optional[bool] = None - team_member_budget: Optional[float] = None - team_member_budget_duration: Optional[str] = None - team_member_rpm_limit: Optional[int] = None - team_member_tpm_limit: Optional[int] = None - team_member_key_duration: Optional[str] = None - allowed_passthrough_routes: Optional[list] = None - secret_manager_settings: Optional[dict] = None - prompts: Optional[List[str]] = None - model_rpm_limit: Optional[Dict[str, int]] = None - model_tpm_limit: Optional[Dict[str, int]] = None - mcp_rpm_limit: Optional[Dict[str, int]] = None - allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None - enforced_batch_output_expires_after: Optional[dict] = None - enforced_file_expires_after: Optional[dict] = None - router_settings: Optional[dict] = None - access_group_ids: Optional[List[str]] = None - budget_limits: Optional[List[BudgetLimitEntry]] = None # multiple concurrent budget windows - default_team_member_models: Optional[List[str]] = None # default allowed_models seeded onto new team members + team_alias: str | None = None + organization_id: str | None = None + metadata: dict | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + max_budget: float | None = None + soft_budget: float | None = None + models: list | None = None + blocked: bool | None = None + budget_duration: str | None = None + tags: list | None = None + model_aliases: dict | None = None + guardrails: list[str] | None = None + policies: list[str] | None = None + object_permission: LiteLLM_ObjectPermissionBase | None = None + disable_global_guardrails: bool | None = None + team_member_budget: float | None = None + team_member_budget_duration: str | None = None + team_member_rpm_limit: int | None = None + team_member_tpm_limit: int | None = None + team_member_key_duration: str | None = None + allowed_passthrough_routes: list | None = None + secret_manager_settings: dict | None = None + prompts: list[str] | None = None + model_rpm_limit: dict[str, int] | None = None + model_tpm_limit: dict[str, int] | None = None + mcp_rpm_limit: dict[str, int] | None = None + allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None + enforced_batch_output_expires_after: dict | None = None + enforced_file_expires_after: dict | None = None + router_settings: dict | None = None + access_group_ids: list[str] | None = None + budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows + default_team_member_models: list[str] | None = None # default allowed_models seeded onto new team members class PatchTeamRequest(UpdateTeamRequest): @@ -1905,7 +1904,7 @@ class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): class DeleteTeamRequest(LiteLLMPydanticObjectBase): - team_ids: List[str] # required + team_ids: list[str] # required class BlockTeamRequest(LiteLLMPydanticObjectBase): @@ -1922,8 +1921,8 @@ class BlockModelRequest(LiteLLMPydanticObjectBase): class AddTeamCallback(LiteLLMPydanticObjectBase): callback_name: str - callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = "success_and_failure" - callback_vars: Dict[str, str] + callback_type: Literal["success", "failure", "success_and_failure"] | None = "success_and_failure" + callback_vars: dict[str, str] @model_validator(mode="before") @classmethod @@ -1939,11 +1938,11 @@ class AddTeamCallback(LiteLLMPydanticObjectBase): class TeamCallbackMetadata(LiteLLMPydanticObjectBase): - success_callback: Optional[List[str]] = [] - failure_callback: Optional[List[str]] = [] - callbacks: Optional[List[str]] = [] + success_callback: list[str] | None = [] + failure_callback: list[str] | None = [] + callbacks: list[str] | None = [] # for now - only supported for langfuse - callback_vars: Optional[Dict[str, str]] = {} + callback_vars: dict[str, str] | None = {} @model_validator(mode="before") @classmethod @@ -1989,7 +1988,7 @@ from litellm.models.team import ( # noqa: E402 class TeamRequest(LiteLLMPydanticObjectBase): - teams: List[str] + teams: list[str] from litellm.models.budget import ( # noqa: E402 @@ -2004,26 +2003,26 @@ from litellm.models.budget import ( # noqa: E402 class NewOrganizationRequest(LiteLLM_BudgetTable): - organization_id: Optional[str] = None + organization_id: str | None = None organization_alias: str - models: List = [] - budget_id: Optional[str] = None - metadata: Optional[dict] = None - model_rpm_limit: Optional[Dict[str, int]] = None - model_tpm_limit: Optional[Dict[str, int]] = None + models: list = [] + budget_id: str | None = None + metadata: dict | None = None + model_rpm_limit: dict[str, int] | None = None + model_tpm_limit: dict[str, int] | None = None ######################################################### # Object Permission - MCP, Vector Stores etc. ######################################################### - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + object_permission: LiteLLM_ObjectPermissionBase | None = None class OrganizationRequest(LiteLLMPydanticObjectBase): - organizations: List[str] + organizations: list[str] class DeleteOrganizationRequest(LiteLLMPydanticObjectBase): - organization_ids: List[str] # required + organization_ids: list[str] # required class TeamDefaultSettings(LiteLLMPydanticObjectBase): @@ -2036,23 +2035,23 @@ class TeamDefaultSettings(LiteLLMPydanticObjectBase): class DynamoDBArgs(LiteLLMPydanticObjectBase): billing_mode: Literal["PROVISIONED_THROUGHPUT", "PAY_PER_REQUEST"] - read_capacity_units: Optional[int] = None - write_capacity_units: Optional[int] = None - ssl_verify: Optional[bool] = None + read_capacity_units: int | None = None + write_capacity_units: int | None = None + ssl_verify: bool | None = None region_name: str user_table_name: str = "LiteLLM_UserTable" key_table_name: str = "LiteLLM_VerificationToken" config_table_name: str = "LiteLLM_Config" spend_table_name: str = "LiteLLM_SpendLogs" - aws_role_name: Optional[str] = None - aws_session_name: Optional[str] = None - aws_web_identity_token: Optional[str] = None - aws_provider_id: Optional[str] = None - aws_policy_arns: Optional[List[str]] = None - aws_policy: Optional[str] = None - aws_duration_seconds: Optional[int] = None - assume_role_aws_role_name: Optional[str] = None - assume_role_aws_session_name: Optional[str] = None + aws_role_name: str | None = None + aws_session_name: str | None = None + aws_web_identity_token: str | None = None + aws_provider_id: str | None = None + aws_policy_arns: list[str] | None = None + aws_policy: str | None = None + aws_duration_seconds: int | None = None + assume_role_aws_role_name: str | None = None + assume_role_aws_session_name: str | None = None class PassThroughGuardrailSettings(LiteLLMPydanticObjectBase): @@ -2062,22 +2061,22 @@ class PassThroughGuardrailSettings(LiteLLMPydanticObjectBase): Allows field-level targeting for guardrail execution. """ - request_fields: Optional[List[str]] = Field( + request_fields: list[str] | None = Field( default=None, description="JSONPath expressions for input field targeting (pre_call). Examples: 'query', 'documents[*].text', 'messages[*].content'. If not specified, guardrail runs on entire request payload.", ) - response_fields: Optional[List[str]] = Field( + response_fields: list[str] | None = Field( default=None, description="JSONPath expressions for output field targeting (post_call). Examples: 'results[*].text', 'output'. If not specified, guardrail runs on entire response payload.", ) # Type alias for the guardrails dict: guardrail_name -> settings (or None for defaults) -PassThroughGuardrailsConfig = Dict[str, Optional[PassThroughGuardrailSettings]] +PassThroughGuardrailsConfig = dict[str, PassThroughGuardrailSettings | None] class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): - id: Optional[str] = Field( + id: str | None = Field( default=None, description="Optional unique identifier for the pass-through endpoint. If not provided, endpoints will be identified by path for backwards compatibility.", ) @@ -2099,7 +2098,7 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): default=0.0, description="The USD cost per request to the target endpoint. This is used to calculate the cost of the request to the target endpoint.", ) - timeout: Optional[float] = Field( + timeout: float | None = Field( default=None, description="Upstream request timeout in seconds for this pass-through endpoint. If unset, uses general_settings.pass_through_request_timeout (default 600).", ) @@ -2107,7 +2106,7 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): default=True, description="Whether authentication is required for the pass-through endpoint. Defaults to True so a pass-through silently created without an explicit value still requires a valid LiteLLM API key — set to False only if the endpoint is meant to be a public forwarder (e.g. an unauthenticated webhook target).", ) - guardrails: Optional[PassThroughGuardrailsConfig] = Field( + guardrails: PassThroughGuardrailsConfig | None = Field( default=None, description="Guardrails configuration for this passthrough endpoint. Dict keys are guardrail names, values are optional settings for field targeting. When set, all org/team/key level guardrails will also execute. Defaults to None (no guardrails execute).", ) @@ -2115,14 +2114,14 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): default=False, description="True if this endpoint is defined in the config file, False if from DB. Config-defined endpoints cannot be edited via the UI.", ) - methods: Optional[List[str]] = Field( + methods: list[str] | None = Field( default=None, description="List of HTTP methods this endpoint handles (e.g., ['GET', 'POST']). If None or empty, all methods (GET, POST, PUT, DELETE, PATCH) are supported for backward compatibility. This allows the same path to have different targets for different HTTP methods.", ) class PassThroughEndpointResponse(LiteLLMPydanticObjectBase): - endpoints: List[PassThroughGenericEndpoint] + endpoints: list[PassThroughGenericEndpoint] class ConfigFieldUpdate(LiteLLMPydanticObjectBase): @@ -2145,7 +2144,7 @@ class FieldDetail(BaseModel): field_type: str field_description: str field_default_value: Any = None - stored_in_db: Optional[bool] + stored_in_db: bool | None class ConfigList(LiteLLMPydanticObjectBase): @@ -2153,12 +2152,12 @@ class ConfigList(LiteLLMPydanticObjectBase): field_type: str field_description: str field_value: Any - stored_in_db: Optional[bool] + stored_in_db: bool | None field_default_value: Any premium_field: bool = False - nested_fields: Optional[List[FieldDetail]] = None # For nested dictionary or Pydantic fields - field_options: Optional[list[str]] = None # Allowed values, for field_type == "Select" - field_tab: Optional[str] = None # Admin UI sub-tab this field renders under; None groups it with the rest + nested_fields: list[FieldDetail] | None = None # For nested dictionary or Pydantic fields + field_options: list[str] | None = None # Allowed values, for field_type == "Select" + field_tab: str | None = None # Admin UI sub-tab this field renders under; None groups it with the rest class UserHeaderMapping(LiteLLMPydanticObjectBase): @@ -2208,20 +2207,20 @@ class CoordinationRedisParams(LiteLLMPydanticObjectBase): model_config = ConfigDict(extra="allow", protected_namespaces=()) - host: Optional[str] = Field(None, description="Redis hostname") - port: Optional[int] = Field(None, description="Redis port") - password: Optional[str] = Field(None, description="Redis password") - username: Optional[str] = Field(None, description="Redis username") - url: Optional[str] = Field(None, description="full Redis connection url, e.g. redis://:pass@host:6379") - ssl: Optional[bool] = Field(None, description="connect over TLS") - startup_nodes: Optional[List[CoordinationRedisNode]] = Field( + host: str | None = Field(None, description="Redis hostname") + port: int | None = Field(None, description="Redis port") + password: str | None = Field(None, description="Redis password") + username: str | None = Field(None, description="Redis username") + url: str | None = Field(None, description="full Redis connection url, e.g. redis://:pass@host:6379") + ssl: bool | None = Field(None, description="connect over TLS") + startup_nodes: list[CoordinationRedisNode] | None = Field( None, description="cluster-mode startup nodes; when set a cluster client is used" ) - sentinel_nodes: Optional[List[List[Union[str, int]]]] = Field( + sentinel_nodes: list[list[str | int]] | None = Field( None, description="sentinel [host, port] pairs; when set a sentinel-managed client is used" ) - sentinel_password: Optional[str] = Field(None, description="password for the sentinel nodes") - service_name: Optional[str] = Field(None, description="sentinel service name") + sentinel_password: str | None = Field(None, description="password for the sentinel nodes") + service_name: str | None = Field(None, description="sentinel service name") def has_connection_target(self) -> bool: return any(value is not None for value in (self.host, self.url, self.startup_nodes, self.sentinel_nodes)) @@ -2232,17 +2231,17 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): Documents all the fields supported by `general_settings` in config.yaml """ - completion_model: Optional[str] = Field(None, description="proxy level default model for all chat completion calls") + completion_model: str | None = Field(None, description="proxy level default model for all chat completion calls") plugins: list[PluginConfig] | None = Field( None, description="external services registered as embeddable UI plugins" ) - key_management_system: Optional[KeyManagementSystem] = Field( + key_management_system: KeyManagementSystem | None = Field( None, description="key manager to load keys from / decrypt keys with" ) - use_google_kms: Optional[bool] = Field(None, description="decrypt keys with google kms") - use_azure_key_vault: Optional[bool] = Field(None, description="load keys from azure key vault") - master_key: Optional[str] = Field(None, description="require a key for all calls to proxy") - coordination_redis: Optional[CoordinationRedisParams] = Field( + use_google_kms: bool | None = Field(None, description="decrypt keys with google kms") + use_azure_key_vault: bool | None = Field(None, description="load keys from azure key vault") + master_key: str | None = Field(None, description="require a key for all calls to proxy") + coordination_redis: CoordinationRedisParams | None = Field( None, description=( "standalone Redis for cross-pod coordination (tpm/rpm rate limits, " @@ -2255,18 +2254,18 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine", ) - database_url: Optional[str] = Field( + database_url: str | None = Field( None, description="connect to a postgres db - needed for generating temporary keys + tracking spend / key", ) - database_connection_pool_limit: Optional[int] = Field( + database_connection_pool_limit: int | None = Field( 10, description="default connection pool for prisma client connecting to postgres db", ) - database_connection_timeout: Optional[float] = Field( + database_connection_timeout: float | None = Field( 60, description="default timeout for a connection to the database" ) - database_connect_timeout: Optional[float] = Field( + database_connect_timeout: float | None = Field( None, description=( "Prisma `connect_timeout` URL param (seconds). Bounds how long the " @@ -2274,7 +2273,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "to Prisma's built-in value when unset." ), ) - database_socket_timeout: Optional[float] = Field( + database_socket_timeout: float | None = Field( None, description=( "Prisma `socket_timeout` URL param (seconds). When set, an idle/slow " @@ -2282,7 +2281,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "This is the main knob for capping idle DB connections from LiteLLM." ), ) - database_extra_connection_params: Optional[Dict[str, Any]] = Field( + database_extra_connection_params: dict[str, Any] | None = Field( None, description=( "Escape hatch: extra key/value pairs appended verbatim to the Prisma " @@ -2290,7 +2289,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "`statement_cache_size`). Keys here override any default LiteLLM sets." ), ) - database_disable_prepared_statements: Optional[bool] = Field( + database_disable_prepared_statements: bool | None = Field( None, description=( "Disable server-side prepared statements by setting Prisma's " @@ -2301,31 +2300,31 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "takes precedence." ), ) - database_type: Optional[Literal["dynamo_db"]] = Field(None, description="to use dynamodb instead of postgres db") - database_args: Optional[DynamoDBArgs] = Field( + database_type: Literal["dynamo_db"] | None = Field(None, description="to use dynamodb instead of postgres db") + database_args: DynamoDBArgs | None = Field( None, description="custom args for instantiating dynamodb client - e.g. billing provision", ) - otel: Optional[bool] = Field( + otel: bool | None = Field( None, description="[BETA] OpenTelemetry support - this might change, use with caution.", ) - custom_auth: Optional[str] = Field( + custom_auth: str | None = Field( None, description="override user_api_key_auth with your own auth script - https://docs.litellm.ai/docs/proxy/virtual_keys#custom-auth", ) - max_parallel_requests: Optional[int] = Field( + max_parallel_requests: int | None = Field( None, description="maximum parallel requests for each api key", ) - global_max_parallel_requests: Optional[int] = Field( + global_max_parallel_requests: int | None = Field( None, description="global max parallel requests to allow for a proxy instance." ) - max_request_size_mb: Optional[int] = Field( + max_request_size_mb: int | None = Field( None, description="max request size in MB, if a request is larger than this size it will be rejected", ) - max_response_size_mb: Optional[int] = Field( + max_response_size_mb: int | None = Field( None, description="max response size in MB, if a response is larger than this size it will be rejected", ) @@ -2334,17 +2333,17 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): gt=0, description="how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup", ) - cancel_on_disconnect: Optional[bool] = Field( + cancel_on_disconnect: bool | None = Field( None, description="cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure", ) - infer_model_from_keys: Optional[bool] = Field( + infer_model_from_keys: bool | None = Field( None, description="for `/models` endpoint, infers available model based on environment keys (e.g. OPENAI_API_KEY)", ) - background_health_checks: Optional[bool] = Field(None, description="run health checks in background") + background_health_checks: bool | None = Field(None, description="run health checks in background") health_check_interval: int = Field(300, description="background health check interval in seconds") - health_check_concurrency: Optional[int] = Field( + health_check_concurrency: int | None = Field( None, description=( "limit concurrent health checks per cycle; when unset, health checks run without a concurrency cap" @@ -2357,28 +2356,26 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "are skipped for on-demand GET /health as well as the background health loop." ), ) - alerting: Optional[List] = Field( + alerting: list | None = Field( None, description="List of alerting integrations. Today, just slack - `alerting: ['slack']`", ) - alert_types: Optional[List[AlertType]] = Field( + alert_types: list[AlertType] | None = Field( None, description="List of alerting types. By default it is all alerts", ) - alert_to_webhook_url: Optional[Dict] = Field( + alert_to_webhook_url: dict | None = Field( None, description="Mapping of alert type to webhook url. e.g. `alert_to_webhook_url: {'budget_alerts': 'https://nothooks.slack.com/services/T00000000/B00000000/XXXXXXXXXXXXXXXXXXXXXXXX'}`", ) - alerting_args: Optional[Dict] = Field( - None, description="Controllable params for slack alerting - e.g. ttl in cache." - ) - alerting_threshold: Optional[int] = Field( + alerting_args: dict | None = Field(None, description="Controllable params for slack alerting - e.g. ttl in cache.") + alerting_threshold: int | None = Field( None, description="sends alerts if requests hang for 5min+", ) - ui_access_mode: Optional[Literal["admin_only", "all"]] = Field("all", description="Control access to the Proxy UI") - allowed_routes: Optional[List] = Field(None, description="Proxy API Endpoints you want users to be able to access") - reject_clientside_metadata_tags: Optional[bool] = Field( + ui_access_mode: Literal["admin_only", "all"] | None = Field("all", description="Control access to the Proxy UI") + allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access") + reject_clientside_metadata_tags: bool | None = Field( None, description="When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.", ) @@ -2386,28 +2383,28 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): default=False, description="Public model hub for users to see what models they have access to, supported openai params, etc.", ) - pass_through_request_timeout: Optional[float] = Field( + pass_through_request_timeout: float | None = Field( default=None, description="Default upstream request timeout in seconds for native and custom pass-through endpoints that use pass_through_request. Defaults to 600 when unset.", ) - pass_through_endpoints: Optional[List[PassThroughGenericEndpoint]] = Field( + pass_through_endpoints: list[PassThroughGenericEndpoint] | None = Field( default=None, description="Set-up pass-through endpoints for provider-specific endpoints. Docs - https://docs.litellm.ai/docs/proxy/pass_through", ) - user_header_name: Optional[str] = Field( + user_header_name: str | None = Field( None, description="[DEPRECATED] Use 'user_header_mappings' instead. When set, the header value is treated as the end user id unless overridden by user_header_mappings.", ) - user_header_mappings: Optional[List[UserHeaderMapping]] = None - supported_db_objects: Optional[List[SupportedDBObjectType]] = Field( + user_header_mappings: list[UserHeaderMapping] | None = None + supported_db_objects: list[SupportedDBObjectType] | None = Field( None, description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map', 'tools', 'config_overrides'. If not set, all objects are loaded (default behavior).", ) - user_mcp_management_mode: Optional[UserMCPManagementMode] = Field( + user_mcp_management_mode: UserMCPManagementMode | None = Field( None, description="Controls how non-admin users interact with MCP servers in the dashboard. 'restricted' shows only accessible servers, 'view_all' lists every server in read-only mode.", ) - store_prompts_in_spend_logs: Optional[bool] = Field( + store_prompts_in_spend_logs: bool | None = Field( None, description="If True, stores request messages and responses in spend logs. Default is False.", ) @@ -2415,44 +2412,44 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.", ) - maximum_spend_logs_retention_period: Optional[str] = Field( + maximum_spend_logs_retention_period: str | None = Field( None, description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.", ) - use_spend_logs_partitioning: Optional[bool] = Field( + use_spend_logs_partitioning: bool | None = Field( None, description="If True and LiteLLM_SpendLogs has been converted to a range-partitioned table (db_scripts/partition_spend_logs.sql), retention cleanup drops expired partitions instead of deleting rows, and pre-creates upcoming partitions. Default is False.", ) - mcp_internal_ip_ranges: Optional[List[str]] = Field( + mcp_internal_ip_ranges: list[str] | None = Field( None, description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).", ) - mcp_trusted_proxy_ranges: Optional[List[str]] = Field( + mcp_trusted_proxy_ranges: list[str] | None = Field( None, description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.", ) - mcp_xff_num_trusted_hops: Optional[int] = Field( + mcp_xff_num_trusted_hops: int | None = Field( None, ge=1, description="Number of trusted reverse proxies/load balancers in front of the gateway that append to X-Forwarded-For. When set (and mcp_trusted_proxy_ranges validates the direct peer), the client IP for MCP access control is read this many entries from the right of the chain instead of the spoofable leftmost value, defeating append-style X-Forwarded-For forgery.", ) - trusted_proxy_ranges: Optional[List[str]] = Field( + trusted_proxy_ranges: list[str] | None = Field( None, description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.", ) - store_model_in_db: Optional[bool] = Field( + store_model_in_db: bool | None = Field( None, description="If True, models and config are stored in and loaded from the database. Default is False.", ) - forward_client_headers_to_llm_api: Optional[bool] = Field( + forward_client_headers_to_llm_api: bool | None = Field( None, description="If True, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription.", ) - mcp_required_fields: Optional[List[str]] = Field( + mcp_required_fields: list[str] | None = Field( None, description="List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).", ) - disable_budget_reservation: Optional[bool] = Field( + disable_budget_reservation: bool | None = Field( None, description=( "If True, disables the optimistic per-request budget reservation " @@ -2470,7 +2467,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "is active as a reminder that hard enforcement is relaxed." ), ) - user_url_validation: Optional[bool] = Field( + user_url_validation: bool | None = Field( None, description=( "Master switch for the SSRF guard applied to user-supplied URLs " @@ -2478,7 +2475,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "Set to False to disable DNS/IP validation entirely (not recommended)." ), ) - user_url_allowed_hosts: Optional[list[str]] = Field( + user_url_allowed_hosts: list[str] | None = Field( None, description=( "SSRF allowlist for user-supplied URLs. Entries are `hostname` or " @@ -2488,7 +2485,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "an internal OpenAPI/MCP server." ), ) - provider_url_destination_allowed_hosts: Optional[list[str]] = Field( + provider_url_destination_allowed_hosts: list[str] | None = Field( None, description="Allowlist of hosts a request may redirect a provider call's destination URL to.", ) @@ -2499,20 +2496,20 @@ class ConfigYAML(LiteLLMPydanticObjectBase): Documents all the fields supported by the config.yaml """ - environment_variables: Optional[dict] = Field( + environment_variables: dict | None = Field( None, description="Object to pass in additional environment variables via POST request", ) - model_list: Optional[List[ModelParams]] = Field( + model_list: list[ModelParams] | None = Field( None, description="List of supported models on the server, with model-specific configs", ) - litellm_settings: Optional[dict] = Field( + litellm_settings: dict | None = Field( None, description="litellm Module settings. See __init__.py for all, example litellm.drop_params=True, litellm.set_verbose=True, litellm.api_base, litellm.cache", ) - general_settings: Optional[ConfigGeneralSettings] = None - router_settings: Optional[UpdateRouterConfig] = Field( + general_settings: ConfigGeneralSettings | None = None + router_settings: UpdateRouterConfig | None = Field( None, description="litellm router object settings. See router.py __init__ for all, example router.num_retries=5, router.timeout=5, router.max_retries=5, router.retry_after=5", ) @@ -2533,45 +2530,45 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): Combined view of litellm verification token + litellm team table (select values) """ - team_spend: Optional[float] = None - team_alias: Optional[str] = None - team_tpm_limit: Optional[int] = None - team_rpm_limit: Optional[int] = None - team_max_budget: Optional[float] = None - team_soft_budget: Optional[float] = None - team_models: List = [] + team_spend: float | None = None + team_alias: str | None = None + team_tpm_limit: int | None = None + team_rpm_limit: int | None = None + team_max_budget: float | None = None + team_soft_budget: float | None = None + team_models: list = [] team_blocked: bool = False - soft_budget: Optional[float] = None - team_model_aliases: Optional[Dict] = None - team_member: Optional[Member] = None - team_metadata: Optional[Dict] = None - team_object_permission_id: Optional[str] = None + soft_budget: float | None = None + team_model_aliases: dict | None = None + team_member: Member | None = None + team_metadata: dict | None = None + team_object_permission_id: str | None = None # Team Member Specific Params - team_member_spend: Optional[float] = None - team_member_tpm_limit: Optional[int] = None - team_member_rpm_limit: Optional[int] = None + team_member_spend: float | None = None + team_member_tpm_limit: int | None = None + team_member_rpm_limit: int | None = None # End User Params - end_user_id: Optional[str] = None - end_user_tpm_limit: Optional[int] = None - end_user_rpm_limit: Optional[int] = None - end_user_max_budget: Optional[float] = None - end_user_model_max_budget: Optional[dict] = None + end_user_id: str | None = None + end_user_tpm_limit: int | None = None + end_user_rpm_limit: int | None = None + end_user_max_budget: float | None = None + end_user_model_max_budget: dict | None = None # Organization Params - organization_alias: Optional[str] = None - organization_max_budget: Optional[float] = None - organization_tpm_limit: Optional[int] = None - organization_rpm_limit: Optional[int] = None - organization_metadata: Optional[dict] = None + organization_alias: str | None = None + organization_max_budget: float | None = None + organization_tpm_limit: int | None = None + organization_rpm_limit: int | None = None + organization_metadata: dict | None = None # Project Params - project_alias: Optional[str] = None - project_metadata: Optional[dict] = None + project_alias: str | None = None + project_metadata: dict | None = None # Time stamps - last_refreshed_at: Optional[float] = None # last time joint view was pulled from db + last_refreshed_at: float | None = None # last time joint view was pulled from db def __init__(self, **kwargs): # Handle litellm_budget_table_* keys (budget table overrides when key value is None or empty) @@ -2603,18 +2600,18 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob Return the row in the db """ - api_key: Optional[str] = None - user_role: Optional[LitellmUserRoles] = None - allowed_model_region: Optional[AllowedModelRegion] = None - parent_otel_span: Optional[Span] = None - rpm_limit_per_model: Optional[Dict[str, int]] = None - tpm_limit_per_model: Optional[Dict[str, int]] = None - user_tpm_limit: Optional[int] = None - user_rpm_limit: Optional[int] = None - user_email: Optional[str] = None - user_spend: Optional[float] = None - user_max_budget: Optional[float] = None - request_route: Optional[str] = None + api_key: str | None = None + user_role: LitellmUserRoles | None = None + allowed_model_region: AllowedModelRegion | None = None + parent_otel_span: Span | None = None + rpm_limit_per_model: dict[str, int] | None = None + tpm_limit_per_model: dict[str, int] | None = None + user_tpm_limit: int | None = None + user_rpm_limit: int | None = None + user_email: str | None = None + user_spend: float | None = None + user_max_budget: float | None = None + request_route: str | None = None is_session_token: bool = False # Server-only marker set exclusively by the MCP gateway admission path # (_reload_admitted_user) for a keyless user-subject admitted via a gateway DCR session @@ -2638,17 +2635,17 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob "user id." ), ) - budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) - budget_throttle_pct: Optional[float] = Field(default=None, exclude=True) - user: Optional[Any] = None # Expanded user object when expand=user is used - created_by_user: Optional[Any] = None # Expanded created_by user when expand=user is used - end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None + budget_reservation: dict[str, Any] | None = Field(default=None, exclude=True) + budget_throttle_pct: float | None = Field(default=None, exclude=True) + user: Any | None = None # Expanded user object when expand=user is used + created_by_user: Any | None = None # Expanded created_by user when expand=user is used + end_user_object_permission: LiteLLM_ObjectPermissionTable | None = None # Team object_permission preloaded in auth (e.g. get_team_object) to avoid # per-request object_permission fetches in downstream checks (vector stores, etc.) - team_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None + team_object_permission: LiteLLM_ObjectPermissionTable | None = None # Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery # and forwarded into outbound tokens by guardrails such as MCPJWTSigner. - jwt_claims: Optional[Dict] = None + jwt_claims: dict | None = None model_config = ConfigDict(arbitrary_types_allowed=True) @@ -2754,10 +2751,10 @@ def user_api_key_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool: class UserInfoResponse(LiteLLMPydanticObjectBase): - user_id: Optional[str] - user_info: Optional[Union[dict, BaseModel]] - keys: List - teams: List + user_id: str | None + user_info: dict | BaseModel | None + keys: list + teams: list class UserInfoV2Response(LiteLLMPydanticObjectBase): @@ -2769,19 +2766,19 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase): """ user_id: str - user_email: Optional[str] = None - user_alias: Optional[str] = None - user_role: Optional[str] = None + user_email: str | None = None + user_alias: str | None = None + user_role: str | None = None spend: float = 0.0 - max_budget: Optional[float] = None - models: List[str] = [] - budget_duration: Optional[str] = None - budget_reset_at: Optional[datetime] = None - metadata: Optional[dict] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None - sso_user_id: Optional[str] = None - teams: List[str] = [] # Just team IDs, not full team objects + max_budget: float | None = None + models: list[str] = [] + budget_duration: str | None = None + budget_reset_at: datetime | None = None + metadata: dict | None = None + created_at: datetime | None = None + updated_at: datetime | None = None + sso_user_id: str | None = None + teams: list[str] = [] # Just team IDs, not full team objects object_permission: LiteLLM_ObjectPermissionTable | None = None @@ -2794,16 +2791,16 @@ from litellm.models.organization_membership import ( # noqa: E402 class LiteLLM_OrganizationTableUpdate(LiteLLM_BudgetTable): """Represents user-controllable params for a LiteLLM_OrganizationTable record""" - organization_id: Optional[str] = None - organization_alias: Optional[str] = None - budget_id: Optional[str] = None - spend: Optional[float] = None - metadata: Optional[dict] = None - models: Optional[List[str]] = None - updated_by: Optional[str] = None - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None - model_tpm_limit: Optional[Dict[str, int]] = None - model_rpm_limit: Optional[Dict[str, int]] = None + organization_id: str | None = None + organization_alias: str | None = None + budget_id: str | None = None + spend: float | None = None + metadata: dict | None = None + models: list[str] | None = None + updated_by: str | None = None + object_permission: LiteLLM_ObjectPermissionBase | None = None + model_tpm_limit: dict[str, int] | None = None + model_rpm_limit: dict[str, int] | None = None @model_validator(mode="before") @classmethod @@ -2851,9 +2848,9 @@ from litellm.models.user import LiteLLM_UserTable as LiteLLM_UserTable # noqa: class LiteLLM_OrganizationTableWithMembers(LiteLLM_OrganizationTable): """Returned by the /organization/info endpoint and /organization/list endpoint""" - members: List[LiteLLM_OrganizationMembershipTable] = [] - teams: List[LiteLLM_TeamTable] = [] - litellm_budget_table: Optional[LiteLLM_BudgetTable] = None + members: list[LiteLLM_OrganizationMembershipTable] = [] + teams: list[LiteLLM_TeamTable] = [] + litellm_budget_table: LiteLLM_BudgetTable | None = None created_at: datetime updated_at: datetime @@ -2870,31 +2867,31 @@ class NewOrganizationResponse(LiteLLM_OrganizationTable): class ProjectBase(LiteLLMPydanticObjectBase): """Base fields shared by project create/update requests""" - project_id: Optional[str] = None - project_alias: Optional[str] = None - team_id: Optional[str] = None - metadata: Optional[dict] = None - models: Optional[List[str]] = None + project_id: str | None = None + project_alias: str | None = None + team_id: str | None = None + metadata: dict | None = None + models: list[str] | None = None blocked: bool = False class NewProjectRequest(LiteLLM_BudgetTable): """Request model for POST /project/new""" - project_id: Optional[str] = None - project_alias: Optional[str] = None - description: Optional[str] = None + project_id: str | None = None + project_alias: str | None = None + description: str | None = None team_id: str - budget_id: Optional[str] = None - metadata: Optional[dict] = None - tags: Optional[List[str]] = None - guardrails: Optional[List[str]] = None - policies: Optional[List[str]] = None - models: List[str] = [] - model_rpm_limit: Optional[dict] = None - model_tpm_limit: Optional[dict] = None + budget_id: str | None = None + metadata: dict | None = None + tags: list[str] | None = None + guardrails: list[str] | None = None + policies: list[str] | None = None + models: list[str] = [] + model_rpm_limit: dict | None = None + model_tpm_limit: dict | None = None blocked: bool = False - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + object_permission: LiteLLM_ObjectPermissionBase | None = None @model_validator(mode="before") @classmethod @@ -2915,19 +2912,19 @@ class UpdateProjectRequest(LiteLLM_BudgetTable): """Request model for POST /project/update""" project_id: str - project_alias: Optional[str] = None - description: Optional[str] = None - team_id: Optional[str] = None - metadata: Optional[dict] = None - tags: Optional[List[str]] = None - guardrails: Optional[List[str]] = None - policies: Optional[List[str]] = None - models: Optional[List[str]] = None - model_rpm_limit: Optional[dict] = None - model_tpm_limit: Optional[dict] = None - blocked: Optional[bool] = None - budget_id: Optional[str] = None - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + project_alias: str | None = None + description: str | None = None + team_id: str | None = None + metadata: dict | None = None + tags: list[str] | None = None + guardrails: list[str] | None = None + policies: list[str] | None = None + models: list[str] | None = None + model_rpm_limit: dict | None = None + model_tpm_limit: dict | None = None + blocked: bool | None = None + budget_id: str | None = None + object_permission: LiteLLM_ObjectPermissionBase | None = None @model_validator(mode="before") @classmethod @@ -2947,7 +2944,7 @@ class UpdateProjectRequest(LiteLLM_BudgetTable): class DeleteProjectRequest(LiteLLMPydanticObjectBase): """Request model for DELETE /project/delete""" - project_ids: List[str] + project_ids: list[str] from litellm.models.project import ( # noqa: E402 @@ -2966,12 +2963,12 @@ class NewProjectResponse(LiteLLM_ProjectTable): class LiteLLM_ProjectTableCachedObj(LiteLLM_ProjectTable): """Cached version for auth checks. Mirrors LiteLLM_TeamTableCachedObj pattern.""" - last_refreshed_at: Optional[float] = None + last_refreshed_at: float | None = None class LiteLLM_UserTableFiltered(BaseModel): # done to avoid exposing sensitive data user_id: str - user_email: Optional[str] = None + user_email: str | None = None class LiteLLM_UserTableWithKeyCount(LiteLLM_UserTable): @@ -2998,13 +2995,13 @@ AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): id: str updated_at: datetime - changed_by: Optional[Any] = None - changed_by_api_key: Optional[str] = None + changed_by: Any | None = None + changed_by_api_key: str | None = None action: AUDIT_ACTIONS table_name: LitellmTableNames object_id: str - before_value: Optional[Json] = None - updated_values: Optional[Json] = None + before_value: Json | None = None + updated_values: Json | None = None @model_validator(mode="before") @classmethod @@ -3020,7 +3017,7 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): masker = SensitiveDataMasker(sensitive_patterns={"key"}) if self.before_value is not None: - json_before_value: Optional[dict] = None + json_before_value: dict | None = None if isinstance(self.before_value, str): json_before_value = json.loads(self.before_value) elif isinstance(self.before_value, dict): @@ -3031,7 +3028,7 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): self.before_value = json.dumps(json_before_value, default=str) if self.updated_values is not None: - json_updated_values: Optional[dict] = None + json_updated_values: dict | None = None if isinstance(self.updated_values, str): json_updated_values = json.loads(self.updated_values) elif isinstance(self.updated_values, dict): @@ -3045,48 +3042,48 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): class LiteLLM_SpendLogs_ResponseObject(LiteLLMPydanticObjectBase): - response: Optional[List[Union[LiteLLM_SpendLogs, Any]]] = None + response: list[LiteLLM_SpendLogs | Any] | None = None class TokenCountRequest(LiteLLMPydanticObjectBase): model: str - prompt: Optional[str] = None - messages: Optional[List[dict]] = None + prompt: str | None = None + messages: list[dict] | None = None """ Anthropic token counting endpoint uses /messages """ - contents: Optional[List[dict]] = None + contents: list[dict] | None = None """ Google /countTokens endpoint expects contents to be a list of dicts with the following structure: """ - tools: Optional[List[dict]] = None - system: Optional[Any] = None + tools: list[dict] | None = None + system: Any | None = None class CallInfo(LiteLLMPydanticObjectBase): """Used for slack budget alerting""" spend: float - max_budget: Optional[float] = None - soft_budget: Optional[float] = None - token: Optional[str] = Field(default=None, description="Hashed value of that key") - customer_id: Optional[str] = None - user_id: Optional[str] = None - team_id: Optional[str] = None - team_alias: Optional[str] = None - organization_id: Optional[str] = None - user_email: Optional[str] = None - key_alias: Optional[str] = None - projected_exceeded_date: Optional[str] = None - projected_spend: Optional[float] = None + max_budget: float | None = None + soft_budget: float | None = None + token: str | None = Field(default=None, description="Hashed value of that key") + customer_id: str | None = None + user_id: str | None = None + team_id: str | None = None + team_alias: str | None = None + organization_id: str | None = None + user_email: str | None = None + key_alias: str | None = None + projected_exceeded_date: str | None = None + projected_spend: float | None = None event_group: Litellm_EntityType - alert_emails: Optional[List[str]] = Field( + alert_emails: list[str] | None = Field( default=None, description="Additional email addresses to send alerts to (e.g., from team metadata)", ) - max_budget_alert_emails: Optional[Dict[str, List[str]]] = Field( + max_budget_alert_emails: dict[str, list[str]] | None = Field( default=None, description="Map of threshold percentage to email recipients (e.g., {'50': ['a@co.com'], '75': ['a@co.com', 'b@co.com']})", ) @@ -3139,7 +3136,7 @@ class InvitationModel(LiteLLMPydanticObjectBase): id: str user_id: str is_accepted: bool - accepted_at: Optional[datetime] + accepted_at: datetime | None expires_at: datetime created_at: datetime created_by: str @@ -3160,7 +3157,7 @@ class ConfigFieldInfo(LiteLLMPydanticObjectBase): class CallbackOnUI(LiteLLMPydanticObjectBase): litellm_callback_name: str - litellm_callback_params: Optional[list] + litellm_callback_params: list | None ui_callback_name: str @@ -3297,36 +3294,36 @@ class SpendLogsMetadata(TypedDict): Specific metadata k,v pairs logged to spendlogs for easier cost tracking """ - additional_usage_values: Optional[dict] # covers provider-specific usage information - e.g. prompt caching - user_api_key: Optional[str] - user_api_key_alias: Optional[str] - user_api_key_team_id: Optional[str] - user_api_key_project_id: Optional[str] - user_api_key_project_alias: Optional[str] - user_api_key_org_id: Optional[str] - user_api_key_user_id: Optional[str] - user_api_key_team_alias: Optional[str] - spend_logs_metadata: Optional[dict] # special param to log k,v pairs to spendlogs for a call - requester_ip_address: Optional[str] - litellm_call_id: Optional[str] - applied_guardrails: Optional[List[str]] - mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] - vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]] + additional_usage_values: dict | None # covers provider-specific usage information - e.g. prompt caching + user_api_key: str | None + user_api_key_alias: str | None + user_api_key_team_id: str | None + user_api_key_project_id: str | None + user_api_key_project_alias: str | None + user_api_key_org_id: str | None + user_api_key_user_id: str | None + user_api_key_team_alias: str | None + spend_logs_metadata: dict | None # special param to log k,v pairs to spendlogs for a call + requester_ip_address: str | None + litellm_call_id: str | None + applied_guardrails: list[str] | None + mcp_tool_call_metadata: StandardLoggingMCPToolCall | None + vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None routing_decision: StandardLoggingRoutingDecision | None internal_call_origin: InternalCallOrigin | None - guardrail_information: Optional[List[StandardLoggingGuardrailInformation]] - eval_information: Optional[Any] + guardrail_information: list[StandardLoggingGuardrailInformation] | None + eval_information: Any | None status: StandardLoggingPayloadStatus - proxy_server_request: Optional[str] - batch_models: Optional[List[str]] - error_information: Optional[StandardLoggingPayloadErrorInformation] - usage_object: Optional[dict] - model_map_information: Optional[StandardLoggingModelInformation] - cold_storage_object_key: Optional[str] # S3/GCS object key for cold storage retrieval - litellm_overhead_time_ms: Optional[float] # LiteLLM overhead time in milliseconds - attempted_retries: Optional[int] # Number of retries attempted (0 = first attempt succeeded) - max_retries: Optional[int] # Max retries configured for this request - cost_breakdown: Optional[CostBreakdown] # Detailed cost breakdown (input_cost, output_cost, margin, discount, etc.) + proxy_server_request: str | None + batch_models: list[str] | None + error_information: StandardLoggingPayloadErrorInformation | None + usage_object: dict | None + model_map_information: StandardLoggingModelInformation | None + cold_storage_object_key: str | None # S3/GCS object key for cold storage retrieval + litellm_overhead_time_ms: float | None # LiteLLM overhead time in milliseconds + attempted_retries: int | None # Number of retries attempted (0 = first attempt succeeded) + max_retries: int | None # Max retries configured for this request + cost_breakdown: CostBreakdown | None # Detailed cost breakdown (input_cost, output_cost, margin, discount, etc.) compression_savings: CompressionSavingsMetadata | None @@ -3338,30 +3335,30 @@ class SpendLogsPayload(TypedDict): total_tokens: int prompt_tokens: int completion_tokens: int - startTime: Union[datetime, str] - endTime: Union[datetime, str] - completionStartTime: Optional[Union[datetime, str]] + startTime: datetime | str + endTime: datetime | str + completionStartTime: datetime | str | None model: str - model_id: Optional[str] - model_group: Optional[str] - mcp_namespaced_tool_name: Optional[str] - agent_id: Optional[str] + model_id: str | None + model_group: str | None + mcp_namespaced_tool_name: str | None + agent_id: str | None api_base: str user: str metadata: str # json str cache_hit: str cache_key: str request_tags: str # json str - team_id: Optional[str] - organization_id: Optional[str] - end_user: Optional[str] - requester_ip_address: Optional[str] - custom_llm_provider: Optional[str] - messages: Optional[Union[str, list, dict]] - response: Optional[Union[str, list, dict]] - proxy_server_request: Optional[str] - session_id: Optional[str] - request_duration_ms: Optional[int] + team_id: str | None + organization_id: str | None + end_user: str | None + requester_ip_address: str | None + custom_llm_provider: str | None + messages: str | list | dict | None + response: str | list | dict | None + proxy_server_request: str | None + session_id: str | None + request_duration_ms: int | None status: Literal["success", "failure"] @@ -3428,10 +3425,10 @@ class SpanAttributes(str, enum.Enum): class ManagementEndpointLoggingPayload(LiteLLMPydanticObjectBase): route: str request_data: dict - response: Optional[dict] = None - exception: Optional[Any] = None - start_time: Optional[datetime] = None - end_time: Optional[datetime] = None + response: dict | None = None + exception: Any | None = None + start_time: datetime | None = None + end_time: datetime | None = None class ProxyException(Exception): @@ -3441,11 +3438,11 @@ class ProxyException(Exception): self, message: str, type: str, - param: Optional[str], - code: Optional[Union[int, str]] = None, # maps to status code - headers: Optional[Dict[str, str]] = None, - openai_code: Optional[str] = None, # maps to 'code' in openai - provider_specific_fields: Optional[dict] = None, + param: str | None, + code: int | str | None = None, # maps to status code + headers: dict[str, str] | None = None, + openai_code: str | None = None, # maps to 'code' in openai + provider_specific_fields: dict | None = None, ): self.message = str(message) super().__init__(self.message) @@ -3472,7 +3469,7 @@ class ProxyException(Exception): def to_dict(self) -> dict: """Converts the ProxyException instance to a dictionary.""" - error_dict: Dict[str, Optional[Union[str, Dict]]] = { + error_dict: dict[str, str | dict | None] = { "message": self.message, "type": self.type, "param": self.param, @@ -3498,9 +3495,9 @@ class CommonProxyErrors(str, enum.Enum): class SpendCalculateRequest(LiteLLMPydanticObjectBase): - model: Optional[str] = None - messages: Optional[List] = None - completion_response: Optional[dict] = None + model: str | None = None + messages: list | None = None + completion_response: dict | None = None class ProxyErrorTypes(str, enum.Enum): @@ -3656,18 +3653,18 @@ DB_RETRY_SAFE_ERROR_TYPES = (httpx.ConnectError,) class SSOUserDefinedValues(TypedDict): - models: List[str] + models: list[str] user_id: str - user_email: Optional[str] - user_role: Optional[str] - max_budget: Optional[float] - budget_duration: Optional[str] + user_email: str | None + user_role: str | None + max_budget: float | None + budget_duration: str | None class VirtualKeyEvent(LiteLLMPydanticObjectBase): created_by_user_id: str created_by_user_role: str - created_by_key_alias: Optional[str] + created_by_key_alias: str | None request_kwargs: dict @@ -3685,7 +3682,7 @@ from litellm.models.team_membership import ( # noqa: E402 class MemberAddRequest(LiteLLMPydanticObjectBase): - member: Union[List[Member], Member] = Field( + member: list[Member] | Member = Field( description="Member object or list of member objects to add. Each member must include either user_id or user_email, and a role" ) @@ -3706,7 +3703,7 @@ class MemberAddRequest(LiteLLMPydanticObjectBase): class OrgMemberAddRequest(LiteLLMPydanticObjectBase): - member: Union[List[OrgMember], OrgMember] + member: list[OrgMember] | OrgMember def __init__(self, **data): member_data = data.get("member") @@ -3728,19 +3725,19 @@ class OrgMemberAddRequest(LiteLLMPydanticObjectBase): class TeamAddMemberResponse(LiteLLM_TeamTable): - updated_users: List[LiteLLM_UserTable] - updated_team_memberships: List[LiteLLM_TeamMembership] + updated_users: list[LiteLLM_UserTable] + updated_team_memberships: list[LiteLLM_TeamMembership] class OrganizationAddMemberResponse(LiteLLMPydanticObjectBase): organization_id: str - updated_users: List[LiteLLM_UserTable] - updated_organization_memberships: List[LiteLLM_OrganizationMembershipTable] + updated_users: list[LiteLLM_UserTable] + updated_organization_memberships: list[LiteLLM_OrganizationMembershipTable] class MemberDeleteRequest(LiteLLMPydanticObjectBase): - user_id: Optional[str] = None - user_email: Optional[str] = None + user_id: str | None = None + user_email: str | None = None @model_validator(mode="before") @classmethod @@ -3752,7 +3749,7 @@ class MemberDeleteRequest(LiteLLMPydanticObjectBase): class MemberUpdateResponse(LiteLLMPydanticObjectBase): user_id: str - user_email: Optional[str] = None + user_email: str | None = None # Team Member Requests @@ -3774,15 +3771,15 @@ class TeamMemberAddRequest(MemberAddRequest): """ team_id: str = Field(description="The ID of the team to add the member to") - max_budget_in_team: Optional[float] = Field( + max_budget_in_team: float | None = Field( default=None, description="Maximum budget allocated to this user within the team. If not set, user has unlimited budget within team limits", ) - budget_duration: Optional[str] = Field( + budget_duration: str | None = Field( default=None, description="Duration after which this team member's budget resets (e.g. '1h', '24h', '7d', '30d'). If not set, the budget never resets.", ) - allowed_models: Optional[List[str]] = Field( + allowed_models: list[str] | None = Field( default=None, description="List of models this team member can access. If not set, inherits the team's default_team_member_models or all team models.", ) @@ -3793,15 +3790,15 @@ class TeamMemberDeleteRequest(MemberDeleteRequest): class TeamMemberUpdateRequest(TeamMemberDeleteRequest): - max_budget_in_team: Optional[float] = None - role: Optional[Literal["admin", "user"]] = None - tpm_limit: Optional[int] = Field(default=None, description="Tokens per minute limit for this team member") - rpm_limit: Optional[int] = Field(default=None, description="Requests per minute limit for this team member") - budget_duration: Optional[str] = Field( + max_budget_in_team: float | None = None + role: Literal["admin", "user"] | None = None + tpm_limit: int | None = Field(default=None, description="Tokens per minute limit for this team member") + rpm_limit: int | None = Field(default=None, description="Requests per minute limit for this team member") + budget_duration: str | None = Field( default=None, description="Duration after which this team member's budget resets (e.g. '1h', '24h', '7d', '30d'). If not set, the budget never resets.", ) - allowed_models: Optional[List[str]] = Field( + allowed_models: list[str] | None = Field( default=None, description="List of models this team member can access. Pass an empty list to remove per-member model restrictions.", ) @@ -3809,31 +3806,31 @@ class TeamMemberUpdateRequest(TeamMemberDeleteRequest): class TeamMemberUpdateResponse(MemberUpdateResponse): team_id: str - max_budget_in_team: Optional[float] = None - tpm_limit: Optional[int] = None - rpm_limit: Optional[int] = None - budget_duration: Optional[str] = None - allowed_models: Optional[List[str]] = None + max_budget_in_team: float | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + budget_duration: str | None = None + allowed_models: list[str] | None = None class TeamModelAddRequest(BaseModel): """Request to add models to a team""" team_id: str - models: List[str] + models: list[str] class TeamModelDeleteRequest(BaseModel): """Request to delete models from a team""" team_id: str - models: List[str] + models: list[str] # Organization Member Requests class OrganizationMemberAddRequest(OrgMemberAddRequest): organization_id: str - max_budget_in_organization: Optional[float] = None # Users max budget within the organization + max_budget_in_organization: float | None = None # Users max budget within the organization class OrganizationMemberDeleteRequest(MemberDeleteRequest): @@ -3848,11 +3845,11 @@ ROLES_WITHIN_ORG = [ class OrganizationMemberUpdateRequest(OrganizationMemberDeleteRequest): - max_budget_in_organization: Optional[float] = None - role: Optional[LitellmUserRoles] = None + max_budget_in_organization: float | None = None + role: LitellmUserRoles | None = None @field_validator("role") - def validate_role(cls, value: Optional[LitellmUserRoles]) -> Optional[LitellmUserRoles]: + def validate_role(cls, value: LitellmUserRoles | None) -> LitellmUserRoles | None: if value is not None and value not in ROLES_WITHIN_ORG: raise ValueError(f"Invalid role. Must be one of: {[role.value for role in ROLES_WITHIN_ORG]}") return value @@ -3867,30 +3864,30 @@ class OrganizationMemberUpdateResponse(MemberUpdateResponse): class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable): - team_member_budget_table: Optional[LiteLLM_BudgetTableFull] = None + team_member_budget_table: LiteLLM_BudgetTableFull | None = None # Resources inherited from access groups (separate from direct assignments) - access_group_models: Optional[List[str]] = None - access_group_mcp_server_ids: Optional[List[str]] = None - access_group_agent_ids: Optional[List[str]] = None + access_group_models: list[str] | None = None + access_group_mcp_server_ids: list[str] | None = None + access_group_agent_ids: list[str] | None = None class TeamInfoResponseObject(TypedDict): team_id: str team_info: TeamInfoResponseObjectTeamTable - keys: List - team_memberships: List[LiteLLM_TeamMembership] + keys: list + team_memberships: list[LiteLLM_TeamMembership] class TeamListResponseObject(LiteLLM_TeamTable): - team_memberships: List[LiteLLM_TeamMembership] - keys: List # list of keys that belong to the team + team_memberships: list[LiteLLM_TeamMembership] + keys: list # list of keys that belong to the team class KeyListResponseObject(TypedDict, total=False): - keys: List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]] - total_count: Optional[int] - current_page: Optional[int] - total_pages: Optional[int] + keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken] + total_count: int | None + current_page: int | None + total_pages: int | None class CurrentItemRateLimit(TypedDict): @@ -3900,28 +3897,28 @@ class CurrentItemRateLimit(TypedDict): class LoggingCallbackStatus(TypedDict, total=False): - callbacks: List[str] + callbacks: list[str] status: Literal["healthy", "unhealthy"] - details: Optional[str] + details: str | None class KeyHealthResponse(TypedDict, total=False): key: Literal["healthy", "unhealthy"] - logging_callbacks: Optional[LoggingCallbackStatus] + logging_callbacks: LoggingCallbackStatus | None class CreateJWTKeyMappingRequest(LiteLLMPydanticObjectBase): jwt_claim_name: str jwt_claim_value: str key: str - description: Optional[str] = None + description: str | None = None class UpdateJWTKeyMappingRequest(LiteLLMPydanticObjectBase): id: str - key: Optional[str] = None - description: Optional[str] = None - is_active: Optional[bool] = None + key: str | None = None + description: str | None = None + is_active: bool | None = None class DeleteJWTKeyMappingRequest(LiteLLMPydanticObjectBase): @@ -3932,12 +3929,12 @@ class JWTKeyMappingResponse(LiteLLMPydanticObjectBase): id: str jwt_claim_name: str jwt_claim_value: str - description: Optional[str] = None + description: str | None = None is_active: bool created_at: datetime updated_at: datetime - created_by: Optional[str] = None - updated_by: Optional[str] = None + created_by: str | None = None + updated_by: str | None = None class SpecialHeaders(enum.Enum): @@ -3979,10 +3976,10 @@ class SpecialHeaders(enum.Enum): class LitellmDataForBackendLLMCall(TypedDict, total=False): headers: dict organization: str - timeout: Optional[float] - stream_timeout: Optional[float] - user: Optional[str] - num_retries: Optional[int] + timeout: float | None + stream_timeout: float | None + user: str | None + num_retries: int | None class LitellmMetadataFromRequestHeaders(TypedDict, total=False): @@ -3990,17 +3987,17 @@ class LitellmMetadataFromRequestHeaders(TypedDict, total=False): Headers a user can pass that will get added to litellm metadata for the request """ - spend_logs_metadata: Optional[dict] - agent_id: Optional[str] - trace_id: Optional[str] - session_id: Optional[str] + spend_logs_metadata: dict | None + agent_id: str | None + trace_id: str | None + session_id: str | None class JWTKeyItem(TypedDict, total=False): kid: str -JWKKeyValue = Union[List[JWTKeyItem], JWTKeyItem] +JWKKeyValue = Union[list[JWTKeyItem], JWTKeyItem] class JWKUrlResponse(TypedDict, total=False): @@ -4055,7 +4052,7 @@ PassThroughEndpointLoggingResultValues = Union[ class PassThroughEndpointLoggingTypedDict(TypedDict): - result: Optional[PassThroughEndpointLoggingResultValues] + result: PassThroughEndpointLoggingResultValues | None kwargs: dict @@ -4099,10 +4096,10 @@ class ProviderBudgetResponseObject(LiteLLMPydanticObjectBase): Configuration for a single provider's budget settings """ - budget_limit: Optional[float] # Budget limit in USD for the time period - time_period: Optional[str] # Time period for budget (e.g., '1d', '30d', '1mo') - spend: Optional[float] = 0.0 # Current spend for this provider - budget_reset_at: Optional[str] = None # When the current budget period resets + budget_limit: float | None # Budget limit in USD for the time period + time_period: str | None # Time period for budget (e.g., '1d', '30d', '1mo') + spend: float | None = 0.0 # Current spend for this provider + budget_reset_at: str | None = None # When the current budget period resets class ProviderBudgetResponse(LiteLLMPydanticObjectBase): @@ -4111,7 +4108,7 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase): Maps provider names to their budget configs. """ - providers: Dict[ + providers: dict[ str, ProviderBudgetResponseObject ] = {} # Dictionary mapping provider names to their budget configurations @@ -4129,17 +4126,17 @@ UI_TEAM_ID = "litellm-dashboard" class JWTAuthBuilderResult(TypedDict): is_proxy_admin: bool - team_object: Optional[LiteLLM_TeamTable] - user_object: Optional[LiteLLM_UserTable] - end_user_object: Optional[LiteLLM_EndUserTable] - org_object: Optional[LiteLLM_OrganizationTable] + team_object: LiteLLM_TeamTable | None + user_object: LiteLLM_UserTable | None + end_user_object: LiteLLM_EndUserTable | None + org_object: LiteLLM_OrganizationTable | None token: str - team_id: Optional[str] - user_id: Optional[str] + team_id: str | None + user_id: str | None user_email: str | None - end_user_id: Optional[str] - org_id: Optional[str] - team_membership: Optional[LiteLLM_TeamMembership] + end_user_id: str | None + org_id: str | None + team_membership: LiteLLM_TeamMembership | None jwt_claims: dict # Decoded JWT token claims (avoids re-decoding) @@ -4149,7 +4146,7 @@ class ClientSideFallbackModel(TypedDict, total=False): """ model: Required[str] - messages: List[AllMessageValues] + messages: list[AllMessageValues] ALL_FALLBACK_MODEL_VALUES = Union[str, ClientSideFallbackModel] @@ -4163,8 +4160,8 @@ RBAC_ROLES = Literal[ class OIDCPermissions(LiteLLMPydanticObjectBase): - models: Optional[List[str]] = None - routes: Optional[List[str]] = None + models: list[str] | None = None + routes: list[str] | None = None class RoleBasedPermissions(OIDCPermissions): @@ -4206,10 +4203,10 @@ class JWTRoutingOverride(BaseModel): scope strings), not to ``iss``, ``aud``, or ``client_id``. """ - iss: Union[str, List[str]] - client_id: Optional[Union[str, List[str]]] = None - scope: Optional[Union[str, List[str]]] = None - aud: Optional[Union[str, List[str]]] = None + iss: str | list[str] + client_id: str | list[str] | None = None + scope: str | list[str] | None = None + aud: str | list[str] | None = None path: Literal["oauth2"] = "oauth2" model_config = { @@ -4248,11 +4245,11 @@ class JWTIssuerConfig(BaseModel): """ issuer: str = Field(description="Exact expected JWT issuer (`iss`) value.") - jwks_url: Optional[str] = Field( + jwks_url: str | None = Field( default=None, description="Issuer JWKS URL. If omitted, LiteLLM uses the issuer's OIDC discovery document.", ) - audience: Optional[Union[str, List[str]]] = Field( + audience: str | list[str] | None = Field( default=None, description="Expected token audience for this issuer.", ) @@ -4260,27 +4257,27 @@ class JWTIssuerConfig(BaseModel): default=False, description="Explicitly disable audience validation for this issuer. Use only when the issuer cannot provide an audience suitable for LiteLLM.", ) - user_id_jwt_field: Optional[str] = Field( + user_id_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's user id.", ) - user_email_jwt_field: Optional[str] = Field( + user_email_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's user email.", ) - team_id_jwt_field: Optional[str] = Field( + team_id_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's team id.", ) - team_ids_jwt_field: Optional[str] = Field( + team_ids_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's team ids.", ) - org_id_jwt_field: Optional[str] = Field( + org_id_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's organization id.", ) - end_user_id_jwt_field: Optional[str] = Field( + end_user_id_jwt_field: str | None = Field( default=None, description="Issuer-specific claim path to normalize into LiteLLM's end-user id.", ) @@ -4328,62 +4325,62 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): """ admin_jwt_scope: str = "litellm_proxy_admin" - admin_allowed_routes: List[str] = [ + admin_allowed_routes: list[str] = [ "management_routes", "spend_tracking_routes", "global_spend_tracking_routes", "info_routes", ] - team_id_jwt_field: Optional[str] = None + team_id_jwt_field: str | None = None team_id_upsert: bool = False - team_ids_jwt_field: Optional[str] = None + team_ids_jwt_field: str | None = None upsert_sso_user_to_team: bool = False - team_allowed_routes: List[str] = [ + team_allowed_routes: list[str] = [ "openai_routes", "info_routes", "mcp_routes", "/v1/messages", "/v1/messages/count_tokens", ] - team_id_default: Optional[str] = Field( + team_id_default: str | None = Field( default=None, description="If no team_id given, default permissions/spend-tracking to this team.s", ) - team_alias_jwt_field: Optional[str] = Field( + team_alias_jwt_field: str | None = Field( default=None, description="The field in the JWT token that stores the team name/alias. Will be resolved to team_id via database lookup.", ) - org_id_jwt_field: Optional[str] = None - org_alias_jwt_field: Optional[str] = Field( + org_id_jwt_field: str | None = None + org_alias_jwt_field: str | None = Field( default=None, description="The field in the JWT token that stores the organization name/alias. Will be resolved to org_id via database lookup.", ) - user_id_jwt_field: Optional[str] = None - user_email_jwt_field: Optional[str] = None - user_allowed_email_domain: Optional[str] = None - user_roles_jwt_field: Optional[str] = None - user_allowed_roles: Optional[List[str]] = None + user_id_jwt_field: str | None = None + user_email_jwt_field: str | None = None + user_allowed_email_domain: str | None = None + user_roles_jwt_field: str | None = None + user_allowed_roles: list[str] | None = None user_id_upsert: bool = Field(default=False, description="If user doesn't exist, upsert them into the db.") - end_user_id_jwt_field: Optional[str] = None + end_user_id_jwt_field: str | None = None public_key_ttl: float = 600 - public_allowed_routes: List[str] = ["public_routes"] + public_allowed_routes: list[str] = ["public_routes"] enforce_rbac: bool = False - roles_jwt_field: Optional[str] = None # v2 on role mappings - role_mappings: Optional[List[RoleMapping]] = None - object_id_jwt_field: Optional[str] = None # can be either user / team, inferred from the role mapping - scope_mappings: Optional[List[ScopeMapping]] = None + roles_jwt_field: str | None = None # v2 on role mappings + role_mappings: list[RoleMapping] | None = None + object_id_jwt_field: str | None = None # can be either user / team, inferred from the role mapping + scope_mappings: list[ScopeMapping] | None = None enforce_scope_based_access: bool = False enforce_team_based_model_access: bool = False - custom_validate: Optional[Callable[..., Literal[True]]] = None + custom_validate: Callable[..., Literal[True]] | None = None ######################################################### # Fields for syncing user team membership and roles with IDP provider - jwt_litellm_role_map: Optional[List[JWTLiteLLMRoleMap]] = None + jwt_litellm_role_map: list[JWTLiteLLMRoleMap] | None = None sync_user_role_and_teams: bool = False ######################################################### ######################################################### # OIDC UserInfo Endpoint Configuration - oidc_userinfo_endpoint: Optional[str] = Field( + oidc_userinfo_endpoint: str | None = Field( default=None, description="OIDC UserInfo endpoint URL. If set, LiteLLM will call this endpoint with the access token to retrieve user identity information.", ) @@ -4396,7 +4393,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): description="TTL (in seconds) for caching UserInfo responses. Default: 300s (5 minutes).", ) # JWT-to-Virtual-Key Mapping - virtual_key_claim_field: Optional[str] = Field( + virtual_key_claim_field: str | None = Field( default=None, description="JWT claim field for virtual key mapping lookup (e.g. 'sub', 'email'). Supports dot notation.", ) @@ -4413,7 +4410,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): "'auto_register': auto-create a virtual key and mapping on first encounter." ), ) - routing_overrides: Optional[List[JWTRoutingOverride]] = Field( + routing_overrides: list[JWTRoutingOverride] | None = Field( default=None, description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.", ) @@ -4438,7 +4435,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): "records exist before the fallback runs." ), ) - issuers: Optional[List[JWTIssuerConfig]] = Field( + issuers: list[JWTIssuerConfig] | None = Field( default=None, description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.", ) @@ -4520,28 +4517,29 @@ class DefaultInternalUserParams(LiteLLMPydanticObjectBase): Default parameters to apply when a new user signs in via SSO or is created on the /user/new API endpoint """ - user_role: Optional[ + user_role: ( Literal[ LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, ] - ] = Field( + | None + ) = Field( default=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, description="Default role assigned to new users created", ) - max_budget: Optional[float] = Field( + max_budget: float | None = Field( default=None, description="Default maximum budget (in USD) for new users created", ) - budget_duration: Optional[str] = Field( + budget_duration: str | None = Field( default=None, description="Default budget duration for new users (e.g. 'daily', 'weekly', 'monthly')", ) - models: Optional[List[str]] = Field(default=None, description="Default list of models that new users can access") + models: list[str] | None = Field(default=None, description="Default list of models that new users can access") - teams: Optional[Union[List[str], List[NewUserRequestTeam]]] = Field( + teams: list[str] | list[NewUserRequestTeam] | None = Field( default=None, description="Default teams for new users created", ) @@ -4550,11 +4548,11 @@ class DefaultInternalUserParams(LiteLLMPydanticObjectBase): class BaseDailySpendTransaction(TypedDict): date: str api_key: str - model: Optional[str] - model_group: Optional[str] - mcp_namespaced_tool_name: Optional[str] - custom_llm_provider: Optional[str] - endpoint: Optional[str] + model: str | None + model_group: str | None + mcp_namespaced_tool_name: str | None + custom_llm_provider: str | None + endpoint: str | None # token count metrics prompt_tokens: int @@ -4591,7 +4589,7 @@ class DailyEndUserSpendTransaction(BaseDailySpendTransaction): class DailyTagSpendTransaction(BaseDailySpendTransaction): - request_id: Optional[str] + request_id: str | None tag: str @@ -4604,30 +4602,30 @@ class DBSpendUpdateTransactions(TypedDict): Internal Data Structure for buffering spend updates in Redis or in memory before committing them to the database """ - user_list_transactions: Optional[Dict[str, float]] - end_user_list_transactions: Optional[Dict[str, float]] - key_list_transactions: Optional[Dict[str, float]] - team_list_transactions: Optional[Dict[str, float]] - team_member_list_transactions: Optional[Dict[str, float]] - org_list_transactions: Optional[Dict[str, float]] - tag_list_transactions: Optional[Dict[str, float]] - agent_list_transactions: Optional[Dict[str, float]] + user_list_transactions: dict[str, float] | None + end_user_list_transactions: dict[str, float] | None + key_list_transactions: dict[str, float] | None + team_list_transactions: dict[str, float] | None + team_member_list_transactions: dict[str, float] | None + org_list_transactions: dict[str, float] | None + tag_list_transactions: dict[str, float] | None + agent_list_transactions: dict[str, float] | None class SpendUpdateQueueItem(TypedDict, total=False): entity_type: Litellm_EntityType entity_id: str - response_cost: Optional[float] + response_cost: float | None class ToolDiscoveryQueueItem(TypedDict, total=False): tool_name: str - origin: Optional[str] # MCP server name or "user_defined" - created_by: Optional[str] - key_hash: Optional[str] # hash of virtual key that triggered discovery - team_id: Optional[str] # team that triggered discovery - key_alias: Optional[str] # human-readable key alias - user_agent: Optional[str] # HTTP User-Agent of the caller + origin: str | None # MCP server name or "user_defined" + created_by: str | None + key_hash: str | None # hash of virtual key that triggered discovery + team_id: str | None # team that triggered discovery + key_alias: str | None # human-readable key alias + user_agent: str | None # HTTP User-Agent of the caller from litellm.models.managed_files import ( # noqa: E402 @@ -4647,7 +4645,7 @@ from litellm.models.managed_files import ( # noqa: E402 class EnterpriseLicenseData(TypedDict, total=False): expiration_date: str user_id: str - allowed_features: List[str] + allowed_features: list[str] max_users: int max_teams: int @@ -4662,8 +4660,8 @@ class CostEstimateRequest(LiteLLMPydanticObjectBase): model: str = Field(description="Model name (from /model_group/info)") input_tokens: int = Field(description="Expected input tokens per request", ge=0) output_tokens: int = Field(description="Expected output tokens per request", ge=0) - num_requests_per_day: Optional[int] = Field(default=None, description="Number of requests per day", ge=0) - num_requests_per_month: Optional[int] = Field(default=None, description="Number of requests per month", ge=0) + num_requests_per_day: int | None = Field(default=None, description="Number of requests per day", ge=0) + num_requests_per_month: int | None = Field(default=None, description="Number of requests per month", ge=0) class CostEstimateResponse(LiteLLMPydanticObjectBase): @@ -4672,24 +4670,24 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase): model: str input_tokens: int output_tokens: int - num_requests_per_day: Optional[int] = None - num_requests_per_month: Optional[int] = None + num_requests_per_day: int | None = None + num_requests_per_month: int | None = None # Per-request costs cost_per_request: float = Field(description="Total cost per request (includes margin)") input_cost_per_request: float = Field(description="Input token cost per request (before margin)") output_cost_per_request: float = Field(description="Output token cost per request (before margin)") margin_cost_per_request: float = Field(default=0.0, description="Margin/fee added per request") # Daily costs (if num_requests_per_day provided) - daily_cost: Optional[float] = Field(default=None, description="Total daily cost (includes margin)") - daily_input_cost: Optional[float] = Field(default=None, description="Daily input token cost") - daily_output_cost: Optional[float] = Field(default=None, description="Daily output token cost") - daily_margin_cost: Optional[float] = Field(default=None, description="Daily margin/fee") + daily_cost: float | None = Field(default=None, description="Total daily cost (includes margin)") + daily_input_cost: float | None = Field(default=None, description="Daily input token cost") + daily_output_cost: float | None = Field(default=None, description="Daily output token cost") + daily_margin_cost: float | None = Field(default=None, description="Daily margin/fee") # Monthly costs (if num_requests_per_month provided) - monthly_cost: Optional[float] = Field(default=None, description="Total monthly cost (includes margin)") - monthly_input_cost: Optional[float] = Field(default=None, description="Monthly input token cost") - monthly_output_cost: Optional[float] = Field(default=None, description="Monthly output token cost") - monthly_margin_cost: Optional[float] = Field(default=None, description="Monthly margin/fee") + monthly_cost: float | None = Field(default=None, description="Total monthly cost (includes margin)") + monthly_input_cost: float | None = Field(default=None, description="Monthly input token cost") + monthly_output_cost: float | None = Field(default=None, description="Monthly output token cost") + monthly_margin_cost: float | None = Field(default=None, description="Monthly margin/fee") # Pricing info - input_cost_per_token: Optional[float] = None - output_cost_per_token: Optional[float] = None - provider: Optional[str] = None + input_cost_per_token: float | None = None + output_cost_per_token: float | None = None + provider: str | None = None diff --git a/litellm/proxy/a2a/__init__.py b/litellm/proxy/a2a/__init__.py index 10fb308f9f6..5ee0c6fa8e7 100644 --- a/litellm/proxy/a2a/__init__.py +++ b/litellm/proxy/a2a/__init__.py @@ -10,8 +10,8 @@ A2A registration helpers for the LiteLLM proxy. from litellm.proxy.a2a.agent_card import ( LITELLM_A2A_PROTOCOL_VERSION, - LITELLM_SECURITY_SCHEMES, LITELLM_SECURITY_REQUIREMENTS, + LITELLM_SECURITY_SCHEMES, merge_agent_card, ) from litellm.proxy.a2a.discovery import ( diff --git a/litellm/proxy/a2a/agent_card.py b/litellm/proxy/a2a/agent_card.py index 7fb81057bc9..de5369d4c03 100644 --- a/litellm/proxy/a2a/agent_card.py +++ b/litellm/proxy/a2a/agent_card.py @@ -10,7 +10,7 @@ and uses LiteLLM auth. import re from collections.abc import Mapping from copy import deepcopy -from typing import Any, Dict, List, Literal +from typing import Any, Literal SupportedA2AVersion = Literal["0.3", "1.0"] @@ -53,7 +53,7 @@ def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str: # Security scheme exposed by the LiteLLM-fronted agent card. Always replaces # whatever upstream advertised — the client must authenticate to the proxy, # not the upstream agent. -LITELLM_SECURITY_SCHEMES: Dict[str, Dict[str, Any]] = { +LITELLM_SECURITY_SCHEMES: dict[str, dict[str, Any]] = { "LiteLLMKey": { "type": "http", "scheme": "bearer", @@ -61,7 +61,7 @@ LITELLM_SECURITY_SCHEMES: Dict[str, Dict[str, Any]] = { }, } -LITELLM_SECURITY_REQUIREMENTS: List[Dict[str, List[str]]] = [{"LiteLLMKey": []}] +LITELLM_SECURITY_REQUIREMENTS: list[dict[str, list[str]]] = [{"LiteLLMKey": []}] # Capabilities LiteLLM can faithfully proxy today. Anything not in this set is # dropped during merge so we don't advertise behavior the proxy can't deliver. @@ -112,7 +112,7 @@ _ALLOWED_TOP_LEVEL_KEYS = { "url", } -_DEFAULT_SKILLS: List[Dict[str, Any]] = [ +_DEFAULT_SKILLS: list[dict[str, Any]] = [ { "id": "chat", "name": "Chat", @@ -121,7 +121,7 @@ _DEFAULT_SKILLS: List[Dict[str, Any]] = [ } ] -_DEFAULT_MODES: List[str] = ["text"] +_DEFAULT_MODES: list[str] = ["text"] # Fallback ``version`` when the upstream card omits the field. The A2A v1.0 # schema requires ``version`` on every card, so without this default the @@ -129,7 +129,7 @@ _DEFAULT_MODES: List[str] = ["text"] _DEFAULT_AGENT_VERSION = "1.0.0" -def _filter_capabilities(upstream_capabilities: Any) -> Dict[str, Any]: +def _filter_capabilities(upstream_capabilities: Any) -> dict[str, Any]: """Return a capabilities dict containing only allowlisted, truthy keys.""" if not isinstance(upstream_capabilities, dict): return {} @@ -138,7 +138,7 @@ def _filter_capabilities(upstream_capabilities: Any) -> Dict[str, Any]: } -def _default_litellm_provider(proxy_base_url: str) -> Dict[str, str]: +def _default_litellm_provider(proxy_base_url: str) -> dict[str, str]: return {"organization": "LiteLLM Proxy", "url": proxy_base_url} @@ -149,7 +149,7 @@ def merge_agent_card( proxy_base_url: str, name: str | None = None, description: str | None = None, -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Build the LiteLLM-fronted agent card. @@ -169,7 +169,7 @@ def merge_agent_card( A dict suitable for serving as the proxy's agent card. Only keys in the v1.0 AgentCard schema (plus ``supportedInterfaces``) are emitted. """ - base: Dict[str, Any] = deepcopy(dict(upstream_card)) if upstream_card else {} + base: dict[str, Any] = deepcopy(dict(upstream_card)) if upstream_card else {} # Keep the upstream ``url`` on the stored card: the runtime A2A # invocation path reads it from ``agent_card_params`` to know where to diff --git a/litellm/proxy/a2a/discovery.py b/litellm/proxy/a2a/discovery.py index a95f1d2dfeb..66c661972a8 100644 --- a/litellm/proxy/a2a/discovery.py +++ b/litellm/proxy/a2a/discovery.py @@ -15,7 +15,7 @@ fetcher dispatches by ``discovery_mode``: """ from enum import Enum -from typing import Any, Dict, Optional, Tuple +from typing import Any from urllib.parse import urlencode from litellm._logging import verbose_proxy_logger @@ -37,7 +37,7 @@ class DiscoveryMode(str, Enum): # Paths the pure-A2A fetcher tries in order. The first two are the current and # previous A2A spec locations; ``/agent.json`` is a non-standard root fallback # some agents still serve. -AGENT_CARD_WELL_KNOWN_PATHS: Tuple[str, ...] = ( +AGENT_CARD_WELL_KNOWN_PATHS: tuple[str, ...] = ( "/.well-known/agent-card.json", "/.well-known/agent.json", "/agent.json", @@ -55,8 +55,8 @@ def _normalize_base_url(base_url: str) -> str: def _build_langgraph_platform_paths( - params: Optional[Dict[str, Any]], -) -> Tuple[str, ...]: + params: dict[str, Any] | None, +) -> tuple[str, ...]: """Build the paths to try for LangGraph Platform discovery. LangGraph serves the card at ``/.well-known/agent-card.json`` with the @@ -71,7 +71,7 @@ def _build_langgraph_platform_paths( return tuple(f"{path}?{query}" for path in AGENT_CARD_WELL_KNOWN_PATHS) -def _paths_for_mode(mode: DiscoveryMode, params: Optional[Dict[str, Any]]) -> Tuple[str, ...]: +def _paths_for_mode(mode: DiscoveryMode, params: dict[str, Any] | None) -> tuple[str, ...]: if mode == DiscoveryMode.WELL_KNOWN_FALLBACK: return AGENT_CARD_WELL_KNOWN_PATHS if mode == DiscoveryMode.LANGGRAPH_PLATFORM: @@ -83,10 +83,10 @@ async def fetch_well_known_card( base_url: str, *, discovery_mode: DiscoveryMode = DiscoveryMode.WELL_KNOWN_FALLBACK, - params: Optional[Dict[str, Any]] = None, + params: dict[str, Any] | None = None, timeout: float = DEFAULT_DISCOVERY_TIMEOUT_SECONDS, - headers: Optional[Dict[str, str]] = None, -) -> Dict[str, Any]: + headers: dict[str, str] | None = None, +) -> dict[str, Any]: """ Fetch an agent card from ``base_url`` using the strategy chosen by ``discovery_mode``. Returns the parsed JSON from the first path that @@ -106,7 +106,7 @@ async def fetch_well_known_card( params={"timeout": timeout}, ) - last_error: Optional[str] = None + last_error: str | None = None for path in paths: url = f"{normalized}{path}" try: diff --git a/litellm/proxy/a2a/endpoints.py b/litellm/proxy/a2a/endpoints.py index a46de73fabc..bcc07629ab1 100644 --- a/litellm/proxy/a2a/endpoints.py +++ b/litellm/proxy/a2a/endpoints.py @@ -9,7 +9,7 @@ admin pick which ones to expose through the proxy. The actual merge into a LiteLLM-fronted card happens when the agent is saved via ``POST /v1/agents``. """ -from typing import Any, Dict, Optional +from typing import Any from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import JSONResponse @@ -49,7 +49,7 @@ class DiscoverAgentRequest(BaseModel): "query parameter." ), ) - params: Optional[Dict[str, Any]] = Field( + params: dict[str, Any] | None = Field( default=None, description=( "Mode-specific parameters. ``langgraph_platform`` requires " @@ -60,7 +60,7 @@ class DiscoverAgentRequest(BaseModel): class DiscoverAgentResponse(BaseModel): url: str - agent_card: Dict[str, Any] + agent_card: dict[str, Any] @router.post( diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 0f0f724ab67..79808c06daa 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -13,7 +13,7 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM import json from collections.abc import AsyncGenerator from copy import deepcopy -from typing import TYPE_CHECKING, Any, Dict, List +from typing import TYPE_CHECKING, Any from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -46,7 +46,7 @@ if TYPE_CHECKING: router = APIRouter() -_PASCAL_TO_WIRE: Dict[str, str] = { +_PASCAL_TO_WIRE: dict[str, str] = { "SendMessage": "message/send", "SendStreamingMessage": "message/stream", "GetTask": "tasks/get", @@ -107,8 +107,8 @@ def _validate_push_notification_url(url: str) -> None: raise HTTPException(status_code=400, detail=str(e)) from e -def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str]: - headers: Dict[str, str] = {} +def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> dict[str, str]: + headers: dict[str, str] = {} if user_api_key_dict.user_id: headers["X-LiteLLM-User-Id"] = user_api_key_dict.user_id if user_api_key_dict.team_id: @@ -119,8 +119,8 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str def _forwarding_headers( user_api_key_dict: UserAPIKeyAuth, request_data: dict[str, Any], - agent_extra_headers: Dict[str, str] | None, -) -> Dict[str, str] | None: + agent_extra_headers: dict[str, str] | None, +) -> dict[str, str] | None: sanitized = ( {k: v for k, v in agent_extra_headers.items() if not k.lower().startswith("x-litellm-")} if agent_extra_headers @@ -182,7 +182,7 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None: async def _forward_jsonrpc( agent_url: str, body: dict[str, Any], - extra_headers: Dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, ) -> dict[str, Any]: from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider @@ -207,7 +207,7 @@ async def _a2a_sse_event_source( agent_url: str, body: dict[str, Any], request_id: Any | None = None, - extra_headers: Dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, served_version: A2AVersion = "0.3", ) -> AsyncGenerator[dict, None]: """Stream an upstream A2A SSE response as parsed JSON-RPC event dicts. @@ -269,7 +269,7 @@ async def _forward_jsonrpc_sse( agent_url: str, body: dict[str, Any], request_id: Any | None = None, - extra_headers: Dict[str, str] | None = None, + extra_headers: dict[str, str] | None = None, proxy_logging_obj: Any | None = None, user_api_key_dict: Any | None = None, request_data: dict[str, Any] | None = None, @@ -338,7 +338,7 @@ async def _handle_stream_message( metadata: dict[str, Any] | None = None, proxy_server_request: dict[str, Any] | None = None, *, - agent_extra_headers: Dict[str, str] | None = None, + agent_extra_headers: dict[str, str] | None = None, user_api_key_dict: UserAPIKeyAuth | None = None, request_data: dict[str, Any] | None = None, proxy_logging_obj: Any | None = None, @@ -491,7 +491,7 @@ async def _handle_stream_message( "id": request_id, "error": { "code": -32603, - "message": f"Streaming error: {str(e)}", + "message": f"Streaming error: {e!s}", }, } ) @@ -609,8 +609,8 @@ async def invoke_agent_a2a( version, ) - body: Dict[str, Any] = {} - request_data: Dict[str, Any] = body + body: dict[str, Any] = {} + request_data: dict[str, Any] = body try: body = await request.json() request_data = body @@ -723,12 +723,12 @@ async def invoke_agent_a2a( request_data = data # Build merged headers for the backend agent - static_headers: Dict[str, str] = dict(agent.static_headers or {}) + static_headers: dict[str, str] = dict(agent.static_headers or {}) raw_headers = dict(request.headers) normalized = {k.lower(): v for k, v in raw_headers.items()} - dynamic_headers: Dict[str, str] = {} + dynamic_headers: dict[str, str] = {} # 1. Admin-configured extra_headers: forward named headers from client request if agent.extra_headers: @@ -772,7 +772,7 @@ async def invoke_agent_a2a( if _agent_guardrails: if not isinstance(_agent_guardrails, list): _agent_guardrails = [_agent_guardrails] - _existing_guardrails: List = data.get("guardrails") or [] + _existing_guardrails: list = data.get("guardrails") or [] if not isinstance(_existing_guardrails, list): _existing_guardrails = [_existing_guardrails] data["guardrails"] = _existing_guardrails + [g for g in _agent_guardrails if g not in _existing_guardrails] @@ -826,7 +826,7 @@ async def invoke_agent_a2a( logging_obj._enqueue_deferred_logging = None # type: ignore[union-attr] _enqueue_fn() - response_dict: Dict[str, Any] = ( + response_dict: dict[str, Any] = ( response.model_dump(mode="json", exclude_none=True) # type: ignore if hasattr(response, "model_dump") else response @@ -974,4 +974,4 @@ async def invoke_agent_a2a( ) except Exception: pass - return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {str(e)}", 500) + return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {e!s}", 500) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index a9b75de9807..0410f067560 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -5,11 +5,11 @@ Handles routing for A2A agents (models with "a2a/" prefix). Looks up agents in the registry and injects their API base URL. """ -from typing import Any, Optional +from typing import Any -import litellm from fastapi import HTTPException +import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth @@ -17,8 +17,8 @@ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth async def route_a2a_agent_request( data: dict, route_type: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> Optional[Any]: + user_api_key_dict: UserAPIKeyAuth | None = None, +) -> Any | None: """ Route A2A agent requests directly to litellm with injected API base. diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index e0ed076bcfc..60367178c7f 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -129,7 +129,7 @@ class AgentRegistry: sequence clears them. """ if agent_config is None: - return None + return self.config_agents = tuple(agent_config) @@ -271,7 +271,7 @@ class AgentRegistry: created_agent_dict["object_permission"] = created_agent.object_permission.dict() return AgentResponse(**created_agent_dict) # type: ignore except Exception as e: - raise Exception(f"Error adding agent to DB: {str(e)}") + raise Exception(f"Error adding agent to DB: {e!s}") async def delete_agent_from_db(self, agent_id: str, prisma_client: PrismaClient) -> Mapping[str, object]: """ @@ -281,7 +281,7 @@ class AgentRegistry: deleted_agent = await agents_table(prisma_client).delete(where={"agent_id": agent_id}) return dict(deleted_agent) except Exception as e: - raise Exception(f"Error deleting agent from DB: {str(e)}") + raise Exception(f"Error deleting agent from DB: {e!s}") async def patch_agent_in_db( self, @@ -363,7 +363,7 @@ class AgentRegistry: patched_agent_dict["object_permission"] = patched_agent.object_permission.dict() return AgentResponse(**patched_agent_dict) # type: ignore except Exception as e: - raise Exception(f"Error patching agent in DB: {str(e)}") + raise Exception(f"Error patching agent in DB: {e!s}") async def update_agent_in_db( self, @@ -450,7 +450,7 @@ class AgentRegistry: updated_agent_dict["object_permission"] = updated_agent.object_permission.dict() return AgentResponse(**updated_agent_dict) # type: ignore except Exception as e: - raise Exception(f"Error updating agent in DB: {str(e)}") + raise Exception(f"Error updating agent in DB: {e!s}") @staticmethod async def get_all_agents_from_db( @@ -478,7 +478,7 @@ class AgentRegistry: return agents except Exception as e: - raise Exception(f"Error getting agents from DB: {str(e)}") + raise Exception(f"Error getting agents from DB: {e!s}") def get_agent_by_id( self, @@ -494,7 +494,7 @@ class AgentRegistry: return None except Exception as e: - raise Exception(f"Error getting agent from DB: {str(e)}") + raise Exception(f"Error getting agent from DB: {e!s}") def get_agent_by_name(self, agent_name: str) -> AgentResponse | None: """ @@ -507,7 +507,7 @@ class AgentRegistry: return None except Exception as e: - raise Exception(f"Error getting agent from DB: {str(e)}") + raise Exception(f"Error getting agent from DB: {e!s}") global_agent_registry = AgentRegistry() diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 24f70351ddc..8acff11b009 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -5,8 +5,6 @@ Handles agent permission checking for keys and teams using object_permission_id. Follows the same pattern as MCP permission handling. """ -from typing import List, Optional, Set - from litellm._logging import verbose_logger from litellm.proxy._types import ( UI_TEAM_ID, @@ -33,8 +31,8 @@ class AgentRequestHandler: @staticmethod async def get_allowed_agents( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get list of allowed agent IDs for the given user/key based on permissions. @@ -42,7 +40,7 @@ class AgentRequestHandler: List[str]: List of allowed agent IDs. Empty list means no restrictions (allow all). """ try: - allowed_agents: List[str] = [] + allowed_agents: list[str] = [] allowed_agents_for_key = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth) allowed_agents_for_team = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth) @@ -61,13 +59,13 @@ class AgentRequestHandler: return list(set(allowed_agents)) except Exception as e: - verbose_logger.warning(f"Failed to get allowed agents: {str(e)}") + verbose_logger.warning(f"Failed to get allowed agents: {e!s}") return [] @staticmethod async def is_agent_allowed( agent_id: str, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> bool: """ Check if a specific agent is allowed for the given user/key. @@ -89,8 +87,8 @@ class AgentRequestHandler: @staticmethod def _get_key_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> Optional[LiteLLM_ObjectPermissionTable]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> LiteLLM_ObjectPermissionTable | None: """ Get key object_permission - already loaded by get_key_object() in main auth flow. @@ -104,8 +102,8 @@ class AgentRequestHandler: @staticmethod async def _get_team_object_permission( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> Optional[LiteLLM_ObjectPermissionTable]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> LiteLLM_ObjectPermissionTable | None: """ Get team object_permission - automatically loaded by get_team_object() in main auth flow. @@ -123,7 +121,7 @@ class AgentRequestHandler: return None # Get the team object (which has object_permission already loaded) - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( + team_obj: LiteLLM_TeamTable | None = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -138,8 +136,8 @@ class AgentRequestHandler: @staticmethod async def _get_allowed_agents_for_key( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get allowed agents for a key. @@ -152,7 +150,7 @@ class AgentRequestHandler: return [] try: - all_agents: List[str] = [] + all_agents: list[str] = [] # 1. Get agents from object_permission (native permissions) key_object_permission = AgentRequestHandler._get_key_object_permission(user_api_key_auth) @@ -181,13 +179,13 @@ class AgentRequestHandler: return list(set(all_agents)) except Exception as e: - verbose_logger.warning(f"Failed to get allowed agents for key: {str(e)}") + verbose_logger.warning(f"Failed to get allowed agents for key: {e!s}") return [] @staticmethod async def _get_allowed_agents_for_team( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get allowed agents for a team. @@ -225,7 +223,7 @@ class AgentRequestHandler: if team_obj is None: return [] - all_agents: List[str] = [] + all_agents: list[str] = [] # 1. Get agents from object_permission (native permissions) object_permissions = team_obj.object_permission @@ -257,15 +255,15 @@ class AgentRequestHandler: # litellm-dashboard is the default UI team and will never have agents; # skip noisy warnings for it. if user_api_key_auth.team_id != UI_TEAM_ID: - verbose_logger.warning(f"Failed to get allowed agents for team: {str(e)}") + verbose_logger.warning(f"Failed to get allowed agents for team: {e!s}") return [] @staticmethod - def _get_config_agent_ids_for_access_groups(config_agents: List, access_groups: List[str]) -> Set[str]: + def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]: """ Helper to get agent_ids from config-loaded agents that match any of the given access groups. """ - server_ids: Set[str] = set() + server_ids: set[str] = set() for agent in config_agents: agent_access_groups = getattr(agent, "agent_access_groups", None) if agent_access_groups: @@ -274,11 +272,11 @@ class AgentRequestHandler: return server_ids @staticmethod - async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: List[str]) -> Set[str]: + async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]: """ Helper to get agent_ids from DB agents that match any of the given access groups. """ - agent_ids: Set[str] = set() + agent_ids: set[str] = set() if access_groups and prisma_client is not None: try: agents = await AgentsRepository(prisma_client).table.find_many( @@ -292,8 +290,8 @@ class AgentRequestHandler: @staticmethod async def _get_agents_from_access_groups( - access_groups: List[str], - ) -> List[str]: + access_groups: list[str], + ) -> list[str]: """ Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents. """ @@ -312,17 +310,17 @@ class AgentRequestHandler: return list(agent_ids) except Exception as e: - verbose_logger.warning(f"Failed to get agents from access groups: {str(e)}") + verbose_logger.warning(f"Failed to get agents from access groups: {e!s}") return [] @staticmethod async def get_agent_access_groups( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """ Get list of agent access groups for the given user/key based on permissions. """ - access_groups: List[str] = [] + access_groups: list[str] = [] access_groups_for_key = await AgentRequestHandler._get_agent_access_groups_for_key(user_api_key_auth) access_groups_for_team = await AgentRequestHandler._get_agent_access_groups_for_team(user_api_key_auth) @@ -338,8 +336,8 @@ class AgentRequestHandler: @staticmethod async def _get_agent_access_groups_for_key( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """Get agent access groups for the key.""" from litellm.proxy.auth.auth_checks import get_object_permission from litellm.proxy.proxy_server import ( @@ -371,13 +369,13 @@ class AgentRequestHandler: return key_object_permission.agent_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get agent access groups for key: {str(e)}") + verbose_logger.warning(f"Failed to get agent access groups for key: {e!s}") return [] @staticmethod async def _get_agent_access_groups_for_team( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: """Get agent access groups for the team.""" from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( @@ -397,7 +395,7 @@ class AgentRequestHandler: return [] try: - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( + team_obj: LiteLLM_TeamTable | None = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -414,5 +412,5 @@ class AgentRequestHandler: return object_permissions.agent_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get agent access groups for team: {str(e)}") + verbose_logger.warning(f"Failed to get agent access groups for team: {e!s}") return [] diff --git a/litellm/proxy/agent_endpoints/databricks_oauth.py b/litellm/proxy/agent_endpoints/databricks_oauth.py index 1c01f789916..eabf2be7f3f 100644 --- a/litellm/proxy/agent_endpoints/databricks_oauth.py +++ b/litellm/proxy/agent_endpoints/databricks_oauth.py @@ -26,7 +26,7 @@ import asyncio import base64 import hashlib from dataclasses import dataclass -from typing import Any, Dict, Optional, Tuple +from typing import Any import httpx @@ -43,7 +43,7 @@ _TOKEN_EXPIRY_BUFFER_SECONDS = 60 _DEFAULT_TTL_SECONDS = 3600 -def _resolve_secret(value: Any) -> Optional[str]: +def _resolve_secret(value: Any) -> str | None: """Resolve a config value, expanding ``os.environ/`` references.""" if not isinstance(value, str): return None @@ -55,8 +55,7 @@ def _resolve_secret(value: Any) -> Optional[str]: def _token_url_from_workspace(workspace_url: str) -> str: """Build the workspace OIDC token endpoint from a workspace URL.""" base = workspace_url.strip().rstrip("/") - if base.endswith("/serving-endpoints"): - base = base[: -len("/serving-endpoints")] + base = base.removesuffix("/serving-endpoints") return f"{base}/oidc/v1/token" @@ -76,8 +75,8 @@ class DatabricksAppOAuthConfig: def parse_databricks_oauth_config( - litellm_params: Optional[Dict[str, Any]], -) -> Optional[DatabricksAppOAuthConfig]: + litellm_params: dict[str, Any] | None, +) -> DatabricksAppOAuthConfig | None: """Build a Databricks App OAuth config from an agent's ``litellm_params``. Returns ``None`` when the agent has no ``databricks_oauth`` block. Raises @@ -129,7 +128,7 @@ class DatabricksAppOAuthTokenCache(InMemoryCache): def __init__(self) -> None: super().__init__(default_ttl=_DEFAULT_TTL_SECONDS) - self._locks: Dict[str, asyncio.Lock] = {} + self._locks: dict[str, asyncio.Lock] = {} def _get_lock(self, cache_key: str) -> asyncio.Lock: return self._locks.setdefault(cache_key, asyncio.Lock()) @@ -167,7 +166,7 @@ class DatabricksAppOAuthTokenCache(InMemoryCache): self._locks.pop(cache_key, None) return token - async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> Tuple[str, int]: + async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> tuple[str, int]: client = get_async_httpx_client(llm_provider=httpxSpecialProvider.A2A) verbose_logger.debug("Fetching Databricks App OAuth token from %s", config.token_url) @@ -216,8 +215,8 @@ databricks_app_oauth_token_cache = DatabricksAppOAuthTokenCache() async def resolve_databricks_app_auth_header( - litellm_params: Optional[Dict[str, Any]], -) -> Optional[Dict[str, str]]: + litellm_params: dict[str, Any] | None, +) -> dict[str, str] | None: """Return ``{"Authorization": "Bearer "}`` for a Databricks App agent. Returns ``None`` when the agent is not configured for Databricks App OAuth. diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index c3308bbfa8c..1efbdeb0132 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -148,9 +148,7 @@ def _check_agent_management_permission(user_api_key_dict: UserAPIKeyAuth) -> Non raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can create, update, or delete agents. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can create, update, or delete agents. Your role={user_api_key_dict.user_role}" }, ) @@ -318,10 +316,8 @@ async def get_agents( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.agent_endpoints.get_agents(): Exception occurred - {}".format(str(e)) - ) - raise HTTPException(status_code=500, detail={"error": f"Internal server error: {str(e)}"}) + verbose_proxy_logger.exception(f"litellm.proxy.agent_endpoints.get_agents(): Exception occurred - {e!s}") + raise HTTPException(status_code=500, detail={"error": f"Internal server error: {e!s}"}) #### CRUD ENDPOINTS FOR AGENTS #### @@ -851,9 +847,7 @@ async def make_agent_public( raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can update public model groups. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can update public model groups. Your role={user_api_key_dict.user_role}" }, ) @@ -964,9 +958,7 @@ async def make_agents_public( raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can update public model groups. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can update public model groups. Your role={user_api_key_dict.user_role}" }, ) diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py index d8e2639521f..56053f59c85 100644 --- a/litellm/proxy/agent_endpoints/model_list_helpers.py +++ b/litellm/proxy/agent_endpoints/model_list_helpers.py @@ -4,8 +4,6 @@ Helper functions for appending A2A agents to model lists. Used by proxy model endpoints to make agents appear in UI alongside models. """ -from typing import List - from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.types.proxy.management_endpoints.model_management_endpoints import ( @@ -14,9 +12,9 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import async def append_agents_to_model_group( - model_groups: List[ModelGroupInfoProxy], + model_groups: list[ModelGroupInfoProxy], user_api_key_dict: UserAPIKeyAuth, -) -> List[ModelGroupInfoProxy]: +) -> list[ModelGroupInfoProxy]: """ Append A2A agents to model groups list for UI display. @@ -48,9 +46,9 @@ async def append_agents_to_model_group( async def append_agents_to_model_info( - models: List[dict], + models: list[dict], user_api_key_dict: UserAPIKeyAuth, -) -> List[dict]: +) -> list[dict]: """ Append A2A agents to model info list for UI display. diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index 7d69e09b8da..03fdede0cf4 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -18,7 +18,7 @@ Endpoints: import json import re from datetime import datetime, timezone -from typing import Any, Dict +from typing import Any from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import JSONResponse @@ -87,7 +87,7 @@ async def get_marketplace(): verbose_proxy_logger.warning(f"Plugin {plugin.name} has no source field, skipping") continue - entry: Dict[str, Any] = { + entry: dict[str, Any] = { "name": plugin.name, "source": manifest["source"], } @@ -121,7 +121,7 @@ async def get_marketplace(): verbose_proxy_logger.exception(f"Error generating marketplace: {e}") raise HTTPException( status_code=500, - detail={"error": f"Failed to generate marketplace: {str(e)}"}, + detail={"error": f"Failed to generate marketplace: {e!s}"}, ) @@ -132,7 +132,7 @@ async def get_marketplace(): _VALID_GIT_SUBDIR_PATH_RE = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9._-]*(/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$") -def _validate_plugin_source(source: Dict[str, Any]) -> None: +def _validate_plugin_source(source: dict[str, Any]) -> None: """Validate plugin source format, raising HTTPException on invalid input.""" source_type = source.get("source") if source_type == "github": @@ -231,7 +231,7 @@ async def register_plugin( _validate_plugin_source(source) # Build manifest for storage - manifest: Dict[str, Any] = { + manifest: dict[str, Any] = { "name": request.name, "source": request.source, } @@ -304,7 +304,7 @@ async def register_plugin( verbose_proxy_logger.exception(f"Error registering plugin: {e}") raise HTTPException( status_code=500, - detail={"error": f"Registration failed: {str(e)}"}, + detail={"error": f"Registration failed: {e!s}"}, ) diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 92279a9e685..4b566caf2b1 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -142,7 +142,7 @@ async def anthropic_response( _usage = _blocked_response_usage(e.original_response) _anthropic_response = AnthropicMessagesResponse( - id=f"msg_{str(uuid.uuid4())}", + id=f"msg_{uuid.uuid4()!s}", type="message", role="assistant", content=[{"type": "text", "text": e.message}], @@ -189,9 +189,7 @@ async def anthropic_response( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.anthropic_response(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.anthropic_response(): Exception occured - {e!s}") # Extract model_id from request metadata (same as success path) litellm_metadata = data.get("litellm_metadata", {}) or {} @@ -211,7 +209,7 @@ async def anthropic_response( litellm_logging_obj=None, ) - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -303,10 +301,8 @@ async def count_tokens( detail=detail, ) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.anthropic_endpoints.count_tokens(): Exception occurred - {}".format(str(e)) - ) - raise HTTPException(status_code=500, detail={"error": f"Internal server error: {str(e)}"}) + verbose_proxy_logger.exception(f"litellm.proxy.anthropic_endpoints.count_tokens(): Exception occurred - {e!s}") + raise HTTPException(status_code=500, detail={"error": f"Internal server error: {e!s}"}) @router.post( diff --git a/litellm/proxy/anthropic_endpoints/skills_endpoints.py b/litellm/proxy/anthropic_endpoints/skills_endpoints.py index 772006a7fa1..e651bbf2b9a 100644 --- a/litellm/proxy/anthropic_endpoints/skills_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/skills_endpoints.py @@ -2,8 +2,6 @@ Anthropic Skills API endpoints - /v1/skills """ -from typing import Optional - import orjson from fastapi import APIRouter, Depends, Request, Response @@ -32,7 +30,7 @@ router = APIRouter() async def create_skill( fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "anthropic", + custom_llm_provider: str | None = "anthropic", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -130,10 +128,10 @@ async def create_skill( async def list_skills( fastapi_response: Response, request: Request, - limit: Optional[int] = 10, - after_id: Optional[str] = None, - before_id: Optional[str] = None, - custom_llm_provider: Optional[str] = "anthropic", + limit: int | None = 10, + after_id: str | None = None, + before_id: str | None = None, + custom_llm_provider: str | None = "anthropic", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -235,7 +233,7 @@ async def get_skill( skill_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "anthropic", + custom_llm_provider: str | None = "anthropic", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -332,7 +330,7 @@ async def delete_skill( skill_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "anthropic", + custom_llm_provider: str | None = "anthropic", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 263fec77d12..f876b303510 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -13,7 +13,7 @@ import asyncio import math import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, Union, cast from fastapi import HTTPException, Request, status from pydantic import BaseModel @@ -59,18 +59,17 @@ from litellm.proxy._types import ( SpecialModelNames, UserAPIKeyAuth, ) -from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.budget_throttle import ( budget_throttle_percentage, should_throttle_budget_exceeded, ) -from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start +from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec -from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, ) +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, get_management_object_ttl, @@ -82,6 +81,7 @@ from litellm.proxy.guardrails.tool_name_extraction import ( extract_request_tool_names, ) from litellm.proxy.route_llm_request import route_request +from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.object_permission_repository import ObjectPermissionRepository @@ -141,7 +141,7 @@ def _log_budget_lookup_failure(entity: str, error: Exception) -> None: ) -def _get_router_zero_cost_cache(llm_router: Router) -> Optional[Dict[str, bool]]: +def _get_router_zero_cost_cache(llm_router: Router) -> dict[str, bool] | None: """ Return the router's per-instance zero-cost cache, or ``None`` for objects that don't expose one (e.g. ``MagicMock`` stand-ins in unit tests). @@ -157,7 +157,7 @@ def _get_router_zero_cost_cache(llm_router: Router) -> Optional[Dict[str, bool]] return cache if isinstance(cache, dict) else None -def _is_model_cost_zero(model: Optional[Union[str, List[str]]], llm_router: Optional[Router]) -> bool: +def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: """ Check if a model has zero cost (no configured pricing). @@ -246,7 +246,7 @@ def _is_model_cost_zero(model: Optional[Union[str, List[str]]], llm_router: Opti except Exception as e: # If we can't determine the cost, assume it has cost (conservative approach) - verbose_proxy_logger.debug(f"Error checking cost for model {model_name}: {str(e)}, assuming it has cost") + verbose_proxy_logger.debug(f"Error checking cost for model {model_name}: {e!s}, assuming it has cost") return False # All models checked have zero cost @@ -276,11 +276,11 @@ def _is_cost_explicitly_configured(model: str, llm_router: "Router") -> bool: async def _run_project_checks( - project_object: Optional[LiteLLM_ProjectTableCachedObj], - _model: Optional[Union[str, List[str]]], - llm_router: Optional[Router], + project_object: LiteLLM_ProjectTableCachedObj | None, + _model: str | list[str] | None, + llm_router: Router | None, skip_budget_checks: bool, - valid_token: Optional[UserAPIKeyAuth], + valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, ) -> None: """ @@ -353,7 +353,7 @@ def _reject_clientside_metadata_tags_check(general_settings: dict, request_body: ) -def _global_proxy_budget_check(global_proxy_spend: Optional[float], skip_budget_checks: bool, route: str) -> None: +def _global_proxy_budget_check(global_proxy_spend: float | None, skip_budget_checks: bool, route: str) -> None: if ( litellm.max_budget > 0 and not skip_budget_checks @@ -378,7 +378,7 @@ _GUARDRAIL_MODIFICATION_KEYS: tuple = ( ) -def _guardrail_modification_check(request_body: dict, team_object: Optional[LiteLLM_TeamTable]) -> None: +def _guardrail_modification_check(request_body: dict, team_object: LiteLLM_TeamTable | None) -> None: """ Reject user-supplied metadata flags that would modify guardrail behavior unless the team has explicit permission. Checked keys include the plural @@ -393,7 +393,7 @@ def _guardrail_modification_check(request_body: dict, team_object: Optional[Lite """ from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails - def _coerce_to_dict(container: Any) -> Optional[dict]: + def _coerce_to_dict(container: Any) -> dict | None: """Accept dict or JSON-string (from multipart/form-data or extra_body). Without this, an attacker can smuggle guardrail keys past the check by @@ -434,8 +434,8 @@ def _guardrail_modification_check(request_body: dict, team_object: Optional[Lite async def check_tools_allowlist( request_body: dict, - valid_token: Optional[UserAPIKeyAuth], - team_object: Optional[LiteLLM_TeamTable], + valid_token: UserAPIKeyAuth | None, + team_object: LiteLLM_TeamTable | None, route: str, ) -> None: """ @@ -497,18 +497,18 @@ BUDGET_ENFORCED_SIDE_EFFECT_ROUTES = frozenset( async def common_checks( request_body: dict, - team_object: Optional[LiteLLM_TeamTable], - user_object: Optional[LiteLLM_UserTable], - end_user_object: Optional[LiteLLM_EndUserTable], - global_proxy_spend: Optional[float], + team_object: LiteLLM_TeamTable | None, + user_object: LiteLLM_UserTable | None, + end_user_object: LiteLLM_EndUserTable | None, + global_proxy_spend: float | None, general_settings: dict, route: str, - llm_router: Optional[Router], + llm_router: Router | None, proxy_logging_obj: ProxyLogging, - valid_token: Optional[UserAPIKeyAuth], + valid_token: UserAPIKeyAuth | None, request: Request, skip_budget_checks: bool = False, - project_object: Optional[LiteLLM_ProjectTableCachedObj] = None, + project_object: LiteLLM_ProjectTableCachedObj | None = None, ) -> bool: """ Common checks across jwt + key-based auth. @@ -531,7 +531,7 @@ async def common_checks( """ from litellm.proxy.proxy_server import prisma_client, user_api_key_cache - _model: Optional[Union[str, List[str]]] = get_model_from_request( + _model: str | list[str] | None = get_model_from_request( request_data=request_body, route=route, request_headers=_safe_get_request_headers(request=request), @@ -729,7 +729,7 @@ async def common_checks( # 10 [OPTIONAL] Organization RBAC checks organization_role_based_access_check(user_object=user_object, route=route, request_body=request_body) - async def _fetch_team_org_id(team_id: str) -> Optional[str]: + async def _fetch_team_org_id(team_id: str) -> str | None: try: team = await get_team_object( team_id=team_id, @@ -777,8 +777,8 @@ async def common_checks( def _get_user_role( - user_obj: Optional[LiteLLM_UserTable], -) -> Optional[LitellmUserRoles]: + user_obj: LiteLLM_UserTable | None, +) -> LitellmUserRoles | None: if user_obj is None: return None @@ -797,8 +797,8 @@ def _is_api_route_allowed( route: str, request: Request, request_data: dict, - valid_token: Optional[UserAPIKeyAuth], - user_obj: Optional[LiteLLM_UserTable] = None, + valid_token: UserAPIKeyAuth | None, + user_obj: LiteLLM_UserTable | None = None, ) -> bool: """ - Route b/w api token check and normal token check @@ -820,7 +820,7 @@ def _is_api_route_allowed( return True -def _is_user_proxy_admin(user_obj: Optional[LiteLLM_UserTable]): +def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None): if user_obj is None: return False @@ -888,7 +888,7 @@ def allowed_routes_check( def allowed_route_check_inside_route( user_api_key_dict: UserAPIKeyAuth, - requested_user_id: Optional[str], + requested_user_id: str | None, ) -> bool: ret_val = True if ( @@ -918,10 +918,10 @@ def get_actual_routes(allowed_routes: list) -> list: async def get_default_end_user_budget( - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, -) -> Optional[LiteLLM_BudgetTable]: + parent_otel_span: Span | None = None, +) -> LiteLLM_BudgetTable | None: """ Fetches the default end user budget from the database if litellm.max_end_user_budget_id is configured. @@ -973,16 +973,16 @@ async def get_default_end_user_budget( return _budget_obj except Exception as e: - verbose_proxy_logger.error(f"Error fetching default end user budget: {str(e)}") + verbose_proxy_logger.error(f"Error fetching default end user budget: {e!s}") return None @log_db_metrics async def get_team_member_default_budget( budget_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, -) -> Optional[LiteLLM_BudgetTable]: +) -> LiteLLM_BudgetTable | None: """ Fetches the team-level default per-member budget referenced by team.metadata["team_member_budget_id"]. @@ -1033,7 +1033,7 @@ async def _apply_default_budget_to_end_user( end_user_obj: LiteLLM_EndUserTable, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> LiteLLM_EndUserTable: """ Helper function to apply default budget to end user if they don't have a budget assigned. @@ -1116,13 +1116,13 @@ async def _check_end_user_budget( @log_db_metrics async def get_end_user_object( - end_user_id: Optional[str], - prisma_client: Optional[PrismaClient], + end_user_id: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - route: Optional[str] = "", - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[LiteLLM_EndUserTable]: + route: str | None = "", + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> LiteLLM_EndUserTable | None: """ Returns end user object from database or cache. @@ -1146,7 +1146,7 @@ async def get_end_user_object( if end_user_id is None: return None - _key = "end_user_id:{}".format(end_user_id) + _key = f"end_user_id:{end_user_id}" # Check cache first cached_user_obj = await user_api_key_cache.async_get_cache( @@ -1188,7 +1188,7 @@ async def get_end_user_object( # Save to cache await user_api_key_cache.async_set_cache( - key="end_user_id:{}".format(end_user_id), + key=f"end_user_id:{end_user_id}", value=_response, model_type=LiteLLM_EndUserTable, ) @@ -1204,13 +1204,13 @@ _END_USER_VALIDATION_POSITIVE_TTL = 300 async def resolve_and_validate_end_user_id( - raw_end_user_id: Optional[str], - prisma_client: Optional[PrismaClient], + raw_end_user_id: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, route: str = "", -) -> Optional[str]: +) -> str | None: """Optionally drop end-user ids that don't resolve to a known DB row. Default: pass-through. LiteLLM's documented pattern is that the `user` @@ -1272,8 +1272,8 @@ async def _end_user_id_exists_in_db( end_user_id: str, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, route: str = "", ) -> bool: """True when the id matches an EndUser, User, or user_email row.""" @@ -1312,12 +1312,12 @@ async def _end_user_id_exists_in_db( @log_db_metrics async def get_tag_objects_batch( - tag_names: List[str], - prisma_client: Optional[PrismaClient], + tag_names: list[str], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Dict[str, LiteLLM_TagTable]: + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> dict[str, LiteLLM_TagTable]: """ Batch fetch multiple tag objects from cache and db. @@ -1383,12 +1383,12 @@ async def get_tag_objects_batch( @log_db_metrics async def get_tag_object( - tag_name: Optional[str], - prisma_client: Optional[PrismaClient], + tag_name: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[LiteLLM_TagTable]: + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> LiteLLM_TagTable | None: """ Returns tag object from cache or db. @@ -1423,10 +1423,10 @@ async def get_tag_object( async def get_team_membership( user_id: str, team_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, ) -> Optional["LiteLLM_TeamMembership"]: """ Returns team membership object if user is member of team. @@ -1441,7 +1441,7 @@ async def get_team_membership( if user_id is None or team_id is None: return None - _key = "team_membership:{}:{}".format(user_id, team_id) + _key = f"team_membership:{user_id}:{team_id}" # check if in cache cached_membership_obj = await user_api_key_cache.async_get_cache( @@ -1478,7 +1478,7 @@ async def get_team_membership( return None -def model_in_access_group(model: str, team_models: Optional[List[str]], llm_router: Optional[Router]) -> bool: +def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool: from collections import defaultdict if team_models is None: @@ -1512,9 +1512,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c """ current_time = time.time() # if key doesn't exist in last_db_access_time -> check db - if key not in last_db_access_time: - return True - elif last_db_access_time[key][0] is not None: # check db for non-null values (for refresh operations) + if key not in last_db_access_time or last_db_access_time[key][0] is not None: return True elif last_db_access_time[key][0] is None: if current_time - last_db_access_time[key][1] >= db_cache_expiry: @@ -1522,7 +1520,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c return False -def _update_last_db_access_time(key: str, value: Optional[Any], last_db_access_time: LimitedSizeOrderedDict): +def _update_last_db_access_time(key: str, value: Any | None, last_db_access_time: LimitedSizeOrderedDict): last_db_access_time[key] = (value, time.time()) @@ -1530,12 +1528,12 @@ def _get_role_based_permissions( rbac_role: RBAC_ROLES, general_settings: dict, key: Literal["models", "routes"], -) -> Optional[List[str]]: +) -> list[str] | None: """ Get the role based permissions from the general settings. """ role_based_permissions = cast( - Optional[List[RoleBasedPermissions]], + list[RoleBasedPermissions] | None, general_settings.get("role_permissions", []), ) if role_based_permissions is None: @@ -1551,7 +1549,7 @@ def _get_role_based_permissions( def get_role_based_models( rbac_role: RBAC_ROLES, general_settings: dict, -) -> Optional[List[str]]: +) -> list[str] | None: """ Get the models allowed for a user role. @@ -1568,7 +1566,7 @@ def get_role_based_models( def get_role_based_routes( rbac_role: RBAC_ROLES, general_settings: dict, -) -> Optional[List[str]]: +) -> list[str] | None: """ Get the routes allowed for a user role. """ @@ -1582,9 +1580,9 @@ def get_role_based_routes( async def _get_fuzzy_user_object( prisma_client: PrismaClient, - sso_user_id: Optional[str] = None, - user_email: Optional[str] = None, -) -> Optional[LiteLLM_UserTable]: + sso_user_id: str | None = None, + user_email: str | None = None, +) -> LiteLLM_UserTable | None: """ Checks if sso user is in db. @@ -1624,16 +1622,16 @@ async def _get_fuzzy_user_object( @log_db_metrics async def get_user_object( - user_id: Optional[str], - prisma_client: Optional[PrismaClient], + user_id: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, user_id_upsert: bool, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, - sso_user_id: Optional[str] = None, - user_email: Optional[str] = None, - check_db_only: Optional[bool] = None, -) -> Optional[LiteLLM_UserTable]: + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, + sso_user_id: str | None = None, + user_email: str | None = None, + check_db_only: bool | None = None, +) -> LiteLLM_UserTable | None: """ - Check if user id in proxy User Table - if valid, return LiteLLM_UserTable object with defined limits @@ -1655,7 +1653,7 @@ async def get_user_object( if prisma_client is None: raise Exception("No db connected") try: - db_access_time_key = "user_id:{}".format(user_id) + db_access_time_key = f"user_id:{user_id}" should_check_db = _should_check_db( key=db_access_time_key, last_db_access_time=last_db_access_time, @@ -1688,7 +1686,7 @@ async def get_user_object( scalar_default_params = { key: value for key, value in default_params.items() if key not in ("teams", "available_teams") } - new_user_params: Dict[str, Any] = { + new_user_params: dict[str, Any] = { "user_id": user_id, **({"user_email": user_email} if user_email is not None else {}), **scalar_default_params, @@ -1761,11 +1759,11 @@ async def get_user_object( async def _cache_management_object( key: str, - value: Union[BaseModel, Dict[str, Any]], + value: BaseModel | dict[str, Any], user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, *, - model_type: Type[BaseModel], + model_type: type[BaseModel], ): """ Persist management objects via ``UserApiKeyCache`` (in-memory + optional Redis). @@ -1784,12 +1782,12 @@ async def _cache_team_object( team_id: str, team_table: LiteLLM_TeamTableCachedObj, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, ): ## CACHE REFRESH TIME! team_table.last_refreshed_at = time.time() - key = "team_id:{}".format(team_id) + key = f"team_id:{team_id}" if proxy_logging_obj is not None: try: @@ -1825,7 +1823,7 @@ async def _cache_team_object( # uniqueness (len(teams) > 1 raises HTTPException) before populating # the cache from a verified single row. if team_table.team_alias: - alias_key = "team_alias:{}".format(team_table.team_alias) + alias_key = f"team_alias:{team_table.team_alias}" try: user_api_key_cache.delete_cache(key=alias_key) if proxy_logging_obj is not None: @@ -1843,7 +1841,7 @@ async def _cache_key_object( hashed_token: str, user_api_key_obj: UserAPIKeyAuth, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, ): key = hashed_token @@ -1863,7 +1861,7 @@ async def _cache_key_object( async def _delete_cache_key_object( hashed_token: str, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, ): key = hashed_token @@ -1875,7 +1873,7 @@ async def _delete_cache_key_object( @log_db_metrics -async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_upsert: Optional[bool] = None): +async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None): response = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if response is None and team_id_upsert: @@ -1905,9 +1903,9 @@ async def _get_team_object_from_user_api_key_cache( user_api_key_cache: UserApiKeyCache, last_db_access_time: LimitedSizeOrderedDict, db_cache_expiry: int, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, key: str, - team_id_upsert: Optional[bool] = None, + team_id_upsert: bool | None = None, ) -> LiteLLM_TeamTableCachedObj: db_access_time_key = key should_check_db = _should_check_db( @@ -1960,10 +1958,10 @@ async def _get_team_object_from_user_api_key_cache( async def _get_team_object_from_cache( key: str, - proxy_logging_obj: Optional[ProxyLogging], + proxy_logging_obj: ProxyLogging | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], -) -> Optional[LiteLLM_TeamTableCachedObj]: + parent_otel_span: Span | None, +) -> LiteLLM_TeamTableCachedObj | None: ## INTERNAL USAGE CACHE (plain DualCache) — checked before UserApiKeyCache stores ## if proxy_logging_obj is not None and proxy_logging_obj.internal_usage_cache.dual_cache: cached_raw = await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache( @@ -1984,13 +1982,13 @@ async def _get_team_object_from_cache( async def get_team_object( team_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, - check_cache_only: Optional[bool] = None, - check_db_only: Optional[bool] = None, - team_id_upsert: Optional[bool] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, + check_cache_only: bool | None = None, + check_db_only: bool | None = None, + team_id_upsert: bool | None = None, ) -> LiteLLM_TeamTableCachedObj: """ - Check if team id in proxy Team Table @@ -2004,7 +2002,7 @@ async def get_team_object( raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") # check if in cache - key = "team_id:{}".format(team_id) + key = f"team_id:{team_id}" if not check_db_only: cached_team_obj = await _get_team_object_from_cache( @@ -2046,9 +2044,9 @@ async def _cache_access_object( access_group_id: str, access_group_table: LiteLLM_AccessGroupTable, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging] = None, + proxy_logging_obj: ProxyLogging | None = None, ): - key = "access_group_id:{}".format(access_group_id) + key = f"access_group_id:{access_group_id}" await user_api_key_cache.async_set_cache( key=key, value=access_group_table, @@ -2060,9 +2058,9 @@ async def _cache_access_object( async def _delete_cache_access_object( access_group_id: str, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging] = None, + proxy_logging_obj: ProxyLogging | None = None, ): - key = "access_group_id:{}".format(access_group_id) + key = f"access_group_id:{access_group_id}" user_api_key_cache.delete_cache(key=key) @@ -2074,9 +2072,9 @@ async def _delete_cache_access_object( @log_db_metrics async def get_access_object( access_group_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging] = None, + proxy_logging_obj: ProxyLogging | None = None, ) -> LiteLLM_AccessGroupTable: """ - Check if access_group_id in proxy AccessGroupTable @@ -2093,7 +2091,7 @@ async def get_access_object( if prisma_client is None: raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") - key = "access_group_id:{}".format(access_group_id) + key = f"access_group_id:{access_group_id}" cached_access_obj = await user_api_key_cache.async_get_cache( key=key, @@ -2141,10 +2139,10 @@ async def get_access_object( @log_db_metrics async def get_team_object_by_alias( team_alias: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, + proxy_logging_obj: ProxyLogging | None = None, ) -> LiteLLM_TeamTableCachedObj: """ Look up a team by its team_alias (name) in the database. @@ -2166,7 +2164,7 @@ async def get_team_object_by_alias( raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") # Check cache first (keyed by alias) - cache_key = "team_alias:{}".format(team_alias) + cache_key = f"team_alias:{team_alias}" cached_team_obj = await _get_team_object_from_cache( key=cache_key, @@ -2224,7 +2222,7 @@ async def get_team_object_by_alias( ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by team_id for consistency - team_id_cache_key = "team_id:{}".format(team_obj.team_id) + team_id_cache_key = f"team_id:{team_obj.team_id}" await user_api_key_cache.async_set_cache( key=team_id_cache_key, value=team_obj, @@ -2240,18 +2238,18 @@ async def get_team_object_by_alias( verbose_proxy_logger.exception("Error looking up team by alias: %s", team_alias) raise HTTPException( status_code=500, - detail={"error": f"Error looking up team by alias '{team_alias}': {str(e)}"}, + detail={"error": f"Error looking up team by alias '{team_alias}': {e!s}"}, ) @log_db_metrics async def get_org_object_by_alias( org_alias: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[LiteLLM_OrganizationTable]: + proxy_logging_obj: ProxyLogging | None = None, +) -> LiteLLM_OrganizationTable | None: """ Look up an organization by its organization_alias in the database. @@ -2272,7 +2270,7 @@ async def get_org_object_by_alias( raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") # Check cache first (keyed by alias) - cache_key = "org_alias:{}".format(org_alias) + cache_key = f"org_alias:{org_alias}" cached_org_obj = await user_api_key_cache.async_get_cache( key=cache_key, model_type=LiteLLM_OrganizationTable, @@ -2312,7 +2310,7 @@ async def get_org_object_by_alias( ) # Also cache by org_id for consistency await user_api_key_cache.async_set_cache( - key="org_id:{}".format(org_obj.organization_id), + key=f"org_id:{org_obj.organization_id}", value=org_obj, model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, @@ -2326,7 +2324,7 @@ async def get_org_object_by_alias( verbose_proxy_logger.exception("Error looking up organization by alias: %s", org_alias) raise HTTPException( status_code=500, - detail={"error": f"Error looking up organization by alias '{org_alias}': {str(e)}"}, + detail={"error": f"Error looking up organization by alias '{org_alias}': {e!s}"}, ) @@ -2367,9 +2365,9 @@ class ExperimentalUIJWTToken: @staticmethod def get_cli_jwt_auth_token( user_info: LiteLLM_UserTable, - team_id: Optional[str] = None, - team_alias: Optional[str] = None, - max_budget: Optional[float] = None, + team_id: str | None = None, + team_alias: str | None = None, + max_budget: float | None = None, ) -> str: """ Generate a JWT token for CLI authentication with configurable expiration. @@ -2430,7 +2428,7 @@ class ExperimentalUIJWTToken: @staticmethod def get_key_object_from_ui_hash_key( hashed_token: str, - ) -> Optional[UserAPIKeyAuth]: + ) -> UserAPIKeyAuth | None: import json from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth @@ -2450,9 +2448,9 @@ class ExperimentalUIJWTToken: async def _fetch_key_object_from_db_with_reconnect( hashed_token: str, prisma_client: PrismaClient, - parent_otel_span: Optional[Span], - proxy_logging_obj: Optional[ProxyLogging], -) -> Optional[BaseModel]: + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, +) -> BaseModel | None: """ Fetch key object from DB and retry once if a DB connection error can be healed. """ @@ -2493,7 +2491,7 @@ async def get_jwt_key_mapping_object( jwt_claim_name: str, jwt_claim_value: str, prisma_client: PrismaClient, -) -> Optional[str]: +) -> str | None: """ Lookup a JWT-to-virtual-key mapping from the database. @@ -2514,11 +2512,11 @@ async def get_jwt_key_mapping_object( @log_db_metrics async def get_key_object( hashed_token: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, - check_cache_only: Optional[bool] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, + check_cache_only: bool | None = None, ) -> UserAPIKeyAuth: """ - Check if team id in proxy Team Table @@ -2544,7 +2542,7 @@ async def get_key_object( raise Exception(f"Key doesn't exist in cache + check_cache_only=True. key={key}.") # else, check db - _valid_token: Optional[BaseModel] = await _fetch_key_object_from_db_with_reconnect( + _valid_token: BaseModel | None = await _fetch_key_object_from_db_with_reconnect( hashed_token=hashed_token, prisma_client=prisma_client, parent_otel_span=parent_otel_span, @@ -2553,9 +2551,7 @@ async def get_key_object( if _valid_token is None: raise ProxyException( - message="Authentication Error, Invalid proxy server token passed. key={}, not found in db. Create key via `/key/generate` call.".format( - hashed_token - ), + message=f"Authentication Error, Invalid proxy server token passed. key={hashed_token}, not found in db. Create key via `/key/generate` call.", type=ProxyErrorTypes.token_not_found_in_db, param="key", code=status.HTTP_401_UNAUTHORIZED, @@ -2603,11 +2599,11 @@ def _copy_user_api_key_auth_for_cache( @log_db_metrics async def get_object_permission( object_permission_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[LiteLLM_ObjectPermissionTable]: + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> LiteLLM_ObjectPermissionTable | None: """ - Check if object permission id in proxy ObjectPermissionTable - if valid, return LiteLLM_ObjectPermissionTable object @@ -2649,12 +2645,12 @@ async def get_object_permission( @log_db_metrics async def get_managed_vector_store_rows_by_uuids( - uuids: List[str], - prisma_client: Optional[PrismaClient], + uuids: list[str], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> List[LiteLLM_ManagedVectorStoresTable]: + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> list[LiteLLM_ManagedVectorStoresTable]: """ Fetch managed vector store rows by their internal UUIDs. @@ -2666,11 +2662,11 @@ async def get_managed_vector_store_rows_by_uuids( if not uuids or prisma_client is None: return [] - result: List[LiteLLM_ManagedVectorStoresTable] = [] - cache_misses: List[str] = [] + result: list[LiteLLM_ManagedVectorStoresTable] = [] + cache_misses: list[str] = [] for uuid in uuids: - key = "managed_vector_store_id:{}".format(uuid) + key = f"managed_vector_store_id:{uuid}" deserialized_vs = await user_api_key_cache.async_get_cache( key=key, model_type=LiteLLM_ManagedVectorStoresTable, @@ -2695,7 +2691,7 @@ async def get_managed_vector_store_rows_by_uuids( if not row_dict: continue cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict) - key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) + key = f"managed_vector_store_id:{cached_obj.vector_store_id}" await user_api_key_cache.async_set_cache( key=key, value=cached_obj, @@ -2719,12 +2715,12 @@ class OrganizationNotFoundError(Exception): @log_db_metrics async def get_org_object( org_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, include_budget_table: bool = False, -) -> Optional[LiteLLM_OrganizationTable]: +) -> LiteLLM_OrganizationTable | None: """ - Check if org id in proxy Org Table - if valid, return LiteLLM_OrganizationTable object @@ -2744,9 +2740,9 @@ async def get_org_object( return None # Use different cache key if budget table is included - cache_key = "org_id:{}".format(org_id) + cache_key = f"org_id:{org_id}" if include_budget_table: - cache_key = "org_id:{}:with_budget".format(org_id) + cache_key = f"org_id:{org_id}:with_budget" # check if in cache deserialized_org = await user_api_key_cache.async_get_cache( @@ -2757,7 +2753,7 @@ async def get_org_object( return deserialized_org # else, check db try: - query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}} + query_kwargs: dict[str, Any] = {"where": {"organization_id": org_id}} if include_budget_table: query_kwargs["include"] = {"litellm_budget_table": True} @@ -2788,12 +2784,12 @@ async def get_org_object( async def _get_resources_from_access_groups( - access_group_ids: List[str], + access_group_ids: list[str], resource_field: Literal["access_model_names", "access_mcp_server_ids", "access_agent_ids"], - prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[UserApiKeyCache] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> List[str]: + prisma_client: PrismaClient | None = None, + user_api_key_cache: UserApiKeyCache | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> list[str]: """ Fetch access groups by their IDs (from cache or DB) and collect the specified resource field across all of them. @@ -2827,7 +2823,7 @@ async def _get_resources_from_access_groups( if user_api_key_cache is None: return [] - resources: List[str] = [] + resources: list[str] = [] for ag_id in access_group_ids: try: ag = await get_access_object( @@ -2847,11 +2843,11 @@ async def _get_resources_from_access_groups( async def _get_models_from_access_groups( - access_group_ids: List[str], - prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[UserApiKeyCache] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> List[str]: + access_group_ids: list[str], + prisma_client: PrismaClient | None = None, + user_api_key_cache: UserApiKeyCache | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> list[str]: """ Collect model names from unified access groups. Models are matched by model name for backwards compatibility. @@ -2866,11 +2862,11 @@ async def _get_models_from_access_groups( async def _get_mcp_server_ids_from_access_groups( - access_group_ids: List[str], - prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[UserApiKeyCache] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> List[str]: + access_group_ids: list[str], + prisma_client: PrismaClient | None = None, + user_api_key_cache: UserApiKeyCache | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> list[str]: """ Collect MCP server IDs from unified access groups. MCPs are matched by server ID. @@ -2885,11 +2881,11 @@ async def _get_mcp_server_ids_from_access_groups( async def _get_agent_ids_from_access_groups( - access_group_ids: List[str], - prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[UserApiKeyCache] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> List[str]: + access_group_ids: list[str], + prisma_client: PrismaClient | None = None, + user_api_key_cache: UserApiKeyCache | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> list[str]: """ Collect agent IDs from unified access groups. Agents are matched by agent ID. @@ -2904,10 +2900,10 @@ async def _get_agent_ids_from_access_groups( def _resolve_all_team_model_sentinel_for_auth_check( - models: List[str], - llm_router: Optional[Router], - team_id: Optional[str], -) -> List[str]: + models: list[str], + llm_router: Router | None, + team_id: str | None, +) -> list[str]: if SpecialModelNames.all_team_models.value not in models or team_id is None or llm_router is None: return models proxy_models = llm_router.get_model_names() @@ -2919,15 +2915,15 @@ def _resolve_all_team_model_sentinel_for_auth_check( def _check_model_access_helper( model: str, - llm_router: Optional[Router], - models: List[str], - team_model_aliases: Optional[Dict[str, str]] = None, - team_id: Optional[str] = None, + llm_router: Router | None, + models: list[str], + team_model_aliases: dict[str, str] | None = None, + team_id: str | None = None, ) -> bool: ## check if model in allowed model names from collections import defaultdict - access_groups: Dict[str, List[str]] = defaultdict(list) + access_groups: dict[str, list[str]] = defaultdict(list) if llm_router: access_groups = llm_router.get_model_access_groups(model_name=model, team_id=team_id) @@ -2966,11 +2962,11 @@ def _check_model_access_helper( def _can_object_call_model( - model: Union[str, List[str]], - llm_router: Optional[Router], - models: List[str], - team_model_aliases: Optional[Dict[str, str]] = None, - team_id: Optional[str] = None, + model: str | list[str], + llm_router: Router | None, + models: list[str], + team_model_aliases: dict[str, str] | None = None, + team_id: str | None = None, object_type: Literal["user", "team", "key", "org", "project"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -2991,7 +2987,7 @@ def _can_object_call_model( - Exception: If token not allowed to call model """ if fallback_depth >= DEFAULT_MAX_RECURSE_DEPTH: - raise Exception("Unable to parse model, max fallback depth exceeded - received model: {}".format(model)) + raise Exception(f"Unable to parse model, max fallback depth exceeded - received model: {model}") if isinstance(model, list): for m in model: _can_object_call_model( @@ -3032,7 +3028,7 @@ def _can_object_call_model( ) -def _model_in_team_aliases(model: str, team_model_aliases: Optional[Dict[str, str]] = None) -> bool: +def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool: """ Returns True if `model` being accessed is an alias of a team model @@ -3050,7 +3046,7 @@ def _model_in_team_aliases(model: str, team_model_aliases: Optional[Dict[str, st return False -def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> List[str]: +def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]: """ Expand key model sentinels before auth checks. @@ -3069,10 +3065,10 @@ def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> List[str] async def can_key_call_model( - model: Union[str, List[str]], - llm_model_list: Optional[list], + model: str | list[str], + llm_model_list: list | None, valid_token: UserAPIKeyAuth, - llm_router: Optional[litellm.Router], + llm_router: litellm.Router | None, ) -> Literal[True]: """ Checks if token can call a given model @@ -3117,9 +3113,9 @@ async def can_key_call_model( async def can_key_call_resolved_model( model: str, - llm_model_list: Optional[list], + llm_model_list: list | None, valid_token: UserAPIKeyAuth, - llm_router: Optional[litellm.Router], + llm_router: litellm.Router | None, ) -> None: from litellm.proxy.proxy_server import ( prisma_client, @@ -3138,7 +3134,7 @@ async def can_key_call_resolved_model( llm_router=llm_router, ) - team_object: Optional[LiteLLM_TeamTableCachedObj] = None + team_object: LiteLLM_TeamTableCachedObj | None = None team_object_from_lookup = False if valid_token.team_id is not None: try: @@ -3208,9 +3204,9 @@ async def can_key_call_resolved_model( def can_org_access_model( model: str, - org_object: Optional[LiteLLM_OrganizationTable], - llm_router: Optional[Router], - team_model_aliases: Optional[Dict[str, str]] = None, + org_object: LiteLLM_OrganizationTable | None, + llm_router: Router | None, + team_model_aliases: dict[str, str] | None = None, ) -> Literal[True]: """ Returns True if the team can access a specific model. @@ -3226,10 +3222,10 @@ def can_org_access_model( async def can_team_access_model( - model: Union[str, List[str]], - team_object: Optional[LiteLLM_TeamTable], - llm_router: Optional[Router], - team_model_aliases: Optional[Dict[str, str]] = None, + model: str | list[str], + team_object: LiteLLM_TeamTable | None, + llm_router: Router | None, + team_model_aliases: dict[str, str] | None = None, ) -> Literal[True]: """ Returns True if the team can access a specific model. @@ -3266,10 +3262,10 @@ async def can_team_access_model( async def get_authorized_resources_from_key_access_groups( - valid_token: Optional[UserAPIKeyAuth], - team_object: Optional[LiteLLM_TeamTable], + valid_token: UserAPIKeyAuth | None, + team_object: LiteLLM_TeamTable | None, resource_field: Literal["access_model_names", "access_mcp_server_ids", "access_agent_ids"], -) -> List[str]: +) -> list[str]: """ For each access_group_id on the key, fetch the LiteLLM_AccessGroupTable row and contribute its `resource_field` only if the group authorizes the caller @@ -3294,7 +3290,7 @@ async def get_authorized_resources_from_key_access_groups( key_team_id = valid_token.team_id or (team_object.team_id if team_object is not None else None) key_token = valid_token.token - authorized_resources: List[str] = [] + authorized_resources: list[str] = [] for ag_id in key_access_group_ids: try: ag = await get_access_object( @@ -3314,10 +3310,10 @@ async def get_authorized_resources_from_key_access_groups( async def _key_access_group_grants_model( - model: Union[str, List[str]], - valid_token: Optional[UserAPIKeyAuth], - team_object: Optional[LiteLLM_TeamTable], - llm_router: Optional[Router], + model: str | list[str], + valid_token: UserAPIKeyAuth | None, + team_object: LiteLLM_TeamTable | None, + llm_router: Router | None, ) -> bool: """ Returns True if the key's `access_group_ids` expand to models that grant @@ -3346,9 +3342,9 @@ async def _key_access_group_grants_model( def can_project_access_model( - model: Union[str, List[str]], + model: str | list[str], project_object: LiteLLM_ProjectTableCachedObj, - llm_router: Optional[Router], + llm_router: Router | None, ) -> Literal[True]: """ Returns True if the project can access a specific model. @@ -3364,9 +3360,9 @@ def can_project_access_model( async def can_user_call_model( - model: Union[str, List[str]], - llm_router: Optional[Router], - user_object: Optional[LiteLLM_UserTable], + model: str | list[str], + llm_router: Router | None, + user_object: LiteLLM_UserTable | None, ) -> Literal[True]: if user_object is None: return True @@ -3388,8 +3384,8 @@ async def can_user_call_model( def _search_tool_names_from_object_permission( - object_permission: Optional[LiteLLM_ObjectPermissionTable], -) -> List[str]: + object_permission: LiteLLM_ObjectPermissionTable | None, +) -> list[str]: """Return allowlisted search tool names from object_permission (empty = unrestricted).""" if object_permission is None: return [] @@ -3401,7 +3397,7 @@ def _search_tool_names_from_object_permission( def _can_object_call_search_tool( search_tool_name: str, - allowed_search_tools: List[str], + allowed_search_tools: list[str], object_type: Literal["key", "team", "project"], ) -> Literal[True]: """ @@ -3466,7 +3462,7 @@ async def can_key_call_search_tool( async def can_team_call_search_tool( search_tool_name: str, - team_object: Optional[LiteLLM_TeamTable], + team_object: LiteLLM_TeamTable | None, ) -> Literal[True]: """ Check if a team can access a specific search tool. @@ -3496,7 +3492,7 @@ async def can_team_call_search_tool( async def can_user_view_search_tool( search_tool_name: str, valid_token: UserAPIKeyAuth, - team_object: Optional[LiteLLM_TeamTable], + team_object: LiteLLM_TeamTable | None, ) -> bool: """ Boolean variant of the key + team authorization enforced on /search, used to @@ -3518,8 +3514,8 @@ async def can_user_view_search_tool( async def is_valid_fallback_model( model: str, - llm_router: Optional[Router], - user_model: Optional[str], + llm_router: Router | None, + user_model: str | None, ) -> Literal[True]: """ Try to route the fallback model request. @@ -3562,7 +3558,7 @@ def _apply_budget_exceeded_throttle(valid_token: UserAPIKeyAuth) -> bool: async def _virtual_key_max_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, - user_obj: Optional[LiteLLM_UserTable] = None, + user_obj: LiteLLM_UserTable | None = None, ): """ Raises: @@ -3680,7 +3676,7 @@ async def _virtual_key_multi_budget_check( async def _virtual_key_soft_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, - user_obj: Optional[LiteLLM_UserTable] = None, + user_obj: LiteLLM_UserTable | None = None, ): """ Triggers a budget alert if the token is over it's soft budget. @@ -3716,7 +3712,7 @@ async def _virtual_key_soft_budget_check( ) -def _parse_email_list(raw: Any) -> List[str]: +def _parse_email_list(raw: Any) -> list[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): return [e.strip() for e in raw if isinstance(e, str) and e.strip()] @@ -3726,8 +3722,8 @@ def _parse_email_list(raw: Any) -> List[str]: def _normalize_alert_emails( - cfg: Optional[Dict[str, Any]], -) -> Dict[str, List[str]]: + cfg: dict[str, Any] | None, +) -> dict[str, list[str]]: """Coerce user-supplied threshold→recipients mapping to Dict[str, List[str]]. Values may legitimately arrive as list, comma-separated string, or None @@ -3739,9 +3735,9 @@ def _normalize_alert_emails( def _merge_budget_alert_email_configs( - global_cfg: Optional[Dict[str, Any]], - per_key_cfg: Optional[Dict[str, Any]], -) -> Optional[Dict[str, List[str]]]: + global_cfg: dict[str, Any] | None, + per_key_cfg: dict[str, Any] | None, +) -> dict[str, list[str]] | None: """ Per-threshold additive merge: each threshold's recipient list is the union of global + per-key entries (deduped, global-first ordering). Missing @@ -3760,7 +3756,7 @@ def _merge_budget_alert_email_configs( async def _virtual_key_max_budget_alert_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, - user_obj: Optional[LiteLLM_UserTable] = None, + user_obj: LiteLLM_UserTable | None = None, ): """ Triggers a budget alert if the token has reached EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE @@ -3771,7 +3767,7 @@ async def _virtual_key_max_budget_alert_check( if valid_token.max_budget is not None and valid_token.spend is not None and valid_token.spend > 0: owner_email = user_obj.user_email if user_obj else None - alert_email_config: Optional[Dict[str, List[str]]] = _merge_budget_alert_email_configs( + alert_email_config: dict[str, list[str]] | None = _merge_budget_alert_email_configs( global_cfg=litellm.default_key_max_budget_alert_emails, per_key_cfg=(valid_token.metadata or {}).get("max_budget_alert_emails"), ) @@ -3840,10 +3836,10 @@ async def _virtual_key_max_budget_alert_check( async def _check_team_member_budget( - team_object: Optional[LiteLLM_TeamTable], - user_object: Optional[LiteLLM_UserTable], - valid_token: Optional[UserAPIKeyAuth], - prisma_client: Optional[PrismaClient], + team_object: LiteLLM_TeamTable | None, + user_object: LiteLLM_UserTable | None, + valid_token: UserAPIKeyAuth | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): @@ -3864,7 +3860,7 @@ async def _check_team_member_budget( # Per-member override wins; otherwise fall back to the team-level # default configured via team.metadata["team_member_budget_id"]. - team_member_budget: Optional[float] = None + team_member_budget: float | None = None if ( team_membership is not None and team_membership.litellm_budget_table is not None @@ -3911,10 +3907,10 @@ async def _check_team_member_budget( async def _check_team_member_model_access( - model: Union[str, List[str]], + model: str | list[str], team_object: LiteLLM_TeamTable, valid_token: UserAPIKeyAuth, - llm_router: Optional[Router], + llm_router: Router | None, prisma_client: Optional["PrismaClient"], user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, @@ -3943,7 +3939,7 @@ async def _check_team_member_model_access( ): return # no per-member restriction — inherit team-level check - member_allowed_models: List[str] = team_membership.litellm_budget_table.allowed_models + member_allowed_models: list[str] = team_membership.litellm_budget_table.allowed_models try: _can_object_call_model( model=model, @@ -3962,8 +3958,8 @@ async def _check_team_member_model_access( async def _team_max_budget_check( - team_object: Optional[LiteLLM_TeamTable], - valid_token: Optional[UserAPIKeyAuth], + team_object: LiteLLM_TeamTable | None, + valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, ): """ @@ -4012,7 +4008,7 @@ async def _team_max_budget_check( async def _team_multi_budget_check( - team_object: Optional[LiteLLM_TeamTable], + team_object: LiteLLM_TeamTable | None, ): """ Raises BudgetExceededError if any budget window in team_object.budget_limits is exceeded. @@ -4051,8 +4047,8 @@ async def _team_multi_budget_check( async def _team_soft_budget_check( - team_object: Optional[LiteLLM_TeamTable], - valid_token: Optional[UserAPIKeyAuth], + team_object: LiteLLM_TeamTable | None, + valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, ): """ @@ -4072,7 +4068,7 @@ async def _team_soft_budget_check( ) if valid_token: # Extract alert emails from team metadata - alert_emails: Optional[List[str]] = None + alert_emails: list[str] | None = None if team_object.metadata is not None and isinstance(team_object.metadata, dict): soft_budget_alert_emails = team_object.metadata.get("soft_budget_alerting_emails") if soft_budget_alert_emails is not None: @@ -4122,8 +4118,8 @@ async def _team_soft_budget_check( async def _project_max_budget_check( - project_object: Optional[LiteLLM_ProjectTableCachedObj], - valid_token: Optional[UserAPIKeyAuth], + project_object: LiteLLM_ProjectTableCachedObj | None, + valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, ): """ @@ -4174,8 +4170,8 @@ async def _project_max_budget_check( async def _project_soft_budget_check( - project_object: Optional[LiteLLM_ProjectTableCachedObj], - valid_token: Optional[UserAPIKeyAuth], + project_object: LiteLLM_ProjectTableCachedObj | None, + valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, ): """ @@ -4219,10 +4215,10 @@ async def _project_soft_budget_check( async def get_project_object( project_id: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[LiteLLM_ProjectTableCachedObj]: + proxy_logging_obj: ProxyLogging | None = None, +) -> LiteLLM_ProjectTableCachedObj | None: """ Fetch project object from cache or DB. @@ -4234,7 +4230,7 @@ async def get_project_object( return None # Check cache first - cache_key = "project_id:{}".format(project_id) + cache_key = f"project_id:{project_id}" deserialized_project = await user_api_key_cache.async_get_cache( key=cache_key, model_type=LiteLLM_ProjectTableCachedObj, @@ -4266,9 +4262,9 @@ async def get_project_object( async def _organization_max_budget_check( - valid_token: Optional[UserAPIKeyAuth], - team_object: Optional[LiteLLM_TeamTable], - prisma_client: Optional[PrismaClient], + valid_token: UserAPIKeyAuth | None, + team_object: LiteLLM_TeamTable | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): @@ -4290,7 +4286,7 @@ async def _organization_max_budget_check( return # Determine organization_id: first try from token, then fallback to team - org_id: Optional[str] = None + org_id: str | None = None if valid_token.org_id is not None: org_id = valid_token.org_id elif team_object is not None and team_object.organization_id is not None: @@ -4317,7 +4313,7 @@ async def _organization_max_budget_check( return # Get max_budget from organization's budget table - org_max_budget: Optional[float] = None + org_max_budget: float | None = None if org_table.litellm_budget_table is not None: org_max_budget = org_table.litellm_budget_table.max_budget @@ -4365,10 +4361,10 @@ async def _organization_max_budget_check( async def _tag_max_budget_check( request_body: dict, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, - valid_token: Optional[UserAPIKeyAuth], + valid_token: UserAPIKeyAuth | None, ): """ Check if any tags in the request are over their max budget. @@ -4511,8 +4507,8 @@ def _is_wildcard_pattern(allowed_model_pattern: str) -> bool: async def vector_store_access_check( request_body: dict, - team_object: Optional[LiteLLM_TeamTable], - valid_token: Optional[UserAPIKeyAuth], + team_object: LiteLLM_TeamTable | None, + valid_token: UserAPIKeyAuth | None, ): """ Checks if the object (key, team, org) has access to the vector store. @@ -4570,8 +4566,8 @@ async def vector_store_access_check( def _can_object_call_vector_stores( object_type: Literal["key", "team", "org"], - vector_store_ids_to_run: List[str], - object_permissions: Optional[LiteLLM_ObjectPermissionTable], + vector_store_ids_to_run: list[str], + object_permissions: LiteLLM_ObjectPermissionTable | None, ): """ Raises ProxyException if the object (key, team, org) cannot access the specific vector store. diff --git a/litellm/proxy/auth/auth_checks_organization.py b/litellm/proxy/auth/auth_checks_organization.py index dc8f2ad758b..472edb54fa1 100644 --- a/litellm/proxy/auth/auth_checks_organization.py +++ b/litellm/proxy/auth/auth_checks_organization.py @@ -3,7 +3,6 @@ Auth Checks for Organizations """ from collections.abc import Awaitable, Callable -from typing import Dict, List, Optional, Tuple from fastapi import status @@ -12,7 +11,7 @@ from litellm.proxy._types import * def organization_role_based_access_check( request_body: dict, - user_object: Optional[LiteLLM_UserTable], + user_object: LiteLLM_UserTable | None, route: str, ): """ @@ -29,7 +28,7 @@ def organization_role_based_access_check( if user_object is None: return - passed_organization_id: Optional[str] = request_body.get("organization_id", None) + passed_organization_id: str | None = request_body.get("organization_id", None) if route == "/organization/new": if user_object.user_role != LitellmUserRoles.PROXY_ADMIN.value: @@ -66,7 +65,7 @@ def organization_role_based_access_check( code=status.HTTP_401_UNAUTHORIZED, ) - user_role: Optional[LitellmUserRoles] = _user_organization_role_mapping.get(passed_organization_id) + user_role: LitellmUserRoles | None = _user_organization_role_mapping.get(passed_organization_id) if user_role is None: raise ProxyException( message=f"You do not have a role within the selected organization. Passed organization_id: {passed_organization_id}. Please contact the organization admin to request access.", @@ -109,7 +108,7 @@ def organization_role_based_access_check( def get_user_organization_info( user_object: LiteLLM_UserTable, -) -> Tuple[List[str], Dict[str, Optional[LitellmUserRoles]]]: +) -> tuple[list[str], dict[str, LitellmUserRoles | None]]: """ Helper function to extract user organization information. @@ -121,8 +120,8 @@ def get_user_organization_info( - List of organization IDs the user is a member of - Dictionary mapping organization IDs to user roles """ - _user_organizations: List[str] = [] - _user_organization_role_mapping: Dict[str, Optional[LitellmUserRoles]] = {} + _user_organizations: list[str] = [] + _user_organization_role_mapping: dict[str, LitellmUserRoles | None] = {} if user_object.organization_memberships is not None: for _membership in user_object.organization_memberships: @@ -135,7 +134,7 @@ def get_user_organization_info( def _user_is_org_admin( request_data: dict, - user_object: Optional[LiteLLM_UserTable] = None, + user_object: LiteLLM_UserTable | None = None, ) -> bool: """ Helper function to check if user is an org admin for all of the passed organizations. @@ -151,7 +150,7 @@ def _user_is_org_admin( return False # Collect candidate org IDs from both fields - candidate_org_ids: List[str] = [] + candidate_org_ids: list[str] = [] singular = request_data.get("organization_id", None) if singular is not None: candidate_org_ids.append(singular) @@ -183,8 +182,8 @@ PATCH_TEAM_ROUTE_TEMPLATE = "/team/{team_id}" async def add_team_org_context_to_request_body( route: str, request_body: dict, - fetch_team_org_id: Callable[[str], Awaitable[Optional[str]]], - route_template: Optional[str] = None, + fetch_team_org_id: Callable[[str], Awaitable[str | None]], + route_template: str | None = None, ) -> dict: """ Return a copy of request_body with organization_id resolved from the target @@ -202,7 +201,7 @@ async def add_team_org_context_to_request_body( return request_body if route in TEAM_ORG_CONTEXT_ROUTES: - team_id: Optional[str] = request_body.get("team_id") + team_id: str | None = request_body.get("team_id") elif route_template == PATCH_TEAM_ROUTE_TEMPLATE: team_id = route.rsplit("/", 1)[-1] else: diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index e19bb3f9f5f..e96d3db65ff 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -2,19 +2,19 @@ Handles Authentication Errors """ -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from fastapi import HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger +from litellm.integrations.otel.runtime import seed_request_identity from litellm.proxy._types import ( LitellmUserRoles, ProxyErrorTypes, ProxyException, UserAPIKeyAuth, ) -from litellm.integrations.otel.runtime import seed_request_identity from litellm.proxy.auth.auth_utils import _get_request_ip_address from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.types.services import ServiceTypes @@ -40,9 +40,9 @@ class UserAPIKeyAuthExceptionHandler: request: Request, request_data: dict, route: str, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, api_key: str, - resolved_identity: Optional[UserAPIKeyAuth] = None, + resolved_identity: UserAPIKeyAuth | None = None, ) -> UserAPIKeyAuth: """ Handles Connection Errors when reading a Virtual Key from LiteLLM DB @@ -95,10 +95,7 @@ class UserAPIKeyAuthExceptionHandler: use_x_forwarded_for=general_settings.get("use_x_forwarded_for", False), ) verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - {}\nRequester IP Address:{}".format( - str(e), - requester_ip, - ), + f"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - {e!s}\nRequester IP Address:{requester_ip}", extra={"requester_ip": requester_ip}, ) @@ -153,7 +150,7 @@ class UserAPIKeyAuthExceptionHandler: ) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_401_UNAUTHORIZED), diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 88b128f7c02..dfa5b22d285 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -4,7 +4,7 @@ import sys from collections.abc import Iterator, Mapping from functools import lru_cache from logging import Logger -from typing import Any, Dict, FrozenSet, List, Optional, Tuple, Union +from typing import Any from fastapi import HTTPException, Request, status @@ -27,7 +27,7 @@ from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS from litellm.types.utils import CustomPricingLiteLLMParams -def _get_request_ip_address(request: Request, use_x_forwarded_for: Optional[bool] = False) -> Optional[str]: +def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None: client_ip = None if use_x_forwarded_for is True and "x-forwarded-for" in request.headers: client_ip = request.headers["x-forwarded-for"] @@ -40,10 +40,10 @@ def _get_request_ip_address(request: Request, use_x_forwarded_for: Optional[bool def _check_valid_ip( - allowed_ips: Optional[List[str]], + allowed_ips: list[str] | None, request: Request, - use_x_forwarded_for: Optional[bool] = False, -) -> Tuple[bool, Optional[str]]: + use_x_forwarded_for: bool | None = False, +) -> tuple[bool, str | None]: """ Returns if ip is allowed or not """ @@ -70,7 +70,7 @@ def check_complete_credentials(request_body: dict) -> bool: be used as an SSRF pivot. Validate any URL fields here so the gate can't be bypassed with ``api_key=anything`` plus a malicious target. """ - given_model: Optional[str] = None + given_model: str | None = None given_model = request_body.get("model") if given_model is None: @@ -131,7 +131,7 @@ def _is_param_allowed( for item in configurable_clientside_auth_params: if isinstance(item, str) and param == item: return True - elif isinstance(item, Dict): + elif isinstance(item, dict): if param == "api_base" and check_regex_or_str_match( request_body_value=request_body_value, regex_str=item["api_base"], @@ -142,7 +142,7 @@ def _is_param_allowed( def _allow_model_level_clientside_configurable_parameters( - model: str, param: str, request_body_value: Any, llm_router: Optional[Router] + model: str, param: str, request_body_value: Any, llm_router: Router | None ) -> bool: """ Check if model is allowed to use configurable client-side params @@ -180,7 +180,7 @@ def _allow_model_level_clientside_configurable_parameters( # ``extra_body.aws_web_identity_token``) without re-validating, so the # banned-key check has to descend into it the same way it descends into # ``litellm_embedding_config``. -_NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config", "extra_body") +_NESTED_CONFIG_KEYS: tuple[str, ...] = ("litellm_embedding_config", "extra_body") # Metadata containers that carry per-request configuration consumed by the # observability callbacks. The same banned-param list applies — a value @@ -188,7 +188,7 @@ _NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config", "extra_body" # leaks the same credentials as the root-level ``langfuse_host``, but the # original check only walked the request-body root, so the metadata path # was an unintentional bypass. -_NESTED_METADATA_KEYS: Tuple[str, ...] = ("metadata", "litellm_metadata") +_NESTED_METADATA_KEYS: tuple[str, ...] = ("metadata", "litellm_metadata") # Banned request-body params. The same list applies to every entry in # ``_NESTED_CONFIG_KEYS`` (dicts spread as ``**kwargs`` into outbound @@ -200,7 +200,7 @@ _NESTED_METADATA_KEYS: Tuple[str, ...] = ("metadata", "litellm_metadata") # without choosing the destination or the credentials, so they don't # contribute to the data-exfil primitive that the rest of # ``_supported_callback_params`` does. -_SAFE_CLIENT_CALLBACK_PARAMS: FrozenSet[str] = frozenset( +_SAFE_CLIENT_CALLBACK_PARAMS: frozenset[str] = frozenset( { "langfuse_prompt_version", "langsmith_sampling_rate", @@ -212,7 +212,7 @@ _SAFE_CLIENT_CALLBACK_PARAMS: FrozenSet[str] = frozenset( # Listed here so the proxy bans them today; the long-term cleanup is to # fold these into the canonical allowlist so they share one source of # truth with the rest. -_EXTRA_BANNED_OBSERVABILITY_PARAMS: FrozenSet[str] = frozenset( +_EXTRA_BANNED_OBSERVABILITY_PARAMS: frozenset[str] = frozenset( { "posthog_api_url", "phoenix_project_name", @@ -228,7 +228,7 @@ _EXTRA_BANNED_OBSERVABILITY_PARAMS: FrozenSet[str] = frozenset( ) -def _build_banned_observability_params() -> FrozenSet[str]: +def _build_banned_observability_params() -> frozenset[str]: """Derive the observability ban list from the canonical allowlist. ``_supported_callback_params`` and ``_request_blocked_callback_params`` in @@ -253,7 +253,7 @@ def _build_banned_observability_params() -> FrozenSet[str]: ) -_BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = ( +_BANNED_REQUEST_BODY_PARAMS: tuple[str, ...] = ( "api_base", "base_url", "user_config", @@ -309,7 +309,7 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = ( def _check_banned_params( body: dict, general_settings: dict, - llm_router: Optional[Router], + llm_router: Router | None, model: str, ) -> None: """Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in. @@ -403,7 +403,7 @@ def _reject_url_valued_fallback_target(value: str) -> None: ) -def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str) -> bool: +def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Router | None, model: str) -> bool: """ Check if the request body is safe. @@ -461,7 +461,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: return True -def _coerce_metadata_to_dict(value: Any) -> Optional[Dict[str, Any]]: +def _coerce_metadata_to_dict(value: Any) -> dict[str, Any] | None: """Return ``value`` as a dict, parsing it from JSON if delivered as a string. Multipart/form-data and ``extra_body`` callers send ``litellm_metadata`` @@ -582,7 +582,7 @@ def route_in_additonal_public_routes(current_route: str): return False except Exception as e: - verbose_proxy_logger.error(f"route_in_additonal_public_routes: {str(e)}") + verbose_proxy_logger.error(f"route_in_additonal_public_routes: {e!s}") return False @@ -619,12 +619,12 @@ def get_request_route(request: Request) -> str: return raw_path except Exception as e: verbose_proxy_logger.debug( - f"error on get_request_route: {str(e)}, defaulting to request.url.path={request.url.path}" + f"error on get_request_route: {e!s}, defaulting to request.url.path={request.url.path}" ) return str(request.url.path) -def get_request_route_template(request: Request) -> Optional[str]: +def get_request_route_template(request: Request) -> str | None: """ Return the low-cardinality route template, e.g. ``/v1/threads/{thread_id}/runs`` (vs. the literal path from @@ -639,7 +639,7 @@ def get_request_route_template(request: Request) -> Optional[str]: template = getattr(route, "path", None) return template if isinstance(template, str) and template else None except Exception as e: - verbose_proxy_logger.debug(f"error on get_request_route_template: {str(e)}") + verbose_proxy_logger.debug(f"error on get_request_route_template: {e!s}") return None @@ -866,7 +866,7 @@ def bytes_to_mb(bytes_value: int): # helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key -def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]: +def _get_deployment_default_limit(model_name: str, field: str) -> int | None: """ Return the minimum value of `field` across all deployments for model_name, or None if no deployment has the field set. @@ -894,18 +894,18 @@ def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]: return min(limits) if limits else None -def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: +def _get_deployment_default_rpm_limit(model_name: str) -> int | None: return _get_deployment_default_limit(model_name, "default_api_key_rpm_limit") -def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]: +def _get_deployment_default_tpm_limit(model_name: str) -> int | None: return _get_deployment_default_limit(model_name, "default_api_key_tpm_limit") def get_key_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, - model_name: Optional[str] = None, -) -> Optional[Dict[str, int]]: + model_name: str | None = None, +) -> dict[str, int] | None: """ Get the model rpm limit for a given api key. @@ -923,7 +923,7 @@ def get_key_model_rpm_limit( # 2. Check model_max_budget if user_api_key_dict.model_max_budget: - model_rpm_limit: Dict[str, Any] = {} + model_rpm_limit: dict[str, Any] = {} for model, budget in user_api_key_dict.model_max_budget.items(): if isinstance(budget, dict) and budget.get("rpm_limit") is not None: model_rpm_limit[model] = budget["rpm_limit"] @@ -947,8 +947,8 @@ def get_key_model_rpm_limit( def get_key_model_tpm_limit( user_api_key_dict: UserAPIKeyAuth, - model_name: Optional[str] = None, -) -> Optional[Dict[str, int]]: + model_name: str | None = None, +) -> dict[str, int] | None: """ Get the model tpm limit for a given api key. @@ -966,7 +966,7 @@ def get_key_model_tpm_limit( # 2. Check model_max_budget (iterate per-model like RPM does) if user_api_key_dict.model_max_budget: - model_tpm_limit: Dict[str, Any] = {} + model_tpm_limit: dict[str, Any] = {} for model, budget in user_api_key_dict.model_max_budget.items(): if isinstance(budget, dict) and budget.get("tpm_limit") is not None: model_tpm_limit[model] = budget["tpm_limit"] @@ -992,7 +992,7 @@ def get_model_rate_limit_from_metadata( user_api_key_dict: UserAPIKeyAuth, metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"], rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"], -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if getattr(user_api_key_dict, metadata_accessor_key): return getattr(user_api_key_dict, metadata_accessor_key).get(rate_limit_key) return None @@ -1000,7 +1000,7 @@ def get_model_rate_limit_from_metadata( def get_team_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if user_api_key_dict.team_metadata: return user_api_key_dict.team_metadata.get("model_rpm_limit") return None @@ -1008,7 +1008,7 @@ def get_team_model_rpm_limit( def get_team_model_tpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if user_api_key_dict.team_metadata: return user_api_key_dict.team_metadata.get("model_tpm_limit") return None @@ -1016,7 +1016,7 @@ def get_team_model_tpm_limit( def get_key_mcp_rpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: """ Get the per-MCP-server rpm limit for a given api key. @@ -1042,7 +1042,7 @@ def get_key_mcp_rpm_limit( def get_team_mcp_rpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if user_api_key_dict.team_metadata: return user_api_key_dict.team_metadata.get("mcp_rpm_limit") return None @@ -1050,7 +1050,7 @@ def get_team_mcp_rpm_limit( def get_key_tag_rpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[dict[str, int]]: +) -> dict[str, int] | None: """ Get the per-request-tag rpm limit configured on a given api key. @@ -1064,7 +1064,7 @@ def get_key_tag_rpm_limit( def get_project_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if user_api_key_dict.project_metadata: return user_api_key_dict.project_metadata.get("model_rpm_limit") return None @@ -1072,7 +1072,7 @@ def get_project_model_rpm_limit( def get_project_model_tpm_limit( user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Dict[str, int]]: +) -> dict[str, int] | None: if user_api_key_dict.project_metadata: return user_api_key_dict.project_metadata.get("model_tpm_limit") return None @@ -1144,7 +1144,7 @@ def _has_user_setup_sso(): return sso_setup -def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[list]: +def get_customer_user_header_from_mapping(user_id_mapping) -> list | None: """Return the header_name mapped to CUSTOMER role, if any (dict-based).""" if not user_id_mapping: return None @@ -1167,8 +1167,8 @@ def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[list]: def _get_customer_id_from_standard_headers( - request_headers: Optional[dict], -) -> Optional[str]: + request_headers: dict | None, +) -> str | None: """ Check standard customer ID headers for a customer/end-user ID. @@ -1193,7 +1193,7 @@ def _get_customer_id_from_standard_headers( return None -def _coerce_user_id_to_str(value: Any) -> Optional[str]: +def _coerce_user_id_to_str(value: Any) -> str | None: """Return a usable end-user identifier string, or None if the value isn't one. Always drops non-string structured values (dict/list/tuple/set) because @@ -1228,7 +1228,7 @@ def _coerce_user_id_to_str(value: Any) -> Optional[str]: return None -def get_end_user_id_from_request_body(request_body: dict, request_headers: Optional[dict] = None) -> Optional[str]: +def get_end_user_id_from_request_body(request_body: dict, request_headers: dict | None = None) -> str | None: # Import general_settings here to avoid potential circular import issues at module level # and to ensure it's fetched at runtime. from litellm.proxy.proxy_server import general_settings @@ -1242,7 +1242,7 @@ def get_end_user_id_from_request_body(request_body: dict, request_headers: Optio # User query: "system not respecting user_header_name property" # This implies the key in general_settings is 'user_header_name'. if request_headers is not None: - custom_header_name_to_check: Optional[Union[list, str]] = None + custom_header_name_to_check: list | str | None = None # Prefer user mappings (new behavior) user_id_mapping = general_settings.get("user_header_mappings", None) @@ -1362,7 +1362,7 @@ _MODEL_ROUTING_ID_FIELDS = ( ) -def _append_model_candidates(candidates: List[str], value: Any) -> None: +def _append_model_candidates(candidates: list[str], value: Any) -> None: if value is None: return @@ -1377,15 +1377,15 @@ def _append_model_candidates(candidates: List[str], value: Any) -> None: candidates.extend(model for model in model_names if model) -def _dedupe_model_candidates(candidates: List[str]) -> List[str]: - deduped: List[str] = [] +def _dedupe_model_candidates(candidates: list[str]) -> list[str]: + deduped: list[str] = [] for model in candidates: if model not in deduped: deduped.append(model) return deduped -def _get_case_insensitive_mapping_value(mapping: Optional[Mapping[str, Any]], key: str) -> Any: +def _get_case_insensitive_mapping_value(mapping: Mapping[str, Any] | None, key: str) -> Any: if not mapping: return None if key in mapping: @@ -1397,7 +1397,7 @@ def _get_case_insensitive_mapping_value(mapping: Optional[Mapping[str, Any]], ke return None -def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool: +def _route_matches_any_marker(route: str, markers: tuple[str, ...]) -> bool: normalized_route = route.lower() return any(marker in normalized_route for marker in markers) @@ -1408,13 +1408,13 @@ def _route_uses_model_routing_sources(route: str) -> bool: def _extract_models_from_managed_resource_id( resource_id: Any, - resource_id_field: Optional[str] = None, - llm_router: Optional[Router] = None, -) -> List[str]: + resource_id_field: str | None = None, + llm_router: Router | None = None, +) -> list[str]: if not isinstance(resource_id, str) or not resource_id: return [] - candidates: List[str] = [] + candidates: list[str] = [] try: from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -1476,7 +1476,7 @@ def _extract_models_from_managed_resource_id( return _dedupe_model_candidates(candidates) -def _resolve_model_id_with_router(model_id: Optional[str], llm_router: Optional[Router]) -> Optional[str]: +def _resolve_model_id_with_router(model_id: str | None, llm_router: Router | None) -> str | None: if model_id is None or llm_router is None: return model_id try: @@ -1489,11 +1489,11 @@ def _resolve_model_id_with_router(model_id: Optional[str], llm_router: Optional[ def _extract_model_candidates_from_request( request_data: dict, route: str, - request_headers: Optional[Mapping[str, Any]] = None, - request_query_params: Optional[Mapping[str, Any]] = None, - llm_router: Optional[Router] = None, -) -> List[str]: - candidates: List[str] = [] + request_headers: Mapping[str, Any] | None = None, + request_query_params: Mapping[str, Any] | None = None, + llm_router: Router | None = None, +) -> list[str]: + candidates: list[str] = [] uses_model_routing_sources = _route_uses_model_routing_sources(route=route) uses_header_or_query_model_sources = _route_matches_any_marker( route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS @@ -1549,8 +1549,8 @@ def _extract_model_candidates_from_request( def _format_model_candidates( - candidates: List[str], -) -> Optional[Union[str, List[str]]]: + candidates: list[str], +) -> str | list[str] | None: if not candidates: return None if len(candidates) == 1: @@ -1582,11 +1582,11 @@ def _request_dispatched_to_pass_through_endpoint(request: Request | None) -> boo def get_model_from_request( request_data: dict, route: str, - request_headers: Optional[Mapping[str, Any]] = None, - request_query_params: Optional[Mapping[str, Any]] = None, - llm_router: Optional[Router] = None, + request_headers: Mapping[str, Any] | None = None, + request_query_params: Mapping[str, Any] | None = None, + llm_router: Router | None = None, request: Request | None = None, -) -> Optional[Union[str, List[str]]]: +) -> str | list[str] | None: """Resolve the model(s) a request targets, for model-access and budget checks. Returns ``None`` when the request was dispatched to a user-defined pass-through diff --git a/litellm/proxy/auth/budget_throttle.py b/litellm/proxy/auth/budget_throttle.py index 19dffee462b..f4e6b78b679 100644 --- a/litellm/proxy/auth/budget_throttle.py +++ b/litellm/proxy/auth/budget_throttle.py @@ -10,13 +10,12 @@ so it never compounds across requests. """ import math -from typing import Optional import litellm from litellm.proxy._types import UserAPIKeyAuth -def budget_throttle_percentage() -> Optional[float]: +def budget_throttle_percentage() -> float | None: """ The global throttle percentage, or None when throttling is disabled / misconfigured (in which case an over-budget key is hard-blocked, the safe @@ -45,7 +44,7 @@ def should_throttle_budget_exceeded(valid_token: UserAPIKeyAuth) -> bool: return budget_throttle_percentage() is not None -def throttled_limit(limit: Optional[int], pct: Optional[float]) -> Optional[int]: +def throttled_limit(limit: int | None, pct: float | None) -> int | None: """ Scale a TPM/RPM limit to ``pct`` of its value, keeping a trickle of at least 1 so a throttled key is slowed rather than fully locked out. An unset limit diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index ff87d0e70da..b1ccdc87830 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -12,7 +12,7 @@ import fnmatch import hashlib import os import re -from typing import Any, List, Literal, NoReturn, Optional, Set, Tuple, Union, cast +from typing import Any, Literal, NoReturn, cast import jwt from cryptography import x509 @@ -82,7 +82,7 @@ class JWTHandler: - if role="litellm_proxy_user" -> allow making calls + info. Can not edit budgets """ - prisma_client: Optional[PrismaClient] + prisma_client: PrismaClient | None user_api_key_cache: UserApiKeyCache # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret @@ -124,7 +124,7 @@ class JWTHandler: def update_environment( self, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, litellm_jwtauth: LiteLLM_JWTAuth, leeway: int = 0, @@ -135,14 +135,14 @@ class JWTHandler: self.leeway = leeway @staticmethod - def is_jwt(token: Optional[str]) -> bool: + def is_jwt(token: str | None) -> bool: if token is None: return False parts = token.split(".") return len(parts) == 3 @staticmethod - def get_unverified_claims(token: str) -> Optional[dict]: + def get_unverified_claims(token: str) -> dict | None: """ Decode JWT claims without signature verification. Used for routing decisions before selecting validation path. @@ -163,7 +163,7 @@ class JWTHandler: verbose_proxy_logger.debug("Failed to decode unverified JWT claims for routing: %s", e) return None - def _rbac_role_from_role_mapping(self, token: dict) -> Optional[RBAC_ROLES]: + def _rbac_role_from_role_mapping(self, token: dict) -> RBAC_ROLES | None: """ Returns the RBAC role the token 'belongs' to based on role mappings. @@ -194,7 +194,7 @@ class JWTHandler: return None - def get_rbac_role(self, token: dict) -> Optional[RBAC_ROLES]: + def get_rbac_role(self, token: dict) -> RBAC_ROLES | None: """ Returns the RBAC role the token 'belongs' to. @@ -219,9 +219,11 @@ class JWTHandler: return LitellmUserRoles.PROXY_ADMIN elif self.get_team_id(token=token, default_value=None) is not None: return LitellmUserRoles.TEAM - elif self.get_user_id(token=token, default_value=None) is not None: - return LitellmUserRoles.INTERNAL_USER - elif user_roles is not None and self.is_allowed_user_role(user_roles=user_roles): + elif ( + self.get_user_id(token=token, default_value=None) is not None + or user_roles is not None + and self.is_allowed_user_role(user_roles=user_roles) + ): return LitellmUserRoles.INTERNAL_USER elif rbac_role := self._rbac_role_from_role_mapping(token=token): return rbac_role @@ -245,7 +247,7 @@ class JWTHandler: def _has_trusted_issuer_normalized_claim(self, token: dict, claim: str) -> bool: return self._is_trusted_issuer_normalized_token(token=token) and claim in token - def get_team_ids_from_jwt(self, token: dict) -> List[str]: + def get_team_ids_from_jwt(self, token: dict) -> list[str]: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_IDS_CLAIM): issuer_team_ids = token.get(self.LITELLM_TEAM_IDS_CLAIM) if isinstance(issuer_team_ids, list): @@ -259,7 +261,7 @@ class JWTHandler: return [] if self.litellm_jwtauth.team_ids_jwt_field is not None: - team_ids: Optional[List[str]] = get_nested_value( + team_ids: list[str] | None = get_nested_value( data=token, key_path=self.litellm_jwtauth.team_ids_jwt_field, default=[], @@ -268,7 +270,7 @@ class JWTHandler: return [] - def get_all_jwt_team_ids(self, token: dict) -> List[str]: + def get_all_jwt_team_ids(self, token: dict) -> list[str]: """ Return team IDs from both the plural ``team_ids_jwt_field`` and the singular ``team_id_jwt_field`` claim (string or list of strings), as a @@ -285,7 +287,7 @@ class JWTHandler: request-bound team, not of the token's claims. Callers that want the default-team behavior should still go through ``get_team_id``. """ - team_ids: List[str] = list(self.get_team_ids_from_jwt(token)) + team_ids: list[str] = list(self.get_team_ids_from_jwt(token)) singular: Any = None if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_ID_CLAIM): singular = token.get(self.LITELLM_TEAM_ID_CLAIM) @@ -307,7 +309,7 @@ class JWTHandler: team_ids.append(str(singular)) return team_ids - def get_end_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_end_user_id(self, token: dict, default_value: str | None) -> str | None: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_END_USER_ID_CLAIM): return token.get(self.LITELLM_END_USER_ID_CLAIM) @@ -348,7 +350,7 @@ class JWTHandler: return True return False - def get_team_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_team_id(self, token: dict, default_value: str | None) -> str | None: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_ID_CLAIM): team_id = token.get(self.LITELLM_TEAM_ID_CLAIM) if isinstance(team_id, list): @@ -390,7 +392,7 @@ class JWTHandler: team_id = default_value return team_id - def get_team_alias(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_team_alias(self, token: dict, default_value: str | None) -> str | None: """ Extract team name/alias from JWT token using the configured team_alias_jwt_field. @@ -415,7 +417,7 @@ class JWTHandler: team_alias = default_value return team_alias - def is_upsert_user_id(self, valid_user_email: Optional[bool] = None) -> bool: + def is_upsert_user_id(self, valid_user_email: bool | None = None) -> bool: """ Returns: - True: if 'user_id_upsert' is set AND valid_user_email is not False @@ -425,7 +427,7 @@ class JWTHandler: return False return self.litellm_jwtauth.user_id_upsert - def get_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_user_id(self, token: dict, default_value: str | None) -> str | None: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_USER_ID_CLAIM): return token.get(self.LITELLM_USER_ID_CLAIM) @@ -442,7 +444,7 @@ class JWTHandler: user_id = default_value return user_id - def get_user_roles(self, token: dict, default_value: Optional[List[str]]) -> Optional[List[str]]: + def get_user_roles(self, token: dict, default_value: list[str] | None) -> list[str] | None: """ Returns the user role from the token. @@ -461,7 +463,7 @@ class JWTHandler: user_roles = default_value return user_roles - def map_jwt_role_to_litellm_role(self, token: dict) -> Optional[LitellmUserRoles]: + def map_jwt_role_to_litellm_role(self, token: dict) -> LitellmUserRoles | None: """Map roles from JWT to LiteLLM user roles""" if not self.litellm_jwtauth.jwt_litellm_role_map: return None @@ -476,7 +478,7 @@ class JWTHandler: return mapping.litellm_role return None - def get_jwt_role(self, token: dict, default_value: Optional[List[str]]) -> Optional[List[str]]: + def get_jwt_role(self, token: dict, default_value: list[str] | None) -> list[str] | None: """ Generic implementation of `get_user_roles` that can be used for both user and team roles. @@ -497,7 +499,7 @@ class JWTHandler: user_roles = default_value return user_roles - def is_allowed_user_role(self, user_roles: Optional[List[str]]) -> bool: + def is_allowed_user_role(self, user_roles: list[str] | None) -> bool: """ Returns the user role from the token. @@ -511,7 +513,7 @@ class JWTHandler: return True return False - def get_user_email(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_user_email(self, token: dict, default_value: str | None) -> str | None: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_USER_EMAIL_CLAIM): return token.get(self.LITELLM_USER_EMAIL_CLAIM) @@ -528,7 +530,7 @@ class JWTHandler: user_email = default_value return user_email - def get_object_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_object_id(self, token: dict, default_value: str | None) -> str | None: try: if self.litellm_jwtauth.object_id_jwt_field is not None: object_id = get_nested_value( @@ -542,7 +544,7 @@ class JWTHandler: object_id = default_value return object_id - def get_org_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_org_id(self, token: dict, default_value: str | None) -> str | None: if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_ORG_ID_CLAIM): return token.get(self.LITELLM_ORG_ID_CLAIM) @@ -559,7 +561,7 @@ class JWTHandler: org_id = default_value return org_id - def get_org_alias(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_org_alias(self, token: dict, default_value: str | None) -> str | None: """ Extract organization name/alias from JWT token using the configured org_alias_jwt_field. @@ -584,7 +586,7 @@ class JWTHandler: org_alias = default_value return org_alias - def get_scopes(self, token: dict) -> List[str]: + def get_scopes(self, token: dict) -> list[str]: try: if isinstance(token["scope"], str): # Assuming the scopes are stored in 'scope' claim and are space-separated @@ -641,7 +643,7 @@ class JWTHandler: return 600 return litellm_jwtauth.public_key_ttl - async def _get_public_key_from_jwks_url(self, jwks_url: str, kid: Optional[str]) -> dict: + async def _get_public_key_from_jwks_url(self, jwks_url: str, kid: str | None) -> dict: resolved_jwks_url = await self._resolve_jwks_url(jwks_url) cache_key = f"litellm_jwt_auth_keys_{resolved_jwks_url}" @@ -675,7 +677,7 @@ class JWTHandler: raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={resolved_jwks_url}, kid={kid}") - async def get_public_key(self, kid: Optional[str]) -> dict: + async def get_public_key(self, kid: str | None) -> dict: keys_url = os.getenv("JWT_PUBLIC_KEY_URL") if keys_url is None: @@ -691,8 +693,8 @@ class JWTHandler: raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={keys_url_list}, kid={kid}") - def parse_keys(self, keys: JWKKeyValue, kid: Optional[str]) -> Optional[JWTKeyItem]: - public_key: Optional[JWTKeyItem] = None + def parse_keys(self, keys: JWKKeyValue, kid: str | None) -> JWTKeyItem | None: + public_key: JWTKeyItem | None = None if len(keys) == 1: if isinstance(keys, dict) and (keys.get("kid", None) == kid or kid is None): public_key = keys @@ -775,8 +777,8 @@ class JWTHandler: return userinfo except Exception as e: - verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {str(e)}") - raise Exception(f"Failed to fetch OIDC UserInfo: {str(e)}") + verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {e!s}") + raise Exception(f"Failed to fetch OIDC UserInfo: {e!s}") _unscoped_jwt_warning_emitted = False @@ -819,7 +821,7 @@ class JWTHandler: "options": options or None, } - def _get_configured_issuer(self, token: str) -> Optional[JWTIssuerConfig]: + def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None: litellm_jwtauth = getattr(self, "litellm_jwtauth", None) if litellm_jwtauth is None: return None @@ -899,10 +901,10 @@ class JWTHandler: def _get_decode_options( self, - audience: Optional[Union[str, List[str]]], - issuer: Optional[str] = None, + audience: str | list[str] | None, + issuer: str | None = None, disable_audience_validation: bool = False, - ) -> Optional[dict]: + ) -> dict | None: # Disabling audience verification must be an explicit choice — never # an implicit consequence of ``audience`` being None. Otherwise a # caller that accidentally constructs a config with ``audience=None`` @@ -921,10 +923,10 @@ class JWTHandler: def _decode_jwt_with_public_key( self, token: str, - public_key: Union[dict, str], - audience: Optional[Union[str, List[str]]], - issuer: Optional[str] = None, - options: Optional[dict] = None, + public_key: dict | str, + audience: str | list[str] | None, + issuer: str | None = None, + options: dict | None = None, disable_audience_validation: bool = False, ) -> dict: decode_options = ( @@ -964,7 +966,7 @@ class JWTHandler: leeway=self.leeway, ) - async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: Optional[str]) -> dict: + async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict: public_key = await self._get_public_key_from_jwks_url( jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config), kid=kid, @@ -985,7 +987,7 @@ class JWTHandler: code=status.HTTP_401_UNAUTHORIZED, ) except Exception as e: - raise Exception(f"Validation fails: {str(e)}") + raise Exception(f"Validation fails: {e!s}") return self._apply_issuer_claim_mappings( token=payload, @@ -1030,7 +1032,7 @@ class JWTHandler: code=status.HTTP_401_UNAUTHORIZED, ) except Exception as e: - raise Exception(f"Validation fails: {str(e)}") + raise Exception(f"Validation fails: {e!s}") raise Exception("Invalid JWT Submitted") @@ -1072,7 +1074,7 @@ class JWTAuthManager: def can_rbac_role_call_model( rbac_role: RBAC_ROLES, general_settings: dict, - model: Optional[str], + model: str | None, ) -> Literal[True]: """ Checks if user is allowed to access the model, based on their role. @@ -1091,8 +1093,8 @@ class JWTAuthManager: @staticmethod def check_scope_based_access( - scope_mappings: List[ScopeMapping], - scopes: List[str], + scope_mappings: list[ScopeMapping], + scopes: list[str], request_data: dict, general_settings: dict, ) -> None: @@ -1100,7 +1102,7 @@ class JWTAuthManager: Check if scope allows access to the requested model """ if not scope_mappings: - return None + return allowed_models = [] for sm in scope_mappings: @@ -1110,14 +1112,14 @@ class JWTAuthManager: requested_model = request_data.get("model") if not requested_model: - return None + return if requested_model not in allowed_models: raise HTTPException( status_code=403, - detail={"error": "model={} not allowed. Allowed_models={}".format(requested_model, allowed_models)}, + detail={"error": f"model={requested_model} not allowed. Allowed_models={allowed_models}"}, ) - return None + return @staticmethod async def check_rbac_role( @@ -1126,7 +1128,7 @@ class JWTAuthManager: general_settings: dict, request_data: dict, route: str, - rbac_role: Optional[RBAC_ROLES], + rbac_role: RBAC_ROLES | None, ) -> None: """Validate RBAC role and model access permissions""" if jwt_handler.litellm_jwtauth.enforce_rbac is True: @@ -1151,12 +1153,12 @@ class JWTAuthManager: jwt_handler: JWTHandler, scopes: list, route: str, - user_id: Optional[str], - org_id: Optional[str], + user_id: str | None, + org_id: str | None, api_key: str, - jwt_valid_token: Optional[dict] = None, + jwt_valid_token: dict | None = None, user_email: str | None = None, - ) -> Optional[JWTAuthBuilderResult]: + ) -> JWTAuthBuilderResult | None: """Check admin status and route access permissions""" if not jwt_handler.is_admin(scopes=scopes): return None @@ -1167,7 +1169,7 @@ class JWTAuthManager: litellm_proxy_roles=jwt_handler.litellm_jwtauth, ) if not is_allowed: - allowed_routes: List[Any] = jwt_handler.litellm_jwtauth.admin_allowed_routes + allowed_routes: list[Any] = jwt_handler.litellm_jwtauth.admin_allowed_routes actual_routes = get_actual_routes(allowed_routes=allowed_routes) raise Exception(f"Admin not allowed to access this route. Route={route}, Allowed Routes={actual_routes}") @@ -1191,11 +1193,11 @@ class JWTAuthManager: async def find_and_validate_specific_team_id( jwt_handler: JWTHandler, jwt_valid_token: dict, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, - ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: + ) -> tuple[str | None, LiteLLM_TeamTable | None]: """Find and validate specific team ID from team_id_jwt_field or team_alias_jwt_field""" individual_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None) team_alias = jwt_handler.get_team_alias(token=jwt_valid_token, default_value=None) @@ -1212,7 +1214,7 @@ class JWTAuthManager: ): individual_team_id = None - team_object: Optional[LiteLLM_TeamTable] = None + team_object: LiteLLM_TeamTable | None = None if individual_team_id: try: @@ -1285,7 +1287,7 @@ class JWTAuthManager: return individual_team_id, team_object @staticmethod - def get_all_team_ids(jwt_handler: JWTHandler, jwt_valid_token: dict) -> Set[str]: + def get_all_team_ids(jwt_handler: JWTHandler, jwt_valid_token: dict) -> set[str]: """Get combined team IDs from groups and individual team_id""" team_ids_from_groups = jwt_handler.get_team_ids_from_jwt(token=jwt_valid_token) @@ -1295,9 +1297,9 @@ class JWTAuthManager: @staticmethod def _team_has_passthrough_route_access( - team_object: Optional[LiteLLM_TeamTable], + team_object: LiteLLM_TeamTable | None, route: str, - request_method: Optional[str] = None, + request_method: str | None = None, ) -> bool: normalized_request_method = request_method.upper() if isinstance(request_method, str) else None if not RouteChecks.is_auth_enforced_pass_through_route( @@ -1325,16 +1327,16 @@ class JWTAuthManager: @staticmethod async def find_team_with_model_access( - team_ids: Set[str], - requested_model: Optional[str], + team_ids: set[str], + requested_model: str | None, route: str, jwt_handler: JWTHandler, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, - request_method: Optional[str] = None, - ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: + request_method: str | None = None, + ) -> tuple[str | None, LiteLLM_TeamTable | None]: """Find first team with access to the requested model""" from litellm.proxy.proxy_server import llm_router @@ -1413,7 +1415,7 @@ class JWTAuthManager: async def get_user_info( jwt_handler: JWTHandler, jwt_valid_token: dict, - ) -> Tuple[Optional[str], Optional[str], Optional[bool]]: + ) -> tuple[str | None, str | None, bool | None]: """Get user email and validation status""" user_email = jwt_handler.get_user_email(token=jwt_valid_token, default_value=None) valid_user_email = None @@ -1424,9 +1426,9 @@ class JWTAuthManager: @staticmethod def _canonical_user_id_from_db( - user_id: Optional[str], - user_object: Optional[LiteLLM_UserTable], - ) -> Optional[str]: + user_id: str | None, + user_object: LiteLLM_UserTable | None, + ) -> str | None: """Id used for spend / team-membership attribution. JWT claim (often email) is only a lookup key. If fuzzy match in @@ -1439,25 +1441,25 @@ class JWTAuthManager: @staticmethod async def get_objects( - user_id: Optional[str], - user_email: Optional[str], - org_id: Optional[str], - end_user_id: Optional[str], - team_id: Optional[str], - valid_user_email: Optional[bool], + user_id: str | None, + user_email: str | None, + org_id: str | None, + end_user_id: str | None, + team_id: str | None, + valid_user_email: bool | None, jwt_handler: JWTHandler, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, route: str, - org_alias: Optional[str] = None, - ) -> Tuple[ - Optional[LiteLLM_UserTable], - Optional[LiteLLM_OrganizationTable], - Optional[LiteLLM_EndUserTable], - Optional[LiteLLM_TeamMembership], - Optional[str], + org_alias: str | None = None, + ) -> tuple[ + LiteLLM_UserTable | None, + LiteLLM_OrganizationTable | None, + LiteLLM_EndUserTable | None, + LiteLLM_TeamMembership | None, + str | None, ]: """Get user, org, end-user, and team-membership objects. @@ -1466,7 +1468,7 @@ class JWTAuthManager: """ # Get org object - first try by ID, then by alias - org_object: Optional[LiteLLM_OrganizationTable] = None + org_object: LiteLLM_OrganizationTable | None = None if org_id: org_object = ( await get_org_object( @@ -1502,7 +1504,7 @@ class JWTAuthManager: code=403, ) - user_object: Optional[LiteLLM_UserTable] = None + user_object: LiteLLM_UserTable | None = None if user_id: user_object = ( await get_user_object( @@ -1519,7 +1521,7 @@ class JWTAuthManager: else None ) - end_user_object: Optional[LiteLLM_EndUserTable] = None + end_user_object: LiteLLM_EndUserTable | None = None if end_user_id: end_user_object = ( await get_end_user_object( @@ -1544,7 +1546,7 @@ class JWTAuthManager: ) user_id = effective_user_id - team_membership_object: Optional[LiteLLM_TeamMembership] = None + team_membership_object: LiteLLM_TeamMembership | None = None if user_id and team_id: team_membership_object = ( await get_team_membership( @@ -1569,8 +1571,8 @@ class JWTAuthManager: @staticmethod def validate_object_id( - user_id: Optional[str], - team_id: Optional[str], + user_id: str | None, + team_id: str | None, enforce_rbac: bool, is_proxy_admin: bool, ) -> Literal[True]: @@ -1584,10 +1586,10 @@ class JWTAuthManager: @staticmethod def get_team_id_from_header( - request_headers: Optional[dict], - allowed_team_ids: Set[str], + request_headers: dict | None, + allowed_team_ids: set[str], fallback_to_db_teams: bool = False, - ) -> Optional[str]: + ) -> str | None: """ Extract team_id from x-litellm-team-id header if present. Validates that the team is in the user's allowed teams from JWT. @@ -1628,8 +1630,8 @@ class JWTAuthManager: @staticmethod async def map_user_to_teams( - user_object: Optional[LiteLLM_UserTable], - team_object: Optional[LiteLLM_TeamTable], + user_object: LiteLLM_UserTable | None, + team_object: LiteLLM_TeamTable | None, ): """ Map user to teams. @@ -1639,15 +1641,15 @@ class JWTAuthManager: from litellm.proxy.management_endpoints.team_endpoints import team_member_add if not user_object: - return None + return if not team_object: - return None + return # check if user is in team for member in team_object.members_with_roles: if member.user_id and member.user_id == user_object.user_id: - return None + return data = TeamMemberAddRequest( member=Member( @@ -1670,18 +1672,18 @@ class JWTAuthManager: verbose_proxy_logger.debug( f"User {user_object.user_id} is already a member of team {team_object.team_id}" ) - return None + return else: raise e - return None + return @staticmethod async def sync_user_role_and_teams( jwt_handler: JWTHandler, jwt_valid_token: dict, - user_object: Optional[LiteLLM_UserTable], - prisma_client: Optional[PrismaClient], - user_api_key_cache: Optional[UserApiKeyCache] = None, + user_object: LiteLLM_UserTable | None, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache | None = None, ) -> None: """ Sync user role and team memberships with JWT claims @@ -1693,10 +1695,10 @@ class JWTAuthManager: This method is only called if sync_user_role_and_teams is set to True in the JWT config. """ if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams: - return None + return if user_object is None or prisma_client is None: - return None + return # Update user role new_role = jwt_handler.map_jwt_role_to_litellm_role(jwt_valid_token) @@ -1746,17 +1748,17 @@ class JWTAuthManager: model_type=LiteLLM_UserTable, ttl=get_management_object_ttl(user_api_key_cache), ) - return None + return @staticmethod async def _attach_team_from_header_for_admin( admin_result: JWTAuthBuilderResult, route: str, - request_headers: Optional[dict], + request_headers: dict | None, jwt_handler: JWTHandler, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, ) -> None: """Attach team context from x-litellm-team-id to an admin result. @@ -1794,13 +1796,13 @@ class JWTAuthManager: @staticmethod async def _resolve_single_team_fallback( - user_object: Optional[LiteLLM_UserTable], - user_id: Optional[str], - prisma_client: Optional[PrismaClient], + user_object: LiteLLM_UserTable | None, + user_id: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, - team_id_upsert: Optional[bool], + team_id_upsert: bool | None, ) -> tuple: """ If JWT did not resolve team_id, but the user belongs to exactly one team @@ -2003,12 +2005,12 @@ class JWTAuthManager: request_data: dict, general_settings: dict, route: str, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, - request_headers: Optional[dict] = None, - request_method: Optional[str] = None, + request_headers: dict | None = None, + request_method: str | None = None, ) -> JWTAuthBuilderResult: """Main authentication and authorization builder""" # Check if OIDC UserInfo endpoint is enabled, but fall back to standard @@ -2058,8 +2060,8 @@ class JWTAuthManager: # Get IDs org_id = jwt_handler.get_org_id(token=jwt_valid_token, default_value=None) end_user_id = jwt_handler.get_end_user_id(token=jwt_valid_token, default_value=None) - team_id: Optional[str] = None - team_object: Optional[LiteLLM_TeamTable] = None + team_id: str | None = None + team_object: LiteLLM_TeamTable | None = None object_id = jwt_handler.get_object_id(token=jwt_valid_token, default_value=None) if rbac_role and object_id: diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index bd1cd596f9c..cf5074aa6ec 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -7,7 +7,7 @@ External callers (public IPs) only see servers with available_on_public_internet import ipaddress from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Union +from typing import Any, Union from fastapi import Request from pydantic import TypeAdapter, ValidationError @@ -62,12 +62,12 @@ class IPAddressUtils: @staticmethod def parse_internal_networks( - configured_ranges: Optional[List[str]], - ) -> List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]]: + configured_ranges: list[str] | None, + ) -> list[ipaddress.IPv4Network | ipaddress.IPv6Network]: """Parse configured CIDR ranges into network objects, falling back to defaults.""" if not configured_ranges: return IPAddressUtils._DEFAULT_INTERNAL_NETWORKS - networks: List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]] = [] + networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = [] for cidr in configured_ranges: try: networks.append(ipaddress.ip_network(cidr, strict=False)) @@ -77,15 +77,15 @@ class IPAddressUtils: @staticmethod def parse_trusted_proxy_networks( - configured_ranges: Optional[List[str]], - ) -> List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]]: + configured_ranges: list[str] | None, + ) -> list[ipaddress.IPv4Network | ipaddress.IPv6Network]: """ Parse trusted proxy CIDR ranges for XFF validation. Returns empty list if not configured (XFF will not be trusted). """ if not configured_ranges: return [] - networks: List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]] = [] + networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = [] for cidr in configured_ranges: try: networks.append(ipaddress.ip_network(cidr, strict=False)) @@ -95,8 +95,8 @@ class IPAddressUtils: @staticmethod def is_trusted_proxy( - proxy_ip: Optional[str], - trusted_networks: List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]], + proxy_ip: str | None, + trusted_networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network], ) -> bool: """Check if the direct connection IP is from a trusted proxy.""" if not proxy_ip or not trusted_networks: @@ -109,8 +109,8 @@ class IPAddressUtils: @staticmethod def is_internal_ip( - client_ip: Optional[str], - internal_networks: Optional[List[Union[ipaddress.IPv4Network, ipaddress.IPv6Network]]] = None, + client_ip: str | None, + internal_networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] | None = None, ) -> bool: """ Check if a client IP is from an internal/private network. @@ -137,7 +137,7 @@ class IPAddressUtils: @staticmethod def is_request_from_trusted_proxy( request: Request, - general_settings: Optional[Dict[str, Any]] = None, + general_settings: dict[str, Any] | None = None, ) -> bool: """ Return True if X-Forwarded-* headers on this request should be trusted. @@ -194,7 +194,7 @@ class IPAddressUtils: def extract_client_ip_from_xff_hops( xff_header: str, num_trusted_hops: int, - ) -> Optional[str]: + ) -> str | None: """ Resolve the originating client IP from an X-Forwarded-For chain by counting ``num_trusted_hops`` entries from the right. @@ -247,8 +247,8 @@ class IPAddressUtils: @staticmethod def get_mcp_client_ip( request: Request, - general_settings: Optional[Dict[str, Any]] = None, - ) -> Optional[str]: + general_settings: dict[str, Any] | None = None, + ) -> str | None: """ Extract client IP from a FastAPI request for MCP access control. diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index 2bb4ae0b354..a25f3e58d2c 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -4,7 +4,7 @@ import base64 import json import os from datetime import datetime -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING import httpx @@ -26,12 +26,12 @@ class LicenseCheck: def __init__(self) -> None: self.license_str = os.getenv("LITELLM_LICENSE", None) - verbose_proxy_logger.debug("License Str value - {}".format(self.license_str)) + verbose_proxy_logger.debug(f"License Str value - {self.license_str}") self.http_handler = HTTPHandler(timeout=NON_LLM_CONNECTION_TIMEOUT) self._premium_check_logged = False self.public_key = None self.read_public_key() - self.airgapped_license_data: Optional["EnterpriseLicenseData"] = None + self.airgapped_license_data: EnterpriseLicenseData | None = None def read_public_key(self): try: @@ -48,17 +48,15 @@ class LicenseCheck: else: self.public_key = None except Exception as e: - verbose_proxy_logger.error(f"Error reading public key: {str(e)}") + verbose_proxy_logger.error(f"Error reading public key: {e!s}") def _verify(self, license_str: str) -> bool: verbose_proxy_logger.debug( - "litellm.proxy.auth.litellm_license.py::_verify - Checking license against {}/verify_license - {}".format( - self.base_url, license_str - ) + f"litellm.proxy.auth.litellm_license.py::_verify - Checking license against {self.base_url}/verify_license - {license_str}" ) - url = "{}/verify_license/{}".format(self.base_url, license_str) + url = f"{self.base_url}/verify_license/{license_str}" - response: Optional[httpx.Response] = None + response: httpx.Response | None = None try: # don't impact user, if call fails num_retries = 3 for i in range(num_retries): @@ -81,14 +79,12 @@ class LicenseCheck: assert isinstance(premium, bool) verbose_proxy_logger.debug( - "litellm.proxy.auth.litellm_license.py::_verify - License={} is premium={}".format(license_str, premium) + f"litellm.proxy.auth.litellm_license.py::_verify - License={license_str} is premium={premium}" ) return premium except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.auth.litellm_license.py::_verify - Unable to verify License={} via api. - {}".format( - license_str, str(e) - ) + f"litellm.proxy.auth.litellm_license.py::_verify - Unable to verify License={license_str} via api. - {e!s}" ) return False @@ -100,9 +96,7 @@ class LicenseCheck: try: if not self._premium_check_logged: verbose_proxy_logger.debug( - "litellm.proxy.auth.litellm_license.py::is_premium() - ENTERING 'IS_PREMIUM' - LiteLLM License={}".format( - self.license_str - ) + f"litellm.proxy.auth.litellm_license.py::is_premium() - ENTERING 'IS_PREMIUM' - LiteLLM License={self.license_str}" ) if self.license_str is None: @@ -110,9 +104,7 @@ class LicenseCheck: if not self._premium_check_logged: verbose_proxy_logger.debug( - "litellm.proxy.auth.litellm_license.py::is_premium() - Updated 'self.license_str' - {}".format( - self.license_str - ) + f"litellm.proxy.auth.litellm_license.py::is_premium() - Updated 'self.license_str' - {self.license_str}" ) self._premium_check_logged = True @@ -121,9 +113,7 @@ class LicenseCheck: elif ( self.verify_license_without_api_request(public_key=self.public_key, license_key=self.license_str) is True - ): - return True - elif self._verify(license_str=self.license_str) is True: + ) or self._verify(license_str=self.license_str) is True: return True return False except Exception: @@ -148,7 +138,7 @@ class LicenseCheck: if self.airgapped_license_data is None: return False - _max_teams_in_license: Optional[int] = self.airgapped_license_data.get("max_teams") + _max_teams_in_license: int | None = self.airgapped_license_data.get("max_teams") if "max_teams" not in self.airgapped_license_data or not isinstance(_max_teams_in_license, int): return False return team_count > _max_teams_in_license @@ -197,8 +187,6 @@ class LicenseCheck: except Exception as e: verbose_proxy_logger.debug( - "litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - {}".format( - str(e) - ) + f"litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - {e!s}" ) return False diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index f35d94c986e..c79956d8c6e 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -8,7 +8,7 @@ login endpoints (e.g., /login and /v2/login). import os import secrets from datetime import datetime, timedelta, timezone -from typing import Literal, Optional, cast +from typing import Literal, cast import jwt from fastapi import HTTPException @@ -55,7 +55,7 @@ async def _rehash_password_if_needed(user_id: str, password: str, stored: str) - ) -def get_ui_credentials(master_key: Optional[str]) -> tuple[str, str]: +def get_ui_credentials(master_key: str | None) -> tuple[str, str]: """ Get UI username and password from environment variables or master key. @@ -87,7 +87,7 @@ class LoginResult: user_id: str key: str - user_email: Optional[str] + user_email: str | None user_role: str login_method: Literal["sso", "username_password"] @@ -95,7 +95,7 @@ class LoginResult: self, user_id: str, key: str, - user_email: Optional[str], + user_email: str | None, user_role: str, login_method: Literal["sso", "username_password"] = "username_password", ): @@ -109,8 +109,8 @@ class LoginResult: async def authenticate_user( username: str, password: str, - master_key: Optional[str], - prisma_client: Optional[PrismaClient], + master_key: str | None, + prisma_client: PrismaClient | None, ) -> LoginResult: """ Authenticate a user and generate an API key for UI access. @@ -142,19 +142,20 @@ async def authenticate_user( ui_username, ui_password = get_ui_credentials(master_key) # Check if we can find the `username` in the db. On the UI, users can enter username=their email - _user_row: Optional[LiteLLM_UserTable] = None - user_role: Optional[ + _user_row: LiteLLM_UserTable | None = None + user_role: ( Literal[ LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - ] = None + | None + ) = None if prisma_client is not None: _user_row = cast( - Optional[LiteLLM_UserTable], + LiteLLM_UserTable | None, await UserRepository(prisma_client).table.find_first( where={"user_email": {"equals": username, "mode": "insensitive"}} ), @@ -221,7 +222,7 @@ async def authenticate_user( if get_secret_bool("EXPERIMENTAL_UI_LOGIN"): from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken - user_info: Optional[LiteLLM_UserTable] = None + user_info: LiteLLM_UserTable | None = None if _user_row is not None: user_info = _user_row elif user_id is not None: # if user_id is not None, we are using the UI_USERNAME and UI_PASSWORD diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 7d31ce1b909..24875dae9ab 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -1,7 +1,7 @@ # What is this? ## Common checks for /v1/models and `/model/info` import copy -from typing import Any, Dict, List, Optional, Set +from typing import Any import litellm from litellm._logging import verbose_proxy_logger @@ -34,7 +34,7 @@ def _check_wildcard_routing(model: str) -> bool: return False -def get_provider_models(provider: str, litellm_params: Optional[LiteLLM_Params] = None) -> Optional[List[str]]: +def get_provider_models(provider: str, litellm_params: LiteLLM_Params | None = None) -> list[str] | None: """ Returns the list of known models by provider """ @@ -48,10 +48,10 @@ def get_provider_models(provider: str, litellm_params: Optional[LiteLLM_Params] def _get_models_from_access_groups( - model_access_groups: Dict[str, List[str]], - all_models: List[str], - include_model_access_groups: Optional[bool] = False, -) -> List[str]: + model_access_groups: dict[str, list[str]], + all_models: list[str], + include_model_access_groups: bool | None = False, +) -> list[str]: idx_to_remove = [] new_models = [] for idx, model in enumerate(all_models): @@ -69,7 +69,7 @@ def _get_models_from_access_groups( async def get_mcp_server_ids( user_api_key_dict: UserAPIKeyAuth, -) -> List[str]: +) -> list[str]: """ Returns the list of MCP server ids for a given key by querying the object_permission table """ @@ -95,11 +95,11 @@ async def get_mcp_server_ids( def get_key_models( user_api_key_dict: UserAPIKeyAuth, - proxy_model_list: List[str], - model_access_groups: Dict[str, List[str]], - include_model_access_groups: Optional[bool] = False, - only_model_access_groups: Optional[bool] = False, -) -> List[str]: + proxy_model_list: list[str], + model_access_groups: dict[str, list[str]], + include_model_access_groups: bool | None = False, + only_model_access_groups: bool | None = False, +) -> list[str]: """ Returns: - List of model name strings @@ -108,7 +108,7 @@ def get_key_models( - If include_model_access_groups is True, it includes the 'keys' of the model_access_groups in the response - {"beta-models": ["gpt-4", "claude-v1"]} -> returns 'beta-models' """ - all_models: List[str] = [] + all_models: list[str] = [] if len(user_api_key_dict.models) > 0: all_models = list(user_api_key_dict.models) # copy to avoid mutating cached objects if SpecialModelNames.all_team_models.value in all_models: @@ -132,23 +132,23 @@ def get_key_models( # deduplicate while preserving order all_models = list(dict.fromkeys(all_models)) - verbose_proxy_logger.debug("ALL KEY MODELS - {}".format(len(all_models))) + verbose_proxy_logger.debug(f"ALL KEY MODELS - {len(all_models)}") return all_models def get_team_models( - team_models: List[str], - proxy_model_list: List[str], - model_access_groups: Dict[str, List[str]], - include_model_access_groups: Optional[bool] = False, -) -> List[str]: + team_models: list[str], + proxy_model_list: list[str], + model_access_groups: dict[str, list[str]], + include_model_access_groups: bool | None = False, +) -> list[str]: """ Returns: - List of model name strings - Empty list if no models set - If model_access_groups is provided, only return models that are in the access groups """ - all_models_set: Set[str] = set() + all_models_set: set[str] = set() if len(team_models) > 0: all_models_set.update(team_models) if SpecialModelNames.all_team_models.value in all_models_set: @@ -173,23 +173,23 @@ def get_team_models( # deduplicate while preserving order all_models = list(dict.fromkeys(all_models)) - verbose_proxy_logger.debug("ALL TEAM MODELS - {}".format(len(all_models))) + verbose_proxy_logger.debug(f"ALL TEAM MODELS - {len(all_models)}") return all_models def get_complete_model_list( - key_models: List[str], - team_models: List[str], - proxy_model_list: List[str], - user_model: Optional[str], - infer_model_from_keys: Optional[bool], - return_wildcard_routes: Optional[bool] = False, - llm_router: Optional[Router] = None, - model_access_groups: Dict[str, List[str]] = {}, - include_model_access_groups: Optional[bool] = False, - only_model_access_groups: Optional[bool] = False, - team_id: Optional[str] = None, -) -> List[str]: + key_models: list[str], + team_models: list[str], + proxy_model_list: list[str], + user_model: str | None, + infer_model_from_keys: bool | None, + return_wildcard_routes: bool | None = False, + llm_router: Router | None = None, + model_access_groups: dict[str, list[str]] = {}, + include_model_access_groups: bool | None = False, + only_model_access_groups: bool | None = False, + team_id: str | None = None, +) -> list[str]: """Logic for returning complete model list for a given key + team pair""" """ @@ -223,7 +223,7 @@ def get_complete_model_list( append_unique(valid_models) if only_model_access_groups: - model_access_groups_to_return: List[str] = [] + model_access_groups_to_return: list[str] = [] for model in unique_models: if model in model_access_groups: model_access_groups_to_return.append(model) @@ -242,8 +242,8 @@ def get_complete_model_list( def _hydrate_litellm_credential_name( - litellm_params: Optional[LiteLLM_Params], -) -> Optional[LiteLLM_Params]: + litellm_params: LiteLLM_Params | None, +) -> LiteLLM_Params | None: if litellm_params is None or litellm_params.litellm_credential_name is None: return litellm_params @@ -259,7 +259,7 @@ def _hydrate_litellm_credential_name( return litellm_params -def get_known_models_from_wildcard(wildcard_model: str, litellm_params: Optional[LiteLLM_Params] = None) -> List[str]: +def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_Params | None = None) -> list[str]: wildcard_model_to_expand = ( litellm_params.model if wildcard_model == "*" @@ -375,11 +375,11 @@ def expand_wildcard_deployments_for_model_info( def _get_wildcard_models( - unique_models: List[str], - return_wildcard_routes: Optional[bool] = False, - llm_router: Optional[Router] = None, - team_id: Optional[str] = None, -) -> List[str]: + unique_models: list[str], + return_wildcard_routes: bool | None = False, + llm_router: Router | None = None, + team_id: str | None = None, +) -> list[str]: models_to_remove = set() all_wildcard_models = [] for model in unique_models: @@ -422,9 +422,9 @@ def _get_wildcard_models( def get_all_fallbacks( model: str, - llm_router: Optional[Router] = None, + llm_router: Router | None = None, fallback_type: str = "general", -) -> List[str]: +) -> list[str]: """ Get all fallbacks for a given model from the router's fallback configuration. diff --git a/litellm/proxy/auth/oauth2_check.py b/litellm/proxy/auth/oauth2_check.py index a77fc510bb7..e16c261ed45 100644 --- a/litellm/proxy/auth/oauth2_check.py +++ b/litellm/proxy/auth/oauth2_check.py @@ -1,6 +1,6 @@ import base64 import os -from typing import Dict, Optional, Tuple, cast +from typing import cast import httpx @@ -20,8 +20,8 @@ class Oauth2Handler: @staticmethod def _is_introspection_endpoint( token_info_endpoint: str, - oauth_client_id: Optional[str], - oauth_client_secret: Optional[str], + oauth_client_id: str | None, + oauth_client_secret: str | None, ) -> bool: """ Determine if this is an introspection endpoint (requires POST) or token info endpoint (uses GET). @@ -43,9 +43,9 @@ class Oauth2Handler: @staticmethod def _prepare_introspection_request( token: str, - oauth_client_id: Optional[str], - oauth_client_secret: Optional[str], - ) -> Tuple[Dict[str, str], Dict[str, str]]: + oauth_client_id: str | None, + oauth_client_secret: str | None, + ) -> tuple[dict[str, str], dict[str, str]]: """ Prepare headers and data for OAuth2 introspection endpoint (RFC 7662). @@ -72,7 +72,7 @@ class Oauth2Handler: return headers, data @staticmethod - def _prepare_token_info_request(token: str) -> Dict[str, str]: + def _prepare_token_info_request(token: str) -> dict[str, str]: """ Prepare headers for generic token info endpoint. @@ -86,11 +86,11 @@ class Oauth2Handler: @staticmethod def _extract_user_info( - response_data: Dict, + response_data: dict, user_id_field_name: str, user_role_field_name: str, user_team_id_field_name: str, - ) -> Tuple[Optional[str], Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None, str | None]: """ Extract user information from OAuth2 response. diff --git a/litellm/proxy/auth/oauth2_proxy_hook.py b/litellm/proxy/auth/oauth2_proxy_hook.py index ca6a7ee4b1d..a7e34072712 100644 --- a/litellm/proxy/auth/oauth2_proxy_hook.py +++ b/litellm/proxy/auth/oauth2_proxy_hook.py @@ -1,5 +1,4 @@ from collections.abc import Mapping -from typing import Dict, FrozenSet, List, Union from fastapi import Request @@ -25,7 +24,7 @@ from litellm.proxy.auth.trusted_proxy_utils import require_trusted_proxy_request # Operators who need a trusted upstream to assert anything beyond # identity should switch to JWT authentication, which validates a # signature on the assertion rather than blindly trusting headers. -ALLOWED_OAUTH2_PROXY_FIELDS: FrozenSet[str] = frozenset( +ALLOWED_OAUTH2_PROXY_FIELDS: frozenset[str] = frozenset( { "user_id", "user_email", @@ -65,7 +64,7 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: feature_name="OAuth2 proxy auth", ) - oauth2_config_mappings: Dict[str, str] = general_settings.get("oauth2_config_mappings") or {} + oauth2_config_mappings: dict[str, str] = general_settings.get("oauth2_config_mappings") or {} verbose_proxy_logger.debug(f"Oauth2 config mappings: {oauth2_config_mappings}") if not oauth2_config_mappings: @@ -84,7 +83,7 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: "(signature-validated) instead of header-trust." ) - auth_data: Mapping[str, Union[str, List[str]]] = { + auth_data: Mapping[str, str | list[str]] = { key: [model.strip() for model in value.split(",")] if key == "models" else value for key, header in oauth2_config_mappings.items() if (value := request.headers.get(header)) diff --git a/litellm/proxy/auth/rds_iam_token.py b/litellm/proxy/auth/rds_iam_token.py index 2ef66d3b7ff..5046c0cb7b5 100644 --- a/litellm/proxy/auth/rds_iam_token.py +++ b/litellm/proxy/auth/rds_iam_token.py @@ -1,18 +1,18 @@ import os -from typing import Any, Optional, Union +from typing import Any import httpx def init_rds_client( - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_region_name: Optional[str] = None, - aws_session_name: Optional[str] = None, - aws_profile_name: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_web_identity_token: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + aws_access_key_id: str | None = None, + aws_secret_access_key: str | None = None, + aws_region_name: str | None = None, + aws_session_name: str | None = None, + aws_profile_name: str | None = None, + aws_role_name: str | None = None, + aws_web_identity_token: str | None = None, + timeout: float | httpx.Timeout | None = None, ): from litellm.secret_managers.main import get_secret @@ -153,7 +153,7 @@ def init_rds_client( return client -def generate_iam_auth_token(db_host, db_port, db_user, client: Optional[Any] = None) -> str: +def generate_iam_auth_token(db_host, db_port, db_user, client: Any | None = None) -> str: from urllib.parse import quote if client is None: diff --git a/litellm/proxy/auth/resolvers/exceptions.py b/litellm/proxy/auth/resolvers/exceptions.py index a2ece57209b..fd57ba94aea 100644 --- a/litellm/proxy/auth/resolvers/exceptions.py +++ b/litellm/proxy/auth/resolvers/exceptions.py @@ -29,9 +29,7 @@ class KeyNotFoundError(IdentityResolutionError, ProxyException): def __init__(self, hashed_token: str) -> None: ProxyException.__init__( self, - message="Authentication Error, Invalid proxy server token passed. key={}, not found in db. Create key via `/key/generate` call.".format( - hashed_token - ), + message=f"Authentication Error, Invalid proxy server token passed. key={hashed_token}, not found in db. Create key via `/key/generate` call.", type=ProxyErrorTypes.token_not_found_in_db, param="key", code=status.HTTP_401_UNAUTHORIZED, diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index dd0a34a7898..1e63d9746a4 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -1,5 +1,4 @@ import re -from typing import List, Optional from fastapi import HTTPException, Request, status @@ -67,7 +66,7 @@ class RouteChecks: def should_call_route( route: str, valid_token: UserAPIKeyAuth, - request: Optional[Request] = None, + request: Request | None = None, ): """ Check if management route is disabled and raise exception @@ -89,7 +88,7 @@ class RouteChecks: def is_virtual_key_allowed_to_call_route( route: str, valid_token: UserAPIKeyAuth, - request: Optional[Request] = None, + request: Request | None = None, ) -> bool: """ Raises Exception if Virtual Key is not allowed to call the route @@ -201,7 +200,7 @@ class RouteChecks: @staticmethod def _raise_admin_only_route_exception( - user_obj: Optional[LiteLLM_UserTable], + user_obj: LiteLLM_UserTable | None, route: str, ) -> None: """ @@ -227,8 +226,8 @@ class RouteChecks: @staticmethod def non_proxy_admin_allowed_routes_check( - user_obj: Optional[LiteLLM_UserTable], - _user_role: Optional[LitellmUserRoles], + user_obj: LiteLLM_UserTable | None, + _user_role: LitellmUserRoles | None, route: str, request: Request, valid_token: UserAPIKeyAuth, @@ -263,9 +262,7 @@ class RouteChecks: if user_id and user_id != valid_token.user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail="key not allowed to access this user's info. user_id={}, key's user_id={}".format( - user_id, valid_token.user_id - ), + detail=f"key not allowed to access this user's info. user_id={user_id}, key's user_id={valid_token.user_id}", ) elif route == "/v2/user/info": # handled by the endpoint itself (full RBAC in handler) @@ -288,23 +285,19 @@ class RouteChecks: request_data=request_data, request=request, ) - elif _user_role == LitellmUserRoles.INTERNAL_USER.value and RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value + elif ( + _user_role == LitellmUserRoles.INTERNAL_USER.value + and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value) + or _user_is_org_admin(request_data=request_data, user_object=user_obj) + and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value) + or _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value + and RouteChecks.check_route_access( + route=route, + allowed_routes=LiteLLMRoutes.internal_user_view_only_routes.value, + ) + or RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.self_managed_routes.value) ): pass - elif _user_is_org_admin(request_data=request_data, user_object=user_obj) and RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value - ): - pass - elif _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value and RouteChecks.check_route_access( - route=route, - allowed_routes=LiteLLMRoutes.internal_user_view_only_routes.value, - ): - pass - elif RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.self_managed_routes.value - ): # routes that manage their own allowed/disallowed logic - pass elif route.startswith("/v1/mcp/") or route.startswith("/mcp-rest/"): pass # authN/authZ handled by api itself elif RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token): @@ -341,7 +334,6 @@ class RouteChecks: status_code=status.HTTP_403_FORBIDDEN, detail=f"user not allowed to access this route. Route={route} is an admin only route", ) - pass @staticmethod def is_llm_api_route(route: str) -> bool: @@ -409,7 +401,7 @@ class RouteChecks: return False @staticmethod - def _is_get_mcp_server_discovery_route(route: str, request: Optional[Request]) -> bool: + def _is_get_mcp_server_discovery_route(route: str, request: Request | None) -> bool: """ Returns True if `request` is a GET against one of the two read-only MCP-server discovery paths: @@ -560,7 +552,7 @@ class RouteChecks: return False @staticmethod - def check_route_access(route: str, allowed_routes: List[str]) -> bool: + def check_route_access(route: str, allowed_routes: list[str]) -> bool: """ Check if a route has access by checking both exact matches and patterns @@ -600,7 +592,7 @@ class RouteChecks: return False @staticmethod - def _get_request_method(request: Optional[Request]) -> Optional[str]: + def _get_request_method(request: Request | None) -> str | None: if request is None: return None @@ -614,7 +606,7 @@ class RouteChecks: return method.upper() @staticmethod - def is_auth_enforced_pass_through_route(route: str, method: Optional[str] = None) -> bool: + def is_auth_enforced_pass_through_route(route: str, method: str | None = None) -> bool: """ True for config/DB pass-through endpoints registered with auth=true. @@ -753,7 +745,7 @@ class RouteChecks: route: str, _user_role: str, request_data: dict, - request: Optional[Request] = None, + request: Request | None = None, ) -> None: """ Check access for PROXY_ADMIN_VIEW_ONLY role. @@ -815,7 +807,7 @@ class RouteChecks: # Allow `/user/update` for self-service email / password change. if route == "/user/update": if request_data is not None and isinstance(request_data, dict): - for param in request_data.keys(): + for param in request_data: if param not in ["user_email", "password"]: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, diff --git a/litellm/proxy/auth/trusted_proxy_utils.py b/litellm/proxy/auth/trusted_proxy_utils.py index 86eb22aee84..7975e66f32b 100644 --- a/litellm/proxy/auth/trusted_proxy_utils.py +++ b/litellm/proxy/auth/trusted_proxy_utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional +from typing import Any from fastapi import Request @@ -12,7 +12,7 @@ from litellm.proxy.auth.network import ( TRUSTED_PROXY_RANGES_KEY = "trusted_proxy_ranges" -def _get_proxy_general_settings() -> Dict[str, Any]: +def _get_proxy_general_settings() -> dict[str, Any]: try: from litellm.proxy.proxy_server import general_settings @@ -37,7 +37,7 @@ def get_trusted_proxy_cidrs( ) -def _get_direct_client_ip(request: Request) -> Optional[str]: +def _get_direct_client_ip(request: Request) -> str | None: client = getattr(request, "client", None) client_host = getattr(client, "host", None) if isinstance(client_host, str): @@ -48,7 +48,7 @@ def _get_direct_client_ip(request: Request) -> Optional[str]: def require_trusted_proxy_request( *, request: Request, - general_settings: Optional[Dict[str, Any]] = None, + general_settings: dict[str, Any] | None = None, feature_name: str, setting_name: str = TRUSTED_PROXY_RANGES_KEY, ) -> None: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index f4d07c1a674..72450453174 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -11,12 +11,11 @@ import asyncio import fnmatch import re import secrets - -import orjson from datetime import datetime, timezone -from typing import Any, Dict, NamedTuple, List, Optional, Protocol, Tuple, Union, cast +from typing import Any, NamedTuple, Protocol, Union, cast import fastapi +import orjson from fastapi import HTTPException, Request, WebSocket, status from fastapi.security.api_key import APIKeyHeader @@ -56,6 +55,7 @@ from litellm.proxy.auth.auth_checks import ( resolve_and_validate_end_user_id, ) from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler +from litellm.proxy.auth.auth_method import AuthMethod from litellm.proxy.auth.auth_utils import ( abbreviate_api_key, get_end_user_id_from_request_body, @@ -68,10 +68,9 @@ from litellm.proxy.auth.auth_utils import ( route_in_additonal_public_routes, ) from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler +from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_context from litellm.proxy.auth.oauth2_check import Oauth2Handler from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request -from litellm.proxy.auth.auth_method import AuthMethod -from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_context from litellm.proxy.auth.resolvers import CredentialRef, Principal from litellm.proxy.auth.resolvers.store import IdentityStore from litellm.proxy.auth.route_checks import RouteChecks @@ -101,7 +100,7 @@ try: enterprise_custom_auth as _enterprise_custom_auth, ) - enterprise_custom_auth: Optional[Callable] = _enterprise_custom_auth + enterprise_custom_auth: Callable | None = _enterprise_custom_auth except ImportError as e: verbose_proxy_logger.debug(f"Error in enterprise custom auth: {e}") enterprise_custom_auth = None @@ -115,7 +114,7 @@ def _normalize_public_auth_route(route: str) -> str: return route -def _route_requires_auth_despite_public(route: str, general_settings: Optional[dict]) -> bool: +def _route_requires_auth_despite_public(route: str, general_settings: dict | None) -> bool: normalized_route = _normalize_public_auth_route(route) if normalized_route == "/metrics": return litellm.require_auth_for_metrics_endpoint is not False @@ -158,9 +157,9 @@ azure_apim_header = APIKeyHeader( def _get_model_from_request_context( request_data: dict, route: str, - request: Optional[Request], - llm_router: Optional[Any] = None, -) -> Optional[Union[str, List[str]]]: + request: Request | None, + llm_router: Any | None = None, +) -> str | list[str] | None: return get_model_from_request( request_data=request_data, route=route, @@ -172,8 +171,8 @@ def _get_model_from_request_context( def _get_model_names_for_budget_checks( - model: Optional[Union[str, List[str]]], -) -> List[str]: + model: str | list[str] | None, +) -> list[str]: if model is None: return [] if isinstance(model, str): @@ -184,9 +183,7 @@ def _get_model_names_for_budget_checks( class _KeyModelBudgetLimiter(Protocol): async def is_key_within_model_budget(self, user_api_key_dict: UserAPIKeyAuth, model: str) -> bool: ... - async def get_fallback_model_within_budget( - self, user_api_key_dict: UserAPIKeyAuth, model: str - ) -> Optional[str]: ... + async def get_fallback_model_within_budget(self, user_api_key_dict: UserAPIKeyAuth, model: str) -> str | None: ... async def _check_key_model_budget_with_fallback( @@ -195,8 +192,8 @@ async def _check_key_model_budget_with_fallback( model_name: str, request_data: dict, request: Request, - llm_model_list: Optional[list] = None, - llm_router: Optional[litellm.Router] = None, + llm_model_list: list | None = None, + llm_router: litellm.Router | None = None, ) -> None: """ Enforce the key's per-model budget for `model_name`. If exceeded and the @@ -286,15 +283,15 @@ def _get_bearer_token_or_received_api_key(api_key: str) -> str: def _routing_selector_matches_claim( - selector_value: Optional[Any], - claim_value: Optional[Any], + selector_value: Any | None, + claim_value: Any | None, *, split_space_delimited: bool = False, ) -> bool: if selector_value is None: return True - selector_list: List[str] = ( + selector_list: list[str] = ( [str(v) for v in selector_value] if isinstance(selector_value, list) else [str(selector_value)] ) @@ -387,7 +384,7 @@ def _get_bearer_token( def _apply_budget_limits_to_end_user_params( end_user_params: dict, budget_info: LiteLLM_BudgetTable, - end_user_id: Optional[str], + end_user_id: str | None, ) -> None: """ Helper function to apply budget limits to end user parameters. @@ -422,7 +419,7 @@ async def user_api_key_auth_websocket(websocket: WebSocket): # ``websocket.url``, which Starlette reconstructs from the (poisonable) # Host header. Carry the ASGI scope's path / root_path so the lookup # never reaches the fallback. - synthetic_scope: Dict[str, Any] = { + synthetic_scope: dict[str, Any] = { "type": "http", "headers": scope_headers, "path": ws_scope.get("path", ""), @@ -500,7 +497,7 @@ async def _fetch_global_spend_with_event_coordination( cache_key: str, user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, -) -> Optional[float]: +) -> float | None: """ Fetch global spend with event-driven coordination to prevent cache stampede. Uses EventDrivenCacheCoordinator: first request queries DB and signals others when done. @@ -509,7 +506,7 @@ async def _fetch_global_spend_with_event_coordination( per request and is zeroed by ResetBudgetJob every ``litellm.budget_duration``. """ - async def _load_global_spend() -> Optional[float]: + async def _load_global_spend() -> float | None: proxy_budget_row = await prisma_client.db.litellm_usertable.find_unique( where={"user_id": LITELLM_PROXY_BUDGET_NAME} ) @@ -525,10 +522,10 @@ async def _fetch_global_spend_with_event_coordination( async def get_global_proxy_spend( litellm_proxy_admin_name: str, user_api_key_cache: UserApiKeyCache, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, token: str, proxy_logging_obj: ProxyLogging, -) -> Optional[float]: +) -> float | None: global_proxy_spend = None if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget # Use event-driven coordination to prevent cache stampede @@ -555,7 +552,7 @@ async def get_global_proxy_spend( return global_proxy_spend -def get_rbac_role(jwt_handler: JWTHandler, scopes: List[str]) -> str: +def get_rbac_role(jwt_handler: JWTHandler, scopes: list[str]) -> str: is_admin = jwt_handler.is_admin(scopes=scopes) if is_admin: return LitellmUserRoles.PROXY_ADMIN @@ -564,16 +561,16 @@ def get_rbac_role(jwt_handler: JWTHandler, scopes: List[str]) -> str: def get_api_key( - custom_litellm_key_header: Optional[str], + custom_litellm_key_header: str | None, api_key: str, - azure_api_key_header: Optional[str], - anthropic_api_key_header: Optional[str], - google_ai_studio_api_key_header: Optional[str], - azure_apim_header: Optional[str], - pass_through_endpoints: Optional[List[dict]], + azure_api_key_header: str | None, + anthropic_api_key_header: str | None, + google_ai_studio_api_key_header: str | None, + azure_apim_header: str | None, + pass_through_endpoints: list[dict] | None, route: str, request: Request, -) -> Tuple[str, Optional[str]]: +) -> tuple[str, str | None]: """ Returns: Tuple[Optional[str], Optional[str]]: Tuple of the api_key and the passed_in_key @@ -584,7 +581,7 @@ def get_api_key( ) api_key = api_key - passed_in_key: Optional[str] = None + passed_in_key: str | None = None if isinstance(custom_litellm_key_header, str): passed_in_key = custom_litellm_key_header api_key = _get_bearer_token_or_received_api_key(custom_litellm_key_header) @@ -614,7 +611,7 @@ def get_api_key( elif pass_through_endpoints is not None: for endpoint in pass_through_endpoints: if endpoint.get("path", "") == route: - headers: Optional[dict] = endpoint.get("headers", None) + headers: dict | None = endpoint.get("headers", None) if headers is not None: header_key: str = headers.get("litellm_user_api_key", "") if request.headers.get(header_key) is not None: @@ -626,9 +623,9 @@ def get_api_key( async def check_api_key_for_custom_headers_or_pass_through_endpoints( request: Request, route: str, - pass_through_endpoints: Optional[List[dict]], + pass_through_endpoints: list[dict] | None, api_key: str, -) -> Union[UserAPIKeyAuth, str]: +) -> UserAPIKeyAuth | str: is_mapped_pass_through_route: bool = False normalized_route = normalize_route_for_root_path(route) if normalized_route is not None: @@ -705,14 +702,14 @@ async def _auto_register_jwt_mapping( jwt_handler: JWTHandler, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, cache_key: str, - team_id: Optional[str] = None, - user_id: Optional[str] = None, - org_id: Optional[str] = None, - end_user_id: Optional[str] = None, -) -> Optional[UserAPIKeyAuth]: + team_id: str | None = None, + user_id: str | None = None, + org_id: str | None = None, + end_user_id: str | None = None, +) -> UserAPIKeyAuth | None: """ Auto-register: create a new virtual key + mapping for an unrecognised JWT claim value. ``team_id`` and ``user_id`` must come from a successful @@ -837,11 +834,11 @@ async def _auto_register_jwt_mapping( async def _resolve_jwt_to_virtual_key( jwt_claims: dict, jwt_handler: JWTHandler, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, -) -> Union[Optional[UserAPIKeyAuth], "_PendingAutoRegister"]: +) -> Union[UserAPIKeyAuth | None, "_PendingAutoRegister"]: """ Returns: - ``UserAPIKeyAuth``: a resolved virtual key (cache hit or DB hit). The @@ -937,7 +934,7 @@ async def _resolve_jwt_to_virtual_key( # Resolve the mapping from DB, or treat prisma_client=None as a definitive # miss (no DB → no mapping can exist → apply no-match policy below). - token_hash: Optional[str] = None + token_hash: str | None = None if prisma_client is not None: token_hash = await get_jwt_key_mapping_object( jwt_claim_name=virtual_key_claim_field, @@ -1044,11 +1041,11 @@ async def _user_api_key_auth_builder( request: Request, api_key: str, azure_api_key_header: str, - anthropic_api_key_header: Optional[str], - google_ai_studio_api_key_header: Optional[str], - azure_apim_header: Optional[str], + anthropic_api_key_header: str | None, + google_ai_studio_api_key_header: str | None, + azure_apim_header: str | None, request_data: dict, - custom_litellm_key_header: Optional[str] = None, + custom_litellm_key_header: str | None = None, ) -> UserAPIKeyAuth: from litellm.proxy.proxy_server import ( general_settings, @@ -1065,7 +1062,7 @@ async def _user_api_key_auth_builder( user_custom_auth, ) - parent_otel_span: Optional[Span] = None + parent_otel_span: Span | None = None # Prefer the receive-instant stamped by the early helper in # user_api_key_auth (before body parse) — overwriting it would shorten # the preprocessing-duration measurement by the body-parse window. @@ -1075,7 +1072,7 @@ async def _user_api_key_auth_builder( except Exception: pass route: str = get_request_route(request=request) - valid_token: Optional[UserAPIKeyAuth] = None + valid_token: UserAPIKeyAuth | None = None custom_auth_api_key: bool = False try: @@ -1085,7 +1082,7 @@ async def _user_api_key_auth_builder( request=request, route=route, ) - pass_through_endpoints: Optional[List[dict]] = general_settings.get("pass_through_endpoints", None) + pass_through_endpoints: list[dict] | None = general_settings.get("pass_through_endpoints", None) ## CHECK IF X-LITELM-API-KEY IS PASSED IN - supercedes Authorization header api_key, passed_in_key = get_api_key( custom_litellm_key_header=custom_litellm_key_header, @@ -1212,10 +1209,10 @@ async def _user_api_key_auth_builder( # Try JWT-to-Virtual-Key mapping first to avoid # unnecessary DB queries in auth_builder do_standard_jwt_auth = True - pending_auto_register: Optional[_PendingAutoRegister] = None + pending_auto_register: _PendingAutoRegister | None = None if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None: # Decode JWT to get claims without running full auth_builder - jwt_claims: Optional[dict] + jwt_claims: dict | None if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt: jwt_claims = await jwt_handler.get_oidc_userinfo(token=api_key) else: @@ -1267,7 +1264,7 @@ async def _user_api_key_auth_builder( user_object = result["user_object"] end_user_id = result["end_user_id"] org_id = result["org_id"] - team_membership: Optional[LiteLLM_TeamMembership] = result.get("team_membership", None) + team_membership: LiteLLM_TeamMembership | None = result.get("team_membership", None) jwt_claims = result.get("jwt_claims", None) if is_proxy_admin: @@ -1484,8 +1481,7 @@ async def _user_api_key_auth_builder( except Exception as e: if isinstance(e, litellm.BudgetExceededError): raise e - verbose_proxy_logger.debug("Unable to find user in db. Error - {}".format(str(e))) - pass + verbose_proxy_logger.debug(f"Unable to find user in db. Error - {e!s}") ### CHECK IF ADMIN ### # note: never string compare api keys, this is vulenerable to a time attack. Use secrets.compare_digest instead @@ -1596,7 +1592,7 @@ async def _user_api_key_auth_builder( if not isinstance(master_key, str): raise HTTPException( status_code=500, - detail={"Master key must be a valid string. Current type={}".format(type(master_key))}, + detail={f"Master key must be a valid string. Current type={type(master_key)}"}, ) if is_master_key_valid: @@ -1651,15 +1647,11 @@ async def _user_api_key_auth_builder( if valid_token is None: if isinstance(api_key, str): # if generated token, make sure it starts with sk-. - _masked_key = "{}****{}".format(api_key[:4], api_key[-4:]) if len(api_key) > 8 else "****" + _masked_key = f"{api_key[:4]}****{api_key[-4:]}" if len(api_key) > 8 else "****" if not api_key.startswith("sk-"): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, - detail=( - "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( - _masked_key - ) - ), + detail=(f"LiteLLM Virtual Key expected. Received={_masked_key}, expected to start with 'sk-'."), ) # prevent token hashes from being used else: verbose_logger.warning( @@ -1683,9 +1675,7 @@ async def _user_api_key_auth_builder( ) except ProxyException as e: if e.code == 401 or e.code == "401": - e.message = "Authentication Error, Invalid proxy server token passed. Received API Key = {}, Key Hash (Token) ={}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`".format( - abbreviated_api_key, api_key - ) + e.message = f"Authentication Error, Invalid proxy server token passed. Received API Key = {abbreviated_api_key}, Key Hash (Token) ={api_key}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`" raise e # update end-user params on valid token # These can change per request - it's important to update them here @@ -1697,7 +1687,7 @@ async def _user_api_key_auth_builder( if valid_token is not None: valid_token = _update_key_budget_with_temp_budget_increase(valid_token) - user_obj: Optional[LiteLLM_UserTable] = None + user_obj: LiteLLM_UserTable | None = None valid_token_dict: dict = {} if valid_token is not None: # Got Valid Token from Cache, DB @@ -1739,9 +1729,7 @@ async def _user_api_key_auth_builder( ) except Exception as e: verbose_logger.debug( - "litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format( - str(e) - ) + f"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {e!s}" ) user_obj = None @@ -2297,7 +2285,7 @@ async def _run_centralized_common_checks( # else branch. After the for-loop above, the only BaseException that # can still appear here is HTTPException (other listed re-raises were # propagated; non-listed exceptions were already swallowed to None). - team_object: Optional[LiteLLM_TeamTableCachedObj] + team_object: LiteLLM_TeamTableCachedObj | None if isinstance(team_result, BaseException): # Token-derived fallback only valid when a team_id is set; # _team_obj_from_token asserts that precondition. @@ -2305,16 +2293,14 @@ async def _run_centralized_common_checks( else: team_object = team_result - user_object: Optional[LiteLLM_UserTable] = None if isinstance(user_result, BaseException) else user_result - project_object: Optional[LiteLLM_ProjectTableCachedObj] = ( + user_object: LiteLLM_UserTable | None = None if isinstance(user_result, BaseException) else user_result + project_object: LiteLLM_ProjectTableCachedObj | None = ( None if isinstance(project_result, BaseException) else project_result ) - end_user_object: Optional[LiteLLM_EndUserTable] = ( + end_user_object: LiteLLM_EndUserTable | None = ( None if isinstance(end_user_result, BaseException) else end_user_result ) - global_proxy_spend: Optional[float] = ( - None if isinstance(global_spend_result, BaseException) else global_spend_result - ) + global_proxy_spend: float | None = None if isinstance(global_spend_result, BaseException) else global_spend_result if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None: user_api_key_auth_obj.org_id = team_object.organization_id @@ -2402,23 +2388,23 @@ async def _run_centralized_common_checks( async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary (e.g. token has no team_id). Keeps the result tuple positional.""" - return None + return async def _reserve_budget_after_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, request_data: dict, route: str, - llm_router: Optional[Any], - team_object: Optional[LiteLLM_TeamTableCachedObj], - user_object: Optional[LiteLLM_UserTable], - prisma_client: Optional[PrismaClient], + llm_router: Any | None, + team_object: LiteLLM_TeamTableCachedObj | None, + user_object: LiteLLM_UserTable | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, skip_budget_checks: bool, general_settings: dict, - end_user_id: Optional[str] = None, - end_user_object: Optional[LiteLLM_EndUserTable] = None, + end_user_id: str | None = None, + end_user_object: LiteLLM_EndUserTable | None = None, ) -> None: user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: @@ -2457,8 +2443,8 @@ async def _reserve_budget_after_common_checks( def _should_skip_budget_checks( request_data: dict, route: str, - request: Optional[Request], - llm_router: Optional[Any], + request: Request | None, + llm_router: Any | None, ) -> bool: model = _get_model_from_request_context( request_data=request_data, @@ -2500,10 +2486,10 @@ async def user_api_key_auth( request: Request, api_key: str = fastapi.Security(api_key_header), azure_api_key_header: str = fastapi.Security(azure_api_key_header), - anthropic_api_key_header: Optional[str] = fastapi.Security(anthropic_api_key_header), - google_ai_studio_api_key_header: Optional[str] = fastapi.Security(google_ai_studio_api_key_header), - azure_apim_header: Optional[str] = fastapi.Security(azure_apim_header), - custom_litellm_key_header: Optional[str] = fastapi.Security(custom_litellm_key_header), + anthropic_api_key_header: str | None = fastapi.Security(anthropic_api_key_header), + google_ai_studio_api_key_header: str | None = fastapi.Security(google_ai_studio_api_key_header), + azure_apim_header: str | None = fastapi.Security(azure_apim_header), + custom_litellm_key_header: str | None = fastapi.Security(custom_litellm_key_header), ) -> UserAPIKeyAuth: """ Parent function to authenticate user api key / jwt token. @@ -2613,13 +2599,13 @@ async def user_api_key_auth( async def _return_user_api_key_auth_obj( - user_obj: Optional[LiteLLM_UserTable], + user_obj: LiteLLM_UserTable | None, api_key: str, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, valid_token_dict: dict, route: str, start_time: datetime, - user_role: Optional[LitellmUserRoles] = None, + user_role: LitellmUserRoles | None = None, ) -> UserAPIKeyAuth: end_time = datetime.now(timezone.utc) @@ -2683,9 +2669,7 @@ def get_api_key_from_custom_header(request: Request, custom_litellm_key_header_n if custom_api_key: api_key = _get_bearer_token(api_key=custom_api_key) verbose_proxy_logger.debug( - "Found custom API key using header: {}, setting api_key={}".format( - custom_litellm_key_header_name, abbreviate_api_key(api_key) - ) + f"Found custom API key using header: {custom_litellm_key_header_name}, setting api_key={abbreviate_api_key(api_key)}" ) else: verbose_proxy_logger.exception( @@ -2719,7 +2703,7 @@ def _update_key_budget_with_temp_budget_increase( async def _lookup_end_user_and_apply_budget( valid_token: UserAPIKeyAuth, route: str, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, prisma_client, user_api_key_cache, proxy_logging_obj, @@ -2770,7 +2754,7 @@ async def _lookup_end_user_and_apply_budget( except Exception as e: if isinstance(e, litellm.BudgetExceededError): raise e - verbose_proxy_logger.debug(f"Unable to find user in db. Error - {str(e)}") + verbose_proxy_logger.debug(f"Unable to find user in db. Error - {e!s}") return valid_token, end_user_object @@ -2779,9 +2763,9 @@ async def _enforce_key_and_fallback_model_access( valid_token: UserAPIKeyAuth, request_data: dict, route: str, - request: Optional[Request], - llm_model_list: Optional[list], - llm_router: Optional[Any], + request: Request | None, + llm_model_list: list | None, + llm_router: Any | None, ) -> None: """ Key-level model allowlist and client fallbacks (same as standard auth). @@ -2846,7 +2830,7 @@ async def _run_post_custom_auth_checks( request: Request, request_data: dict, route: str, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, ) -> UserAPIKeyAuth: from litellm.proxy.proxy_server import ( general_settings, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index a91b29002e3..8b2a437009d 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -5,7 +5,7 @@ ###################################################################### import asyncio -from typing import Any, Dict, Optional, cast +from typing import Any, cast from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response @@ -13,9 +13,9 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest from litellm.proxy._types import * -from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, @@ -27,8 +27,8 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( decode_model_from_file_id, encode_batch_response_ids, encode_file_id_with_model, - get_batch_id_from_unified_batch_id, get_batch_from_database, + get_batch_id_from_unified_batch_id, get_credentials_for_model, get_model_id_from_unified_batch_id, get_models_from_unified_file_id, @@ -88,7 +88,7 @@ async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str | async def create_batch( request: Request, fastapi_response: Response, - provider: Optional[str] = None, + provider: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -116,11 +116,11 @@ async def create_batch( version, ) - data: Dict = {} + data: dict = {} try: data = await _read_request_body(request=request) verbose_proxy_logger.debug( - "Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)), + f"Request received by LiteLLM:\n{json.dumps(data, indent=4)}", ) base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) ( @@ -138,7 +138,7 @@ async def create_batch( data["metadata"] = sanitize_openai_provider_metadata(data.get("metadata")) ## check if model is a loadbalanced model - router_model: Optional[str] = None + router_model: str | None = None is_router_model = False if litellm.enable_loadbalancing_on_batch_endpoints is True: router_model = data.get("model", None) @@ -247,7 +247,7 @@ async def create_batch( if len(target_model_names) != 1: raise HTTPException( status_code=400, - detail={"error": "Expected 1 model, got {}".format(len(target_model_names))}, + detail={"error": f"Expected 1 model, got {len(target_model_names)}"}, ) model = target_model_names[0] _create_batch_data["model"] = model @@ -340,9 +340,7 @@ async def create_batch( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.create_batch(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.create_batch(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -365,7 +363,7 @@ async def retrieve_batch( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - provider: Optional[str] = None, + provider: str | None = None, batch_id: str = Path(title="Batch ID to retrieve", description="The ID of the batch to retrieve"), ): """ @@ -389,7 +387,7 @@ async def retrieve_batch( version, ) - data: Dict = {} + data: dict = {} try: model_from_id = decode_model_from_file_id(batch_id) _retrieve_batch_request = RetrieveBatchRequest( @@ -594,9 +592,7 @@ async def retrieve_batch( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.retrieve_batch(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.retrieve_batch(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -618,11 +614,11 @@ async def retrieve_batch( async def list_batches( request: Request, fastapi_response: Response, - provider: Optional[str] = None, - limit: Optional[int] = None, - after: Optional[str] = None, + provider: str | None = None, + limit: int | None = None, + after: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - target_model_names: Optional[str] = None, + target_model_names: str | None = None, ): """ Lists @@ -645,7 +641,7 @@ async def list_batches( version, ) - verbose_proxy_logger.debug("GET /v1/batches after={} limit={}".format(after, limit)) + verbose_proxy_logger.debug(f"GET /v1/batches after={after} limit={limit}") try: if llm_router is None: raise HTTPException( @@ -777,7 +773,7 @@ async def list_batches( original_exception=e, request_data={"after": after, "limit": limit}, ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.retrieve_batch(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.retrieve_batch(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -800,7 +796,7 @@ async def cancel_batch( request: Request, batch_id: str, fastapi_response: Response, - provider: Optional[str] = None, + provider: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -827,7 +823,7 @@ async def cancel_batch( version, ) - data: Dict = {} + data: dict = {} try: # Check for encoded batch ID with model info model_from_id = decode_model_from_file_id(batch_id) @@ -986,9 +982,7 @@ async def cancel_batch( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.create_batch(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.create_batch(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) diff --git a/litellm/proxy/caching_routes.py b/litellm/proxy/caching_routes.py index 09d68ee3a0c..e64f3e9e7e3 100644 --- a/litellm/proxy/caching_routes.py +++ b/litellm/proxy/caching_routes.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Tuple +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request @@ -19,7 +19,7 @@ router = APIRouter( ) -def _extract_cache_params() -> Dict[str, Any]: +def _extract_cache_params() -> dict[str, Any]: """ Safely extracts and cleans cache parameters. @@ -43,7 +43,7 @@ def _extract_cache_params() -> Dict[str, Any]: cleaned_params = HealthCheckCacheParams(**cache_params).model_dump() if cache_params else {} return masker.mask_dict(cleaned_params) except (AttributeError, TypeError) as e: - verbose_proxy_logger.debug(f"Error extracting cache params: {str(e)}") + verbose_proxy_logger.debug(f"Error extracting cache params: {e!s}") return {} @@ -56,8 +56,8 @@ async def cache_ping(): """ Endpoint for checking if cache can be pinged """ - litellm_cache_params: Dict[str, Any] = {} - cleaned_cache_params: Dict[str, Any] = {} + litellm_cache_params: dict[str, Any] = {} + cleaned_cache_params: dict[str, Any] = {} if litellm.cache is None: raise ProxyException( message=safe_dumps( @@ -158,11 +158,11 @@ async def cache_delete(request: Request): except Exception as e: raise HTTPException( status_code=500, - detail=f"Cache Delete Failed({str(e)})", + detail=f"Cache Delete Failed({e!s})", ) -def _get_redis_client_info(cache_instance) -> Tuple[List, int]: +def _get_redis_client_info(cache_instance) -> tuple[list, int]: """ Helper function to safely get Redis client list information. @@ -173,7 +173,7 @@ def _get_redis_client_info(cache_instance) -> Tuple[List, int]: client_list = cache_instance.client_list() return client_list, len(client_list) except Exception as e: - verbose_proxy_logger.warning(f"CLIENT LIST command failed (likely restricted on managed Redis): {str(e)}") + verbose_proxy_logger.warning(f"CLIENT LIST command failed (likely restricted on managed Redis): {e!s}") return ["CLIENT LIST command not available on this Redis instance"], -1 @@ -209,7 +209,7 @@ async def cache_redis_info(): except Exception as e: raise HTTPException( status_code=503, - detail=f"Service Unhealthy ({str(e)})", + detail=f"Service Unhealthy ({e!s})", ) @@ -245,5 +245,5 @@ async def cache_flushall(): except Exception as e: raise HTTPException( status_code=503, - detail=f"Service Unhealthy ({str(e)})", + detail=f"Service Unhealthy ({e!s})", ) diff --git a/litellm/proxy/client/__init__.py b/litellm/proxy/client/__init__.py index 370585728b8..dff3dc226ab 100644 --- a/litellm/proxy/client/__init__.py +++ b/litellm/proxy/client/__init__.py @@ -1,17 +1,17 @@ -from .client import Client from .chat import ChatClient -from .models import ModelsManagementClient -from .model_groups import ModelGroupsManagementClient +from .client import Client from .exceptions import UnauthorizedError -from .users import UsersManagementClient from .health import HealthManagementClient +from .model_groups import ModelGroupsManagementClient +from .models import ModelsManagementClient +from .users import UsersManagementClient __all__ = [ - "Client", "ChatClient", - "ModelsManagementClient", - "ModelGroupsManagementClient", - "UsersManagementClient", - "UnauthorizedError", + "Client", "HealthManagementClient", + "ModelGroupsManagementClient", + "ModelsManagementClient", + "UnauthorizedError", + "UsersManagementClient", ] diff --git a/litellm/proxy/client/chat.py b/litellm/proxy/client/chat.py index 0bc42685b88..3739354ae3b 100644 --- a/litellm/proxy/client/chat.py +++ b/litellm/proxy/client/chat.py @@ -1,6 +1,6 @@ import json from collections.abc import Iterator -from typing import Any, Dict, List, Optional, Union +from typing import Any import requests @@ -8,7 +8,7 @@ from .exceptions import UnauthorizedError class ChatClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the ChatClient. @@ -19,7 +19,7 @@ class ChatClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -34,16 +34,16 @@ class ChatClient: def completions( self, model: str, - messages: List[Dict[str, str]], - temperature: Optional[float] = None, - top_p: Optional[float] = None, - n: Optional[int] = None, - max_tokens: Optional[int] = None, - presence_penalty: Optional[float] = None, - frequency_penalty: Optional[float] = None, - user: Optional[str] = None, + messages: list[dict[str, str]], + temperature: float | None = None, + top_p: float | None = None, + n: int | None = None, + max_tokens: int | None = None, + presence_penalty: float | None = None, + frequency_penalty: float | None = None, + user: str | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Create a chat completion. @@ -70,7 +70,7 @@ class ChatClient: url = f"{self._base_url}/chat/completions" # Build request data with required fields - data: Dict[str, Any] = {"model": model, "messages": messages} + data: dict[str, Any] = {"model": model, "messages": messages} # Add optional parameters if provided if temperature is not None: @@ -107,15 +107,15 @@ class ChatClient: def completions_stream( self, model: str, - messages: List[Dict[str, str]], - temperature: Optional[float] = None, - top_p: Optional[float] = None, - n: Optional[int] = None, - max_tokens: Optional[int] = None, - presence_penalty: Optional[float] = None, - frequency_penalty: Optional[float] = None, - user: Optional[str] = None, - ) -> Iterator[Dict[str, Any]]: + messages: list[dict[str, str]], + temperature: float | None = None, + top_p: float | None = None, + n: int | None = None, + max_tokens: int | None = None, + presence_penalty: float | None = None, + frequency_penalty: float | None = None, + user: str | None = None, + ) -> Iterator[dict[str, Any]]: """ Create a streaming chat completion. @@ -140,7 +140,7 @@ class ChatClient: url = f"{self._base_url}/chat/completions" # Build request data with required fields - data: Dict[str, Any] = {"model": model, "messages": messages, "stream": True} + data: dict[str, Any] = {"model": model, "messages": messages, "stream": True} # Add optional parameters if provided if temperature is not None: diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py index e055b294855..3d34d27e82a 100644 --- a/litellm/proxy/client/cli/commands/agents.py +++ b/litellm/proxy/client/cli/commands/agents.py @@ -2,7 +2,6 @@ import os import shutil import sys from collections.abc import Callable, Mapping, Sequence -from typing import Dict, FrozenSet, List, Optional, Tuple import click import requests @@ -18,13 +17,13 @@ OPENAI_API_KEY_ENV = "OPENAI_API_KEY" PROFILE_ANTHROPIC = "anthropic" PROFILE_OPENAI = "openai" -_KNOWN_AGENTS: Dict[str, Tuple[str, FrozenSet[str]]] = { +_KNOWN_AGENTS: dict[str, tuple[str, frozenset[str]]] = { "claude": ("Claude Code", frozenset({PROFILE_ANTHROPIC})), "codex": ("Codex", frozenset({PROFILE_OPENAI})), "opencode": ("OpenCode", frozenset({PROFILE_OPENAI})), } -_INSTALL_DOCS: Dict[str, str] = { +_INSTALL_DOCS: dict[str, str] = { "claude": "https://docs.claude.com/en/docs/claude-code/setup", "codex": "https://developers.openai.com/codex/cli", "opencode": "https://opencode.ai/docs", @@ -37,7 +36,7 @@ class AgentRunError(Exception): """Raised for any user-actionable failure while preparing to run an agent.""" -def agent_profile(command: str) -> Tuple[str, FrozenSet[str]]: +def agent_profile(command: str) -> tuple[str, frozenset[str]]: """Return the (display name, env profiles) for a wrapped command. Known agents map to the API family they speak. Anything else gets both @@ -53,8 +52,8 @@ def build_agent_env( base_env: Mapping[str, str], base_url: str, api_key: str, - profiles: FrozenSet[str], -) -> Dict[str, str]: + profiles: frozenset[str], +) -> dict[str, str]: """Return a copy of base_env wired to route the agent through the proxy. Anthropic clients (Claude Code) append /v1/messages to ANTHROPIC_BASE_URL, @@ -74,7 +73,7 @@ def build_agent_env( return env -def _codex_proxy_args(base_url: str) -> List[str]: +def _codex_proxy_args(base_url: str) -> list[str]: """Codex `-c` overrides that point it at the proxy. Codex ignores OPENAI_BASE_URL (it always dials api.openai.com), so the env @@ -101,12 +100,12 @@ def _codex_proxy_args(base_url: str) -> List[str]: ] -_PROXY_ARGS: Dict[str, Callable[[str], List[str]]] = { +_PROXY_ARGS: dict[str, Callable[[str], list[str]]] = { "codex": _codex_proxy_args, } -def agent_launch_args(command: str, base_url: str) -> List[str]: +def agent_launch_args(command: str, base_url: str) -> list[str]: """Extra CLI args an agent needs to actually honor the proxy. Claude Code and OpenCode respect the exported env vars, so they get nothing @@ -172,11 +171,11 @@ def run_agent( command: Sequence[str], *, skip_verify: bool = False, - base_env: Optional[Mapping[str, str]] = None, - which: Callable[[str], Optional[str]] = shutil.which, + base_env: Mapping[str, str] | None = None, + which: Callable[[str], str | None] = shutil.which, verify: Callable[[str, str], None] = verify_proxy_key, launcher: Callable[[str, Sequence[str], Mapping[str, str]], None] = _exec, - reattach_terminal: Optional[Callable[[], None]] = None, + reattach_terminal: Callable[[], None] | None = None, ) -> None: """Validate, wire the environment, and hand off to the agent. @@ -277,18 +276,18 @@ def _make_agent_command(binary: str, display_name: str) -> click.Command: return _command -def agent_commands() -> List[click.Command]: +def agent_commands() -> list[click.Command]: """Build one top-level command per known agent, e.g. `lite claude`.""" return [_make_agent_command(binary, name) for binary, (name, _profiles) in _KNOWN_AGENTS.items()] __all__ = [ - "agent_commands", - "run_agent", - "build_agent_env", - "agent_launch_args", - "verify_proxy_key", - "agent_profile", - "resolve_api_key", "AgentRunError", + "agent_commands", + "agent_launch_args", + "agent_profile", + "build_agent_env", + "resolve_api_key", + "run_agent", + "verify_proxy_key", ] diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 970d801dc6d..3179435db35 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -4,7 +4,7 @@ import sys import time import webbrowser from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any from urllib.parse import urlencode import click @@ -27,12 +27,12 @@ def get_token_file_path() -> str: return str(config_dir / "token.json") -def save_token(token_data: Dict[str, Any]) -> None: +def save_token(token_data: dict[str, Any]) -> None: """Save token data to file""" write_private_json(get_token_file_path(), token_data) -def load_token() -> Optional[Dict[str, Any]]: +def load_token() -> dict[str, Any] | None: """Load token data from file""" token_file = get_token_file_path() if not os.path.exists(token_file): @@ -41,7 +41,7 @@ def load_token() -> Optional[Dict[str, Any]]: try: with open(token_file, "r") as f: return json.load(f) - except (json.JSONDecodeError, IOError): + except (OSError, json.JSONDecodeError): return None @@ -52,7 +52,7 @@ def clear_token() -> None: os.remove(token_file) -def get_stored_api_key(expected_base_url: Optional[str] = None) -> Optional[str]: +def get_stored_api_key(expected_base_url: str | None = None) -> str | None: """Get the stored API key from token file. If expected_base_url is provided, the key is only returned when it was @@ -65,7 +65,7 @@ def get_stored_api_key(expected_base_url: Optional[str] = None) -> Optional[str] # Team selection utilities -def display_teams_table(teams: List[Dict[str, Any]]) -> None: +def display_teams_table(teams: list[dict[str, Any]]) -> None: """Display teams in a formatted table""" console = Console() @@ -153,7 +153,7 @@ def get_key_input(): return None -def display_interactive_team_selection(teams: List[Dict[str, Any]], selected_index: int = 0) -> None: +def display_interactive_team_selection(teams: list[dict[str, Any]], selected_index: int = 0) -> None: """Display teams with one highlighted for selection""" console = Console() @@ -191,7 +191,7 @@ def display_interactive_team_selection(teams: List[Dict[str, Any]], selected_ind console.print(f" Budget: [dim]{budget_str}[/dim]\n") -def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]: +def prompt_team_selection(teams: list[dict[str, Any]]) -> dict[str, Any] | None: """Interactive team selection with arrow keys""" if not teams: return None @@ -241,8 +241,8 @@ def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any def prompt_team_selection_fallback( - teams: List[Dict[str, Any]], -) -> Optional[Dict[str, Any]]: + teams: list[dict[str, Any]], +) -> dict[str, Any] | None: """Fallback team selection for non-interactive environments""" if not teams: return None @@ -299,20 +299,20 @@ def _is_permanent_polling_error(status_code: int) -> bool: def _poll_for_ready_data( url: str, *, - headers: Optional[Dict[str, str]] = None, + headers: dict[str, str] | None = None, total_timeout: int = 300, poll_interval: int = 2, request_timeout: int = 10, - pending_message: Optional[str] = None, + pending_message: str | None = None, pending_log_every: int = 10, - other_status_message: Optional[str] = None, + other_status_message: str | None = None, other_status_log_every: int = 10, http_error_log_every: int = 10, connection_error_log_every: int = 10, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: for attempt in range(total_timeout // poll_interval): try: - request_kwargs: Dict[str, Any] = {"timeout": request_timeout} + request_kwargs: dict[str, Any] = {"timeout": request_timeout} if headers is not None: request_kwargs["headers"] = headers response = requests.get(url, **request_kwargs) @@ -365,7 +365,7 @@ def _normalize_teams(teams, team_details): return [] -def _start_cli_sso_flow(base_url: str) -> Dict[str, Any]: +def _start_cli_sso_flow(base_url: str) -> dict[str, Any]: start_url = f"{base_url}/sso/cli/start" try: response = requests.post(start_url, timeout=10) @@ -408,11 +408,11 @@ def _start_cli_sso_flow(base_url: str) -> Dict[str, Any]: return data -def _get_cli_sso_poll_headers(poll_secret: str) -> Dict[str, str]: +def _get_cli_sso_poll_headers(poll_secret: str) -> dict[str, str]: return {"x-litellm-cli-poll-secret": poll_secret} -def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Optional[dict]: +def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> dict | None: """ Poll the server for authentication completion and handle team selection. @@ -431,7 +431,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Op teams = data.get("teams", []) team_details = data.get("team_details") user_id = data.get("user_id") - normalized_teams: List[Dict[str, Any]] = _normalize_teams(teams, team_details) + normalized_teams: list[dict[str, Any]] = _normalize_teams(teams, team_details) if not normalized_teams: click.echo("Warning: No teams available for selection.") return None @@ -478,8 +478,8 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Op def _handle_team_selection_during_polling( - base_url: str, key_id: str, poll_secret: str, teams: List[Dict[str, Any]] -) -> Optional[str]: + base_url: str, key_id: str, poll_secret: str, teams: list[dict[str, Any]] +) -> str | None: """ Handle team selection and re-poll with selected team_id. @@ -522,7 +522,7 @@ def _handle_team_selection_during_polling( return None -def _render_and_prompt_for_team_selection(teams: List[Dict[str, Any]]) -> Optional[str]: +def _render_and_prompt_for_team_selection(teams: list[dict[str, Any]]) -> str | None: """Render teams table and prompt user for a team selection. Returns the selected team_id as a string, or None if selection was @@ -725,7 +725,7 @@ auth_group.add_command(print_token) # Export functions for use by other CLI commands -__all__ = ["login", "logout", "print_token", "auth_group", "whoami", "prompt_team_selection"] +__all__ = ["auth_group", "login", "logout", "print_token", "prompt_team_selection", "whoami"] # Export individual commands instead of grouping them # login, logout, and whoami will be added as top-level commands diff --git a/litellm/proxy/client/cli/commands/autoroute/config.py b/litellm/proxy/client/cli/commands/autoroute/config.py index 237705564ff..21afb472f85 100644 --- a/litellm/proxy/client/cli/commands/autoroute/config.py +++ b/litellm/proxy/client/cli/commands/autoroute/config.py @@ -230,11 +230,11 @@ def master_key_from_config(config: dict[str, JsonValue]) -> str | None: __all__ = [ "AUTOROUTER_MODEL_NAME", + "DEFAULT_KEYWORD_TIER_RULES", "TIER_NAMES", "AutorouteConfig", "ClassifierChoice", "ConfigGenerationError", - "DEFAULT_KEYWORD_TIER_RULES", "DiscoveredModel", "HeuristicClassifier", "KeywordTierRule", diff --git a/litellm/proxy/client/cli/commands/chat.py b/litellm/proxy/client/cli/commands/chat.py index d78feb84bd6..f0c91be686d 100644 --- a/litellm/proxy/client/cli/commands/chat.py +++ b/litellm/proxy/client/cli/commands/chat.py @@ -1,6 +1,6 @@ import json import sys -from typing import Any, Dict, List, Optional +from typing import Any import click import requests @@ -13,7 +13,7 @@ from ... import Client from ...chat import ChatClient -def _get_available_models(ctx: click.Context) -> List[Dict[str, Any]]: +def _get_available_models(ctx: click.Context) -> list[dict[str, Any]]: """Get list of available models from the proxy server""" try: client = Client(base_url=ctx.obj["base_url"], api_key=ctx.obj["api_key"]) @@ -28,7 +28,7 @@ def _get_available_models(ctx: click.Context) -> List[Dict[str, Any]]: return [] -def _select_model(console: Console, available_models: List[Dict[str, Any]]) -> Optional[str]: +def _select_model(console: Console, available_models: list[dict[str, Any]]) -> str | None: """Interactive model selection""" if not available_models: console.print("[yellow]No models available or could not fetch models list.[/yellow]") @@ -42,7 +42,7 @@ def _select_model(console: Console, available_models: List[Dict[str, Any]]) -> O table.add_column("Owned By", style="yellow") MAX_MODELS_TO_DISPLAY = 200 - models_to_display: List[Dict[str, Any]] = available_models[:MAX_MODELS_TO_DISPLAY] + models_to_display: list[dict[str, Any]] = available_models[:MAX_MODELS_TO_DISPLAY] for i, model in enumerate(models_to_display): # Limit to first 200 models table.add_row(str(i + 1), str(model.get("id", "")), str(model.get("owned_by", ""))) @@ -104,10 +104,10 @@ def _select_model(console: Console, available_models: List[Dict[str, Any]]) -> O @click.pass_context def chat( ctx: click.Context, - model: Optional[str], + model: str | None, temperature: float, - max_tokens: Optional[int] = None, - system: Optional[str] = None, + max_tokens: int | None = None, + system: str | None = None, ): """Interactive chat with streaming responses @@ -135,7 +135,7 @@ def chat( client = ChatClient(ctx.obj["base_url"], ctx.obj["api_key"]) # Initialize conversation history - messages: List[Dict[str, Any]] = [] + messages: list[dict[str, Any]] = [] # Add system message if provided if system: @@ -238,7 +238,7 @@ def _show_help(console: Console): console.print(Panel(help_text, title="Help")) -def _show_history(console: Console, messages: List[Dict[str, Any]]): +def _show_history(console: Console, messages: list[dict[str, Any]]): """Show conversation history""" if not messages: console.print("[yellow]No conversation history.[/yellow]") @@ -260,7 +260,7 @@ def _show_history(console: Console, messages: List[Dict[str, Any]]): ) -def _save_conversation(console: Console, messages: List[Dict[str, Any]], command: str): +def _save_conversation(console: Console, messages: list[dict[str, Any]], command: str): """Save conversation to a file""" parts = command.split() if len(parts) < 2: @@ -279,7 +279,7 @@ def _save_conversation(console: Console, messages: List[Dict[str, Any]], command console.print(f"[red]Error saving conversation: {e}[/red]") -def _load_conversation(console: Console, command: str, system: Optional[str]) -> List[Dict[str, Any]]: +def _load_conversation(console: Console, command: str, system: str | None) -> list[dict[str, Any]]: """Load conversation from a file""" parts = command.split() if len(parts) < 2: @@ -309,10 +309,10 @@ def _load_conversation(console: Console, command: str, system: Optional[str]) -> def _handle_special_commands( console: Console, user_input: str, - messages: List[Dict[str, Any]], - system: Optional[str], + messages: list[dict[str, Any]], + system: str | None, ctx: click.Context, -) -> tuple[bool, List[Dict[str, Any]], Optional[str]]: +) -> tuple[bool, list[dict[str, Any]], str | None]: """Handle special chat commands. Returns (should_exit, updated_messages, updated_model)""" if user_input.lower() in ["/quit", "/exit", "/q"]: console.print("[yellow]Chat session ended.[/yellow]") @@ -353,10 +353,10 @@ def _stream_response( console: Console, client: ChatClient, model: str, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], temperature: float, - max_tokens: Optional[int], -) -> Optional[str]: + max_tokens: int | None, +) -> str | None: """Stream the model response and return the complete content""" try: assistant_content = "" @@ -386,5 +386,5 @@ def _stream_response( console.print(f"[red]{e.response.text}[/red]") return None except Exception as e: - console.print(f"\n[red]Error: {str(e)}[/red]") + console.print(f"\n[red]Error: {e!s}[/red]") return None diff --git a/litellm/proxy/client/cli/commands/credentials.py b/litellm/proxy/client/cli/commands/credentials.py index 44f4a112ce4..8187f811778 100644 --- a/litellm/proxy/client/cli/commands/credentials.py +++ b/litellm/proxy/client/cli/commands/credentials.py @@ -2,8 +2,8 @@ import json from typing import Literal import click -import rich import requests +import rich from rich.table import Table from ...credentials import CredentialsManagementClient @@ -12,7 +12,6 @@ from ...credentials import CredentialsManagementClient @click.group() def credentials(): """Manage credentials for the LiteLLM proxy server""" - pass @credentials.command() @@ -72,7 +71,7 @@ def create(ctx: click.Context, credential_name: str, info: str, values: str): credential_info = json.loads(info) credential_values = json.loads(values) except json.JSONDecodeError as e: - raise click.BadParameter(f"Invalid JSON: {str(e)}") + raise click.BadParameter(f"Invalid JSON: {e!s}") try: response = client.create(credential_name, credential_info, credential_values) diff --git a/litellm/proxy/client/cli/commands/encryption.py b/litellm/proxy/client/cli/commands/encryption.py index 4b460bac19c..a9bed6c5ba6 100644 --- a/litellm/proxy/client/cli/commands/encryption.py +++ b/litellm/proxy/client/cli/commands/encryption.py @@ -9,7 +9,6 @@ from ...http_client import HTTPClient @click.group() def encryption(): """Migrate at-rest credentials to AES-256-GCM and attest residual state.""" - pass @encryption.command(name="migrate") diff --git a/litellm/proxy/client/cli/commands/http.py b/litellm/proxy/client/cli/commands/http.py index dba36f9d92c..96bc94e8131 100644 --- a/litellm/proxy/client/cli/commands/http.py +++ b/litellm/proxy/client/cli/commands/http.py @@ -1,9 +1,8 @@ import json as json_lib -from typing import Optional import click -import rich import requests +import rich from ...http_client import HTTPClient @@ -11,7 +10,6 @@ from ...http_client import HTTPClient @click.group() def http(): """Make HTTP requests to the LiteLLM proxy server""" - pass @http.command() @@ -40,8 +38,8 @@ def request( ctx: click.Context, method: str, uri: str, - data: Optional[str] = None, - json: Optional[str] = None, + data: str | None = None, + json: str | None = None, header: tuple[str, ...] = (), ): """Make an HTTP request to the LiteLLM proxy server diff --git a/litellm/proxy/client/cli/commands/keys.py b/litellm/proxy/client/cli/commands/keys.py index afbaa3702c1..ec5dca25518 100644 --- a/litellm/proxy/client/cli/commands/keys.py +++ b/litellm/proxy/client/cli/commands/keys.py @@ -1,10 +1,11 @@ +import builtins import json from datetime import datetime -from typing import Literal, Optional, List, Dict, Any +from typing import Any, Literal import click -import rich import requests +import rich from rich.table import Table from ...keys import KeysManagementClient @@ -13,7 +14,6 @@ from ...keys import KeysManagementClient @click.group() def keys(): """Manage API keys for the LiteLLM proxy server""" - pass @keys.command() @@ -41,13 +41,13 @@ def keys(): @click.pass_context def list( ctx: click.Context, - page: Optional[int], - size: Optional[int], - user_id: Optional[str], - team_id: Optional[str], - organization_id: Optional[str], - key_hash: Optional[str], - key_alias: Optional[str], + page: int | None, + size: int | None, + user_id: str | None, + team_id: str | None, + organization_id: str | None, + key_hash: str | None, + key_alias: str | None, include_team_keys: bool, output_format: Literal["table", "json"], return_full_object: bool, @@ -105,15 +105,15 @@ def list( @click.pass_context def generate( ctx: click.Context, - models: Optional[str], - aliases: Optional[str], - spend: Optional[float], - duration: Optional[str], - key_alias: Optional[str], - team_id: Optional[str], - user_id: Optional[str], - budget_id: Optional[str], - config: Optional[str], + models: str | None, + aliases: str | None, + spend: float | None, + duration: str | None, + key_alias: str | None, + team_id: str | None, + user_id: str | None, + budget_id: str | None, + config: str | None, ): """Generate a new API key""" client = KeysManagementClient(ctx.obj["base_url"], ctx.obj["api_key"]) @@ -122,7 +122,7 @@ def generate( aliases_dict = json.loads(aliases) if aliases else None config_dict = json.loads(config) if config else None except json.JSONDecodeError as e: - raise click.BadParameter(f"Invalid JSON: {str(e)}") + raise click.BadParameter(f"Invalid JSON: {e!s}") try: response = client.generate( models=models_list, @@ -150,7 +150,7 @@ def generate( @click.option("--keys", type=str, help="Comma-separated list of API keys to delete") @click.option("--key-aliases", type=str, help="Comma-separated list of key aliases to delete") @click.pass_context -def delete(ctx: click.Context, keys: Optional[str], key_aliases: Optional[str]): +def delete(ctx: click.Context, keys: str | None, key_aliases: str | None): """Delete API keys by key or alias""" client = KeysManagementClient(ctx.obj["base_url"], ctx.obj["api_key"]) keys_list = [k.strip() for k in keys.split(",")] if keys else None @@ -168,7 +168,7 @@ def delete(ctx: click.Context, keys: Optional[str], key_aliases: Optional[str]): raise click.Abort() -def _parse_created_since_filter(created_since: Optional[str]) -> Optional[datetime]: +def _parse_created_since_filter(created_since: str | None) -> datetime | None: """Parse and validate the created_since date filter.""" if not created_since: return None @@ -187,7 +187,9 @@ def _parse_created_since_filter(created_since: Optional[str]) -> Optional[dateti raise click.Abort() -def _fetch_all_keys_with_pagination(source_client: KeysManagementClient, source_base_url: str) -> List[Dict[str, Any]]: +def _fetch_all_keys_with_pagination( + source_client: KeysManagementClient, source_base_url: str +) -> builtins.list[dict[str, Any]]: """Fetch all keys from source instance using pagination.""" click.echo(f"Fetching keys from source server: {source_base_url}") source_keys = [] @@ -216,10 +218,10 @@ def _fetch_all_keys_with_pagination(source_client: KeysManagementClient, source_ def _filter_keys_by_created_since( - source_keys: List[Dict[str, Any]], - created_since_dt: Optional[datetime], + source_keys: builtins.list[dict[str, Any]], + created_since_dt: datetime | None, created_since: str, -) -> List[Dict[str, Any]]: +) -> builtins.list[dict[str, Any]]: """Filter keys by created_since date if specified.""" if not created_since_dt: return source_keys @@ -246,7 +248,7 @@ def _filter_keys_by_created_since( return filtered_keys -def _display_dry_run_table(source_keys: List[Dict[str, Any]]) -> None: +def _display_dry_run_table(source_keys: builtins.list[dict[str, Any]]) -> None: """Display a table of keys that would be imported in dry-run mode.""" click.echo("\n--- DRY RUN MODE ---") table = Table(title="Keys that would be imported") @@ -269,7 +271,7 @@ def _display_dry_run_table(source_keys: List[Dict[str, Any]]) -> None: rich.print(table) -def _prepare_key_import_data(key: Dict[str, Any]) -> Dict[str, Any]: +def _prepare_key_import_data(key: dict[str, Any]) -> dict[str, Any]: """Prepare key data for import by extracting relevant fields.""" import_data = {} @@ -291,7 +293,7 @@ def _prepare_key_import_data(key: Dict[str, Any]) -> Dict[str, Any]: def _import_keys_to_destination( - source_keys: List[Dict[str, Any]], dest_client: KeysManagementClient + source_keys: builtins.list[dict[str, Any]], dest_client: KeysManagementClient ) -> tuple[int, int]: """Import each key to the destination instance and return counts.""" imported_count = 0 @@ -314,7 +316,7 @@ def _import_keys_to_destination( except Exception as e: failed_count += 1 key_alias = key.get("key_alias", "N/A") - click.echo(f"Failed to import key {key_alias}: {str(e)}", err=True) + click.echo(f"Failed to import key {key_alias}: {e!s}", err=True) return imported_count, failed_count @@ -339,9 +341,9 @@ def _import_keys_to_destination( def import_keys( ctx: click.Context, source_base_url: str, - source_api_key: Optional[str], + source_api_key: str | None, dry_run: bool, - created_since: Optional[str], + created_since: str | None, ): """Import API keys from another LiteLLM instance""" # Parse created_since filter if provided @@ -387,5 +389,5 @@ def import_keys( click.echo(e.response.text, err=True) raise click.Abort() except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) raise click.Abort() diff --git a/litellm/proxy/client/cli/commands/models.py b/litellm/proxy/client/cli/commands/models.py index 15266488c84..3982b1f7b17 100644 --- a/litellm/proxy/client/cli/commands/models.py +++ b/litellm/proxy/client/cli/commands/models.py @@ -1,14 +1,14 @@ # stdlib imports -from datetime import datetime import re -from typing import Optional, Literal, Any -import yaml -from dataclasses import dataclass from collections import defaultdict +from dataclasses import dataclass +from datetime import datetime +from typing import Any, Literal # third party imports import click import rich +import yaml # local imports from ... import Client @@ -46,7 +46,7 @@ def _get_model_info_obj_from_yaml(model: dict[str, Any]) -> ModelYamlInfo: ) -def format_iso_datetime_str(iso_datetime_str: Optional[str]) -> str: +def format_iso_datetime_str(iso_datetime_str: str | None) -> str: """Format an ISO format datetime string to human-readable date with minute resolution.""" if not iso_datetime_str: return "" @@ -58,7 +58,7 @@ def format_iso_datetime_str(iso_datetime_str: Optional[str]) -> str: return str(iso_datetime_str) -def format_timestamp(timestamp: Optional[int]) -> str: +def format_timestamp(timestamp: int | None) -> str: """Format a Unix timestamp (integer) to human-readable date with minute resolution.""" if timestamp is None: return "" @@ -69,7 +69,7 @@ def format_timestamp(timestamp: Optional[int]) -> str: return str(timestamp) -def format_cost_per_1k_tokens(cost: Optional[float]) -> str: +def format_cost_per_1k_tokens(cost: float | None) -> str: """Format a per-token cost to cost per 1000 tokens.""" if cost is None: return "" @@ -90,7 +90,6 @@ def create_client(ctx: click.Context) -> Client: @click.group() def models() -> None: """Manage models on your LiteLLM proxy server""" - pass @models.command("list") @@ -180,7 +179,7 @@ def delete_model(ctx: click.Context, model_id: str) -> None: @click.option("--id", "model_id", help="ID of the model to retrieve") @click.option("--name", "model_name", help="Name of the model to retrieve") @click.pass_context -def get_model(ctx: click.Context, model_id: Optional[str], model_name: Optional[str]) -> None: +def get_model(ctx: click.Context, model_id: str | None, model_name: str | None) -> None: """Get information about a specific model""" if not model_id and not model_name: raise click.UsageError("Either --id or --name must be provided") @@ -411,8 +410,8 @@ def import_models( ctx: click.Context, yaml_file: str, dry_run: bool, - only_models_matching_regex: Optional[str], - only_access_groups_matching_regex: Optional[str], + only_models_matching_regex: str | None, + only_access_groups_matching_regex: str | None, ) -> None: """Import models from a YAML file and add them to the proxy.""" provider_counts: dict[str, int] = defaultdict(int) diff --git a/litellm/proxy/client/cli/commands/teams.py b/litellm/proxy/client/cli/commands/teams.py index b0ccdc8f9bf..442ac40a775 100644 --- a/litellm/proxy/client/cli/commands/teams.py +++ b/litellm/proxy/client/cli/commands/teams.py @@ -1,6 +1,6 @@ """Team management commands for LiteLLM CLI.""" -from typing import Any, Dict, List, Optional +from typing import Any import click import requests @@ -13,10 +13,9 @@ from litellm.proxy.client import Client @click.group() def teams(): """Manage teams and team assignments""" - pass -def display_teams_table(teams: List[Dict[str, Any]]) -> None: +def display_teams_table(teams: list[dict[str, Any]]) -> None: """Display teams in a formatted table""" console = Console() @@ -77,7 +76,7 @@ def list(ctx: click.Context): click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True) raise click.Abort() except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) raise click.Abort() @@ -100,14 +99,14 @@ def available(ctx: click.Context): error_body = e.response.json() click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True) except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) raise click.Abort() @teams.command() @click.option("--team-id", type=str, help="Team ID to assign the key to") @click.pass_context -def assign_key(ctx: click.Context, team_id: Optional[str]): +def assign_key(ctx: click.Context, team_id: str | None): """Assign your current CLI key to a team""" client = Client(ctx.obj["base_url"], ctx.obj["api_key"]) api_key = ctx.obj["api_key"] @@ -159,5 +158,5 @@ def assign_key(ctx: click.Context, team_id: Optional[str]): click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True) raise click.Abort() except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) raise click.Abort() diff --git a/litellm/proxy/client/cli/commands/users.py b/litellm/proxy/client/cli/commands/users.py index 1671b92c089..f4a7c4d82fe 100644 --- a/litellm/proxy/client/cli/commands/users.py +++ b/litellm/proxy/client/cli/commands/users.py @@ -1,12 +1,12 @@ import click import rich + from ... import UsersManagementClient @click.group() def users(): """Manage users on your LiteLLM proxy server""" - pass @users.command("list") @@ -20,8 +20,8 @@ def list_users(ctx: click.Context): if not users: click.echo("No users found.") return - from rich.table import Table from rich.console import Console + from rich.table import Table table = Table(title="Users") table.add_column("User ID", style="cyan") diff --git a/litellm/proxy/client/cli/interface.py b/litellm/proxy/client/cli/interface.py index e953742f412..9f42c3b06bb 100644 --- a/litellm/proxy/client/cli/interface.py +++ b/litellm/proxy/client/cli/interface.py @@ -1,16 +1,12 @@ # stdlib imports import os import sys -from typing import TYPE_CHECKING # third party imports import click from litellm._logging import verbose_logger -if TYPE_CHECKING: - pass - def styled_prompt(): """Create a styled blue box prompt for user input.""" diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 24e5cdf747b..153666eaac5 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -1,5 +1,4 @@ # stdlib imports -from typing import Optional # third party imports import click @@ -26,7 +25,7 @@ from .commands.users import users from .interface import interactive_shell -def print_version(base_url: str, api_key: Optional[str]): +def print_version(base_url: str, api_key: str | None): """Print CLI and server version info.""" click.echo(f"LiteLLM Proxy CLI Version: {litellm_version}") if base_url: @@ -65,7 +64,7 @@ def print_version(base_url: str, api_key: Optional[str]): help="API key for authentication", ) @click.pass_context -def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None: +def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: str | None) -> None: """LiteLLM Proxy CLI - Manage your LiteLLM proxy server""" ctx.ensure_object(dict) diff --git a/litellm/proxy/client/client.py b/litellm/proxy/client/client.py index f481a61c328..d71802e06c8 100644 --- a/litellm/proxy/client/client.py +++ b/litellm/proxy/client/client.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key from .chat import ChatClient @@ -17,7 +15,7 @@ class Client: def __init__( self, base_url: str, - api_key: Optional[str] = None, + api_key: str | None = None, timeout: int = 30, ): """ diff --git a/litellm/proxy/client/credentials.py b/litellm/proxy/client/credentials.py index e0b41a91326..1364e816b47 100644 --- a/litellm/proxy/client/credentials.py +++ b/litellm/proxy/client/credentials.py @@ -1,11 +1,12 @@ +from typing import Any + import requests -from typing import Dict, Any, Optional, Union from .exceptions import UnauthorizedError class CredentialsManagementClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the CredentialsManagementClient. @@ -16,7 +17,7 @@ class CredentialsManagementClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -31,7 +32,7 @@ class CredentialsManagementClient: def list( self, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ List all credentials. @@ -66,10 +67,10 @@ class CredentialsManagementClient: def create( self, credential_name: str, - credential_info: Dict[str, Any], - credential_values: Dict[str, Any], + credential_info: dict[str, Any], + credential_values: dict[str, Any], return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Create a new credential. @@ -114,7 +115,7 @@ class CredentialsManagementClient: self, credential_name: str, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Delete a credential by name. @@ -151,7 +152,7 @@ class CredentialsManagementClient: self, credential_name: str, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Get a credential by name. diff --git a/litellm/proxy/client/exceptions.py b/litellm/proxy/client/exceptions.py index 34884c0b995..27de7031e3c 100644 --- a/litellm/proxy/client/exceptions.py +++ b/litellm/proxy/client/exceptions.py @@ -1,13 +1,11 @@ -from typing import Union - import requests from litellm.litellm_core_utils.secret_redaction import redact_string def _redact_orig_exception( - orig_exception: Union[requests.exceptions.HTTPError, str], -) -> Union[requests.exceptions.HTTPError, str]: + orig_exception: requests.exceptions.HTTPError | str, +) -> requests.exceptions.HTTPError | str: if isinstance(orig_exception, requests.exceptions.HTTPError): return requests.exceptions.HTTPError(redact_string(str(orig_exception)), response=orig_exception.response) return redact_string(str(orig_exception)) @@ -16,7 +14,7 @@ def _redact_orig_exception( class UnauthorizedError(Exception): """Exception raised when the API returns a 401 Unauthorized response.""" - def __init__(self, orig_exception: Union[requests.exceptions.HTTPError, str]): + def __init__(self, orig_exception: requests.exceptions.HTTPError | str): self.orig_exception = _redact_orig_exception(orig_exception) super().__init__(str(self.orig_exception)) @@ -24,6 +22,6 @@ class UnauthorizedError(Exception): class NotFoundError(Exception): """Exception raised when the API returns a 404 Not Found response or indicates a resource was not found.""" - def __init__(self, orig_exception: Union[requests.exceptions.HTTPError, str]): + def __init__(self, orig_exception: requests.exceptions.HTTPError | str): self.orig_exception = _redact_orig_exception(orig_exception) super().__init__(str(self.orig_exception)) diff --git a/litellm/proxy/client/health.py b/litellm/proxy/client/health.py index 3cfcd151d6b..85b6767b327 100644 --- a/litellm/proxy/client/health.py +++ b/litellm/proxy/client/health.py @@ -1,4 +1,5 @@ -from typing import Optional, Dict, Any +from typing import Any + from .http_client import HTTPClient @@ -7,7 +8,7 @@ class HealthManagementClient: Client for interacting with the health endpoints of the LiteLLM proxy server. """ - def __init__(self, base_url: str, api_key: Optional[str] = None, timeout: int = 30): + def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 30): """ Initialize the HealthManagementClient. @@ -18,7 +19,7 @@ class HealthManagementClient: """ self._http = HTTPClient(base_url=base_url, api_key=api_key, timeout=timeout) - def get_readiness(self) -> Dict[str, Any]: + def get_readiness(self) -> dict[str, Any]: """ Check the readiness of the LiteLLM proxy server. @@ -31,7 +32,7 @@ class HealthManagementClient: """ return self._http.request("GET", "/health/readiness") - def get_server_version(self) -> Optional[str]: + def get_server_version(self) -> str | None: """ Get the LiteLLM server version from the readiness endpoint. diff --git a/litellm/proxy/client/http_client.py b/litellm/proxy/client/http_client.py index 4357f6e35b5..3d714678426 100644 --- a/litellm/proxy/client/http_client.py +++ b/litellm/proxy/client/http_client.py @@ -1,13 +1,14 @@ """HTTP client for making requests to the LiteLLM proxy server.""" -from typing import Any, Dict, Optional, Union +from typing import Any + import requests class HTTPClient: """HTTP client for making requests to the LiteLLM proxy server.""" - def __init__(self, base_url: str, api_key: Optional[str] = None, timeout: int = 30): + def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 30): """Initialize the HTTP client. Args: @@ -24,9 +25,9 @@ class HTTPClient: method: str, uri: str, *, - data: Optional[Union[Dict[str, Any], list, bytes]] = None, - json: Optional[Union[Dict[str, Any], list]] = None, - headers: Optional[Dict[str, str]] = None, + data: dict[str, Any] | list | bytes | None = None, + json: dict[str, Any] | list | None = None, + headers: dict[str, str] | None = None, **kwargs: Any, ) -> Any: """Make an HTTP request to the LiteLLM proxy server. diff --git a/litellm/proxy/client/keys.py b/litellm/proxy/client/keys.py index dcf915d812a..a2852764728 100644 --- a/litellm/proxy/client/keys.py +++ b/litellm/proxy/client/keys.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, List, Optional, Union +import builtins +from typing import Any import requests @@ -8,7 +9,7 @@ from .exceptions import UnauthorizedError class KeysManagementClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the KeysManagementClient. @@ -19,7 +20,7 @@ class KeysManagementClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -33,17 +34,17 @@ class KeysManagementClient: def list( self, - page: Optional[int] = None, - size: Optional[int] = None, - user_id: Optional[str] = None, - team_id: Optional[str] = None, - organization_id: Optional[str] = None, - key_hash: Optional[str] = None, - key_alias: Optional[str] = None, - return_full_object: Optional[bool] = None, - include_team_keys: Optional[bool] = None, + page: int | None = None, + size: int | None = None, + user_id: str | None = None, + team_id: str | None = None, + organization_id: str | None = None, + key_hash: str | None = None, + key_alias: str | None = None, + return_full_object: bool | None = None, + include_team_keys: bool | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ List all API keys with optional filtering and pagination. @@ -69,7 +70,7 @@ class KeysManagementClient: requests.exceptions.RequestException: If the request fails with any other error """ url = f"{self._base_url}/key/list" - params: Dict[str, Any] = {} + params: dict[str, Any] = {} # Add optional query parameters if page is not None: @@ -108,17 +109,17 @@ class KeysManagementClient: def generate( self, - models: Optional[List[str]] = None, - aliases: Optional[Dict[str, str]] = None, - spend: Optional[float] = None, - duration: Optional[str] = None, - key_alias: Optional[str] = None, - team_id: Optional[str] = None, - user_id: Optional[str] = None, - budget_id: Optional[str] = None, - config: Optional[Dict[str, Any]] = None, + models: builtins.list[str] | None = None, + aliases: dict[str, str] | None = None, + spend: float | None = None, + duration: str | None = None, + key_alias: str | None = None, + team_id: str | None = None, + user_id: str | None = None, + budget_id: str | None = None, + config: dict[str, Any] | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Generate an API key based on the provided data. @@ -146,7 +147,7 @@ class KeysManagementClient: """ url = f"{self._base_url}/key/generate" - data: Dict[str, Any] = {} + data: dict[str, Any] = {} if models is not None: data["models"] = models if aliases is not None: @@ -183,10 +184,10 @@ class KeysManagementClient: def delete( self, - keys: Optional[List[str]] = None, - key_aliases: Optional[List[str]] = None, + keys: builtins.list[str] | None = None, + key_aliases: builtins.list[str] | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Delete existing keys @@ -228,14 +229,14 @@ class KeysManagementClient: def update( self, key: str, - models: Optional[List[str]] = None, - aliases: Optional[Dict[str, str]] = None, - spend: Optional[float] = None, - duration: Optional[str] = None, - key_alias: Optional[str] = None, - team_id: Optional[str] = None, - user_id: Optional[str] = None, - ) -> Union[Dict[str, Any], requests.Request]: + models: builtins.list[str] | None = None, + aliases: dict[str, str] | None = None, + spend: float | None = None, + duration: str | None = None, + key_alias: str | None = None, + team_id: str | None = None, + user_id: str | None = None, + ) -> dict[str, Any] | requests.Request: """ Update an existing API key's parameters. @@ -258,7 +259,7 @@ class KeysManagementClient: """ url = f"{self._base_url}/key/update" - data: Dict[str, Any] = {"key": key} + data: dict[str, Any] = {"key": key} if key_alias is not None: data["key_alias"] = key_alias @@ -276,7 +277,7 @@ class KeysManagementClient: data["aliases"] = aliases request = requests.Request("POST", url, headers=self._get_headers(), json=data) session = requests.Session() - response_text: Optional[str] = None + response_text: str | None = None try: response = session.send(request.prepare()) response_text = response.text @@ -285,7 +286,7 @@ class KeysManagementClient: except Exception: raise Exception(f"Error updating key: {response_text}") - def info(self, key: str, return_request: bool = False) -> Union[Dict[str, Any], requests.Request]: + def info(self, key: str, return_request: bool = False) -> dict[str, Any] | requests.Request: """ Get information about API keys. diff --git a/litellm/proxy/client/model_groups.py b/litellm/proxy/client/model_groups.py index 2be6e10e542..578604642e7 100644 --- a/litellm/proxy/client/model_groups.py +++ b/litellm/proxy/client/model_groups.py @@ -1,10 +1,12 @@ +from typing import Any + import requests -from typing import List, Dict, Any, Optional, Union + from .exceptions import UnauthorizedError class ModelGroupsManagementClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the ModelGroupsManagementClient. @@ -15,7 +17,7 @@ class ModelGroupsManagementClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -27,7 +29,7 @@ class ModelGroupsManagementClient: headers["Authorization"] = f"Bearer {self._api_key}" return headers - def info(self, return_request: bool = False) -> Union[List[Dict[str, Any]], requests.Request]: + def info(self, return_request: bool = False) -> list[dict[str, Any]] | requests.Request: """ Get detailed information about all model groups from the server. diff --git a/litellm/proxy/client/models.py b/litellm/proxy/client/models.py index bb375426295..9e49175ba1d 100644 --- a/litellm/proxy/client/models.py +++ b/litellm/proxy/client/models.py @@ -1,10 +1,13 @@ +import builtins +from typing import Any + import requests -from typing import List, Dict, Any, Optional, Union -from .exceptions import UnauthorizedError, NotFoundError + +from .exceptions import NotFoundError, UnauthorizedError class ModelsManagementClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the ModelsManagementClient. @@ -15,7 +18,7 @@ class ModelsManagementClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -27,7 +30,7 @@ class ModelsManagementClient: headers["Authorization"] = f"Bearer {self._api_key}" return headers - def list(self, return_request: bool = False) -> Union[List[Dict[str, Any]], requests.Request]: + def list(self, return_request: bool = False) -> list[dict[str, Any]] | requests.Request: """ Get the list of models supported by the server. @@ -63,10 +66,10 @@ class ModelsManagementClient: def new( self, model_name: str, - model_params: Dict[str, Any], - model_info: Optional[Dict[str, Any]] = None, + model_params: dict[str, Any], + model_info: dict[str, Any] | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Add a new model to the proxy. @@ -109,7 +112,7 @@ class ModelsManagementClient: raise UnauthorizedError(e) raise - def delete(self, model_id: str, return_request: bool = False) -> Union[Dict[str, Any], requests.Request]: + def delete(self, model_id: str, return_request: bool = False) -> dict[str, Any] | requests.Request: """ Delete a model from the proxy. @@ -149,10 +152,10 @@ class ModelsManagementClient: def get( self, - model_id: Optional[str] = None, - model_name: Optional[str] = None, + model_id: str | None = None, + model_name: str | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Get information about a specific model by its ID or name. @@ -182,7 +185,7 @@ class ModelsManagementClient: # Get all models and filter models = self.info() - assert isinstance(models, List) + assert isinstance(models, list) # Find the matching model for model in models: @@ -205,7 +208,7 @@ class ModelsManagementClient: ) ) - def info(self, return_request: bool = False) -> Union[List[Dict[str, Any]], requests.Request]: + def info(self, return_request: bool = False) -> builtins.list[dict[str, Any]] | requests.Request: """ Get detailed information about all models from the server. @@ -240,10 +243,10 @@ class ModelsManagementClient: def update( self, model_id: str, - model_params: Dict[str, Any], - model_info: Optional[Dict[str, Any]] = None, + model_params: dict[str, Any], + model_info: dict[str, Any] | None = None, return_request: bool = False, - ) -> Union[Dict[str, Any], requests.Request]: + ) -> dict[str, Any] | requests.Request: """ Update an existing model's configuration. diff --git a/litellm/proxy/client/teams.py b/litellm/proxy/client/teams.py index 017d0744857..f48b7360c38 100644 --- a/litellm/proxy/client/teams.py +++ b/litellm/proxy/client/teams.py @@ -1,6 +1,7 @@ """Teams management client for LiteLLM proxy.""" -from typing import Any, Dict, List, Optional, Union +import builtins +from typing import Any import requests @@ -10,7 +11,7 @@ from .exceptions import UnauthorizedError class TeamsManagementClient: """Client for managing teams in LiteLLM proxy.""" - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): """ Initialize the TeamsManagementClient. @@ -21,7 +22,7 @@ class TeamsManagementClient: self._base_url = base_url.rstrip("/") # Remove trailing slash if present self._api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: """ Get the headers for API requests, including authorization if api_key is set. @@ -35,9 +36,9 @@ class TeamsManagementClient: def list( self, - user_id: Optional[str] = None, - organization_id: Optional[str] = None, - ) -> List[Dict[str, Any]]: + user_id: str | None = None, + organization_id: str | None = None, + ) -> list[dict[str, Any]]: """ List teams that the user belongs to. @@ -69,15 +70,15 @@ class TeamsManagementClient: def list_v2( self, - user_id: Optional[str] = None, - organization_id: Optional[str] = None, - team_id: Optional[str] = None, - team_alias: Optional[str] = None, + user_id: str | None = None, + organization_id: str | None = None, + team_id: str | None = None, + team_alias: str | None = None, page: int = 1, page_size: int = 10, - sort_by: Optional[str] = None, + sort_by: str | None = None, sort_order: str = "asc", - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Get a paginated list of teams with filtering and sorting options. @@ -99,7 +100,7 @@ class TeamsManagementClient: UnauthorizedError: If authentication fails """ url = f"{self._base_url}/v2/team/list" - params: Dict[str, Union[str, int]] = { + params: dict[str, str | int] = { "page": page, "page_size": page_size, "sort_order": sort_order, @@ -124,7 +125,7 @@ class TeamsManagementClient: response.raise_for_status() return response.json() - def get_available(self) -> List[Dict[str, Any]]: + def get_available(self) -> builtins.list[dict[str, Any]]: """ Get list of available teams that the user can join. diff --git a/litellm/proxy/client/users.py b/litellm/proxy/client/users.py index 9f2d53c6d7d..133edfd33be 100644 --- a/litellm/proxy/client/users.py +++ b/litellm/proxy/client/users.py @@ -1,20 +1,22 @@ +from typing import Any + import requests -from typing import List, Dict, Any, Optional -from .exceptions import UnauthorizedError, NotFoundError + +from .exceptions import NotFoundError, UnauthorizedError class UsersManagementClient: - def __init__(self, base_url: str, api_key: Optional[str] = None): + def __init__(self, base_url: str, api_key: str | None = None): self.base_url = base_url.rstrip("/") self.api_key = api_key - def _get_headers(self) -> Dict[str, str]: + def _get_headers(self) -> dict[str, str]: headers = {"Content-Type": "application/json"} if self.api_key: headers["Authorization"] = f"Bearer {self.api_key}" return headers - def list_users(self, params: Optional[Dict[str, Any]] = None) -> List[Dict[str, Any]]: + def list_users(self, params: dict[str, Any] | None = None) -> list[dict[str, Any]]: """List users (GET /user/list)""" url = f"{self.base_url}/user/list" response = requests.get(url, headers=self._get_headers(), params=params) @@ -23,7 +25,7 @@ class UsersManagementClient: response.raise_for_status() return response.json().get("users", response.json()) - def get_user(self, user_id: Optional[str] = None) -> Dict[str, Any]: + def get_user(self, user_id: str | None = None) -> dict[str, Any]: """Get user info (GET /user/info)""" url = f"{self.base_url}/user/info" params = {"user_id": user_id} if user_id else {} @@ -35,7 +37,7 @@ class UsersManagementClient: response.raise_for_status() return response.json() - def get_user_v2(self, user_id: Optional[str] = None) -> Dict[str, Any]: + def get_user_v2(self, user_id: str | None = None) -> dict[str, Any]: """Get user info v2 - lightweight, returns only user object (GET /v2/user/info)""" url = f"{self.base_url}/v2/user/info" params = {"user_id": user_id} if user_id else {} @@ -47,7 +49,7 @@ class UsersManagementClient: response.raise_for_status() return response.json() - def create_user(self, user_data: Dict[str, Any]) -> Dict[str, Any]: + def create_user(self, user_data: dict[str, Any]) -> dict[str, Any]: """Create a new user (POST /user/new)""" url = f"{self.base_url}/user/new" response = requests.post(url, headers=self._get_headers(), json=user_data) @@ -56,7 +58,7 @@ class UsersManagementClient: response.raise_for_status() return response.json() - def delete_user(self, user_ids: List[str]) -> Dict[str, Any]: + def delete_user(self, user_ids: list[str]) -> dict[str, Any]: """Delete users (POST /user/delete)""" url = f"{self.base_url}/user/delete" response = requests.post(url, headers=self._get_headers(), json={"user_ids": user_ids}) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 3f3da6dc258..4d4f459a080 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -10,11 +10,7 @@ from functools import lru_cache from typing import ( TYPE_CHECKING, Any, - Dict, Literal, - Optional, - Tuple, - Union, ) import anyio @@ -102,7 +98,7 @@ def _should_return_raw_model_name(request_data: dict[str, object]) -> bool: ) -def _apply_client_disconnect_metadata(target_metadata: Optional[dict[str, object]]) -> None: +def _apply_client_disconnect_metadata(target_metadata: dict[str, object] | None) -> None: if target_metadata is None: return target_metadata["client_disconnected"] = True @@ -314,7 +310,7 @@ def _stream_usage_tracking_updates( def _serialize_http_exception_detail( detail: Any, -) -> Tuple[str, Optional[dict]]: +) -> tuple[str, dict | None]: """ Convert an HTTPException.detail value into (message, structured_fields) for ProxyException / SSE error frames. @@ -343,7 +339,7 @@ def _serialize_http_exception_detail( return str(detail), None -def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]: +def _collect_response_file_search_vector_store_ids(data: dict[str, Any]) -> set[str]: vector_store_ids: set[str] = set() tools = data.get("tools") if not isinstance(tools, list): @@ -370,7 +366,7 @@ def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[ async def _authorize_response_file_search_vector_stores( - data: Dict[str, Any], + data: dict[str, Any], user_api_key_dict: UserAPIKeyAuth, ) -> None: vector_store_ids = _collect_response_file_search_vector_store_ids(data) @@ -388,7 +384,7 @@ async def _authorize_response_file_search_vector_stores( ) -async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]: +async def _parse_event_data_for_error(event_line: str | bytes) -> int | None: """Parses an event line and returns an error code if present, else None.""" event_line = event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line if event_line.startswith("data: "): @@ -399,7 +395,7 @@ async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional data = orjson.loads(json_str) if isinstance(data, dict) and "error" in data and isinstance(data["error"], dict): error_code_raw = data["error"].get("code") - error_code: Optional[int] = None + error_code: int | None = None if isinstance(error_code_raw, int): error_code = error_code_raw @@ -411,7 +407,6 @@ async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional f"Error code is a string but not a valid integer: {error_code_raw}" ) # Not a valid integer string, treat as if no valid code was found for this check - pass # Ensure error_code is a valid HTTP status code if error_code is not None and 100 <= error_code <= 599: @@ -424,7 +419,7 @@ async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional return None -def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict: +def _extract_error_from_sse_chunk(event_line: str | bytes) -> dict: """ Extract error dictionary from SSE format chunk. @@ -478,10 +473,10 @@ class _UpstreamClosingStreamingResponse(StreamingResponse): self, content: AsyncGenerator[str, None], *, - media_type: Optional[str] = None, - headers: Optional[dict] = None, + media_type: str | None = None, + headers: dict | None = None, status_code: int = status.HTTP_200_OK, - upstream_generator: Optional[AsyncGenerator[str, None]] = None, + upstream_generator: AsyncGenerator[str, None] | None = None, ) -> None: super().__init__(content, status_code=status_code, headers=headers, media_type=media_type) self._upstream_generator = upstream_generator @@ -527,7 +522,7 @@ async def _wait_for_http_disconnect(request: Request) -> None: async def _buffer_first_chunk_honoring_disconnect( generator: AsyncGenerator[str, None], - request: Optional[Request], + request: Request | None, ) -> str: """Fetch the first streamed chunk, cancelling the upstream LLM call if the client disconnects before it arrives. @@ -581,8 +576,8 @@ async def create_response( media_type: str, headers: dict, default_status_code: int = status.HTTP_200_OK, - request: Optional[Request] = None, -) -> Union[StreamingResponse, JSONResponse]: + request: Request | None = None, +) -> StreamingResponse | JSONResponse: """ Create streaming response, checking if the first chunk is an error. If the first chunk is an error, return a standard JSON error response. @@ -595,7 +590,7 @@ async def create_response( "Cache-Control": "no-cache", "X-Accel-Buffering": "no", } - first_chunk_value: Optional[str] = None + first_chunk_value: str | None = None final_status_code = default_status_code try: @@ -673,13 +668,13 @@ async def create_response( existing_fields = getattr(e, "provider_specific_fields", None) or {} if structured_fields: - merged_fields: Optional[dict] = {**existing_fields, **structured_fields} + merged_fields: dict | None = {**existing_fields, **structured_fields} else: merged_fields = existing_fields or None # Match ProxyException.to_dict() shape so streaming and non-streaming # error frames are byte-identical. - error_obj: Dict[str, Any] = { + error_obj: dict[str, Any] = { "message": message, "type": getattr(e, "type", "None"), "param": getattr(e, "param", "None"), @@ -855,8 +850,8 @@ def _override_openai_response_model( def _get_cost_breakdown_from_logging_obj( - litellm_logging_obj: Optional[LiteLLMLoggingObj], -) -> Tuple[Optional[float], Optional[float], Optional[float], Optional[float]]: + litellm_logging_obj: LiteLLMLoggingObj | None, +) -> tuple[float | None, float | None, float | None, float | None]: """ Extract discount and margin information from logging object's cost breakdown. @@ -913,9 +908,7 @@ def _log_llm_api_exception(e: Exception) -> None: "litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, upstream LLM request cancelled" ) return - verbose_proxy_logger.exception( - f"litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - {str(e)}" - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - {e!s}") async def _cancel_llm_call_on_client_disconnect( @@ -964,18 +957,18 @@ class ProxyBaseLLMRequestProcessing: def get_custom_headers( *, user_api_key_dict: UserAPIKeyAuth, - call_id: Optional[str] = None, - model_id: Optional[str] = None, - cache_key: Optional[str] = None, - api_base: Optional[str] = None, - version: Optional[str] = None, - model_region: Optional[str] = None, - response_cost: Optional[Union[float, str]] = None, - hidden_params: Optional[dict] = None, - fastest_response_batch_completion: Optional[bool] = None, - request_data: Optional[dict] = {}, - timeout: Optional[Union[float, int, httpx.Timeout]] = None, - litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, + call_id: str | None = None, + model_id: str | None = None, + cache_key: str | None = None, + api_base: str | None = None, + version: str | None = None, + model_region: str | None = None, + response_cost: float | str | None = None, + hidden_params: dict | None = None, + fastest_response_batch_completion: bool | None = None, + request_data: dict | None = {}, + timeout: float | httpx.Timeout | None = None, + litellm_logging_obj: LiteLLMLoggingObj | None = None, **kwargs, ) -> dict: exclude_values = {"", None, "None"} @@ -1066,9 +1059,9 @@ class ProxyBaseLLMRequestProcessing: request: Request, user_api_key_dict: UserAPIKeyAuth, logging_obj: LiteLLMLoggingObj, - version: Optional[str], + version: str | None, proxy_logging_obj: ProxyLogging, - ) -> Dict[str, str]: + ) -> dict[str, str]: """ Build LiteLLM proxy response headers for routes that call the LLM directly (e.g. Google native :generateContent) instead of base_process_llm_request. @@ -1213,15 +1206,15 @@ class ProxyBaseLLMRequestProcessing: "adelete_run", "apply_guardrail", ], - version: Optional[str] = None, - user_model: Optional[str] = None, - user_temperature: Optional[float] = None, - user_request_timeout: Optional[float] = None, - user_max_tokens: Optional[int] = None, - user_api_base: Optional[str] = None, - model: Optional[str] = None, - llm_router: Optional[Router] = None, - ) -> Tuple[dict, LiteLLMLoggingObj]: + version: str | None = None, + user_model: str | None = None, + user_temperature: float | None = None, + user_request_timeout: float | None = None, + user_max_tokens: int | None = None, + user_api_base: str | None = None, + model: str | None = None, + llm_router: Router | None = None, + ) -> tuple[dict, LiteLLMLoggingObj]: start_time = datetime.now() # start before calling guardrail hooks self.data = await add_litellm_data_to_request( @@ -1374,16 +1367,16 @@ class ProxyBaseLLMRequestProcessing: general_settings: dict, proxy_logging_obj: ProxyLogging, user_api_key_dict: UserAPIKeyAuth, - version: Optional[str], + version: str | None, proxy_config: ProxyConfig, - user_model: Optional[str], - user_temperature: Optional[float], - user_request_timeout: Optional[float], - user_max_tokens: Optional[int], - user_api_base: Optional[str], - model: Optional[str], + user_model: str | None, + user_temperature: float | None, + user_request_timeout: float | None, + user_max_tokens: int | None, + user_api_base: str | None, + model: str | None, route_type: str, - llm_router: Optional[Router], + llm_router: Router | None, ) -> tuple[dict, LiteLLMLoggingObj]: from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError @@ -1459,7 +1452,7 @@ class ProxyBaseLLMRequestProcessing: model: str, llm_router: Router, user_api_key_dict: UserAPIKeyAuth, - ) -> Optional[list]: + ) -> list | None: from litellm.router_utils.fallback_event_handlers import get_fallback_model_group fallbacks = None @@ -1649,23 +1642,23 @@ class ProxyBaseLLMRequestProcessing: proxy_logging_obj: ProxyLogging, general_settings: dict, proxy_config: ProxyConfig, - select_data_generator: Optional[Callable] = None, - llm_router: Optional[Router] = None, - model: Optional[str] = None, - user_model: Optional[str] = None, - user_temperature: Optional[float] = None, - user_request_timeout: Optional[float] = None, - user_max_tokens: Optional[int] = None, - user_api_base: Optional[str] = None, - version: Optional[str] = None, - is_streaming_request: Optional[bool] = False, - contents: Optional[list] = None, # Add contents parameter + select_data_generator: Callable | None = None, + llm_router: Router | None = None, + model: str | None = None, + user_model: str | None = None, + user_temperature: float | None = None, + user_request_timeout: float | None = None, + user_max_tokens: int | None = None, + user_api_base: str | None = None, + version: str | None = None, + is_streaming_request: bool | None = False, + contents: list | None = None, # Add contents parameter skip_pre_call_logic: bool = False, ) -> Any: """ Common request processing logic for both chat completions and responses API endpoints """ - requested_model_from_client: Optional[str] = ( + requested_model_from_client: str | None = ( self.data.get("model") if isinstance(self.data.get("model"), str) else None ) self._debug_log_request_payload() @@ -2188,14 +2181,14 @@ class ProxyBaseLLMRequestProcessing: general_settings: dict, proxy_config: ProxyConfig, select_data_generator: Callable, - llm_router: Optional[Router] = None, - model: Optional[str] = None, - user_model: Optional[str] = None, - user_temperature: Optional[float] = None, - user_request_timeout: Optional[float] = None, - user_max_tokens: Optional[int] = None, - user_api_base: Optional[str] = None, - version: Optional[str] = None, + llm_router: Router | None = None, + model: str | None = None, + user_model: str | None = None, + user_temperature: float | None = None, + user_request_timeout: float | None = None, + user_max_tokens: int | None = None, + user_api_base: str | None = None, + version: str | None = None, ): from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( HttpPassThroughEndpointHelpers, @@ -2259,7 +2252,7 @@ class ProxyBaseLLMRequestProcessing: return False - def _is_streaming_request(self, data: dict, is_streaming_request: Optional[bool] = False) -> bool: + def _is_streaming_request(self, data: dict, is_streaming_request: bool | None = False) -> bool: """ Check if the request is a streaming request. @@ -2332,7 +2325,7 @@ class ProxyBaseLLMRequestProcessing: self.data.get("endpoint"), ) - def _passthrough_event_stream_media_type(self) -> Optional[str]: + def _passthrough_event_stream_media_type(self) -> str | None: """ Content-type for a buffered passthrough event-stream response, resolved from the provider handler so the proxy stays provider-agnostic. Mirrors @@ -2351,8 +2344,8 @@ class ProxyBaseLLMRequestProcessing: proxy_logging_obj: "ProxyLogging", user_api_key_dict: "UserAPIKeyAuth", custom_headers: dict, - request_headers: Dict[str, str], - ) -> Optional[Response]: + request_headers: dict[str, str], + ) -> Response | None: if not self._has_post_call_guardrails_for_passthrough(): return None @@ -2599,7 +2592,7 @@ class ProxyBaseLLMRequestProcessing: e: Exception, user_api_key_dict: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, - version: Optional[str] = None, + version: str | None = None, ): """Raises ProxyException (OpenAI API compatible) if an exception is raised""" _log_llm_api_exception(e) @@ -2622,7 +2615,7 @@ class ProxyBaseLLMRequestProcessing: timeout = getattr( e, "timeout", None ) # returns the timeout set by the wrapper. Used for testing if model-specific timeout are set correctly - _litellm_logging_obj: Optional[LiteLLMLoggingObj] = self.data.get("litellm_logging_obj", None) + _litellm_logging_obj: LiteLLMLoggingObj | None = self.data.get("litellm_logging_obj", None) # Attempt to get model_id from logging object # @@ -2681,7 +2674,7 @@ class ProxyBaseLLMRequestProcessing: message, structured_fields = _serialize_http_exception_detail(raw_detail) existing_fields = getattr(e, "provider_specific_fields", None) or {} if structured_fields: - merged_fields: Optional[dict] = {**existing_fields, **structured_fields} + merged_fields: dict | None = {**existing_fields, **structured_fields} else: merged_fields = existing_fields or None raise ProxyException( @@ -2703,7 +2696,7 @@ class ProxyBaseLLMRequestProcessing: status_code=http_status_error.response.status_code, detail={"error": error_text}, ) - error_msg = f"{str(e)}" + error_msg = f"{e!s}" # Check for AttributeError in the exception chain. # The AttributeError may be wrapped in multiple layers # (e.g. AttributeError -> OpenAIException -> APIConnectionError), @@ -2905,7 +2898,7 @@ class ProxyBaseLLMRequestProcessing: raise except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(str(e)) + f"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {e!s}" ) transformed_exception = await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, @@ -2921,7 +2914,7 @@ class ProxyBaseLLMRequestProcessing: if isinstance(e, HTTPException): raise e error_traceback = _redact_string(traceback.format_exc()) - error_msg = f"{str(e)}\n\n{error_traceback}" + error_msg = f"{e!s}\n\n{error_traceback}" proxy_exception = ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -3011,7 +3004,7 @@ class ProxyBaseLLMRequestProcessing: return chunk @staticmethod - def _inject_cost_into_sse_frame_str(frame_str: str, model_name: str) -> Optional[str]: + def _inject_cost_into_sse_frame_str(frame_str: str, model_name: str) -> str | None: """ Inject cost information into an SSE frame string by modifying the JSON in the 'data:' line. @@ -3041,7 +3034,7 @@ class ProxyBaseLLMRequestProcessing: return None @staticmethod - def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> Optional[dict]: + def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> dict | None: """ Inject cost information into a usage dictionary for message_delta events. @@ -3104,7 +3097,7 @@ class ProxyBaseLLMRequestProcessing: return obj return None - def maybe_get_model_id(self, _logging_obj: Optional[LiteLLMLoggingObj]) -> Optional[str]: + def maybe_get_model_id(self, _logging_obj: LiteLLMLoggingObj | None) -> str | None: """ Get model_id from logging object or request metadata. diff --git a/litellm/proxy/common_utils/cache_coordinator.py b/litellm/proxy/common_utils/cache_coordinator.py index 60d9e273f62..a37316d22fd 100644 --- a/litellm/proxy/common_utils/cache_coordinator.py +++ b/litellm/proxy/common_utils/cache_coordinator.py @@ -13,7 +13,7 @@ pattern: global spend, feature flags, config, or other shared read-through data. import asyncio import time from collections.abc import Awaitable, Callable -from typing import Any, Optional, Protocol, TypeVar +from typing import Any, Protocol, TypeVar from litellm._logging import verbose_proxy_logger @@ -60,11 +60,11 @@ class EventDrivenCacheCoordinator: def __init__(self, log_prefix: str = "[CACHE]"): self._lock = asyncio.Lock() - self._event: Optional[asyncio.Event] = None + self._event: asyncio.Event | None = None self._query_in_progress = False self._log_prefix = log_prefix - async def _get_cached(self, cache_key: str, cache: AsyncCacheProtocol) -> Optional[Any]: + async def _get_cached(self, cache_key: str, cache: AsyncCacheProtocol) -> Any | None: """Return value from cache if present, else None.""" return await cache.async_get_cache(key=cache_key) @@ -76,7 +76,7 @@ class EventDrivenCacheCoordinator: if self._log_prefix: verbose_proxy_logger.debug("%s Cache miss", self._log_prefix) - async def _claim_role(self) -> Optional[asyncio.Event]: + async def _claim_role(self) -> asyncio.Event | None: """ Under lock: return event to wait on if load is in progress, else set us as loader and return None. """ @@ -99,12 +99,12 @@ class EventDrivenCacheCoordinator: event: asyncio.Event, cache_key: str, cache: AsyncCacheProtocol, - ) -> Optional[T]: + ) -> T | None: """Wait for loader to finish, then read from cache.""" await event.wait() if self._log_prefix: verbose_proxy_logger.debug("%s Signal received, reading from cache", self._log_prefix) - value: Optional[T] = await cache.async_get_cache(key=cache_key) + value: T | None = await cache.async_get_cache(key=cache_key) if value is not None and self._log_prefix: verbose_proxy_logger.debug( "%s Cache filled by other request, value: %s", @@ -120,7 +120,7 @@ class EventDrivenCacheCoordinator: cache_key: str, cache: AsyncCacheProtocol, load_fn: Callable[[], Awaitable[T]], - ) -> Optional[T]: + ) -> T | None: """Double-check cache, run load_fn, set cache, return value. Caller must call _signal_done in finally.""" value = await cache.async_get_cache(key=cache_key) if value is not None: @@ -165,7 +165,7 @@ class EventDrivenCacheCoordinator: cache_key: str, cache: AsyncCacheProtocol, load_fn: Callable[[], Awaitable[T]], - ) -> Optional[T]: + ) -> T | None: """ Return cached value or load it once and signal waiters. diff --git a/litellm/proxy/common_utils/cache_pydantic_utils.py b/litellm/proxy/common_utils/cache_pydantic_utils.py index f57f6a299ae..725c2b61145 100644 --- a/litellm/proxy/common_utils/cache_pydantic_utils.py +++ b/litellm/proxy/common_utils/cache_pydantic_utils.py @@ -16,7 +16,7 @@ are encoded; pass a Pydantic model or convert with e.g. ``dataclasses.asdict`` f from __future__ import annotations -from typing import Any, Optional, Type, TypeVar +from typing import Any, TypeVar from pydantic import BaseModel, ValidationError @@ -37,7 +37,7 @@ class CacheCodec: """ @staticmethod - def serialize(value: Any, model_type: Optional[Type[T]] = None) -> Any: + def serialize(value: Any, model_type: type[T] | None = None) -> Any: """ Encode a value for DualCache / Redis (``json.dumps``-safe). @@ -63,7 +63,7 @@ class CacheCodec: return value @staticmethod - def deserialize(cached: Any, model_type: Type[T]) -> Optional[T]: + def deserialize(cached: Any, model_type: type[T]) -> T | None: """ Decode a cache entry to ``model_type``. diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 4988edd2810..1eca0eb768c 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -1,7 +1,7 @@ import copy import os from collections.abc import Callable, Iterable -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Optional import litellm from litellm import get_secret @@ -51,7 +51,7 @@ def initialize_callbacks_on_proxy( premium_user: bool, config_file_path: str, litellm_settings: dict, - callback_specific_params: Optional[dict] = None, + callback_specific_params: dict | None = None, ): if not isinstance(callback_specific_params, dict): callback_specific_params = {} @@ -63,7 +63,7 @@ def initialize_callbacks_on_proxy( verbose_proxy_logger.debug(f"{blue_color_code}initializing callbacks={value} on proxy{reset_color_code}") if isinstance(value, list): - imported_list: List[Any] = [] + imported_list: list[Any] = [] for callback in value: # ["presidio", ] if isinstance(callback, str) and callback == "compression_interception": from litellm.integrations.compression_interception.handler import ( @@ -99,7 +99,7 @@ def initialize_callbacks_on_proxy( _OPTIONAL_PresidioPIIMasking, ) - presidio_logging_only: Optional[bool] = litellm_settings.get("presidio_logging_only", None) + presidio_logging_only: bool | None = litellm_settings.get("presidio_logging_only", None) if presidio_logging_only is not None: presidio_logging_only = bool(presidio_logging_only) # validate boolean given @@ -107,7 +107,7 @@ def initialize_callbacks_on_proxy( if "presidio" in callback_specific_params and isinstance(callback_specific_params["presidio"], dict): _presidio_params = callback_specific_params["presidio"] - params: Dict[str, Any] = { + params: dict[str, Any] = { "logging_only": presidio_logging_only, **_presidio_params, } @@ -325,7 +325,7 @@ def initialize_callbacks_on_proxy( verbose_proxy_logger.debug(f"{blue_color_code} Initialized Callbacks - {litellm.callbacks} {reset_color_code}") -def get_model_group_from_litellm_kwargs(kwargs: dict) -> Optional[str]: +def get_model_group_from_litellm_kwargs(kwargs: dict) -> str | None: _litellm_params = kwargs.get("litellm_params", None) or {} _metadata = _litellm_params.get(get_metadata_variable_name_from_kwargs(kwargs)) or {} _model_group = _metadata.get("model_group", None) @@ -335,7 +335,7 @@ def get_model_group_from_litellm_kwargs(kwargs: dict) -> Optional[str]: return None -def get_model_group_from_request_data(data: dict) -> Optional[str]: +def get_model_group_from_request_data(data: dict) -> str | None: _metadata = data.get("metadata", None) or {} _model_group = _metadata.get("model_group", None) if _model_group is not None: @@ -344,7 +344,7 @@ def get_model_group_from_request_data(data: dict) -> Optional[str]: return None -def get_remaining_tokens_and_requests_from_request_data(data: Dict) -> Dict[str, str]: +def get_remaining_tokens_and_requests_from_request_data(data: dict) -> dict[str, str]: """ Helper function to return x-litellm-key-remaining-tokens-{model_group} and x-litellm-key-remaining-requests-{model_group} @@ -373,8 +373,8 @@ def get_remaining_tokens_and_requests_from_request_data(data: Dict) -> Dict[str, return headers -def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]: - _metadata: Dict = {} +def get_logging_caching_headers(request_data: dict) -> dict | None: + _metadata: dict = {} metadata_bucket = request_data.get("metadata") litellm_metadata_bucket = request_data.get("litellm_metadata") if isinstance(metadata_bucket, dict): @@ -443,8 +443,8 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( def sanitize_openai_provider_metadata( - metadata: Optional[Dict[str, Any]], -) -> Optional[Dict[str, str]]: + metadata: dict[str, Any] | None, +) -> dict[str, str] | None: """ Keep only provider-safe OpenAI metadata entries (string keys -> string values). @@ -453,7 +453,7 @@ def sanitize_openai_provider_metadata( """ if not metadata: return metadata - sanitized: Dict[str, str] = {} + sanitized: dict[str, str] = {} for key, value in metadata.items(): if key in LITELLM_PROXY_INTERNAL_METADATA_KEYS: continue @@ -468,7 +468,7 @@ def sanitize_openai_provider_metadata( return sanitized or None -def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_name: Optional[str]): +def add_guardrail_to_applied_guardrails_header(request_data: dict, guardrail_name: str | None): if guardrail_name is None: return _, _metadata = get_or_create_metadata_bucket(request_data) @@ -479,7 +479,7 @@ def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_nam _metadata["applied_guardrails"] = [guardrail_name] -def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optional[str]): +def add_policy_to_applied_policies_header(request_data: dict, policy_name: str | None): """ Add a policy name to the applied_policies list in request metadata. @@ -496,7 +496,7 @@ def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optio _metadata["applied_policies"] = [policy_name] -def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str, str]): +def add_policy_sources_to_metadata(request_data: dict, policy_sources: dict[str, str]): """ Store policy match reasons in metadata for x-litellm-policy-sources header. @@ -520,7 +520,7 @@ def add_guardrail_response_to_standard_logging_object( ): if litellm_logging_obj is None: return - standard_logging_object: Optional[StandardLoggingPayload] = litellm_logging_obj.model_call_details.get( + standard_logging_object: StandardLoggingPayload | None = litellm_logging_obj.model_call_details.get( "standard_logging_object" ) if standard_logging_object is None: @@ -546,7 +546,7 @@ def process_callback(_callback: str, callback_type: str, environment_variables: return {"name": _callback, "variables": env_vars_dict, "type": callback_type} -def normalize_callback_names(callbacks: Iterable[Any]) -> List[Any]: +def normalize_callback_names(callbacks: Iterable[Any]) -> list[Any]: if callbacks is None: return [] return [c.lower() if isinstance(c, str) else c for c in callbacks] @@ -593,7 +593,7 @@ def _transform_callback_vars(metadata: Any, transform: Callable[[str, Any], Any] def is_sensitive_callback_key( key: str, - extra: Optional[set[str]] = None, + extra: set[str] | None = None, ) -> bool: """Return ``True`` if ``key`` is present in ``extra`` (checked as-is), or if its lowercase form is in ``_EXTRA_SENSITIVE_CALLBACK_KEYS``, or if diff --git a/litellm/proxy/common_utils/custom_openapi_spec.py b/litellm/proxy/common_utils/custom_openapi_spec.py index 92cd8741b32..55d6f083fdd 100644 --- a/litellm/proxy/common_utils/custom_openapi_spec.py +++ b/litellm/proxy/common_utils/custom_openapi_spec.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm._logging import verbose_proxy_logger @@ -26,7 +26,7 @@ class CustomOpenAPISpec: RESPONSES_API_PATHS = ["/v1/responses", "/responses"] @staticmethod - def get_pydantic_schema(model_class) -> Optional[Dict[str, Any]]: + def get_pydantic_schema(model_class) -> dict[str, Any] | None: """ Get JSON schema from a Pydantic model, handling both v1 and v2 APIs. @@ -53,7 +53,7 @@ class CustomOpenAPISpec: return None @staticmethod - def add_schema_to_components(openapi_schema: Dict[str, Any], schema_name: str, schema_def: Dict[str, Any]) -> None: + def add_schema_to_components(openapi_schema: dict[str, Any], schema_name: str, schema_def: dict[str, Any]) -> None: """ Add a schema definition to the OpenAPI components/schemas section. @@ -72,7 +72,7 @@ class CustomOpenAPISpec: CustomOpenAPISpec._move_defs_to_components(openapi_schema, {schema_name: schema_def}) @staticmethod - def add_request_body_to_paths(openapi_schema: Dict[str, Any], paths: List[str], schema_ref: str) -> None: + def add_request_body_to_paths(openapi_schema: dict[str, Any], paths: list[str], schema_ref: str) -> None: """ Add request body with expanded form fields for better Swagger UI display. This keeps the request body but expands it to show individual fields in the UI. @@ -130,7 +130,7 @@ class CustomOpenAPISpec: openapi_schema["paths"][path]["post"]["parameters"] = filtered_params @staticmethod - def _move_defs_to_components(openapi_schema: Dict[str, Any], defs: Dict[str, Any]) -> None: + def _move_defs_to_components(openapi_schema: dict[str, Any], defs: dict[str, Any]) -> None: """ Move $defs from Pydantic v2 schema to OpenAPI components/schemas. This makes the definitions resolvable in Swagger/OpenAPI viewers. @@ -190,7 +190,7 @@ class CustomOpenAPISpec: return schema @staticmethod - def _extract_field_schema(field_def: Dict[str, Any]) -> Dict[str, Any]: + def _extract_field_schema(field_def: dict[str, Any]) -> dict[str, Any]: """ Extract a simple schema from a Pydantic field definition for parameter display. @@ -218,7 +218,7 @@ class CustomOpenAPISpec: return {"type": "string"} @staticmethod - def _expand_field_definition(field_def: Dict[str, Any]) -> Dict[str, Any]: + def _expand_field_definition(field_def: dict[str, Any]) -> dict[str, Any]: """ Expand a Pydantic field definition for inline use in OpenAPI schema. This creates a full field definition that Swagger UI can render as individual form fields. @@ -234,12 +234,12 @@ class CustomOpenAPISpec: @staticmethod def add_request_schema( - openapi_schema: Dict[str, Any], - model_class: Type, + openapi_schema: dict[str, Any], + model_class: type, schema_name: str, - paths: List[str], + paths: list[str], operation_name: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Generic method to add a request schema to OpenAPI specification. @@ -273,14 +273,14 @@ class CustomOpenAPISpec: except Exception as e: # If schema addition fails, continue without it - verbose_proxy_logger.debug(f"Failed to add {operation_name} request schema: {str(e)}") + verbose_proxy_logger.debug(f"Failed to add {operation_name} request schema: {e!s}") return openapi_schema @staticmethod def add_chat_completion_request_schema( - openapi_schema: Dict[str, Any], - ) -> Dict[str, Any]: + openapi_schema: dict[str, Any], + ) -> dict[str, Any]: """ Add ProxyChatCompletionRequest schema to chat completion endpoints for documentation. This shows the request body in Swagger without runtime validation. @@ -302,11 +302,11 @@ class CustomOpenAPISpec: operation_name="chat completion", ) except ImportError as e: - verbose_proxy_logger.debug(f"Failed to import ProxyChatCompletionRequest: {str(e)}") + verbose_proxy_logger.debug(f"Failed to import ProxyChatCompletionRequest: {e!s}") return openapi_schema @staticmethod - def add_embedding_request_schema(openapi_schema: Dict[str, Any]) -> Dict[str, Any]: + def add_embedding_request_schema(openapi_schema: dict[str, Any]) -> dict[str, Any]: """ Add EmbeddingRequest schema to embedding endpoints for documentation. This shows the request body in Swagger without runtime validation. @@ -328,13 +328,13 @@ class CustomOpenAPISpec: operation_name="embedding", ) except ImportError as e: - verbose_proxy_logger.debug(f"Failed to import EmbeddingRequest: {str(e)}") + verbose_proxy_logger.debug(f"Failed to import EmbeddingRequest: {e!s}") return openapi_schema @staticmethod def add_responses_api_request_schema( - openapi_schema: Dict[str, Any], - ) -> Dict[str, Any]: + openapi_schema: dict[str, Any], + ) -> dict[str, Any]: """ Add ResponsesAPIRequestParams schema to responses API endpoints for documentation. This shows the request body in Swagger without runtime validation. @@ -356,13 +356,13 @@ class CustomOpenAPISpec: operation_name="responses API", ) except ImportError as e: - verbose_proxy_logger.debug(f"Failed to import ResponsesAPIRequestParams: {str(e)}") + verbose_proxy_logger.debug(f"Failed to import ResponsesAPIRequestParams: {e!s}") return openapi_schema @staticmethod def add_llm_api_request_schema_body( - openapi_schema: Dict[str, Any], - ) -> Dict[str, Any]: + openapi_schema: dict[str, Any], + ) -> dict[str, Any]: """ Add LLM API request schema bodies to OpenAPI specification for documentation. diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py index aed3e1228f1..7d3150a3e72 100644 --- a/litellm/proxy/common_utils/debug_utils.py +++ b/litellm/proxy/common_utils/debug_utils.py @@ -6,7 +6,7 @@ import os import sys import tracemalloc from collections import Counter -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Query @@ -194,7 +194,7 @@ async def memory_usage_in_mem_cache_items( @router.get("/debug/memory/summary", include_in_schema=False) async def get_memory_summary( _: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Get simplified memory usage summary for the proxy. @@ -249,7 +249,7 @@ async def get_memory_summary( process_memory["error"] = str(e) # Get cache information - caches: Dict[str, Any] = {} + caches: dict[str, Any] = {} total_cache_items = 0 try: @@ -310,7 +310,7 @@ async def get_memory_summary( } -def _get_gc_statistics() -> Dict[str, Any]: +def _get_gc_statistics() -> dict[str, Any]: """Get garbage collector statistics.""" return { "enabled": gc.isenabled(), @@ -338,7 +338,7 @@ def _get_gc_statistics() -> Dict[str, Any]: } -def _get_object_type_counts(top_n: int) -> Tuple[int, List[Dict[str, Any]]]: +def _get_object_type_counts(top_n: int) -> tuple[int, list[dict[str, Any]]]: """Count objects by type and return total count and top N types.""" type_counts: Counter = Counter() total_objects = 0 @@ -356,7 +356,7 @@ def _get_object_type_counts(top_n: int) -> Tuple[int, List[Dict[str, Any]]]: return total_objects, top_object_types -def _get_uncollectable_objects_info() -> Dict[str, Any]: +def _get_uncollectable_objects_info() -> dict[str, Any]: """Get information about uncollectable objects (potential memory leaks).""" uncollectable = gc.garbage return { @@ -370,9 +370,9 @@ def _get_uncollectable_objects_info() -> Dict[str, Any]: } -def _get_cache_memory_stats(user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache) -> Dict[str, Any]: +def _get_cache_memory_stats(user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache) -> dict[str, Any]: """Calculate memory usage for all caches.""" - cache_stats: Dict[str, Any] = {} + cache_stats: dict[str, Any] = {} try: # User API key cache user_cache_size = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict) @@ -436,9 +436,9 @@ def _get_cache_memory_stats(user_api_key_cache, llm_router, proxy_logging_obj, r return cache_stats -def _get_router_memory_stats(llm_router) -> Dict[str, Any]: +def _get_router_memory_stats(llm_router) -> dict[str, Any]: """Get memory usage statistics for LiteLLM router.""" - litellm_router_memory: Dict[str, Any] = {} + litellm_router_memory: dict[str, Any] = {} try: if llm_router is not None: # Model list memory size @@ -502,7 +502,7 @@ def _get_router_memory_stats(llm_router) -> Dict[str, Any]: return litellm_router_memory -def _get_process_memory_info(worker_pid: int, include_process_info: bool) -> Optional[Dict[str, Any]]: +def _get_process_memory_info(worker_pid: int, include_process_info: bool) -> dict[str, Any] | None: """Get process-level memory information using psutil.""" if not include_process_info: return None @@ -555,7 +555,7 @@ async def get_memory_details( _: UserAPIKeyAuth = Depends(user_api_key_auth), top_n: int = Query(20, description="Number of top object types to return"), include_process_info: bool = Query(True, description="Include process memory info"), -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Get detailed memory diagnostics for deep debugging. @@ -580,8 +580,8 @@ async def get_memory_details( from litellm.proxy.proxy_server import ( llm_router, proxy_logging_obj, - user_api_key_cache, redis_usage_cache, + user_api_key_cache, ) worker_pid = os.getpid() @@ -615,7 +615,7 @@ async def configure_gc_thresholds_endpoint( generation_0: int = Query(700, description="Generation 0 threshold (default: 700)"), generation_1: int = Query(10, description="Generation 1 threshold (default: 10)"), generation_2: int = Query(10, description="Generation 2 threshold (default: 10)"), -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Configure Python garbage collection thresholds. @@ -653,7 +653,7 @@ async def configure_gc_thresholds_endpoint( ) except Exception as e: verbose_proxy_logger.error(f"Failed to set GC thresholds: {e}") - raise HTTPException(status_code=500, detail=f"Failed to set GC thresholds: {str(e)}") + raise HTTPException(status_code=500, detail=f"Failed to set GC thresholds: {e!s}") # Get current object count to show immediate impact current_count = gc.get_count()[0] @@ -783,4 +783,4 @@ def init_verbose_loggers(): except Exception as e: import logging - logging.warning(f"Failed to init verbose loggers: {str(e)}") + logging.warning(f"Failed to init verbose loggers: {e!s}") diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index 9a0b4f8b982..b7b8bfd1eea 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -1,6 +1,6 @@ import base64 import os -from typing import Literal, Optional, cast +from typing import Literal, cast from litellm._logging import verbose_proxy_logger @@ -91,7 +91,7 @@ def _decrypt_aes_gcm(value: str, signing_key: str) -> str: return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8") -def encrypt_value_helper(value: str, new_encryption_key: Optional[str] = None): +def encrypt_value_helper(value: str, new_encryption_key: str | None = None): signing_key = new_encryption_key or _get_salt_key() try: @@ -145,7 +145,7 @@ def decrypt_value_helper( # if it's not str - do not decrypt it, return the value return value except Exception as e: - error_message = f"Error decrypting value for key: {key}, Did your master_key/salt key change recently? \nError: {str(e)}\nSet permanent salt key - https://docs.litellm.ai/docs/proxy/prod#5-set-litellm-salt-key" + error_message = f"Error decrypting value for key: {key}, Did your master_key/salt key change recently? \nError: {e!s}\nSet permanent salt key - https://docs.litellm.ai/docs/proxy/prod#5-set-litellm-salt-key" if exception_type == "debug": verbose_proxy_logger.debug(error_message) return value if return_original_value else None diff --git a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py index 85ec6b839a4..a8acc28d9de 100644 --- a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py +++ b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py @@ -5,7 +5,7 @@ Deletes expired virtual keys created for LiteLLM dashboard sessions. """ from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.constants import ( @@ -111,8 +111,8 @@ class ExpiredUISessionKeyCleanupManager: @staticmethod def _get_deleted_token_count( - tokens: List[str], - response: Optional[Dict[str, Any]], + tokens: list[str], + response: dict[str, Any] | None, ) -> int: """ Return the number of tokens actually deleted from the delete helper response. @@ -138,7 +138,7 @@ class ExpiredUISessionKeyCleanupManager: return len(tokens) - async def _find_expired_ui_session_keys(self) -> List[LiteLLM_VerificationToken]: + async def _find_expired_ui_session_keys(self) -> list[LiteLLM_VerificationToken]: """ Find expired LiteLLM dashboard session keys. """ diff --git a/litellm/proxy/common_utils/get_routes.py b/litellm/proxy/common_utils/get_routes.py index 64a1332ccb0..4e3c908a8bb 100644 --- a/litellm/proxy/common_utils/get_routes.py +++ b/litellm/proxy/common_utils/get_routes.py @@ -2,7 +2,7 @@ Utility class for getting routes from a FastAPI app. """ -from typing import Any, Dict, List, Optional +from typing import Any from starlette.routing import BaseRoute @@ -14,11 +14,11 @@ class GetRoutes: def get_app_routes( route: BaseRoute, endpoint_route: Any, - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """ Get routes for a regular route. """ - routes: List[Dict[str, Any]] = [] + routes: list[dict[str, Any]] = [] route_info = { "path": getattr(route, "path", None), "methods": getattr(route, "methods", None), @@ -31,11 +31,11 @@ class GetRoutes: @staticmethod def get_routes_for_mounted_app( route: BaseRoute, - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """ Get routes for a mounted sub-application. """ - routes: List[Dict[str, Any]] = [] + routes: list[dict[str, Any]] = [] mount_path = getattr(route, "path", "") sub_app = getattr(route, "app", None) if sub_app and hasattr(sub_app, "routes"): @@ -58,7 +58,7 @@ class GetRoutes: return routes @staticmethod - def _safe_get_endpoint_name(endpoint_function: Any) -> Optional[str]: + def _safe_get_endpoint_name(endpoint_function: Any) -> str | None: """ Safely get the name of the endpoint function. """ diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 9195f02bfbe..0dd910e1901 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -1,7 +1,7 @@ import json import re from collections.abc import Collection -from typing import Any, Dict, List, Optional +from typing import Any import orjson from fastapi import Request, UploadFile, status @@ -40,7 +40,7 @@ def _is_json_content_type(content_type: str) -> bool: return _normalize_media_type(content_type) == "application/json" -async def _read_request_body(request: Optional[Request]) -> Dict: +async def _read_request_body(request: Request | None) -> dict: """ Safely read the request body and parse it as JSON. @@ -55,7 +55,7 @@ async def _read_request_body(request: Optional[Request]) -> Dict: return {} # Check if we already read and parsed the body - _cached_request_body: Optional[dict] = _safe_get_request_parsed_body(request=request) + _cached_request_body: dict | None = _safe_get_request_parsed_body(request=request) if _cached_request_body is not None: return _cached_request_body @@ -98,9 +98,9 @@ async def _read_request_body(request: Optional[Request]) -> Dict: # Above the configured size, skip the repair and raise the 400 now. repair_limit_bytes = MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB * 1024 * 1024 if repair_limit_bytes > 0 and len(body) > repair_limit_bytes: - verbose_proxy_logger.error(f"Invalid JSON payload received: {str(e)}") + verbose_proxy_logger.error(f"Invalid JSON payload received: {e!s}") raise ProxyException( - message=f"Invalid JSON payload: {str(e)}", + message=f"Invalid JSON payload: {e!s}", type="invalid_request_error", param="request_body", code=status.HTTP_400_BAD_REQUEST, @@ -120,9 +120,9 @@ async def _read_request_body(request: Optional[Request]) -> Dict: parsed_body = json.loads(body_str) except json.JSONDecodeError: # If both orjson and json.loads fail, throw a proper error - verbose_proxy_logger.error(f"Invalid JSON payload received: {str(e)}") + verbose_proxy_logger.error(f"Invalid JSON payload received: {e!s}") raise ProxyException( - message=f"Invalid JSON payload: {str(e)}", + message=f"Invalid JSON payload: {e!s}", type="invalid_request_error", param="request_body", code=status.HTTP_400_BAD_REQUEST, @@ -134,15 +134,15 @@ async def _read_request_body(request: Optional[Request]) -> Dict: except (json.JSONDecodeError, orjson.JSONDecodeError, ProxyException) as e: # Re-raise ProxyException as-is - verbose_proxy_logger.error(f"Invalid JSON payload received: {str(e)}") + verbose_proxy_logger.error(f"Invalid JSON payload received: {e!s}") raise except Exception as e: # Catch unexpected errors to avoid crashes - verbose_proxy_logger.exception("Unexpected error reading request body - {}".format(e)) + verbose_proxy_logger.exception(f"Unexpected error reading request body - {e}") return {} -def _safe_get_request_parsed_body(request: Optional[Request]) -> Optional[dict]: +def _safe_get_request_parsed_body(request: Request | None) -> dict | None: if request is None: return None if hasattr(request, "scope") and "parsed_body" in request.scope and isinstance(request.scope["parsed_body"], tuple): @@ -151,7 +151,7 @@ def _safe_get_request_parsed_body(request: Optional[Request]) -> Optional[dict]: return None -def _safe_get_request_query_params(request: Optional[Request]) -> Dict: +def _safe_get_request_query_params(request: Request | None) -> dict: if request is None: return {} try: @@ -159,12 +159,12 @@ def _safe_get_request_query_params(request: Optional[Request]) -> Dict: return dict(request.query_params) return {} except Exception as e: - verbose_proxy_logger.debug("Unexpected error reading request query params - {}".format(e)) + verbose_proxy_logger.debug(f"Unexpected error reading request query params - {e}") return {} def _safe_set_request_parsed_body( - request: Optional[Request], + request: Request | None, parsed_body: dict, ) -> None: try: @@ -172,10 +172,10 @@ def _safe_set_request_parsed_body( return request.scope["parsed_body"] = (tuple(parsed_body.keys()), parsed_body) except Exception as e: - verbose_proxy_logger.debug("Unexpected error setting request parsed body - {}".format(e)) + verbose_proxy_logger.debug(f"Unexpected error setting request parsed body - {e}") -def _safe_get_request_headers(request: Optional[Request]) -> dict: +def _safe_get_request_headers(request: Request | None) -> dict: """ [Non-Blocking] Safely get the request headers. Caches the result on request.state to avoid re-creating dict(request.headers) per call. @@ -190,11 +190,11 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict: if isinstance(cached, dict): return cached if cached is not None: - verbose_proxy_logger.debug("Unexpected cached request headers type - {}".format(type(cached))) + verbose_proxy_logger.debug(f"Unexpected cached request headers type - {type(cached)}") try: headers = dict(request.headers) except Exception as e: - verbose_proxy_logger.debug("Unexpected error reading request headers - {}".format(e)) + verbose_proxy_logger.debug(f"Unexpected error reading request headers - {e}") headers = {} try: if state is not None: @@ -231,7 +231,7 @@ def check_file_size_under_limit( if llm_router is not None and request_data["model"] in router_model_names: try: - deployment: Optional[Deployment] = llm_router.get_deployment_by_model_group_name( + deployment: Deployment | None = llm_router.get_deployment_by_model_group_name( model_group_name=request_data["model"] ) if ( @@ -267,7 +267,7 @@ def check_file_size_under_limit( return True -async def get_form_data(request: Request) -> Dict[str, Any]: +async def get_form_data(request: Request) -> dict[str, Any]: """ Read form data from request @@ -287,8 +287,8 @@ async def get_form_data(request: Request) -> Dict[str, Any]: async def convert_upload_files_to_file_data( - form_data: Dict[str, Any], -) -> Dict[str, Any]: + form_data: dict[str, Any], +) -> dict[str, Any]: """ Convert FastAPI UploadFile objects to file data tuples for litellm. @@ -331,7 +331,7 @@ async def convert_upload_files_to_file_data( return data -async def get_request_body(request: Request) -> Dict[str, Any]: +async def get_request_body(request: Request) -> dict[str, Any]: """ Read the request body and parse it as JSON. """ @@ -346,7 +346,7 @@ async def get_request_body(request: Request) -> Dict[str, Any]: return {} -def extract_nested_form_metadata(form_data: Dict[str, Any], prefix: str = "litellm_metadata[") -> Dict[str, Any]: +def extract_nested_form_metadata(form_data: dict[str, Any], prefix: str = "litellm_metadata[") -> dict[str, Any]: """ Extract nested metadata from form data with bracket notation. @@ -384,7 +384,7 @@ def extract_nested_form_metadata(form_data: Dict[str, Any], prefix: str = "litel if not form_data: return {} - metadata: Dict[str, Any] = {} + metadata: dict[str, Any] = {} for key, value in form_data.items(): # Skip keys that don't start with the prefix @@ -426,13 +426,13 @@ def extract_nested_form_metadata(form_data: Dict[str, Any], prefix: str = "litel verbose_proxy_logger.warning(f"Cannot set value - parent is not a dict for key: {key}") except Exception as e: - verbose_proxy_logger.error(f"Error parsing metadata key '{key}': {str(e)}") + verbose_proxy_logger.error(f"Error parsing metadata key '{key}': {e!s}") continue return metadata -def get_tags_from_request_body(request_body: dict) -> List[str]: +def get_tags_from_request_body(request_body: dict) -> list[str]: """ Extract tags from request body metadata. @@ -455,7 +455,7 @@ def get_tags_from_request_body(request_body: dict) -> List[str]: metadata = {} tags_in_metadata: Any = metadata.get("tags", []) tags_in_request_body: Any = request_body.get("tags", []) - combined_tags: List[str] = [] + combined_tags: list[str] = [] ###################################### # Only combine tags if they are lists diff --git a/litellm/proxy/common_utils/key_rotation_manager.py b/litellm/proxy/common_utils/key_rotation_manager.py index 51f86e514ca..8e065a06979 100644 --- a/litellm/proxy/common_utils/key_rotation_manager.py +++ b/litellm/proxy/common_utils/key_rotation_manager.py @@ -5,7 +5,6 @@ Handles finding keys that need rotation based on their individual schedules. """ from datetime import datetime, timezone -from typing import List from litellm._logging import verbose_proxy_logger from litellm.constants import ( @@ -106,7 +105,7 @@ class KeyRotationManager: cronjob_id=KEY_ROTATION_JOB_NAME, ) - async def _find_keys_needing_rotation(self) -> List[LiteLLM_VerificationToken]: + async def _find_keys_needing_rotation(self) -> list[LiteLLM_VerificationToken]: """ Find keys that are due for rotation based on their key_rotation_at timestamp. diff --git a/litellm/proxy/common_utils/load_config_utils.py b/litellm/proxy/common_utils/load_config_utils.py index 71e8238db27..225f7cfdf6c 100644 --- a/litellm/proxy/common_utils/load_config_utils.py +++ b/litellm/proxy/common_utils/load_config_utils.py @@ -34,10 +34,9 @@ def get_file_contents_from_s3(bucket_name, object_key): except ImportError as e: # this is most likely if a user is not using the litellm docker container - verbose_proxy_logger.error(f"ImportError: {str(e)}") - pass + verbose_proxy_logger.error(f"ImportError: {e!s}") except Exception as e: - verbose_proxy_logger.error(f"Error retrieving file contents: {str(e)}") + verbose_proxy_logger.error(f"Error retrieving file contents: {e!s}") return None @@ -58,7 +57,7 @@ async def get_config_file_contents_from_gcs(bucket_name, object_key): return config except Exception as e: - verbose_proxy_logger.error(f"Error retrieving file contents: {str(e)}") + verbose_proxy_logger.error(f"Error retrieving file contents: {e!s}") return None @@ -112,10 +111,10 @@ def download_python_file_from_s3( return True except ImportError as e: - verbose_proxy_logger.error(f"ImportError: {str(e)}") + verbose_proxy_logger.error(f"ImportError: {e!s}") return False except Exception as e: - verbose_proxy_logger.exception(f"Error downloading Python file: {str(e)}") + verbose_proxy_logger.exception(f"Error downloading Python file: {e!s}") return False @@ -159,7 +158,7 @@ async def download_python_file_from_gcs( return True except Exception as e: - verbose_proxy_logger.exception(f"Error downloading Python file from GCS: {str(e)}") + verbose_proxy_logger.exception(f"Error downloading Python file from GCS: {e!s}") return False diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 2835ecfc751..41fbb804cc0 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -53,7 +53,7 @@ class TeamModelNameTranslator: @staticmethod def build_internal_to_public_map( - llm_router: "Router | None", + llm_router: Router | None, general_settings: Mapping[str, object], ) -> dict[str, str]: """Internal team routing key -> public `team_public_model_name`. @@ -92,7 +92,7 @@ class TeamModelNameTranslator: @staticmethod def listing_entries( model_names: list[str], - llm_router: "Router | None", + llm_router: Router | None, general_settings: Mapping[str, object], ) -> list[tuple[str, str]]: """`(response_id, metadata_lookup_id)` for each listed model, de-duplicated @@ -114,7 +114,7 @@ class TeamModelNameTranslator: @staticmethod def translate_listing( model_names: list[str], - llm_router: "Router | None", + llm_router: Router | None, general_settings: Mapping[str, object], ) -> list[str]: """Public-name view of `model_names` (the `response_id` of each listing @@ -129,7 +129,7 @@ class TeamModelNameTranslator: def resolve_public_name( model_id: str, available_models: list[str], - llm_router: "Router | None", + llm_router: Router | None, general_settings: Mapping[str, object], ) -> str: """Resolve a public team name back to the internal routing key the router diff --git a/litellm/proxy/common_utils/openai_endpoint_utils.py b/litellm/proxy/common_utils/openai_endpoint_utils.py index 905967fa465..d8d5b8d86b7 100644 --- a/litellm/proxy/common_utils/openai_endpoint_utils.py +++ b/litellm/proxy/common_utils/openai_endpoint_utils.py @@ -2,8 +2,6 @@ Contains utils used by OpenAI compatible endpoints """ -from typing import Optional, Set - from fastapi import Request from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker @@ -14,7 +12,7 @@ SENSITIVE_DATA_MASKER = SensitiveDataMasker() def remove_sensitive_info_from_deployment( deployment_dict: dict, - excluded_keys: Optional[Set[str]] = None, + excluded_keys: set[str] | None = None, ) -> dict: """ Removes sensitive information from a deployment dictionary. @@ -49,7 +47,7 @@ def remove_sensitive_info_from_deployment( return deployment_dict -async def get_custom_llm_provider_from_request_body(request: Request) -> Optional[str]: +async def get_custom_llm_provider_from_request_body(request: Request) -> str | None: """ Get the `custom_llm_provider` from the request body @@ -61,7 +59,7 @@ async def get_custom_llm_provider_from_request_body(request: Request) -> Optiona return None -def get_custom_llm_provider_from_request_query(request: Request) -> Optional[str]: +def get_custom_llm_provider_from_request_query(request: Request) -> str | None: """ Get the `custom_llm_provider` from the request query parameters @@ -72,7 +70,7 @@ def get_custom_llm_provider_from_request_query(request: Request) -> Optional[str return None -def get_custom_llm_provider_from_request_headers(request: Request) -> Optional[str]: +def get_custom_llm_provider_from_request_headers(request: Request) -> str | None: """ Get the `custom_llm_provider` from the request header `custom-llm-provider` """ diff --git a/litellm/proxy/common_utils/openapi_schema_compat.py b/litellm/proxy/common_utils/openapi_schema_compat.py index cd96bd2d2a4..06a18524733 100644 --- a/litellm/proxy/common_utils/openapi_schema_compat.py +++ b/litellm/proxy/common_utils/openapi_schema_compat.py @@ -5,7 +5,7 @@ FastAPI 0.120+ has stricter schema generation that fails on certain types like o This module provides a compatibility layer to handle these cases gracefully. """ -from typing import Any, Dict +from typing import Any from litellm._logging import verbose_proxy_logger @@ -16,7 +16,7 @@ def get_openapi_schema_with_compat( version: str, description: str, routes: list, -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Generate OpenAPI schema with compatibility handling for FastAPI 0.120+. diff --git a/litellm/proxy/common_utils/performance_utils.py b/litellm/proxy/common_utils/performance_utils.py index 38c866c86f7..50de40480fd 100644 --- a/litellm/proxy/common_utils/performance_utils.py +++ b/litellm/proxy/common_utils/performance_utils.py @@ -15,7 +15,7 @@ import inspect import threading from collections.abc import Callable from pathlib import Path as PathLib -from typing import Any, Optional +from typing import Any from litellm._logging import verbose_proxy_logger @@ -27,7 +27,7 @@ _sample_counter = 0 _sample_counter_lock = threading.Lock() # Global line_profiler state -_line_profiler: Optional[Any] = None +_line_profiler: Any | None = None _line_profiler_lock = threading.Lock() _wrapped_functions: dict[str, Callable] = {} # Store original functions @@ -230,7 +230,7 @@ def wrap_function_directly(func: Callable) -> Callable: return profiled_function -def collect_line_profiler_stats(output_file: Optional[str] = None) -> None: +def collect_line_profiler_stats(output_file: str | None = None) -> None: """Collect and save line_profiler statistics. This can be called manually to collect stats at any time, or it's @@ -264,7 +264,7 @@ def collect_line_profiler_stats(output_file: Optional[str] = None) -> None: verbose_proxy_logger.error(f"Error collecting line profiler stats: {e}") -def register_shutdown_handler(output_file: Optional[str] = None) -> None: +def register_shutdown_handler(output_file: str | None = None) -> None: """Register a shutdown handler to collect line_profiler stats. This registers an atexit handler that will automatically save profiling diff --git a/litellm/proxy/common_utils/proxy_rate_limit_error.py b/litellm/proxy/common_utils/proxy_rate_limit_error.py index 6b3a3856f52..28c954e9721 100644 --- a/litellm/proxy/common_utils/proxy_rate_limit_error.py +++ b/litellm/proxy/common_utils/proxy_rate_limit_error.py @@ -37,7 +37,7 @@ This module provides a single proxy-side error class that: import json from collections.abc import Mapping -from typing import Any, Dict, Optional, Union +from typing import Any from fastapi import HTTPException @@ -45,8 +45,8 @@ from litellm.exceptions import RateLimitError, RateLimitErrorCategory, RateLimit def map_v3_rate_limit_type( - v3_value: Optional[str], -) -> Optional[RateLimitType]: + v3_value: str | None, +) -> RateLimitType | None: """ Map the v3 rate limiter's internal `status["rate_limit_type"]` strings onto the public :class:`RateLimitType` enum. @@ -144,11 +144,11 @@ class ProxyRateLimitError(HTTPException, RateLimitError): # type: ignore[misc] def __init__( self, detail: Any, - headers: Optional[Mapping[str, Any]] = None, - category: Union[str, RateLimitErrorCategory] = RateLimitErrorCategory.LITELLM_RATE_LIMIT, - rate_limit_type: Optional[Union[str, RateLimitType]] = None, - model: Optional[str] = None, - llm_provider: Optional[str] = "litellm_proxy", + headers: Mapping[str, Any] | None = None, + category: str | RateLimitErrorCategory = RateLimitErrorCategory.LITELLM_RATE_LIMIT, + rate_limit_type: str | RateLimitType | None = None, + model: str | None = None, + llm_provider: str | None = "litellm_proxy", ): # Normalize None → safe defaults so callers (and the resolver helper # in `rate_limiter_utils`) can pass `None` without producing an @@ -158,7 +158,7 @@ class ProxyRateLimitError(HTTPException, RateLimitError): # type: ignore[misc] model = model or "" llm_provider = llm_provider or "litellm_proxy" message = _coerce_message(detail) - stringified_headers: Optional[Dict[str, str]] = {k: str(v) for k, v in headers.items()} if headers else None + stringified_headers: dict[str, str] | None = {k: str(v) for k, v in headers.items()} if headers else None # Initialize the FastAPI HTTPException portion first so its attributes # (status_code, detail, headers) are already on the instance before diff --git a/litellm/proxy/common_utils/realtime_utils.py b/litellm/proxy/common_utils/realtime_utils.py index ee31a902edd..ff039754555 100644 --- a/litellm/proxy/common_utils/realtime_utils.py +++ b/litellm/proxy/common_utils/realtime_utils.py @@ -1,11 +1,10 @@ from functools import lru_cache -from typing import Optional from litellm.constants import _REALTIME_BODY_CACHE_SIZE @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) -def _realtime_request_body(model: Optional[str]) -> bytes: +def _realtime_request_body(model: str | None) -> bytes: """ Generate the realtime websocket request body. Cached with LRU semantics to avoid repeated string formatting work while keeping memory usage bounded. diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index e77e112d120..3087e356f99 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -3,7 +3,7 @@ import json import time from collections.abc import Callable from datetime import datetime, timezone -from typing import Any, List, Literal, Optional, Union +from typing import Any, Literal import litellm from litellm._logging import verbose_proxy_logger @@ -133,12 +133,12 @@ class ResetBudgetJob: async def _cascade_reset_spend_for_budget_link( self, - budgets_to_reset: List[LiteLLM_BudgetTableFull], + budgets_to_reset: list[LiteLLM_BudgetTableFull], table: Any, counter_key_fn: Callable[[Any], str], log_subject: str, - extra_where: Optional[dict] = None, - cache_key_fn: Optional[Callable[[Any], Union[str, List[str]]]] = None, + extra_where: dict | None = None, + cache_key_fn: Callable[[Any], str | list[str]] | None = None, ): """ Generic cascade: zero spend on rows whose budget_id is in the reset set. @@ -174,7 +174,7 @@ class ResetBudgetJob: return update_result - async def reset_budget_for_litellm_team_members(self, budgets_to_reset: List[LiteLLM_BudgetTableFull]): + async def reset_budget_for_litellm_team_members(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): """ Resets the budget for all LiteLLM Team Members if their budget has expired """ @@ -186,7 +186,7 @@ class ResetBudgetJob: cache_key_fn=lambda m: f"{m.team_id}_{m.user_id}", ) - async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: List[LiteLLM_BudgetTableFull]): + async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): """ Resets the spend for keys linked to budget tiers that are being reset. @@ -202,7 +202,7 @@ class ResetBudgetJob: cache_key_fn=lambda k: k.token, ) - async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: List[LiteLLM_BudgetTableFull]): + async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): """ Resets the spend for orgs linked to budget tiers that are being reset. """ @@ -218,7 +218,7 @@ class ResetBudgetJob: ], ) - async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: List[LiteLLM_BudgetTableFull]): + async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): """ Resets the spend for tags linked to budget tiers that are being reset. @@ -247,9 +247,9 @@ class ResetBudgetJob: now = datetime.now(timezone.utc) start_time = time.time() - endusers_to_reset: Optional[List[LiteLLM_EndUserTable]] = None - budgets_to_reset: Optional[List[LiteLLM_BudgetTableFull]] = None - updated_endusers: List[LiteLLM_EndUserTable] = [] + endusers_to_reset: list[LiteLLM_EndUserTable] | None = None + budgets_to_reset: list[LiteLLM_BudgetTableFull] | None = None + updated_endusers: list[LiteLLM_EndUserTable] = [] failed_endusers = [] try: budgets_to_reset = await self.prisma_client.get_data( @@ -369,7 +369,7 @@ class ResetBudgetJob: async def _get_endusers_with_no_budget_id( self, - ) -> List[LiteLLM_EndUserTable]: + ) -> list[LiteLLM_EndUserTable]: """ Fetch end users that have no explicit budget_id set (NULL) and have accumulated spend > 0. These are implicitly-created end users that @@ -384,7 +384,7 @@ class ResetBudgetJob: ) return [LiteLLM_EndUserTable(**row.dict()) for row in rows] - async def _write_key_reset_updates(self, updated_keys: List[LiteLLM_VerificationToken]) -> None: + async def _write_key_reset_updates(self, updated_keys: list[LiteLLM_VerificationToken]) -> None: """ Write per-row {spend, budget_reset_at} updates for keys. @@ -406,7 +406,7 @@ class ResetBudgetJob: ) await batcher.commit() - async def _write_user_reset_updates(self, updated_users: List[LiteLLM_UserTable]) -> None: + async def _write_user_reset_updates(self, updated_users: list[LiteLLM_UserTable]) -> None: """ Write per-row {spend, budget_reset_at} updates for users. @@ -425,7 +425,7 @@ class ResetBudgetJob: ) await batcher.commit() - async def _write_team_reset_updates(self, updated_teams: List[LiteLLM_TeamTable]) -> None: + async def _write_team_reset_updates(self, updated_teams: list[LiteLLM_TeamTable]) -> None: """ Write per-row {spend, budget_reset_at} updates for teams. @@ -452,13 +452,13 @@ class ResetBudgetJob: """ now = datetime.utcnow() start_time = time.time() - keys_to_reset: Optional[List[LiteLLM_VerificationToken]] = None + keys_to_reset: list[LiteLLM_VerificationToken] | None = None try: keys_to_reset = await self.prisma_client.get_data( table_name="key", query_type="find_all", expires=now, reset_at=now ) verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str)) - updated_keys: List[LiteLLM_VerificationToken] = [] + updated_keys: list[LiteLLM_VerificationToken] = [] failed_keys = [] if keys_to_reset is not None and len(keys_to_reset) > 0: for key in keys_to_reset: @@ -530,10 +530,10 @@ class ResetBudgetJob: """ now = datetime.utcnow() start_time = time.time() - users_to_reset: Optional[List[LiteLLM_UserTable]] = None + users_to_reset: list[LiteLLM_UserTable] | None = None try: users_to_reset = await self.prisma_client.get_data(table_name="user", query_type="find_all", reset_at=now) - updated_users: List[LiteLLM_UserTable] = [] + updated_users: list[LiteLLM_UserTable] = [] failed_users = [] if users_to_reset is not None and len(users_to_reset) > 0: for user in users_to_reset: @@ -611,10 +611,10 @@ class ResetBudgetJob: """ now = datetime.utcnow() start_time = time.time() - teams_to_reset: Optional[List[LiteLLM_TeamTable]] = None + teams_to_reset: list[LiteLLM_TeamTable] | None = None try: teams_to_reset = await self.prisma_client.get_data(table_name="team", query_type="find_all", reset_at=now) - updated_teams: List[LiteLLM_TeamTable] = [] + updated_teams: list[LiteLLM_TeamTable] = [] failed_teams = [] if teams_to_reset is not None and len(teams_to_reset) > 0: for team in teams_to_reset: @@ -786,7 +786,7 @@ class ResetBudgetJob: @staticmethod async def _reset_budget_common( - item: Union[LiteLLM_TeamTable, LiteLLM_UserTable, LiteLLM_VerificationToken], + item: LiteLLM_TeamTable | LiteLLM_UserTable | LiteLLM_VerificationToken, current_time: datetime, item_type: Literal["key", "team", "user"], reset_settings: BudgetResetSettings, @@ -843,7 +843,7 @@ class ResetBudgetJob: @staticmethod async def _reset_budget_for_enduser( enduser: LiteLLM_EndUserTable, - ) -> Optional[LiteLLM_EndUserTable]: + ) -> LiteLLM_EndUserTable | None: try: enduser.spend = 0.0 except Exception as e: diff --git a/litellm/proxy/common_utils/resource_ownership.py b/litellm/proxy/common_utils/resource_ownership.py index 936c4e18bde..6f8300248ce 100644 --- a/litellm/proxy/common_utils/resource_ownership.py +++ b/litellm/proxy/common_utils/resource_ownership.py @@ -1,9 +1,7 @@ -from typing import List, Optional - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -def is_proxy_admin(user_api_key_dict: Optional[UserAPIKeyAuth]) -> bool: +def is_proxy_admin(user_api_key_dict: UserAPIKeyAuth | None) -> bool: if user_api_key_dict is None: return False @@ -14,8 +12,8 @@ def is_proxy_admin(user_api_key_dict: Optional[UserAPIKeyAuth]) -> bool: def get_resource_owner_scopes( - user_api_key_dict: Optional[UserAPIKeyAuth], -) -> List[str]: + user_api_key_dict: UserAPIKeyAuth | None, +) -> list[str]: """ Return ownership scopes that may access a user-created proxy resource. @@ -32,9 +30,9 @@ def get_resource_owner_scopes( if user_api_key_dict is None: return [] - scopes: List[str] = [] + scopes: list[str] = [] - def _add(scope: Optional[str]) -> None: + def _add(scope: str | None) -> None: if scope and scope not in scopes: scopes.append(scope) @@ -54,8 +52,8 @@ def get_resource_owner_scopes( def get_primary_resource_owner_scope( - user_api_key_dict: Optional[UserAPIKeyAuth], -) -> Optional[str]: + user_api_key_dict: UserAPIKeyAuth | None, +) -> str | None: """Return the canonical owner scope to stamp on newly-created rows. ``None`` for identity-less callers — callers that depend on a primary @@ -80,8 +78,8 @@ def get_primary_resource_owner_scope( def user_can_access_resource_owner( - owner: Optional[str], - user_api_key_dict: Optional[UserAPIKeyAuth], + owner: str | None, + user_api_key_dict: UserAPIKeyAuth | None, ) -> bool: if user_api_key_dict is None: return True diff --git a/litellm/proxy/common_utils/static_asset_utils.py b/litellm/proxy/common_utils/static_asset_utils.py index c108af2b475..21565545d83 100644 --- a/litellm/proxy/common_utils/static_asset_utils.py +++ b/litellm/proxy/common_utils/static_asset_utils.py @@ -1,14 +1,13 @@ """Helpers for unauthenticated logo / favicon endpoints.""" import os -from typing import Optional, Tuple from litellm._logging import verbose_proxy_logger LOCAL_IMAGE_HEADER_BYTES = 512 -def detect_local_image_media_type(header: bytes) -> Optional[str]: +def detect_local_image_media_type(header: bytes) -> str | None: """Return a browser image media type for supported local image signatures.""" if header[0:8] == b"\x89PNG\r\n\x1a\n": return "image/png" @@ -23,7 +22,7 @@ def detect_local_image_media_type(header: bytes) -> Optional[str]: return None -def resolve_validated_local_image_path(candidate: str) -> Optional[Tuple[str, str]]: +def resolve_validated_local_image_path(candidate: str) -> tuple[str, str] | None: """Resolve ``candidate`` only when it is an existing supported image file.""" if not candidate: return None diff --git a/litellm/proxy/common_utils/swagger_utils.py b/litellm/proxy/common_utils/swagger_utils.py index 75a64707cd4..1057ee9b658 100644 --- a/litellm/proxy/common_utils/swagger_utils.py +++ b/litellm/proxy/common_utils/swagger_utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict +from typing import Any from pydantic import BaseModel, Field @@ -6,7 +6,7 @@ from litellm.exceptions import LITELLM_EXCEPTION_TYPES class ErrorResponse(BaseModel): - detail: Dict[str, Any] = Field( + detail: dict[str, Any] = Field( ..., example={ # type: ignore "error": { diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index dbb5b2c24d0..fc270a59f3e 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Optional, Type, TypeVar, Union, cast, overload +from typing import Any, TypeVar, cast, overload from pydantic import BaseModel @@ -44,9 +44,9 @@ class UserApiKeyCache(DualCache): parent_otel_span: Any = None, local_only: bool = False, *, - model_type: Type[T], + model_type: type[T], **kwargs: Any, - ) -> Optional[T]: ... + ) -> T | None: ... @overload def get_cache( @@ -62,11 +62,11 @@ class UserApiKeyCache(DualCache): key, parent_otel_span=None, local_only: bool = False, - model_type: Optional[Type[BaseModel]] = None, + model_type: type[BaseModel] | None = None, **kwargs, - ) -> Union[Any, Optional[BaseModel]]: + ) -> Any | BaseModel | None: if model_type is None and "model_type" in kwargs: - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) cached = super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs) if model_type is None: return cached @@ -89,9 +89,9 @@ class UserApiKeyCache(DualCache): parent_otel_span: Any = None, local_only: bool = False, *, - model_type: Type[T], + model_type: type[T], **kwargs: Any, - ) -> Optional[T]: ... + ) -> T | None: ... @overload async def async_get_cache( @@ -107,11 +107,11 @@ class UserApiKeyCache(DualCache): key, parent_otel_span=None, local_only: bool = False, - model_type: Optional[Type[BaseModel]] = None, + model_type: type[BaseModel] | None = None, **kwargs, - ) -> Union[Any, Optional[BaseModel]]: + ) -> Any | BaseModel | None: if model_type is None and "model_type" in kwargs: - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) cached = await super().async_get_cache( key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs ) @@ -130,12 +130,12 @@ class UserApiKeyCache(DualCache): return decoded def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) payload = CacheCodec.serialize(value, model_type=model_type) return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs) async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) payload = CacheCodec.serialize(value, model_type=model_type) return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs) @@ -180,7 +180,7 @@ def get_management_object_ttl(cache: DualCache) -> float: propagates onto ``default_in_memory_ttl`` at startup, and falls back to ``DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL`` when no default is configured. """ - configured: Optional[float] = getattr(cache, "default_in_memory_ttl", None) + configured: float | None = getattr(cache, "default_in_memory_ttl", None) if configured is not None: return configured return DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL diff --git a/litellm/proxy/compliance_checks.py b/litellm/proxy/compliance_checks.py index 445257a1e01..e843fcd5fe1 100644 --- a/litellm/proxy/compliance_checks.py +++ b/litellm/proxy/compliance_checks.py @@ -5,8 +5,6 @@ Provides guardrail-agnostic compliance validation based on guardrail modes and execution results rather than specific guardrail names. """ -from typing import Dict, List - from litellm.types.proxy.compliance_endpoints import ( ComplianceCheckRequest, ComplianceCheckResult, @@ -28,7 +26,7 @@ class ComplianceChecker: self.data = data self.guardrails = data.guardrail_information or [] - def _get_guardrails_by_mode(self, mode: str) -> List[Dict]: + def _get_guardrails_by_mode(self, mode: str) -> list[dict]: """ Get all guardrails that ran in a specific mode. @@ -79,7 +77,7 @@ class ComplianceChecker: return all(_branch_runs_in_mode(branch) for branch in [default, *tag_branches]) return False - def _has_guardrail_intervention(self, guardrails: List[Dict]) -> bool: + def _has_guardrail_intervention(self, guardrails: list[dict]) -> bool: """Check if any guardrail intervened (blocked/masked content).""" for g in guardrails: status = g.get("guardrail_status", "") @@ -87,7 +85,7 @@ class ComplianceChecker: return True return False - def _all_guardrails_passed(self, guardrails: List[Dict]) -> bool: + def _all_guardrails_passed(self, guardrails: list[dict]) -> bool: """Check if all guardrails passed (no issues detected).""" if not guardrails: return False @@ -210,7 +208,7 @@ class ComplianceChecker: # ── Main Compliance Check Methods ──────────────────────────────────────── - def check_eu_ai_act(self) -> List[ComplianceCheckResult]: + def check_eu_ai_act(self) -> list[ComplianceCheckResult]: """ Check EU AI Act compliance. @@ -226,7 +224,7 @@ class ComplianceChecker: self._check_art_12_audit_complete(), ] - def check_gdpr(self) -> List[ComplianceCheckResult]: + def check_gdpr(self) -> list[ComplianceCheckResult]: """ Check GDPR compliance. diff --git a/litellm/proxy/container_endpoints/endpoints.py b/litellm/proxy/container_endpoints/endpoints.py index fc1f77bb684..c19573ed9b1 100644 --- a/litellm/proxy/container_endpoints/endpoints.py +++ b/litellm/proxy/container_endpoints/endpoints.py @@ -1,6 +1,6 @@ #### Container Endpoints ##### -from typing import Any, Dict +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import ORJSONResponse @@ -208,7 +208,7 @@ async def list_containers( # Read query parameters query_params = dict(request.query_params) - data: Dict[str, Any] = {"query_params": query_params} + data: dict[str, Any] = {"query_params": query_params} # Extract custom_llm_provider using priority chain custom_llm_provider = ( @@ -312,7 +312,7 @@ async def retrieve_container( ) # Include container_id in request data - data: Dict[str, Any] = {"container_id": container_id} + data: dict[str, Any] = {"container_id": container_id} # Extract custom_llm_provider using priority chain custom_llm_provider = ( @@ -417,7 +417,7 @@ async def delete_container( ) # Include container_id in request data - data: Dict[str, Any] = {"container_id": container_id} + data: dict[str, Any] = {"container_id": container_id} # Extract custom_llm_provider using priority chain custom_llm_provider = ( diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index bc871479356..c879aa72e58 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -7,7 +7,7 @@ FastAPI route handlers for ALL container file endpoints. import json from pathlib import Path -from typing import Any, Dict, List +from typing import Any from fastapi import APIRouter, Depends, Request, Response from fastapi.responses import ORJSONResponse @@ -25,14 +25,14 @@ from litellm.proxy.container_endpoints.ownership import ( ) -def _load_endpoints_config() -> Dict: +def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" config_path = Path(__file__).parent.parent.parent / "containers" / "endpoints.json" with open(config_path) as f: return json.load(f) -def get_all_route_types() -> List[str]: +def get_all_route_types() -> list[str]: """Get all async route types for registration in route_llm_request.py""" config = _load_endpoints_config() return [endpoint["async_name"] for endpoint in config["endpoints"]] @@ -52,7 +52,7 @@ def _get_container_provider_config(custom_llm_provider: str): def _create_handler_for_path_params( - path_params: List[str], + path_params: list[str], route_type: str, returns_binary: bool = False, is_multipart: bool = False, @@ -194,7 +194,7 @@ async def _process_binary_request( user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, ) - data: Dict[str, Any] = { + data: dict[str, Any] = { "file_id": file_id, **( await get_container_forwarding_params( @@ -356,7 +356,7 @@ async def _process_request( fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth, route_type: str, - path_params: Dict[str, str], + path_params: dict[str, str], ): """Common request processing logic.""" from litellm.proxy.proxy_server import ( @@ -374,7 +374,7 @@ async def _process_request( ) query_params = dict(request.query_params) - data: Dict[str, Any] = { + data: dict[str, Any] = { "query_params": query_params, **path_params, } diff --git a/litellm/proxy/container_endpoints/ownership.py b/litellm/proxy/container_endpoints/ownership.py index 583d8db7d30..5b168ff5a26 100644 --- a/litellm/proxy/container_endpoints/ownership.py +++ b/litellm/proxy/container_endpoints/ownership.py @@ -1,5 +1,5 @@ import json -from typing import Any, Dict, List, Optional, Set, Tuple +from typing import Any from fastapi import HTTPException @@ -39,7 +39,7 @@ _CONTAINER_STORED_ID_CACHE = InMemoryCache(max_size_in_memory=10000, default_ttl _ALLOWED_CONTAINER_IDS_CACHE = InMemoryCache(max_size_in_memory=2048, default_ttl=60) -def _allowed_container_ids_cache_key(owner_scopes: List[str]) -> str: +def _allowed_container_ids_cache_key(owner_scopes: list[str]) -> str: """JSON-encode the sorted scope list — using a separator like ``|`` would collide for any tenant whose user_id / team_id / org_id / api_key happens to contain the separator. JSON quoting escapes @@ -51,7 +51,7 @@ def _container_model_object_id(original_container_id: str, custom_llm_provider: return f"{CONTAINER_OBJECT_PURPOSE}:{custom_llm_provider}:{original_container_id}" -def decode_container_id_for_ownership(container_id: str, custom_llm_provider: str) -> Tuple[str, str]: +def decode_container_id_for_ownership(container_id: str, custom_llm_provider: str) -> tuple[str, str]: decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) original_container_id = decoded.get("response_id", container_id) decoded_provider = decoded.get("custom_llm_provider") @@ -62,7 +62,7 @@ def decode_container_id_for_ownership(container_id: str, custom_llm_provider: st async def get_container_forwarding_params( container_id: str, original_container_id: str, custom_llm_provider: str -) -> Dict[str, str]: +) -> dict[str, str]: params = { "container_id": original_container_id, "custom_llm_provider": custom_llm_provider, @@ -86,7 +86,7 @@ async def get_container_forwarding_params( return params -def _get_response_id(response: Any) -> Optional[str]: +def _get_response_id(response: Any) -> str | None: if response is None: return None if isinstance(response, dict): @@ -96,7 +96,7 @@ def _get_response_id(response: Any) -> Optional[str]: return value if isinstance(value, str) else None -def _dump_response(response: Any) -> Dict[str, Any]: +def _dump_response(response: Any) -> dict[str, Any]: if isinstance(response, dict): return dict(response) if hasattr(response, "model_dump"): @@ -116,7 +116,7 @@ def _custom_llm_provider_from_responses_response( response: Any, default: str = "openai", ) -> str: - hidden_params: Dict[str, Any] = {} + hidden_params: dict[str, Any] = {} if isinstance(response, dict): hidden_params = response.get("_hidden_params") or {} else: @@ -131,7 +131,7 @@ def _custom_llm_provider_from_responses_response( async def record_container_owners_from_responses_response( response: Any, user_api_key_dict: UserAPIKeyAuth, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, ) -> None: """Track containers created implicitly by code interpreter in /v1/responses.""" container_ids = ResponsesAPIRequestUtils.collect_container_ids_from_responses_response(response) @@ -234,7 +234,7 @@ async def record_container_owner( return response -async def _get_container_owner(original_container_id: str, custom_llm_provider: str) -> Optional[str]: +async def _get_container_owner(original_container_id: str, custom_llm_provider: str) -> str | None: model_object_id = _container_model_object_id(original_container_id, custom_llm_provider) cached = _CONTAINER_OWNER_CACHE.get_cache(model_object_id) @@ -263,7 +263,7 @@ async def _get_container_owner(original_container_id: str, custom_llm_provider: return owner -async def _get_stored_container_id(original_container_id: str, custom_llm_provider: str) -> Optional[str]: +async def _get_stored_container_id(original_container_id: str, custom_llm_provider: str) -> str | None: """Return the ``unified_object_id`` stored at create time, if any. Used by :func:`get_container_forwarding_params` to recover the @@ -301,7 +301,7 @@ async def assert_user_can_access_container( container_id: str, user_api_key_dict: UserAPIKeyAuth, custom_llm_provider: str, -) -> Tuple[str, str]: +) -> tuple[str, str]: original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider) if is_proxy_admin(user_api_key_dict): @@ -317,7 +317,7 @@ async def assert_user_can_access_container( return original_container_id, resolved_provider -def _get_container_list_data(response: Any) -> Optional[List[Any]]: +def _get_container_list_data(response: Any) -> list[Any] | None: if response is None: return None if isinstance(response, dict): @@ -327,7 +327,7 @@ def _get_container_list_data(response: Any) -> Optional[List[Any]]: return data if isinstance(data, list) else None -def _set_container_list_data(response: Any, data: List[Any], removed_filtered_items: bool = False) -> Any: +def _set_container_list_data(response: Any, data: list[Any], removed_filtered_items: bool = False) -> Any: if isinstance(response, dict): response["data"] = data if data: @@ -353,7 +353,7 @@ def _set_container_list_data(response: Any, data: List[Any], removed_filtered_it async def _get_allowed_container_ids( user_api_key_dict: UserAPIKeyAuth, -) -> Set[str]: +) -> set[str]: owner_scopes = get_resource_owner_scopes(user_api_key_dict) if not owner_scopes: return set() @@ -394,7 +394,7 @@ async def filter_container_list_response( return response allowed_container_ids = await _get_allowed_container_ids(user_api_key_dict) - filtered: List[Any] = [] + filtered: list[Any] = [] for item in data: container_id = _get_response_id(item) if container_id is None: diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index e94016c555e..b89c3c2f33f 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -2,9 +2,7 @@ CRUD endpoints for storing reusable credentials. """ -from typing import Optional - -from fastapi import APIRouter, Depends, HTTPException, Request, Response, Path +from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response import litellm from litellm._logging import verbose_proxy_logger @@ -22,9 +20,7 @@ router = APIRouter() class CredentialHelperUtils: @staticmethod - def encrypt_credential_values( - credential: CredentialItem, new_encryption_key: Optional[str] = None - ) -> CredentialItem: + def encrypt_credential_values(credential: CredentialItem, new_encryption_key: str | None = None) -> CredentialItem: """Encrypt values in credential.credential_values and add to DB""" encrypted_credential_values = {} for key, value in (credential.credential_values or {}).items(): @@ -204,7 +200,7 @@ async def get_credential_by_model( number_of_asterisks=4, ) credential = CredentialItem( - credential_name="{}-credential-{}".format(model.model_name, model_id), + credential_name=f"{model.model_name}-credential-{model_id}", credential_values=masked_credential_values, credential_info={}, ) @@ -248,7 +244,7 @@ async def delete_credential( def update_db_credential( db_credential: CredentialItem, updated_patch: CredentialItem, - new_encryption_key: Optional[str] = None, + new_encryption_key: str | None = None, ) -> CredentialItem: """ Update a credential in the DB. @@ -323,7 +319,7 @@ async def update_credential( # Sync in-memory credential_list (skip if not in memory - e.g., proxy restarted) new_name = merged_credential.credential_name - existing_in_memory: Optional[CredentialItem] = None + existing_in_memory: CredentialItem | None = None for cred in litellm.credential_list: if cred.credential_name == credential_name: existing_in_memory = cred diff --git a/litellm/proxy/custom_auth_auto.py b/litellm/proxy/custom_auth_auto.py index 97e928b6546..8ac8443adf6 100644 --- a/litellm/proxy/custom_auth_auto.py +++ b/litellm/proxy/custom_auth_auto.py @@ -4,14 +4,12 @@ Example custom auth function. This will allow all keys starting with "my-custom-key" to pass through. """ -from typing import Union - from fastapi import Request from litellm.proxy._types import ProxyException, UserAPIKeyAuth -async def user_api_key_auth(request: Request, api_key: str) -> Union[UserAPIKeyAuth, str]: +async def user_api_key_auth(request: Request, api_key: str) -> UserAPIKeyAuth | str: try: if api_key.startswith("my-custom-key"): return "sk-P1zJMdsqCPNN54alZd_ETw" diff --git a/litellm/proxy/custom_prompt_management.py b/litellm/proxy/custom_prompt_management.py index 62a51911409..355edb69897 100644 --- a/litellm/proxy/custom_prompt_management.py +++ b/litellm/proxy/custom_prompt_management.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Tuple - from litellm._logging import verbose_logger from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.types.llms.openai import AllMessageValues @@ -11,17 +9,17 @@ class X42PromptManagement(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, - ignore_prompt_manager_model: Optional[bool] = False, - ignore_prompt_manager_optional_params: Optional[bool] = False, - ) -> Tuple[str, List[AllMessageValues], dict]: + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: """ Returns: - model: str - the model to use (can be pulled from prompt management tool) diff --git a/litellm/proxy/custom_sso.py b/litellm/proxy/custom_sso.py index 4f621a9659e..53d087c1964 100644 --- a/litellm/proxy/custom_sso.py +++ b/litellm/proxy/custom_sso.py @@ -14,8 +14,8 @@ Flow: from fastapi_sso.sso.base import OpenID -from litellm.proxy._types import LitellmUserRoles, SSOUserDefinedValues from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles, SSOUserDefinedValues async def custom_sso_handler(userIDPInfo: OpenID) -> SSOUserDefinedValues: diff --git a/litellm/proxy/db/base_client.py b/litellm/proxy/db/base_client.py index 6e5fff44c79..127581866bb 100644 --- a/litellm/proxy/db/base_client.py +++ b/litellm/proxy/db/base_client.py @@ -1,4 +1,4 @@ -from typing import Any, Literal, List +from typing import Any, Literal class CustomDB: @@ -13,21 +13,18 @@ class CustomDB: """ Check if key valid """ - pass def insert_data(self, value: Any, table_name: Literal["user", "key", "config"]): """ For new key / user logic """ - pass def update_data(self, key: str, value: Any, table_name: Literal["user", "key", "config"]): """ For cost tracking logic """ - pass - def delete_data(self, keys: List[str], table_name: Literal["user", "key", "config"]): + def delete_data(self, keys: list[str], table_name: Literal["user", "key", "config"]): """ For /key/delete endpoint s """ @@ -38,7 +35,6 @@ class CustomDB: """ For connecting to db and creating / updating any tables """ - pass def disconnect( self, @@ -46,4 +42,3 @@ class CustomDB: """ For closing connection on server shutdown """ - pass diff --git a/litellm/proxy/db/check_migration.py b/litellm/proxy/db/check_migration.py index 2aacaed8aff..5e53e118e45 100644 --- a/litellm/proxy/db/check_migration.py +++ b/litellm/proxy/db/check_migration.py @@ -2,12 +2,11 @@ import os import subprocess -from typing import List, Optional, Tuple from litellm._logging import verbose_logger -def extract_sql_commands(diff_output: str) -> List[str]: +def extract_sql_commands(diff_output: str) -> list[str]: """ Extract SQL commands from the Prisma migrate diff output. Args: @@ -44,7 +43,7 @@ def extract_sql_commands(diff_output: str) -> List[str]: return sql_commands -def check_prisma_schema_diff_helper(db_url: str) -> Tuple[bool, List[str]]: +def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]: """Checks for differences between current database and Prisma schema. Returns: A tuple containing: @@ -89,7 +88,7 @@ def check_prisma_schema_diff_helper(db_url: str) -> Tuple[bool, List[str]]: return False, [] -def check_prisma_schema_diff(db_url: Optional[str] = None) -> None: +def check_prisma_schema_diff(db_url: str | None = None) -> None: """Main function to run the Prisma schema diff check.""" if db_url is None: db_url = os.getenv("DATABASE_URL") @@ -98,7 +97,5 @@ def check_prisma_schema_diff(db_url: Optional[str] = None) -> None: has_diff, message = check_prisma_schema_diff_helper(db_url) if has_diff: verbose_logger.exception( - "🚨🚨🚨 prisma schema out of sync with db. Consider running these sql_commands to sync the two - {}".format( - message - ) + f"🚨🚨🚨 prisma schema out of sync with db. Consider running these sql_commands to sync the two - {message}" ) diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 411faa7e4e0..3ced1589757 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -222,8 +222,6 @@ async def create_missing_views(db: _db): verbose_logger.debug("Last30dTopEndUsersSpend Created!") - return - async def should_create_missing_views(db: _db) -> bool: """ @@ -240,7 +238,7 @@ async def should_create_missing_views(db: _db) -> bool: result = await db.query_raw(query=sql_query) - verbose_logger.debug("Estimated Row count of LiteLLM_SpendLogs = {}".format(result)) + verbose_logger.debug(f"Estimated Row count of LiteLLM_SpendLogs = {result}") if ( result and isinstance(result, list) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index fd8132fef22..17410698aed 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -16,11 +16,7 @@ from datetime import datetime, timedelta, timezone from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Union, cast, overload, ) @@ -29,8 +25,8 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache from litellm.constants import ( - DB_SPEND_UPDATE_JOB_NAME, DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, + DB_SPEND_UPDATE_JOB_NAME, ) from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy._types import ( @@ -107,7 +103,7 @@ class DBSpendUpdateWriter: def __init__( self, - redis_cache: Optional[RedisCache] = None, + redis_cache: RedisCache | None = None, ): self.redis_cache = redis_cache self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache) @@ -124,17 +120,17 @@ class DBSpendUpdateWriter: async def update_database( # LiteLLM management object fields self, - token: Optional[str], - user_id: Optional[str], - end_user_id: Optional[str], - team_id: Optional[str], - org_id: Optional[str], + token: str | None, + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + org_id: str | None, # Completion object fields - kwargs: Optional[dict], - completion_response: Optional[Union[litellm.ModelResponse, Any, Exception]], - start_time: Optional[datetime], - end_time: Optional[datetime], - response_cost: Optional[float], + kwargs: dict | None, + completion_response: litellm.ModelResponse | Any | Exception | None, + start_time: datetime | None, + end_time: datetime | None, + response_cost: float | None, ): from litellm.proxy.proxy_server import ( disable_spend_logs, @@ -261,10 +257,10 @@ class DBSpendUpdateWriter: def _enqueue_tool_registry_upsert( self, - kwargs: Optional[dict], - completion_response: Optional[Any], - hashed_token: Optional[str] = None, - team_id: Optional[str] = None, + kwargs: dict | None, + completion_response: Any | None, + hashed_token: str | None = None, + team_id: str | None = None, ) -> None: """ Extract tool names from the LLM request and response and enqueue them @@ -283,7 +279,7 @@ class DBSpendUpdateWriter: return # Extract key_alias from kwargs metadata if available - key_alias: Optional[str] = None + key_alias: str | None = None _litellm_params = kwargs.get("litellm_params") or {} _metadata = _litellm_params.get("metadata") or {} key_alias = _metadata.get("user_api_key_alias") or None @@ -345,14 +341,14 @@ class DBSpendUpdateWriter: async def _batch_database_updates( self, *, - response_cost: Optional[float], - user_id: Optional[str], - hashed_token: Optional[str], - team_id: Optional[str], - org_id: Optional[str], - end_user_id: Optional[str], - prisma_client: Optional[PrismaClient], - litellm_proxy_budget_name: Optional[str], + response_cost: float | None, + user_id: str | None, + hashed_token: str | None, + team_id: str | None, + org_id: str | None, + end_user_id: str | None, + prisma_client: PrismaClient | None, + litellm_proxy_budget_name: str | None, payload: SpendLogsPayload, ): """ @@ -510,9 +506,9 @@ class DBSpendUpdateWriter: async def _update_key_db( self, - response_cost: Optional[float], - hashed_token: Optional[str], - prisma_client: Optional[PrismaClient], + response_cost: float | None, + hashed_token: str | None, + prisma_client: PrismaClient | None, ): try: if hashed_token is None or prisma_client is None: @@ -531,11 +527,11 @@ class DBSpendUpdateWriter: async def _update_user_db( self, - response_cost: Optional[float], - user_id: Optional[str], - prisma_client: Optional[PrismaClient], - litellm_proxy_budget_name: Optional[str], - end_user_id: Optional[str] = None, + response_cost: float | None, + user_id: str | None, + prisma_client: PrismaClient | None, + litellm_proxy_budget_name: str | None, + end_user_id: str | None = None, ): """ - Update that user's row @@ -578,10 +574,10 @@ class DBSpendUpdateWriter: async def _update_team_db( self, - response_cost: Optional[float], - team_id: Optional[str], - user_id: Optional[str], - prisma_client: Optional[PrismaClient], + response_cost: float | None, + team_id: str | None, + user_id: str | None, + prisma_client: PrismaClient | None, ): try: if team_id is None or prisma_client is None: @@ -632,9 +628,9 @@ class DBSpendUpdateWriter: async def _update_org_db( self, - response_cost: Optional[float], - org_id: Optional[str], - prisma_client: Optional[PrismaClient], + response_cost: float | None, + org_id: str | None, + prisma_client: PrismaClient | None, ): try: if org_id is None or prisma_client is None: @@ -662,9 +658,9 @@ class DBSpendUpdateWriter: async def _update_agent_db( self, - response_cost: Optional[float], - agent_id: Optional[str], - prisma_client: Optional[PrismaClient], + response_cost: float | None, + agent_id: str | None, + prisma_client: PrismaClient | None, ): try: if agent_id is None or prisma_client is None: @@ -689,9 +685,9 @@ class DBSpendUpdateWriter: async def _update_tag_db( self, - response_cost: Optional[float], - request_tags: Optional[str], - prisma_client: Optional[PrismaClient], + response_cost: float | None, + request_tags: str | None, + prisma_client: PrismaClient | None, ): """ Update spend for all tags in the request. @@ -739,19 +735,16 @@ class DBSpendUpdateWriter: async def _insert_spend_log_to_db( self, - payload: Union[dict, SpendLogsPayload], - prisma_client: Optional[PrismaClient] = None, - spend_logs_url: Optional[str] = os.getenv("SPEND_LOGS_URL"), - ) -> Optional[PrismaClient]: + payload: dict | SpendLogsPayload, + prisma_client: PrismaClient | None = None, + spend_logs_url: str | None = os.getenv("SPEND_LOGS_URL"), + ) -> PrismaClient | None: verbose_proxy_logger.debug( "Writing spend log to db - request_id: {}, spend: {}".format( payload.get("request_id"), payload.get("spend") ) ) - if prisma_client is not None and spend_logs_url is not None: - async with prisma_client._spend_log_transactions_lock: - prisma_client.spend_log_transactions.append(payload) - elif prisma_client is not None: + if prisma_client is not None and spend_logs_url is not None or prisma_client is not None: async with prisma_client._spend_log_transactions_lock: prisma_client.spend_log_transactions.append(payload) else: @@ -934,7 +927,7 @@ class DBSpendUpdateWriter: ################## Daily Spend Update Transactions ################## # Aggregate all in memory daily spend transactions and commit to db daily_spend_update_transactions = cast( - Dict[str, DailyUserSpendTransaction], + dict[str, DailyUserSpendTransaction], await self.daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -948,7 +941,7 @@ class DBSpendUpdateWriter: ################## Daily Team Spend Update Transactions ################## # Aggregate all in memory daily team spend transactions and commit to db daily_team_spend_update_transactions = cast( - Dict[str, DailyTeamSpendTransaction], + dict[str, DailyTeamSpendTransaction], await self.daily_team_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -962,7 +955,7 @@ class DBSpendUpdateWriter: ################## Daily Organization Spend Update Transactions ################## # Aggregate all in memory daily org spend transactions and commit to db daily_org_spend_update_transactions = cast( - Dict[str, DailyOrganizationSpendTransaction], + dict[str, DailyOrganizationSpendTransaction], await self.daily_org_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -978,7 +971,7 @@ class DBSpendUpdateWriter: ################## Daily End-User Spend Update Transactions ################## # Aggregate all in memory daily end-user spend transactions and commit to db daily_end_user_spend_update_transactions = cast( - Dict[str, DailyEndUserSpendTransaction], + dict[str, DailyEndUserSpendTransaction], await self.daily_end_user_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -992,7 +985,7 @@ class DBSpendUpdateWriter: ################## Daily Agent Spend Update Transactions ################## # Aggregate all in memory daily agent spend transactions and commit to db daily_agent_spend_update_transactions = cast( - Dict[str, DailyAgentSpendTransaction], + dict[str, DailyAgentSpendTransaction], await self.daily_agent_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -1017,7 +1010,7 @@ class DBSpendUpdateWriter: This is called by a separate scheduler job at a longer interval. """ daily_tag_spend_update_transactions = cast( - Dict[str, DailyTagSpendTransaction], + dict[str, DailyTagSpendTransaction], await self.daily_tag_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), ) @@ -1105,7 +1098,7 @@ class DBSpendUpdateWriter: ### UPDATE USER TABLE ### user_list_transactions = db_spend_update_transactions["user_list_transactions"] - verbose_proxy_logger.debug("User Spend transactions: {}".format(user_list_transactions)) + verbose_proxy_logger.debug(f"User Spend transactions: {user_list_transactions}") if user_list_transactions is not None and len(user_list_transactions.keys()) > 0: for i in range(n_retry_times + 1): start_time = time.time() @@ -1137,7 +1130,7 @@ class DBSpendUpdateWriter: ### UPDATE END-USER TABLE ### end_user_list_transactions = db_spend_update_transactions["end_user_list_transactions"] - verbose_proxy_logger.debug("End-User Spend transactions: {}".format(end_user_list_transactions)) + verbose_proxy_logger.debug(f"End-User Spend transactions: {end_user_list_transactions}") if end_user_list_transactions is not None and len(end_user_list_transactions.keys()) > 0: await ProxyUpdateSpend.update_end_user_spend( n_retry_times=n_retry_times, @@ -1147,7 +1140,7 @@ class DBSpendUpdateWriter: ) ### UPDATE KEY TABLE ### key_list_transactions = db_spend_update_transactions["key_list_transactions"] - verbose_proxy_logger.debug("KEY Spend transactions: {}".format(key_list_transactions)) + verbose_proxy_logger.debug(f"KEY Spend transactions: {key_list_transactions}") if key_list_transactions is not None and len(key_list_transactions.keys()) > 0: for i in range(n_retry_times + 1): start_time = time.time() @@ -1180,7 +1173,7 @@ class DBSpendUpdateWriter: ### UPDATE TEAM TABLE ### team_list_transactions = db_spend_update_transactions["team_list_transactions"] - verbose_proxy_logger.debug("Team Spend transactions: {}".format(team_list_transactions)) + verbose_proxy_logger.debug(f"Team Spend transactions: {team_list_transactions}") if team_list_transactions is not None and len(team_list_transactions.keys()) > 0: for i in range(n_retry_times + 1): start_time = time.time() @@ -1189,9 +1182,7 @@ class DBSpendUpdateWriter: async with transaction.batch_() as batcher: # Sort by team_id for consistent lock ordering across pods to prevent deadlocks. for team_id, response_cost in sorted(team_list_transactions.items()): - verbose_proxy_logger.debug( - "Updating spend for team id={} by {}".format(team_id, response_cost) - ) + verbose_proxy_logger.debug(f"Updating spend for team id={team_id} by {response_cost}") batcher.litellm_teamtable.update_many( # 'update_many' prevents error from being raised if no row exists where={"team_id": team_id}, data={"spend": {"increment": response_cost}}, @@ -1213,10 +1204,10 @@ class DBSpendUpdateWriter: ### UPDATE TEAM Membership TABLE with spend ### team_member_list_transactions = db_spend_update_transactions["team_member_list_transactions"] - verbose_proxy_logger.debug("Team Membership Spend transactions: {}".format(team_member_list_transactions)) + verbose_proxy_logger.debug(f"Team Membership Spend transactions: {team_member_list_transactions}") if team_member_list_transactions is not None and len(team_member_list_transactions.keys()) > 0: # Track which team memberships will be updated for cache invalidation - team_memberships_to_invalidate: List[tuple[str, str]] = [] + team_memberships_to_invalidate: list[tuple[str, str]] = [] for key in team_member_list_transactions.keys(): # key is "team_id::::user_id::" team_id = key.split("::")[1] @@ -1264,7 +1255,7 @@ class DBSpendUpdateWriter: user_api_key_cache = proxy_logging_obj.call_details.get("user_api_key_cache") if user_api_key_cache is not None: for user_id, team_id in team_memberships_to_invalidate: - cache_key = "team_membership:{}:{}".format(user_id, team_id) + cache_key = f"team_membership:{user_id}:{team_id}" await user_api_key_cache.async_delete_cache(key=cache_key) verbose_proxy_logger.debug( f"Invalidated team membership cache for user_id={user_id}, team_id={team_id}" @@ -1272,7 +1263,7 @@ class DBSpendUpdateWriter: ### UPDATE ORG TABLE ### org_list_transactions = db_spend_update_transactions["org_list_transactions"] - verbose_proxy_logger.debug("Org Spend transactions: {}".format(org_list_transactions)) + verbose_proxy_logger.debug(f"Org Spend transactions: {org_list_transactions}") if org_list_transactions is not None and len(org_list_transactions.keys()) > 0: for i in range(n_retry_times + 1): start_time = time.time() @@ -1334,7 +1325,7 @@ class DBSpendUpdateWriter: @staticmethod async def _update_entity_spend_in_db( entity_name: str, - transactions: Optional[Dict[str, float]], + transactions: dict[str, float] | None, table_accessor: Any, where_field: str, n_retry_times: int, @@ -1393,7 +1384,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyUserSpendTransaction], + daily_spend_transactions: dict[str, DailyUserSpendTransaction], entity_type: Literal["user"], entity_id_field: str, table_name: str, @@ -1407,7 +1398,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyTeamSpendTransaction], + daily_spend_transactions: dict[str, DailyTeamSpendTransaction], entity_type: Literal["team"], entity_id_field: str, table_name: str, @@ -1421,7 +1412,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyOrganizationSpendTransaction], + daily_spend_transactions: dict[str, DailyOrganizationSpendTransaction], entity_type: Literal["org"], entity_id_field: str, table_name: str, @@ -1435,7 +1426,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyEndUserSpendTransaction], + daily_spend_transactions: dict[str, DailyEndUserSpendTransaction], entity_type: Literal["end_user"], entity_id_field: str, table_name: str, @@ -1449,7 +1440,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyAgentSpendTransaction], + daily_spend_transactions: dict[str, DailyAgentSpendTransaction], entity_type: Literal["agent"], entity_id_field: str, table_name: str, @@ -1463,7 +1454,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyTagSpendTransaction], + daily_spend_transactions: dict[str, DailyTagSpendTransaction], entity_type: Literal["tag"], entity_id_field: str, table_name: str, @@ -1477,14 +1468,12 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Union[ - Dict[str, DailyUserSpendTransaction], - Dict[str, DailyTeamSpendTransaction], - Dict[str, DailyTagSpendTransaction], - Dict[str, DailyOrganizationSpendTransaction], - Dict[str, DailyEndUserSpendTransaction], - Dict[str, DailyAgentSpendTransaction], - ], + daily_spend_transactions: dict[str, DailyUserSpendTransaction] + | dict[str, DailyTeamSpendTransaction] + | dict[str, DailyTagSpendTransaction] + | dict[str, DailyOrganizationSpendTransaction] + | dict[str, DailyEndUserSpendTransaction] + | dict[str, DailyAgentSpendTransaction], entity_type: Literal["user", "team", "org", "tag", "end_user", "agent"], entity_id_field: str, table_name: str, @@ -1664,7 +1653,7 @@ class DBSpendUpdateWriter: ) # Remove processed transactions - for key in transactions_to_process.keys(): + for key in transactions_to_process: daily_spend_transactions.pop(key, None) break @@ -1687,7 +1676,7 @@ class DBSpendUpdateWriter: except Exception as e: if "transactions_to_process" in locals(): - for key in transactions_to_process.keys(): # type: ignore + for key in transactions_to_process: # type: ignore daily_spend_transactions.pop(key, None) _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) @@ -1696,7 +1685,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyUserSpendTransaction], + daily_spend_transactions: dict[str, DailyUserSpendTransaction], ): """ Batch job to update LiteLLM_DailyUserSpend table using in-memory daily_spend_transactions @@ -1717,7 +1706,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyTeamSpendTransaction], + daily_spend_transactions: dict[str, DailyTeamSpendTransaction], ): """ Batch job to update LiteLLM_DailyTeamSpend table using in-memory daily_spend_transactions @@ -1738,7 +1727,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyOrganizationSpendTransaction], + daily_spend_transactions: dict[str, DailyOrganizationSpendTransaction], ): """ Batch job to update LiteLLM_DailyOrganizationSpend table using in-memory daily_spend_transactions @@ -1759,7 +1748,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyEndUserSpendTransaction], + daily_spend_transactions: dict[str, DailyEndUserSpendTransaction], ): """ Batch job to update LiteLLM_DailyEndUserSpend table using in-memory daily_spend_transactions @@ -1780,7 +1769,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyAgentSpendTransaction], + daily_spend_transactions: dict[str, DailyAgentSpendTransaction], ): """ Batch job to update LiteLLM_DailyAgentSpend table using in-memory daily_spend_transactions @@ -1801,7 +1790,7 @@ class DBSpendUpdateWriter: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - daily_spend_transactions: Dict[str, DailyTagSpendTransaction], + daily_spend_transactions: dict[str, DailyTagSpendTransaction], ): """ Batch job to update LiteLLM_DailyTagSpend table using in-memory daily_spend_transactions @@ -1819,10 +1808,10 @@ class DBSpendUpdateWriter: async def _common_add_spend_log_transaction_to_daily_transaction( self, - payload: Union[dict, SpendLogsPayload], + payload: dict | SpendLogsPayload, prisma_client: PrismaClient, type: Literal["user", "team", "org", "request_tags", "end_user", "agent"] = "user", - ) -> Optional[BaseDailySpendTransaction]: + ) -> BaseDailySpendTransaction | None: common_expected_keys = ["startTime", "api_key"] if type == "user": expected_keys = ["user", *common_expected_keys] @@ -1914,8 +1903,8 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_user_transaction( self, - payload: Union[dict, SpendLogsPayload], - prisma_client: Optional[PrismaClient] = None, + payload: dict | SpendLogsPayload, + prisma_client: PrismaClient | None = None, ): """ Add a spend log transaction to the `daily_spend_update_queue` @@ -1942,7 +1931,7 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_team_transaction( self, payload: SpendLogsPayload, - prisma_client: Optional[PrismaClient] = None, + prisma_client: PrismaClient | None = None, ) -> None: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") @@ -1965,8 +1954,8 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_org_transaction( self, payload: SpendLogsPayload, - prisma_client: Optional[PrismaClient] = None, - org_id: Optional[str] = None, + prisma_client: PrismaClient | None = None, + org_id: str | None = None, ) -> None: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") @@ -1998,7 +1987,7 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_end_user_transaction( self, payload: SpendLogsPayload, - prisma_client: Optional[PrismaClient] = None, + prisma_client: PrismaClient | None = None, ) -> None: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") @@ -2031,7 +2020,7 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_agent_transaction( self, payload: SpendLogsPayload, - prisma_client: Optional[PrismaClient] = None, + prisma_client: PrismaClient | None = None, ) -> None: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") @@ -2058,7 +2047,7 @@ class DBSpendUpdateWriter: async def add_spend_log_transaction_to_daily_tag_transaction( self, payload: SpendLogsPayload, - prisma_client: Optional[PrismaClient] = None, + prisma_client: PrismaClient | None = None, ) -> None: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/db/db_transaction_queue/base_update_queue.py b/litellm/proxy/db/db_transaction_queue/base_update_queue.py index fb6010e7d21..3bd15e89ee5 100644 --- a/litellm/proxy/db/db_transaction_queue/base_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/base_update_queue.py @@ -3,7 +3,6 @@ Base class for in memory buffer for database transactions """ import asyncio -from typing import Optional from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging @@ -51,7 +50,6 @@ class BaseUpdateQueue: async def _emit_new_item_added_to_queue_event( self, - queue_size: Optional[int] = None, + queue_size: int | None = None, ): """placeholder, emit event when a new item is added to the queue""" - pass diff --git a/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py index b6462636393..18df25093f1 100644 --- a/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py @@ -1,6 +1,5 @@ import asyncio from copy import deepcopy -from typing import Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm.constants import LITELLM_ASYNCIO_QUEUE_MAXSIZE @@ -54,11 +53,11 @@ class DailySpendUpdateQueue(BaseUpdateQueue): def __init__(self): super().__init__() - self.update_queue: asyncio.Queue[Dict[str, BaseDailySpendTransaction]] = asyncio.Queue( + self.update_queue: asyncio.Queue[dict[str, BaseDailySpendTransaction]] = asyncio.Queue( maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE ) - async def add_update(self, update: Dict[str, BaseDailySpendTransaction]): + async def add_update(self, update: dict[str, BaseDailySpendTransaction]): """Enqueue an update.""" verbose_proxy_logger.debug("Adding update to queue: %s", update) await self.update_queue.put(update) @@ -73,13 +72,13 @@ class DailySpendUpdateQueue(BaseUpdateQueue): Combine all updates in the queue into a single update. This is used to reduce the size of the in-memory queue. """ - updates: List[Dict[str, BaseDailySpendTransaction]] = await self.flush_all_updates_from_in_memory_queue() + updates: list[dict[str, BaseDailySpendTransaction]] = await self.flush_all_updates_from_in_memory_queue() aggregated_updates = self.get_aggregated_daily_spend_update_transactions(updates) await self.update_queue.put(aggregated_updates) async def flush_and_get_aggregated_daily_spend_update_transactions( self, - ) -> Dict[str, BaseDailySpendTransaction]: + ) -> dict[str, BaseDailySpendTransaction]: """Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates.""" updates = await self.flush_all_updates_from_in_memory_queue() if len(updates) > 0: @@ -98,10 +97,10 @@ class DailySpendUpdateQueue(BaseUpdateQueue): @staticmethod def get_aggregated_daily_spend_update_transactions( - updates: List[Dict[str, BaseDailySpendTransaction]], - ) -> Dict[str, BaseDailySpendTransaction]: + updates: list[dict[str, BaseDailySpendTransaction]], + ) -> dict[str, BaseDailySpendTransaction]: """Aggregate updates by daily_transaction_key.""" - aggregated_daily_spend_update_transactions: Dict[str, BaseDailySpendTransaction] = {} + aggregated_daily_spend_update_transactions: dict[str, BaseDailySpendTransaction] = {} for _update in updates: for _key, payload in _update.items(): if _key in aggregated_daily_spend_update_transactions: @@ -140,7 +139,7 @@ class DailySpendUpdateQueue(BaseUpdateQueue): async def _emit_new_item_added_to_queue_event( self, - queue_size: Optional[int] = None, + queue_size: int | None = None, ): asyncio.create_task( service_logger_obj.async_service_success_hook( diff --git a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py index e04c2ba9a4e..01f6a92485a 100644 --- a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py +++ b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py @@ -1,9 +1,9 @@ import asyncio import json -from litellm._uuid import uuid -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS from litellm.proxy.db.db_transaction_queue.base_update_queue import service_logger_obj @@ -30,10 +30,10 @@ else end """ - def __init__(self, redis_cache: Optional[RedisCache] = None): + def __init__(self, redis_cache: RedisCache | None = None): self.pod_id = str(uuid.uuid4()) self.redis_cache = redis_cache - self._release_lock_script: Optional[Any] = None + self._release_lock_script: Any | None = None @staticmethod def get_redis_lock_key(cronjob_id: str) -> str: @@ -42,8 +42,8 @@ end async def acquire_lock( self, cronjob_id: str, - ttl: Optional[int] = None, - ) -> Optional[bool]: + ttl: int | None = None, + ) -> bool | None: """ Attempt to acquire the lock for a specific cron job using Redis. Uses the SET command with NX and EX options to ensure atomicity. diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index c924448669d..69614defb9e 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -6,7 +6,7 @@ This is to prevent deadlocks and improve reliability import asyncio import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache @@ -60,7 +60,7 @@ class RedisUpdateBuffer: def __init__( self, - redis_cache: Optional[RedisCache] = None, + redis_cache: RedisCache | None = None, ): self.redis_cache = redis_cache @@ -74,9 +74,7 @@ class RedisUpdateBuffer: """ from litellm.proxy.proxy_server import general_settings - _use_redis_transaction_buffer: Optional[Union[bool, str]] = general_settings.get( - "use_redis_transaction_buffer", False - ) + _use_redis_transaction_buffer: bool | str | None = general_settings.get("use_redis_transaction_buffer", False) if isinstance(_use_redis_transaction_buffer, str): _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) if _use_redis_transaction_buffer is None: @@ -203,7 +201,7 @@ class RedisUpdateBuffer: verbose_proxy_logger.debug("ALL DAILY SPEND UPDATE TRANSACTIONS: %s", daily_spend_update_transactions) # Build a list of rpush operations, skipping empty/None transaction sets - _queue_configs: List[Tuple[Any, str, ServiceTypes]] = [ + _queue_configs: list[tuple[Any, str, ServiceTypes]] = [ ( db_spend_update_transactions, REDIS_UPDATE_BUFFER_KEY, @@ -236,8 +234,8 @@ class RedisUpdateBuffer: ), ] - rpush_list: List[RedisPipelineRpushOperation] = [] - service_types: List[ServiceTypes] = [] + rpush_list: list[RedisPipelineRpushOperation] = [] + service_types: list[ServiceTypes] = [] for transactions, redis_key, service_type in _queue_configs: if transactions is None or len(transactions) == 0: continue @@ -293,12 +291,12 @@ class RedisUpdateBuffer: @staticmethod async def _restore_spend_updates_to_in_memory_queues( - db_spend_update_transactions: Optional[DBSpendUpdateTransactions], - daily_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]], - daily_team_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]], - daily_org_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]], - daily_end_user_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]], - daily_agent_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]], + db_spend_update_transactions: DBSpendUpdateTransactions | None, + daily_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None, + daily_team_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None, + daily_org_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None, + daily_end_user_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None, + daily_agent_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None, spend_update_queue: SpendUpdateQueue, daily_spend_update_queue: DailySpendUpdateQueue, daily_team_spend_update_queue: DailySpendUpdateQueue, @@ -314,7 +312,7 @@ class RedisUpdateBuffer: because the source queues were already drained before the rpush. """ if db_spend_update_transactions is not None: - entity_entries: List[Tuple[Litellm_EntityType, Optional[Dict[str, float]]]] = [ + entity_entries: list[tuple[Litellm_EntityType, dict[str, float] | None]] = [ ( Litellm_EntityType.USER, db_spend_update_transactions.get("user_list_transactions"), @@ -360,7 +358,7 @@ class RedisUpdateBuffer: ) ) - daily_pairs: List[Tuple[Optional[Dict[str, BaseDailySpendTransaction]], DailySpendUpdateQueue]] = [ + daily_pairs: list[tuple[dict[str, BaseDailySpendTransaction] | None, DailySpendUpdateQueue]] = [ (daily_spend_update_transactions, daily_spend_update_queue), (daily_team_spend_update_transactions, daily_team_spend_update_queue), (daily_org_spend_update_transactions, daily_org_spend_update_queue), @@ -388,7 +386,7 @@ class RedisUpdateBuffer: return num_transactions @staticmethod - def _remove_prefix_from_keys(data: Dict[str, Any], prefix: str) -> Dict[str, Any]: + def _remove_prefix_from_keys(data: dict[str, Any], prefix: str) -> dict[str, Any]: """ Removes the specified prefix from the keys of a dictionary. """ @@ -396,7 +394,7 @@ class RedisUpdateBuffer: async def get_all_update_transactions_from_redis_buffer( self, - ) -> Optional[DBSpendUpdateTransactions]: + ) -> DBSpendUpdateTransactions | None: """ Gets all the update transactions from Redis @@ -463,13 +461,13 @@ class RedisUpdateBuffer: async def get_all_transactions_from_redis_buffer_pipeline( self, - ) -> Tuple[ - Optional[DBSpendUpdateTransactions], - Optional[Dict[str, DailyUserSpendTransaction]], - Optional[Dict[str, DailyTeamSpendTransaction]], - Optional[Dict[str, DailyOrganizationSpendTransaction]], - Optional[Dict[str, DailyEndUserSpendTransaction]], - Optional[Dict[str, DailyAgentSpendTransaction]], + ) -> tuple[ + DBSpendUpdateTransactions | None, + dict[str, DailyUserSpendTransaction] | None, + dict[str, DailyTeamSpendTransaction] | None, + dict[str, DailyOrganizationSpendTransaction] | None, + dict[str, DailyEndUserSpendTransaction] | None, + dict[str, DailyAgentSpendTransaction] | None, ]: """ Drains the main 6 Redis buffer queues in a single pipeline round-trip. @@ -485,7 +483,7 @@ class RedisUpdateBuffer: if self.redis_cache is None: return None, None, None, None, None, None - lpop_list: List[RedisPipelineLpopOperation] = [ + lpop_list: list[RedisPipelineLpopOperation] = [ RedisPipelineLpopOperation(key=REDIS_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), RedisPipelineLpopOperation( key=REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, @@ -516,14 +514,14 @@ class RedisUpdateBuffer: raw_results.append(None) # Slot 0: DBSpendUpdateTransactions - db_spend: Optional[DBSpendUpdateTransactions] = None + db_spend: DBSpendUpdateTransactions | None = None if raw_results[0] is not None: parsed = self._parse_list_of_transactions(raw_results[0]) if len(parsed) > 0: db_spend = self._combine_list_of_transactions(parsed) # Slots 1-5: daily spend categories - daily_results: List[Optional[Dict[str, Any]]] = [] + daily_results: list[dict[str, Any] | None] = [] for slot in range(1, 6): slot_result = raw_results[slot] if slot_result is None: @@ -535,11 +533,11 @@ class RedisUpdateBuffer: return ( db_spend, - cast(Optional[Dict[str, DailyUserSpendTransaction]], daily_results[0]), - cast(Optional[Dict[str, DailyTeamSpendTransaction]], daily_results[1]), - cast(Optional[Dict[str, DailyOrganizationSpendTransaction]], daily_results[2]), - cast(Optional[Dict[str, DailyEndUserSpendTransaction]], daily_results[3]), - cast(Optional[Dict[str, DailyAgentSpendTransaction]], daily_results[4]), + cast(dict[str, DailyUserSpendTransaction] | None, daily_results[0]), + cast(dict[str, DailyTeamSpendTransaction] | None, daily_results[1]), + cast(dict[str, DailyOrganizationSpendTransaction] | None, daily_results[2]), + cast(dict[str, DailyEndUserSpendTransaction] | None, daily_results[3]), + cast(dict[str, DailyAgentSpendTransaction] | None, daily_results[4]), ) async def store_in_memory_daily_tag_spend_updates_in_redis( @@ -560,7 +558,7 @@ class RedisUpdateBuffer: async def get_all_daily_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyUserSpendTransaction]]: + ) -> dict[str, DailyUserSpendTransaction] | None: """ Gets all the daily spend update transactions from Redis """ @@ -574,7 +572,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyUserSpendTransaction], + dict[str, DailyUserSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -582,7 +580,7 @@ class RedisUpdateBuffer: async def get_all_daily_team_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyTeamSpendTransaction]]: + ) -> dict[str, DailyTeamSpendTransaction] | None: """ Gets all the daily team spend update transactions from Redis """ @@ -596,7 +594,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyTeamSpendTransaction], + dict[str, DailyTeamSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -604,7 +602,7 @@ class RedisUpdateBuffer: async def get_all_daily_org_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyOrganizationSpendTransaction]]: + ) -> dict[str, DailyOrganizationSpendTransaction] | None: """ Gets all the daily organization spend update transactions from Redis """ @@ -618,7 +616,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyOrganizationSpendTransaction], + dict[str, DailyOrganizationSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -626,7 +624,7 @@ class RedisUpdateBuffer: async def get_all_daily_end_user_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyEndUserSpendTransaction]]: + ) -> dict[str, DailyEndUserSpendTransaction] | None: """ Gets all the daily end-user spend update transactions from Redis """ @@ -640,7 +638,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyEndUserSpendTransaction], + dict[str, DailyEndUserSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -648,7 +646,7 @@ class RedisUpdateBuffer: async def get_all_daily_agent_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyAgentSpendTransaction]]: + ) -> dict[str, DailyAgentSpendTransaction] | None: """ Gets all the daily agent spend update transactions from Redis """ @@ -662,7 +660,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyAgentSpendTransaction], + dict[str, DailyAgentSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -670,7 +668,7 @@ class RedisUpdateBuffer: async def get_all_daily_tag_spend_update_transactions_from_redis_buffer( self, - ) -> Optional[Dict[str, DailyTagSpendTransaction]]: + ) -> dict[str, DailyTagSpendTransaction] | None: """ Gets all the daily tag spend update transactions from Redis """ @@ -684,7 +682,7 @@ class RedisUpdateBuffer: return None list_of_daily_spend_update_transactions = [json.loads(transaction) for transaction in list_of_transactions] return cast( - Dict[str, DailyTagSpendTransaction], + dict[str, DailyTagSpendTransaction], DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( list_of_daily_spend_update_transactions ), @@ -692,8 +690,8 @@ class RedisUpdateBuffer: @staticmethod def _parse_list_of_transactions( - list_of_transactions: Union[Any, List[Any]], - ) -> List[DBSpendUpdateTransactions]: + list_of_transactions: Any | list[Any], + ) -> list[DBSpendUpdateTransactions]: """ Parses the list of transactions from Redis """ @@ -704,7 +702,7 @@ class RedisUpdateBuffer: @staticmethod def _combine_list_of_transactions( - list_of_transactions: List[DBSpendUpdateTransactions], + list_of_transactions: list[DBSpendUpdateTransactions], ) -> DBSpendUpdateTransactions: """ Combines the list of transactions into a single DBSpendUpdateTransactions object diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 9bf3da0066d..b5ff2ffa9cc 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -1,6 +1,5 @@ import asyncio from datetime import datetime, timedelta, timezone -from typing import Optional from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache @@ -31,11 +30,11 @@ class SpendLogCleanup: def __init__( self, general_settings=None, - redis_cache: Optional[RedisCache] = None, - partition_manager: Optional[SpendLogsPartitionManager] = None, + redis_cache: RedisCache | None = None, + partition_manager: SpendLogsPartitionManager | None = None, ): self.batch_size = SPEND_LOG_CLEANUP_BATCH_SIZE - self.retention_seconds: Optional[int] = None + self.retention_seconds: int | None = None self.partition_manager = partition_manager or SpendLogsPartitionManager() from litellm.proxy.proxy_server import general_settings as default_settings @@ -69,7 +68,7 @@ class SpendLogCleanup: return True except ValueError as e: verbose_proxy_logger.warning( - f"Invalid maximum_spend_logs_retention_period value: {retention_setting}, error: {str(e)}" + f"Invalid maximum_spend_logs_retention_period value: {retention_setting}, error: {e!s}" ) return False diff --git a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py index 932675a6ac3..57321bf8bf3 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py +++ b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py @@ -15,7 +15,6 @@ keeps the batched-DELETE path, so existing deployments are untouched. import re from datetime import date, datetime, timedelta, timezone -from typing import List, Optional, Tuple from litellm._logging import verbose_proxy_logger from litellm.constants import ( @@ -59,12 +58,12 @@ def partition_name(start: date) -> str: return f"{SPEND_LOGS_TABLE}_p{start.strftime('%Y%m%d')}" -def upcoming_partitions(today: date, interval: PartitionInterval, ahead: int) -> List[Tuple[str, date, date]]: +def upcoming_partitions(today: date, interval: PartitionInterval, ahead: int) -> list[tuple[str, date, date]]: """ Specs (name, lower_inclusive, upper_exclusive) for the current period plus the next `ahead` periods, so writes always have a partition to land in. """ - specs: List[Tuple[str, date, date]] = [] + specs: list[tuple[str, date, date]] = [] start = period_start(today, interval) for _ in range(ahead + 1): upper = next_period_start(start, interval) @@ -73,7 +72,7 @@ def upcoming_partitions(today: date, interval: PartitionInterval, ahead: int) -> return specs -def parse_partition_upper_bound(bound_expr: str) -> Optional[datetime]: +def parse_partition_upper_bound(bound_expr: str) -> datetime | None: """ Upper bound of a Postgres partition from its `pg_get_expr(relpartbound)` string, e.g. "FOR VALUES FROM ('2026-06-01 00:00:00') TO ('2026-06-02 00:00:00')". @@ -91,7 +90,7 @@ def parse_partition_upper_bound(bound_expr: str) -> Optional[datetime]: return None -def select_partitions_to_drop(partitions: List[Tuple[str, Optional[datetime]]], cutoff: datetime) -> List[str]: +def select_partitions_to_drop(partitions: list[tuple[str, datetime | None]], cutoff: datetime) -> list[str]: """ Names of partitions whose entire range is older than `cutoff` (upper bound <= cutoff). `cutoff` and the bounds are UTC-naive. Partitions without a @@ -140,13 +139,13 @@ class SpendLogsPartitionManager: return False return bool(rows and rows[0].get("partitioned")) - async def ensure_partitions(self, prisma_client) -> List[str]: + async def ensure_partitions(self, prisma_client) -> list[str]: """ Ensure the current and upcoming partitions exist, returning the names now present. CREATE TABLE IF NOT EXISTS is a no-op for partitions that already exist, so this list is "ensured present", not "newly created". """ - ensured: List[str] = [] + ensured: list[str] = [] for name, lower, upper in upcoming_partitions( datetime.now(timezone.utc).date(), self.interval, self.precreate_ahead ): @@ -161,7 +160,7 @@ class SpendLogsPartitionManager: verbose_proxy_logger.warning("Failed to ensure spend-log partition %s: %s", name, e) return ensured - async def _list_partitions(self, prisma_client) -> List[Tuple[str, Optional[datetime]]]: + async def _list_partitions(self, prisma_client) -> list[tuple[str, datetime | None]]: rows = await prisma_client.db.query_raw( """ SELECT c.relname AS name, @@ -177,12 +176,12 @@ class SpendLogsPartitionManager: ) return [(row["name"], parse_partition_upper_bound(row.get("bound") or "")) for row in rows] - async def drop_partitions_older_than(self, prisma_client, cutoff: datetime) -> List[str]: + async def drop_partitions_older_than(self, prisma_client, cutoff: datetime) -> list[str]: """DROP every partition whose whole range is older than `cutoff`.""" cutoff_naive = cutoff.astimezone(timezone.utc).replace(tzinfo=None) partitions = await self._list_partitions(prisma_client) to_drop = select_partitions_to_drop(partitions, cutoff_naive) - dropped: List[str] = [] + dropped: list[str] = [] for name in to_drop: try: await prisma_client.db.execute_raw(f'DROP TABLE IF EXISTS "{name}"') diff --git a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py index 0689fc00b02..43383e7b5d5 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py @@ -1,5 +1,4 @@ import asyncio -from typing import Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm.constants import LITELLM_ASYNCIO_QUEUE_MAXSIZE @@ -51,15 +50,14 @@ class SpendUpdateQueue(BaseUpdateQueue): async def aggregate_queue_updates(self): """Concatenate all updates in the queue to reduce the size of in-memory queue""" - updates: List[SpendUpdateQueueItem] = await self.flush_all_updates_from_in_memory_queue() + updates: list[SpendUpdateQueueItem] = await self.flush_all_updates_from_in_memory_queue() aggregated_updates = self._get_aggregated_spend_update_queue_item(updates) for update in aggregated_updates: await self.update_queue.put(update) - return def _get_aggregated_spend_update_queue_item( - self, updates: List[SpendUpdateQueueItem] - ) -> List[SpendUpdateQueueItem]: + self, updates: list[SpendUpdateQueueItem] + ) -> list[SpendUpdateQueueItem]: """ This is used to reduce the size of the in-memory queue by aggregating updates by entity type + id @@ -102,9 +100,9 @@ class SpendUpdateQueue(BaseUpdateQueue): "Aggregating spend updates, current queue size: %s", self.update_queue.qsize(), ) - aggregated_spend_updates: List[SpendUpdateQueueItem] = [] + aggregated_spend_updates: list[SpendUpdateQueueItem] = [] - _in_memory_map: Dict[str, SpendUpdateQueueItem] = {} + _in_memory_map: dict[str, SpendUpdateQueueItem] = {} """ Used for combining several updates into a single update Key=entity_type:entity_id @@ -127,7 +125,7 @@ class SpendUpdateQueue(BaseUpdateQueue): return aggregated_spend_updates def get_aggregated_db_spend_update_transactions( - self, updates: List[SpendUpdateQueueItem] + self, updates: list[SpendUpdateQueueItem] ) -> DBSpendUpdateTransactions: """Aggregate updates by entity type.""" # Initialize all transaction lists as empty dicts @@ -209,7 +207,7 @@ class SpendUpdateQueue(BaseUpdateQueue): async def _emit_new_item_added_to_queue_event( self, - queue_size: Optional[int] = None, + queue_size: int | None = None, ): asyncio.create_task( service_logger_obj.async_service_success_hook( diff --git a/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py b/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py index 5d23bcaa944..b05e2a325e9 100644 --- a/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py +++ b/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py @@ -7,8 +7,6 @@ cycle (~30s). The seen-set is cleared on every flush so that call_count increments in subsequent cycles rather than stopping after the first flush. """ -from typing import List, Set - from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ToolDiscoveryQueueItem @@ -25,8 +23,8 @@ class ToolDiscoveryQueue: """ def __init__(self) -> None: - self._seen_tool_names: Set[str] = set() - self._pending: List[ToolDiscoveryQueueItem] = [] + self._seen_tool_names: set[str] = set() + self._pending: list[ToolDiscoveryQueueItem] = [] def add_update(self, item: ToolDiscoveryQueueItem) -> None: """Enqueue a tool discovery item if tool_name has not been seen before.""" @@ -44,7 +42,7 @@ class ToolDiscoveryQueue: item.get("origin"), ) - def flush(self) -> List[ToolDiscoveryQueueItem]: + def flush(self) -> list[ToolDiscoveryQueueItem]: """Return and clear all pending items. Resets seen-set so the next flush cycle can re-count the same tools.""" items, self._pending = self._pending, [] diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 04256be50b3..28551f49031 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -1,5 +1,5 @@ from collections.abc import Awaitable, Callable -from typing import Any, Optional, Union +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( @@ -26,9 +26,7 @@ class PrismaDBExceptionHandler: """ from litellm.proxy.proxy_server import general_settings - _allow_requests_on_db_unavailable: Union[bool, str] = general_settings.get( - "allow_requests_on_db_unavailable", False - ) + _allow_requests_on_db_unavailable: bool | str = general_settings.get("allow_requests_on_db_unavailable", False) if isinstance(_allow_requests_on_db_unavailable, bool): return _allow_requests_on_db_unavailable if str_to_bool(_allow_requests_on_db_unavailable) is True: @@ -262,7 +260,7 @@ class PrismaDBExceptionHandler: PrismaDBExceptionHandler.is_database_connection_error(e) and PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() ): - return None + return raise e @@ -287,8 +285,8 @@ async def call_with_db_reconnect_retry( coro_factory: Callable[[], Awaitable[Any]], *, reason: str, - timeout_seconds: Optional[float] = None, - lock_timeout_seconds: Optional[float] = None, + timeout_seconds: float | None = None, + lock_timeout_seconds: float | None = None, ) -> Any: """Run a Prisma read coroutine with one transport-reconnect-and-retry. diff --git a/litellm/proxy/db/log_db_metrics.py b/litellm/proxy/db/log_db_metrics.py index 241f7db733f..29a438476d8 100644 --- a/litellm/proxy/db/log_db_metrics.py +++ b/litellm/proxy/db/log_db_metrics.py @@ -8,13 +8,12 @@ import asyncio from collections.abc import Callable from datetime import datetime from functools import wraps -from typing import Dict, Optional, Tuple from litellm._service_logger import ServiceTypes from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs -def _safe_db_event_metadata(kwargs: Dict) -> Optional[Dict[str, str]]: +def _safe_db_event_metadata(kwargs: dict) -> dict[str, str] | None: """Minimal, non-sensitive ``event_metadata`` for a DB service log. The raw ``kwargs``/``args`` carry live objects (Prisma client, OTel spans) @@ -117,8 +116,8 @@ def _is_exception_related_to_db(e: Exception) -> bool: async def _handle_logging_db_exception( e: Exception, func: Callable, - kwargs: Dict, - args: Tuple, + kwargs: dict, + args: tuple, start_time: datetime, end_time: datetime, ) -> None: diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index f8a99d2c120..79f86d548c5 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -13,7 +13,7 @@ import urllib.parse from collections.abc import Callable from dataclasses import dataclass from datetime import datetime, timedelta -from typing import Any, Protocol, Union +from typing import Any, Protocol from litellm._logging import verbose_proxy_logger from litellm.secret_managers.main import str_to_bool @@ -912,7 +912,7 @@ class PrismaManager: def should_update_prisma_schema( - disable_updates: Union[bool, str] | None = None, + disable_updates: bool | str | None = None, ) -> bool: """ Determines if Prisma Schema updates should be applied during startup. diff --git a/litellm/proxy/db/query_engine_reaper.py b/litellm/proxy/db/query_engine_reaper.py index 0e5f0e68910..996513877d4 100644 --- a/litellm/proxy/db/query_engine_reaper.py +++ b/litellm/proxy/db/query_engine_reaper.py @@ -26,7 +26,6 @@ import signal import sys import threading import time -from typing import Optional from litellm._logging import verbose_proxy_logger @@ -53,7 +52,7 @@ def set_child_subreaper() -> bool: return False -def _read_comm_and_ppid(pid: int, proc_root: str) -> Optional[tuple[str, int]]: +def _read_comm_and_ppid(pid: int, proc_root: str) -> tuple[str, int] | None: try: with open(f"{proc_root}/{pid}/stat", encoding="ascii", errors="replace") as stat_file: data = stat_file.read() @@ -180,7 +179,7 @@ def _reaper_loop(parent_pid: int) -> None: REAPER_THREAD_NAME = "litellm-orphan-query-engine-reaper" -def start_query_engine_reaper() -> Optional[threading.Thread]: +def start_query_engine_reaper() -> threading.Thread | None: """Start the reaper daemon thread in the supervisor process. Must only be called from a process that never hosts the proxy app diff --git a/litellm/proxy/db/routing_prisma_wrapper.py b/litellm/proxy/db/routing_prisma_wrapper.py index bca36dcb712..01b1d12f5ea 100644 --- a/litellm/proxy/db/routing_prisma_wrapper.py +++ b/litellm/proxy/db/routing_prisma_wrapper.py @@ -39,7 +39,7 @@ class _RoutedActions: fails a recreate) is observed without re-fetching the actions accessor. """ - __slots__ = ("_writer_actions", "_reader_actions", "_should_use_reader") + __slots__ = ("_reader_actions", "_should_use_reader", "_writer_actions") def __init__( self, diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 079cbd163dc..a4e80a32066 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -54,7 +54,7 @@ class SpendCounterReseed: """ _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() - _registry_lock: ClassVar[Optional[asyncio.Lock]] = None + _registry_lock: ClassVar[asyncio.Lock | None] = None @staticmethod async def _get_lock(counter_key: str) -> asyncio.Lock: @@ -72,7 +72,7 @@ class SpendCounterReseed: return lock @staticmethod - async def from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> Optional[float]: + async def from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None: """ Read the authoritative spend for a counter from the DB. @@ -106,9 +106,7 @@ class SpendCounterReseed: elif counter_key.startswith("spend:user:"): user_id = counter_key[len("spend:user:") :] row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - elif counter_key.startswith("spend:end_user:"): - return None - elif counter_key.startswith("spend:tag:"): + elif counter_key.startswith("spend:end_user:") or counter_key.startswith("spend:tag:"): return None elif counter_key.startswith("spend:org:"): org_id = counter_key[len("spend:org:") :] @@ -143,7 +141,7 @@ class SpendCounterReseed: spend_counter_cache: "DualCache", counter_key: str, require_cache_warm: bool = False, - ) -> Optional[float]: + ) -> float | None: """ Reseed a cold spend counter from the DB and warm the cache, coalesced via a per-counter lock so concurrent callers (read path @@ -213,7 +211,7 @@ class SpendCounterReseed: entity_type: str, entity_id: str, window_start: datetime, - ) -> Optional[float]: + ) -> float | None: if prisma_client is None: return None @@ -261,7 +259,7 @@ class SpendCounterReseed: entity_type: str, entity_id: str, window_start: datetime, - ) -> Optional[float]: + ) -> float | None: lock = await SpendCounterReseed._get_lock(counter_key) async with lock: redis_clean_miss = False diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 5f7de772a2c..20ff62c8100 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -7,7 +7,7 @@ Admins use the management endpoints to read and update input_policy / output_pol import uuid from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ToolDiscoveryQueueItem @@ -23,7 +23,7 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: +def _row_to_model(row: dict | Any) -> LiteLLM_ToolTableRow: """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow.""" model_dump = getattr(row, "model_dump", None) if callable(model_dump): @@ -72,7 +72,7 @@ def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: async def batch_upsert_tools( prisma_client: "PrismaClient", - items: List[ToolDiscoveryQueueItem], + items: list[ToolDiscoveryQueueItem], ) -> None: """ Batch-upsert tool registry rows via Prisma. @@ -128,8 +128,8 @@ async def batch_upsert_tools( async def list_tools( prisma_client: "PrismaClient", - input_policy: Optional[str] = None, -) -> List[LiteLLM_ToolTableRow]: + input_policy: str | None = None, +) -> list[LiteLLM_ToolTableRow]: """Return all tools, optionally filtered by input_policy.""" try: where = {"input_policy": input_policy} if input_policy is not None else {} @@ -146,7 +146,7 @@ async def list_tools( async def get_tool( prisma_client: "PrismaClient", tool_name: str, -) -> Optional[LiteLLM_ToolTableRow]: +) -> LiteLLM_ToolTableRow | None: """Return a single tool row by tool_name.""" try: row = await ToolRepository(prisma_client).table.find_unique( @@ -163,10 +163,10 @@ async def get_tool( async def update_tool_policy( prisma_client: "PrismaClient", tool_name: str, - updated_by: Optional[str], - input_policy: Optional[str] = None, - output_policy: Optional[str] = None, -) -> Optional[LiteLLM_ToolTableRow]: + updated_by: str | None, + input_policy: str | None = None, + output_policy: str | None = None, +) -> LiteLLM_ToolTableRow | None: """Update input_policy and/or output_policy for a tool. Upserts the row if it does not exist yet.""" try: _updated_by = updated_by or "system" @@ -206,8 +206,8 @@ async def update_tool_policy( async def get_tools_by_names( prisma_client: "PrismaClient", - tool_names: List[str], -) -> Dict[str, Tuple[str, str]]: + tool_names: list[str], +) -> dict[str, tuple[str, str]]: """ Return a {tool_name: (input_policy, output_policy)} map for the given tool names. """ @@ -232,12 +232,12 @@ async def get_tools_by_names( async def list_overrides_for_tool( prisma_client: "PrismaClient", tool_name: str, -) -> List[ToolPolicyOverrideRow]: +) -> list[ToolPolicyOverrideRow]: """ Return override-like rows for a tool by finding object permissions that have this tool in blocked_tools, then resolving each permission to key/team scope for display. """ - out: List[ToolPolicyOverrideRow] = [] + out: list[ToolPolicyOverrideRow] = [] try: perms = await ObjectPermissionRepository(prisma_client).table.find_many( where={"blocked_tools": {"has": tool_name}}, @@ -289,9 +289,9 @@ class ToolPolicyRegistry: """ def __init__(self) -> None: - self._tool_input_policies: Dict[str, str] = {} - self._tool_output_policies: Dict[str, str] = {} - self._blocked_tools_by_op_id: Dict[str, List[str]] = {} + self._tool_input_policies: dict[str, str] = {} + self._tool_output_policies: dict[str, str] = {} + self._blocked_tools_by_op_id: dict[str, list[str]] = {} self._initialized: bool = False def is_initialized(self) -> bool: @@ -342,10 +342,10 @@ class ToolPolicyRegistry: def get_effective_policies( self, - tool_names: List[str], - object_permission_id: Optional[str] = None, - team_object_permission_id: Optional[str] = None, - ) -> Dict[str, str]: + tool_names: list[str], + object_permission_id: str | None = None, + team_object_permission_id: str | None = None, + ) -> dict[str, str]: """ Return effective input_policy per tool from in-memory state. If tool is in key or team blocked_tools -> "blocked", else global input_policy or "untrusted". @@ -356,7 +356,7 @@ class ToolPolicyRegistry: for op_id in (object_permission_id, team_object_permission_id): if op_id and op_id.strip(): blocked.update(self._blocked_tools_by_op_id.get(op_id.strip(), [])) - result: Dict[str, str] = {} + result: dict[str, str] = {} for name in tool_names: if name in blocked: result[name] = "blocked" @@ -365,7 +365,7 @@ class ToolPolicyRegistry: return result -_tool_policy_registry: Optional[ToolPolicyRegistry] = None +_tool_policy_registry: ToolPolicyRegistry | None = None def get_tool_policy_registry() -> ToolPolicyRegistry: diff --git a/litellm/proxy/dd_span_tagger.py b/litellm/proxy/dd_span_tagger.py index 08b7d928d0e..55f3ecf4da3 100644 --- a/litellm/proxy/dd_span_tagger.py +++ b/litellm/proxy/dd_span_tagger.py @@ -1,5 +1,3 @@ -from typing import Optional - from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.dd_tracing import set_active_span_tag from litellm.proxy._types import UserAPIKeyAuth @@ -9,7 +7,7 @@ class DDSpanTagger: """Best-effort helpers for tagging the active Datadog APM span with LiteLLM request metadata.""" @staticmethod - def tag_call_id(litellm_call_id: Optional[str]) -> None: + def tag_call_id(litellm_call_id: str | None) -> None: """ Attach LiteLLM call id to the active Datadog APM span. @@ -29,7 +27,7 @@ class DDSpanTagger: @staticmethod def tag_request( user_api_key_dict: UserAPIKeyAuth, - requested_model: Optional[str], + requested_model: str | None, ) -> None: """ Attach key and model tags to the active Datadog APM span. diff --git a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py index ba2e0a39de3..41698a3906e 100644 --- a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -14,9 +14,8 @@ router = APIRouter() @router.get("/litellm/.well-known/litellm-ui-config", response_model=UiDiscoveryEndpoints) # if mounted at root path async def get_ui_config(): from litellm.proxy.auth.auth_utils import _has_user_setup_sso - from litellm.proxy.utils import get_proxy_base_url, get_server_root_path - from litellm.proxy.proxy_server import general_settings + from litellm.proxy.utils import get_proxy_base_url, get_server_root_path auto_redirect_ui_login_to_sso = ( os.getenv("AUTO_REDIRECT_UI_LOGIN_TO_SSO", "false").lower() == "true" diff --git a/litellm/proxy/enterprise_billing/billing_metrics.py b/litellm/proxy/enterprise_billing/billing_metrics.py index f9f8ceaf721..ce6ddf3e7d4 100644 --- a/litellm/proxy/enterprise_billing/billing_metrics.py +++ b/litellm/proxy/enterprise_billing/billing_metrics.py @@ -61,10 +61,10 @@ class BillingMetricsConfig: endpoint: str client_cert_path: str client_key_path: str - ca_cert_path: Optional[str] + ca_cert_path: str | None export_interval_ms: int litellm_version: str - license_id: Optional[str] + license_id: str | None def _metrics_endpoint(endpoint: str) -> str: @@ -83,7 +83,7 @@ def _resource_attributes(config: BillingMetricsConfig) -> dict[str, AttributeVal def _billable_attributes( - category: BillableCategory, route: str, status_code: int, model_id: Optional[str] + category: BillableCategory, route: str, status_code: int, model_id: str | None ) -> dict[str, AttributeValue]: base: dict[str, AttributeValue] = { "litellm.endpoint.category": category.value, @@ -124,7 +124,7 @@ class BillingMetricsRecorder: description="Count of 2xx HTTP requests to billable LLM/MCP/A2A endpoints", ) - def record(self, *, category: BillableCategory, route: str, status_code: int, model_id: Optional[str]) -> None: + def record(self, *, category: BillableCategory, route: str, status_code: int, model_id: str | None) -> None: self._counter.add(1, _billable_attributes(category, route, status_code, model_id)) def shutdown(self) -> None: @@ -150,7 +150,7 @@ def _export_interval_ms() -> int: class _CredentialPaths: client_cert_path: str client_key_path: str - ca_cert_path: Optional[str] + ca_cert_path: str | None def _is_pem_content(value: str) -> bool: @@ -165,7 +165,7 @@ def _write_pem(directory: str, filename: str, pem: str) -> str: return path -def _resolve_credential_paths(*, client_cert: str, client_key: str, ca_cert: Optional[str]) -> _CredentialPaths: +def _resolve_credential_paths(*, client_cert: str, client_key: str, ca_cert: str | None) -> _CredentialPaths: """ Accept either a filesystem path or inline PEM content for each credential. @@ -197,7 +197,7 @@ def _resolve_credential_paths(*, client_cert: str, client_key: str, ca_cert: Opt def load_billing_metrics_config( *, license_data: Optional["EnterpriseLicenseData"], litellm_version: str -) -> Optional[BillingMetricsConfig]: +) -> BillingMetricsConfig | None: endpoint = os.getenv(ENDPOINT_ENV) client_cert = os.getenv(CLIENT_CERT_ENV) client_key = os.getenv(CLIENT_KEY_ENV) @@ -267,12 +267,12 @@ class _ActiveRecorderRegistry: proxy_shutdown_event.""" def __init__(self) -> None: - self._recorder: Optional[BillingMetricsRecorder] = None + self._recorder: BillingMetricsRecorder | None = None def set(self, recorder: BillingMetricsRecorder) -> None: self._recorder = recorder - def pop(self) -> Optional[BillingMetricsRecorder]: + def pop(self) -> BillingMetricsRecorder | None: recorder = self._recorder self._recorder = None return recorder @@ -283,7 +283,7 @@ _ACTIVE_RECORDER = _ActiveRecorderRegistry() def build_billing_metrics_recorder( *, premium: bool, license_data: Optional["EnterpriseLicenseData"], litellm_version: str -) -> Optional[BillingMetricsRecorder]: +) -> BillingMetricsRecorder | None: """Build the recorder, or None when the deployment is not licensed or metering is unconfigured.""" if not premium: # Debug, not warning: unlicensed is the common case and a warning here diff --git a/litellm/proxy/fine_tuning_endpoints/endpoints.py b/litellm/proxy/fine_tuning_endpoints/endpoints.py index ef2943df76d..778ea729e32 100644 --- a/litellm/proxy/fine_tuning_endpoints/endpoints.py +++ b/litellm/proxy/fine_tuning_endpoints/endpoints.py @@ -6,7 +6,7 @@ ########################################################################## import asyncio -from typing import Optional, cast +from typing import cast from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response @@ -112,7 +112,7 @@ async def create_fine_tuning_job( # Convert Pydantic model to dict verbose_proxy_logger.debug( - "Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)), + f"Request received by LiteLLM:\n{json.dumps(data, indent=4)}", ) # Include original request and headers in the data @@ -133,7 +133,7 @@ async def create_fine_tuning_job( ## CHECK IF MANAGED FILE ID unified_file_id: Union[str, Literal[False]] = False training_file = fine_tuning_request.training_file - response: Optional[LiteLLMFineTuningJob] = None + response: LiteLLMFineTuningJob | None = None if training_file: unified_file_id = _is_base64_encoded_unified_file_id(training_file) ## IF SO, Route based on that @@ -200,7 +200,7 @@ async def create_fine_tuning_job( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.create_fine_tuning_job(): Exception occurred - {}".format(str(e)) + f"litellm.proxy.proxy_server.create_fine_tuning_job(): Exception occurred - {e!s}" ) raise handle_exception_on_proxy(e) @@ -221,7 +221,7 @@ async def retrieve_fine_tuning_job( request: Request, fastapi_response: Response, fine_tuning_job_id: str, - custom_llm_provider: Optional[Literal["openai", "azure"]] = None, + custom_llm_provider: Literal["openai", "azure"] | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -269,7 +269,7 @@ async def retrieve_fine_tuning_job( ## CHECK IF MANAGED FILE ID unified_finetuning_job_id: Union[str, Literal[False]] = False - response: Optional[LiteLLMFineTuningJob] = None + response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) if unified_finetuning_job_id: @@ -340,7 +340,7 @@ async def retrieve_fine_tuning_job( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.retrieve_fine_tuning_job(): Exception occurred - {}".format(str(e)) + f"litellm.proxy.proxy_server.retrieve_fine_tuning_job(): Exception occurred - {e!s}" ) raise handle_exception_on_proxy(e) @@ -360,13 +360,13 @@ async def retrieve_fine_tuning_job( async def list_fine_tuning_jobs( request: Request, fastapi_response: Response, - custom_llm_provider: Optional[Literal["openai", "azure"]] = None, - target_model_names: Optional[str] = Query( + custom_llm_provider: Literal["openai", "azure"] | None = None, + target_model_names: str | None = Query( default=None, description="Comma separated list of model names to filter by. Example: 'gpt-4o,gpt-4o-mini'", ), - after: Optional[str] = None, - limit: Optional[int] = None, + after: str | None = None, + limit: int | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -406,7 +406,7 @@ async def list_fine_tuning_jobs( route_type=CallTypes.alist_fine_tuning_jobs.value, ) - response: Optional[Any] = None + response: Any | None = None if target_model_names and isinstance(target_model_names, str): target_model_names_list = target_model_names.split(",") if len(target_model_names_list) != 1: @@ -469,7 +469,7 @@ async def list_fine_tuning_jobs( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.list_fine_tuning_jobs(): Exception occurred - {}".format(str(e)) + f"litellm.proxy.proxy_server.list_fine_tuning_jobs(): Exception occurred - {e!s}" ) raise handle_exception_on_proxy(e) @@ -538,7 +538,7 @@ async def cancel_fine_tuning_job( ## CHECK IF MANAGED FILE ID unified_finetuning_job_id: Union[str, Literal[False]] = False - response: Optional[LiteLLMFineTuningJob] = None + response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) if unified_finetuning_job_id: @@ -609,6 +609,6 @@ async def cancel_fine_tuning_job( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.cancel_fine_tuning_job(): Exception occurred - {}".format(str(e)) + f"litellm.proxy.proxy_server.cancel_fine_tuning_job(): Exception occurred - {e!s}" ) raise handle_exception_on_proxy(e) diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py index 8e09064562e..1e06bb9f345 100644 --- a/litellm/proxy/guardrails/_content_utils.py +++ b/litellm/proxy/guardrails/_content_utils.py @@ -9,7 +9,7 @@ every text fragment. """ from collections.abc import Callable, Iterator -from typing import Any, Dict, FrozenSet, List +from typing import Any # Call types whose body carries free-form chat / prompt text that # text-content guardrails (banned keywords, content moderation, secret @@ -24,7 +24,7 @@ from typing import Any, Dict, FrozenSet, List # ``pre_call_hook`` directly with the sync name. Embedding, moderation, # audio, and transcription endpoints are deliberately excluded — text # guardrails on those paths are a separate scope. -TEXT_CONTENT_CALL_TYPES: FrozenSet[str] = frozenset({"completion", "acompletion", "aresponses"}) +TEXT_CONTENT_CALL_TYPES: frozenset[str] = frozenset({"completion", "acompletion", "aresponses"}) def is_text_content_call_type(call_type: str) -> bool: @@ -33,7 +33,7 @@ def is_text_content_call_type(call_type: str) -> bool: return call_type in TEXT_CONTENT_CALL_TYPES -TEXT_PART_TYPES: FrozenSet[str] = frozenset({"text", "input_text", "output_text"}) +TEXT_PART_TYPES: frozenset[str] = frozenset({"text", "input_text", "output_text"}) # Responses-API item types whose ``output`` field carries user/tool text # that guardrails should inspect. ``function_call_output`` is the @@ -64,13 +64,13 @@ def _iter_text_parts_in_content(content: Any) -> Iterator[str]: yield text -def _coerce_input_to_messages(input_value: Any) -> List[Dict[str, Any]]: +def _coerce_input_to_messages(input_value: Any) -> list[dict[str, Any]]: """Coerce a Responses-API ``data["input"]`` value into chat-style messages.""" if isinstance(input_value, str): return [{"role": "user", "content": input_value}] if not isinstance(input_value, list): return [] - messages: List[Dict[str, Any]] = [] + messages: list[dict[str, Any]] = [] for item in input_value: if isinstance(item, str): messages.append({"role": "user", "content": item}) @@ -84,7 +84,7 @@ def _coerce_input_to_messages(input_value: Any) -> List[Dict[str, Any]]: return messages -def _iter_inspection_messages(data: Dict[str, Any]) -> Iterator[Dict[str, Any]]: +def _iter_inspection_messages(data: dict[str, Any]) -> Iterator[dict[str, Any]]: """Yield every message-like dict, walking ``messages`` AND ``input``.""" messages = data.get("messages") if isinstance(messages, list): @@ -92,7 +92,7 @@ def _iter_inspection_messages(data: Dict[str, Any]) -> Iterator[Dict[str, Any]]: yield from _coerce_input_to_messages(data.get("input")) -def iter_message_text(data: Dict[str, Any]) -> Iterator[str]: +def iter_message_text(data: dict[str, Any]) -> Iterator[str]: """Yield every text fragment from ``messages`` AND ``input``. Walks every role (user, assistant, system, …) — guardrails inspect @@ -104,7 +104,7 @@ def iter_message_text(data: Dict[str, Any]) -> Iterator[str]: yield from _iter_text_parts_in_content(message.get("content")) -def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int: +def walk_user_text(data: dict[str, Any], visit: Callable[[str], str]) -> int: """Rewrite every text fragment in place via ``visit``. Mutates ``data["messages"]`` and ``data["input"]``. Returns the number @@ -121,7 +121,7 @@ def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int: return visit(content) return content if isinstance(content, list): - new_parts: List[Any] = [] + new_parts: list[Any] = [] for part in content: if isinstance(part, str) and part: visited += 1 @@ -171,7 +171,7 @@ def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int: return visited -def apply_redacted_messages_back(data: Dict[str, Any], redacted_messages: List[Dict[str, Any]]) -> None: +def apply_redacted_messages_back(data: dict[str, Any], redacted_messages: list[dict[str, Any]]) -> None: """Write redacted messages back to whichever field(s) the caller used. Mask/anonymize paths take a synthesised messages list (from @@ -185,7 +185,7 @@ def apply_redacted_messages_back(data: Dict[str, Any], redacted_messages: List[D if "messages" in data: data["messages"] = redacted_messages if isinstance(data.get("input"), str): - text_parts: List[str] = [] + text_parts: list[str] = [] for msg in redacted_messages: if not isinstance(msg, dict): continue @@ -193,7 +193,7 @@ def apply_redacted_messages_back(data: Dict[str, Any], redacted_messages: List[D data["input"] = "\n".join(text_parts) -def has_non_string_content(data: Dict[str, Any]) -> bool: +def has_non_string_content(data: dict[str, Any]) -> bool: """Return True if any inspected content is not a plain string. Used by hooks whose mask/redact path operates on string offsets and @@ -213,7 +213,7 @@ def has_non_string_content(data: Dict[str, Any]) -> bool: return False -def build_inspection_messages(data: Dict[str, Any]) -> List[Dict[str, str]]: +def build_inspection_messages(data: dict[str, Any]) -> list[dict[str, str]]: """Synthesize a chat-style messages list for posting to a guardrail API. Each returned message has a plain-string ``content`` — multimodal text @@ -224,7 +224,7 @@ def build_inspection_messages(data: Dict[str, Any]) -> List[Dict[str, str]]: call this instead of ``data.get("messages", [])`` so the Responses API and multimodal content are covered. """ - flattened: List[Dict[str, str]] = [] + flattened: list[dict[str, str]] = [] for message in _iter_inspection_messages(data): if not isinstance(message, dict): continue diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index a29899a3965..3d8ed8dc1e2 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -11,11 +11,7 @@ from datetime import datetime, timezone from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Type, TypeVar, Union, cast, @@ -100,14 +96,14 @@ async def _find_team_guardrail_rows( def _get_guardrails_list_response( - guardrails_config: List[Dict], + guardrails_config: list[dict], ) -> ListGuardrailsResponse: """ Helper function to get the guardrails list response """ from litellm.litellm_core_utils.litellm_logging import _get_masked_values - guardrail_configs: List[GuardrailInfoResponse] = [] + guardrail_configs: list[GuardrailInfoResponse] = [] for guardrail in guardrails_config: litellm_params = guardrail.get("litellm_params") or {} masked_params = _get_masked_values( @@ -171,7 +167,7 @@ async def list_guardrails(): config = proxy_config.config - _guardrails_config = cast(Optional[list[dict]], config.get("guardrails")) + _guardrails_config = cast(list[dict] | None, config.get("guardrails")) if _guardrails_config is None: return _get_guardrails_list_response([]) @@ -235,7 +231,7 @@ async def list_guardrails_v2( excluded_guardrail_ids: set = set() if not is_admin: caller_team_ids = await _get_user_team_ids(user_api_key_dict) - allowed: List[Guardrail] = [] + allowed: list[Guardrail] = [] for g in guardrails: g_team_id = g.get("team_id") if g_team_id is None or g_team_id in caller_team_ids: @@ -246,10 +242,10 @@ async def list_guardrails_v2( excluded_guardrail_ids.add(gid) guardrails = allowed - guardrail_configs: List[GuardrailInfoResponse] = [] + guardrail_configs: list[GuardrailInfoResponse] = [] seen_guardrail_ids: set = excluded_guardrail_ids.copy() for guardrail in guardrails: - litellm_params: Optional[Union[LitellmParams, dict]] = guardrail.get("litellm_params") + litellm_params: LitellmParams | dict | None = guardrail.get("litellm_params") litellm_params_dict = ( litellm_params.model_dump(exclude_none=True) if isinstance(litellm_params, LitellmParams) @@ -614,11 +610,11 @@ class RegisterGuardrailRequest(BaseModel): """Request body for POST /guardrails/register. Follows Generic Guardrail API config.""" guardrail_name: str - litellm_params: Dict[str, Any] # guardrail, mode, api_base required; api_key, headers, etc. optional - guardrail_info: Optional[Dict[str, object]] = None - team_id: Optional[str] = None + litellm_params: dict[str, Any] # guardrail, mode, api_base required; api_key, headers, etc. optional + guardrail_info: dict[str, object] | None = None + team_id: str | None = None - def get_litellm_params_dict(self) -> Dict[str, Any]: + def get_litellm_params_dict(self) -> dict[str, Any]: return dict(self.litellm_params) @@ -626,7 +622,7 @@ class RegisterGuardrailResponse(BaseModel): guardrail_id: str guardrail_name: str status: str - submitted_at: Optional[datetime] = None + submitted_at: datetime | None = None class GuardrailSubmissionSummary(BaseModel): @@ -640,22 +636,22 @@ class GuardrailSubmissionItem(BaseModel): guardrail_id: str guardrail_name: str status: str # pending_review | active | rejected - team_id: Optional[str] = None + team_id: str | None = None team_guardrail: bool = ( False # True when submitted via team (team_id set); use to distinguish team vs regular guardrails ) - litellm_params: Optional[Dict[str, object]] = None - guardrail_info: Optional[Dict[str, object]] = None - submitted_by_user_id: Optional[str] = None - submitted_by_email: Optional[str] = None - submitted_at: Optional[datetime] = None - reviewed_at: Optional[datetime] = None - created_at: Optional[datetime] = None - updated_at: Optional[datetime] = None + litellm_params: dict[str, object] | None = None + guardrail_info: dict[str, object] | None = None + submitted_by_user_id: str | None = None + submitted_by_email: str | None = None + submitted_at: datetime | None = None + reviewed_at: datetime | None = None + created_at: datetime | None = None + updated_at: datetime | None = None class ListGuardrailSubmissionsResponse(BaseModel): - submissions: List[GuardrailSubmissionItem] + submissions: list[GuardrailSubmissionItem] summary: GuardrailSubmissionSummary @@ -774,7 +770,7 @@ async def register_guardrail( raise HTTPException(status_code=500, detail=str(e)) -def _parse_json_field(value: object) -> Optional[Dict[str, Any]]: +def _parse_json_field(value: object) -> dict[str, Any] | None: if value is None: return None if isinstance(value, dict): @@ -787,7 +783,7 @@ def _parse_json_field(value: object) -> Optional[Dict[str, Any]]: return None -async def _get_user_team_ids(user_api_key_dict: UserAPIKeyAuth) -> List[str]: +async def _get_user_team_ids(user_api_key_dict: UserAPIKeyAuth) -> list[str]: """Return the list of team_ids the caller belongs to (empty list if none).""" from litellm.proxy.auth.auth_checks import get_user_object from litellm.proxy.proxy_server import ( @@ -841,9 +837,9 @@ def _row_to_submission_item(row: "LiteLLM_GuardrailsTable") -> GuardrailSubmissi response_model=ListGuardrailSubmissionsResponse, ) async def list_guardrail_submissions( - status: Optional[str] = None, - team_id: Optional[str] = None, - search: Optional[str] = None, + status: str | None = None, + team_id: str | None = None, + search: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -868,7 +864,7 @@ async def list_guardrail_submissions( # Proxy Admin would (no writes — registration / approval still gated # elsewhere by their own per-action checks). is_admin = _user_has_admin_view(user_api_key_dict) - visible_team_ids: Optional[List[str]] = None + visible_team_ids: list[str] | None = None if not is_admin: visible_team_ids = await _get_user_team_ids(user_api_key_dict) if team_id is not None and team_id not in visible_team_ids: @@ -878,7 +874,7 @@ async def list_guardrail_submissions( ) try: - where_clause: Dict[str, object] = {"team_id": {"not": None}} + where_clause: dict[str, object] = {"team_id": {"not": None}} if visible_team_ids is not None: if not visible_team_ids: # Non-admin with no team memberships: nothing visible. @@ -1289,7 +1285,7 @@ async def get_guardrail_info(guardrail_id: str): if result is None: raise HTTPException(status_code=404, detail=f"Guardrail with ID {guardrail_id} not found") - litellm_params: Optional[Union[LitellmParams, dict]] = result.get("litellm_params") + litellm_params: LitellmParams | dict | None = result.get("litellm_params") result_litellm_params_dict = ( litellm_params.model_dump(exclude_none=True) if isinstance(litellm_params, LitellmParams) @@ -1424,7 +1420,7 @@ async def get_category_yaml(category_name: str): "file_type": file_type, } except Exception as e: - raise HTTPException(status_code=500, detail=f"Error reading category file: {str(e)}") + raise HTTPException(status_code=500, detail=f"Error reading category file: {e!s}") @router.get( @@ -1456,7 +1452,7 @@ async def get_major_airlines(): airlines = json.load(f) return {"airlines": airlines} except Exception as e: - raise HTTPException(status_code=500, detail=f"Error reading major_airlines.json: {str(e)}") from e + raise HTTPException(status_code=500, detail=f"Error reading major_airlines.json: {e!s}") from e @router.post( @@ -1464,7 +1460,7 @@ async def get_major_airlines(): tags=["Guardrails"], dependencies=[Depends(user_api_key_auth)], ) -async def validate_blocked_words_file(request: Dict[str, str]): +async def validate_blocked_words_file(request: dict[str, str]): """ Validate a blocked_words YAML file content. @@ -1544,10 +1540,10 @@ async def validate_blocked_words_file(request: Dict[str, str]): "message": f"Valid YAML file with {len(blocked_words_list)} blocked word(s)", } except yaml.YAMLError as e: - return {"valid": False, "error": f"Invalid YAML syntax: {str(e)}"} + return {"valid": False, "error": f"Invalid YAML syntax: {e!s}"} except Exception as e: verbose_proxy_logger.exception("Error validating blocked words file") - return {"valid": False, "error": f"Validation error: {str(e)}"} + return {"valid": False, "error": f"Validation error: {e!s}"} def _get_field_type_from_annotation(field_annotation: Any) -> str: @@ -1584,9 +1580,7 @@ def _get_field_type_from_annotation(field_annotation: Any) -> str: # Handle basic types if field_annotation is str: return "string" - elif field_annotation is int: - return "number" - elif field_annotation is float: + elif field_annotation is int or field_annotation is float: return "number" elif field_annotation is bool: return "boolean" @@ -1599,7 +1593,7 @@ def _get_field_type_from_annotation(field_annotation: Any) -> str: return "string" -def _extract_literal_values(annotation: Any) -> List[str]: +def _extract_literal_values(annotation: Any) -> list[str]: """ Extract literal values from a Literal type annotation """ @@ -1610,7 +1604,7 @@ def _extract_literal_values(annotation: Any) -> List[str]: return [] -def _get_dict_key_options(field_annotation: Any) -> Optional[List[str]]: +def _get_dict_key_options(field_annotation: Any) -> list[str] | None: """ Extract key options from Dict[Literal[...], T] types """ @@ -1642,7 +1636,7 @@ def _get_dict_value_type(field_annotation: Any) -> str: return "string" -def _get_list_element_options(field_annotation: Any) -> Optional[List[str]]: +def _get_list_element_options(field_annotation: Any) -> list[str] | None: """ Extract element options from List[Literal[...]] types """ @@ -1708,7 +1702,7 @@ def _build_field_dict( field_annotation: Any, description: str, required: bool, -) -> Dict[str, Any]: +) -> dict[str, Any]: """Build field dictionary for non-nested fields.""" # Determine the field type from annotation field_type = _get_field_type_from_annotation(field_annotation) @@ -1769,9 +1763,9 @@ def _build_field_dict( def _extract_fields_recursive( - model: Type[BaseModel], + model: type[BaseModel], depth: int = 0, -) -> Dict[str, Any]: +) -> dict[str, Any]: # Check if we've exceeded the maximum recursion depth if depth > DEFAULT_MAX_RECURSE_DEPTH: raise HTTPException( @@ -1807,7 +1801,7 @@ def _extract_fields_recursive( if is_basemodel_subclass: # Recursively get fields from the nested model - nested_fields = _extract_fields_recursive(cast(Type[BaseModel], field_annotation), depth + 1) + nested_fields = _extract_fields_recursive(cast(type[BaseModel], field_annotation), depth + 1) fields[field_name] = { "description": description, "required": required, @@ -1825,7 +1819,7 @@ def _extract_fields_recursive( return fields -def _get_fields_from_model(model_class: Type[BaseModel]) -> Dict[str, Any]: +def _get_fields_from_model(model_class: type[BaseModel]) -> dict[str, Any]: """ Get the fields from a Pydantic model as a nested dictionary structure """ @@ -1926,13 +1920,13 @@ class TestCustomCodeGuardrailRequest(BaseModel): custom_code: str """The Python-like code containing the apply_guardrail function.""" - test_input: Dict[str, object] + test_input: dict[str, object] """The test input to pass to the guardrail. Should contain 'texts', optionally 'images', 'tools', etc.""" input_type: str = "request" """Whether this is a 'request' or 'response' input type.""" - request_data: Optional[Dict[str, object]] = None + request_data: dict[str, object] | None = None """Optional mock request_data (model, user_id, team_id, metadata, etc.).""" @@ -1942,13 +1936,13 @@ class TestCustomCodeGuardrailResponse(BaseModel): success: bool """Whether the test executed successfully (no errors).""" - result: Optional[Dict[str, object]] = None + result: dict[str, object] | None = None """The guardrail result: action (allow/block/modify), reason, modified_texts, etc.""" - error: Optional[str] = None + error: str | None = None """Error message if execution failed.""" - error_type: Optional[str] = None + error_type: str | None = None """Type of error: 'compilation' or 'execution'.""" @@ -2253,7 +2247,7 @@ async def apply_guardrail( start_time = datetime.now(timezone.utc) try: - active_guardrail: Optional[CustomGuardrail] = GUARDRAIL_REGISTRY.get_initialized_guardrail_callback( + active_guardrail: CustomGuardrail | None = GUARDRAIL_REGISTRY.get_initialized_guardrail_callback( guardrail_name=request.guardrail_name ) if active_guardrail is None: diff --git a/litellm/proxy/guardrails/guardrail_helpers.py b/litellm/proxy/guardrails/guardrail_helpers.py index 677ac66fcd0..609e2d833e3 100644 --- a/litellm/proxy/guardrails/guardrail_helpers.py +++ b/litellm/proxy/guardrails/guardrail_helpers.py @@ -1,6 +1,5 @@ import os import sys -from typing import Dict import litellm from litellm._logging import verbose_proxy_logger @@ -16,7 +15,7 @@ def can_modify_guardrails(team_obj: Optional[LiteLLM_TeamTable]) -> bool: team_metadata = team_obj.metadata or {} - if team_metadata.get("guardrails", None) is not None and isinstance(team_metadata.get("guardrails"), Dict): + if team_metadata.get("guardrails", None) is not None and isinstance(team_metadata.get("guardrails"), dict): if team_metadata.get("guardrails", {}).get("modify_guardrails", None) is False: return False diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 601cc25e774..076cdfc8ecd 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -8,7 +8,7 @@ import asyncio import json import os from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, List, Optional, Type, Union +from typing import TYPE_CHECKING, Any from pydantic import BaseModel from websockets.asyncio.client import ClientConnection, connect @@ -47,14 +47,14 @@ class AimGuardrailMissingSecrets(Exception): class AimGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, GuardrailEventHooks.post_call, ] - def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None, **kwargs): + def __init__(self, api_key: str | None = None, api_base: str | None = None, **kwargs): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) ssl_verify = kwargs.pop("ssl_verify", None) self.async_handler = get_async_httpx_client( @@ -80,7 +80,7 @@ class AimGuardrail(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Inside AIM Pre-Call Hook") return await self.call_aim_guardrail(data, hook="pre_call", key_alias=user_api_key_dict.key_alias) @@ -89,13 +89,13 @@ class AimGuardrail(CustomGuardrail): data: dict, user_api_key_dict: UserAPIKeyAuth, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Inside AIM Moderation Hook") await self.call_aim_guardrail(data, hook="moderation", key_alias=user_api_key_dict.key_alias) return data - async def call_aim_guardrail(self, data: dict, hook: str, key_alias: Optional[str]) -> dict: + async def call_aim_guardrail(self, data: dict, hook: str, key_alias: str | None) -> dict: user_email = data.get("metadata", {}).get("headers", {}).get("x-aim-user-email") call_id = data.get("litellm_call_id") headers = self._build_aim_headers( @@ -184,8 +184,8 @@ class AimGuardrail(CustomGuardrail): return data async def call_aim_guardrail_on_output( - self, request_data: dict, output: str, hook: str, key_alias: Optional[str] - ) -> Optional[dict]: + self, request_data: dict, output: str, hook: str, key_alias: str | None + ) -> dict | None: user_email = request_data.get("metadata", {}).get("headers", {}).get("x-aim-user-email") call_id = request_data.get("litellm_call_id") response = await self.async_handler.post( @@ -227,9 +227,9 @@ class AimGuardrail(CustomGuardrail): self, *, hook: str, - key_alias: Optional[str], - user_email: Optional[str], - litellm_call_id: Optional[str], + key_alias: str | None, + user_email: str | None, + litellm_call_id: str | None, ): """ A helper function to build the http headers that are required by AIM guardrails. @@ -260,7 +260,7 @@ class AimGuardrail(CustomGuardrail): self, data: dict, user_api_key_dict: UserAPIKeyAuth, - response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse], + response: Any | ModelResponse | EmbeddingResponse | ImageResponse, ) -> Any: if not (isinstance(response, ModelResponse) and response.choices): return response @@ -345,7 +345,7 @@ class AimGuardrail(CustomGuardrail): await websocket.send(json.dumps({"done": True})) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.aim import ( AimGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py index daae74ae8e0..acb107e509b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -11,11 +11,10 @@ import asyncio import json import os from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Type - -from fastapi import HTTPException +from typing import TYPE_CHECKING, Any, Literal import httpx +from fastapi import HTTPException from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -44,7 +43,7 @@ class AktoGuardrail(CustomGuardrail): HOOK_TO_INPUT = {"pre_call": "request", "post_call": "response"} @staticmethod - def get_config_model() -> Type["GuardrailConfigModel"]: + def get_config_model() -> type["GuardrailConfigModel"]: """Return the Pydantic config model for YAML-based initialization.""" from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( AktoConfigModel, @@ -53,7 +52,7 @@ class AktoGuardrail(CustomGuardrail): return AktoConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -61,12 +60,12 @@ class AktoGuardrail(CustomGuardrail): def __init__( self, - akto_base_url: Optional[str] = None, - akto_api_key: Optional[str] = None, - akto_account_id: Optional[str] = None, - akto_vxlan_id: Optional[str] = None, + akto_base_url: str | None = None, + akto_api_key: str | None = None, + akto_account_id: str | None = None, + akto_vxlan_id: str | None = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", - guardrail_timeout: Optional[int] = None, + guardrail_timeout: int | None = None, **kwargs: Any, ) -> None: """Initialize the Akto guardrail. @@ -107,7 +106,7 @@ class AktoGuardrail(CustomGuardrail): ) @staticmethod - def resolve_metadata_value(request_data: Optional[dict], key: str) -> Optional[str]: + def resolve_metadata_value(request_data: dict | None, key: str) -> str | None: """Look up a metadata value from litellm_metadata or metadata dicts.""" if request_data is None: return None @@ -128,7 +127,7 @@ class AktoGuardrail(CustomGuardrail): route = metadata.get("user_api_key_request_route") return route if route else "/v1/chat/completions" - def prepare_headers(self) -> Dict[str, str]: + def prepare_headers(self) -> dict[str, str]: """Build HTTP headers for the Akto API call.""" return { "content-type": "application/json", @@ -136,9 +135,9 @@ class AktoGuardrail(CustomGuardrail): } @staticmethod - def build_query_params(*, guardrails: bool, ingest_data: bool) -> Dict[str, str]: + def build_query_params(*, guardrails: bool, ingest_data: bool) -> dict[str, str]: """Build query params that control Akto backend behavior (guardrail check and/or data ingestion).""" - params: Dict[str, str] = {"akto_connector": AKTO_CONNECTOR_NAME} + params: dict[str, str] = {"akto_connector": AKTO_CONNECTOR_NAME} if guardrails: params["guardrails"] = "true" if ingest_data: @@ -146,9 +145,9 @@ class AktoGuardrail(CustomGuardrail): return params @staticmethod - def build_request_headers(request_data: dict) -> Dict[str, str]: + def build_request_headers(request_data: dict) -> dict[str, str]: """Build the requestHeaders field from proxy request headers.""" - headers: Dict[str, str] = {"content-type": "application/json"} + headers: dict[str, str] = {"content-type": "application/json"} proxy_req = request_data.get("proxy_server_request", {}) if not isinstance(proxy_req, dict): return headers @@ -162,11 +161,11 @@ class AktoGuardrail(CustomGuardrail): @staticmethod def build_request_body( inputs: GenericGuardrailAPIInputs, - request_data: Optional[dict] = None, - ) -> Dict[str, Any]: + request_data: dict | None = None, + ) -> dict[str, Any]: """Build the LLM request body from guardrail inputs (messages, model, tools).""" model = inputs.get("model", "") or "" - body: Dict[str, Any] = {"model": model} + body: dict[str, Any] = {"model": model} structured = inputs.get("structured_messages") if structured: @@ -194,8 +193,8 @@ class AktoGuardrail(CustomGuardrail): @staticmethod def build_response_body( inputs: GenericGuardrailAPIInputs, - request_data: Optional[dict] = None, - ) -> Dict[str, Any]: + request_data: dict | None = None, + ) -> dict[str, Any]: """Build the LLM response body, preferring the actual model response if available.""" model_response = request_data.get("response") if request_data else None if model_response is not None and hasattr(model_response, "model_dump"): @@ -207,9 +206,9 @@ class AktoGuardrail(CustomGuardrail): return {} @staticmethod - def build_tag_metadata(request_data: dict) -> Dict[str, str]: + def build_tag_metadata(request_data: dict) -> dict[str, str]: """Build tag/metadata dict with user_id and team_id for Akto tracking.""" - tag: Dict[str, str] = {"gen-ai": "Gen AI"} + tag: dict[str, str] = {"gen-ai": "Gen AI"} user_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_user_id") team_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_team_id") if user_id: @@ -225,7 +224,7 @@ class AktoGuardrail(CustomGuardrail): *, status_code: int = 200, include_response: bool = False, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Build the flat MIRRORING payload sent to Akto's HTTP proxy endpoint. All body fields use double-encoding: json.dumps({"body": json.dumps(actual_body)}) @@ -237,7 +236,7 @@ class AktoGuardrail(CustomGuardrail): tag = self.build_tag_metadata(request_data) response_payload = json.dumps({}) # Empty body wrapper when no response yet - response_headers: Dict[str, str] = {} + response_headers: dict[str, str] = {} if include_response: response_body = self.build_response_body(inputs, request_data) response_payload = json.dumps({"body": json.dumps(response_body)}) # Double-encoded @@ -299,7 +298,7 @@ class AktoGuardrail(CustomGuardrail): ) @staticmethod - def handle_guardrail_response(response: httpx.Response) -> Tuple[bool, str]: + def handle_guardrail_response(response: httpx.Response) -> tuple[bool, str]: """Parse the Akto guardrail response. Returns (allowed, reason).""" if response.status_code != 200: verbose_proxy_logger.error("Akto returned HTTP %d", response.status_code) diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py index dc3fc40625c..0ed0140c371 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py @@ -11,7 +11,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import json import sys -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal from fastapi import HTTPException @@ -38,13 +38,13 @@ if TYPE_CHECKING: class AporiaGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.during_call, GuardrailEventHooks.post_call, ] - def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None, **kwargs): + def __init__(self, api_key: str | None = None, api_base: str | None = None, **kwargs): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) self.aporia_api_key = api_key or os.environ["APORIO_API_KEY"] @@ -52,7 +52,7 @@ class AporiaGuardrail(CustomGuardrail): super().__init__(**kwargs) #### CALL HOOKS - proxy only #### - def transform_messages(self, messages: List[dict]) -> List[dict]: + def transform_messages(self, messages: list[dict]) -> list[dict]: supported_openai_roles = ["system", "user", "assistant"] default_role = "other" # for unsupported roles - e.g. tool new_messages = [] @@ -69,7 +69,7 @@ class AporiaGuardrail(CustomGuardrail): return new_messages - async def prepare_aporia_request(self, new_messages: List[dict], response_string: Optional[str] = None) -> dict: + async def prepare_aporia_request(self, new_messages: list[dict], response_string: str | None = None) -> dict: data: dict[str, Any] = {} if new_messages is not None: data["messages"] = new_messages @@ -90,8 +90,8 @@ class AporiaGuardrail(CustomGuardrail): async def make_aporia_api_request( self, request_data: dict, - new_messages: List[dict], - response_string: Optional[str] = None, + new_messages: list[dict], + response_string: str | None = None, ): data = await self.prepare_aporia_request(new_messages=new_messages, response_string=response_string) @@ -156,7 +156,7 @@ class AporiaGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return - response_str: Optional[str] = convert_litellm_response_object_to_str(response) + response_str: str | None = convert_litellm_response_object_to_str(response) if response_str is not None: await self.make_aporia_api_request( request_data=data, @@ -166,8 +166,6 @@ class AporiaGuardrail(CustomGuardrail): add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) - pass - @log_guardrail_information async def async_moderation_hook( self, @@ -206,7 +204,7 @@ class AporiaGuardrail(CustomGuardrail): ): return - new_messages: Optional[List[dict]] = None + new_messages: list[dict] | None = None if "messages" in data and isinstance(data["messages"], list): new_messages = self.transform_messages(messages=data["messages"]) @@ -218,10 +216,9 @@ class AporiaGuardrail(CustomGuardrail): add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) else: verbose_proxy_logger.warning("Aporia AI: not running guardrail. No messages in data") - pass @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.aporia_ai import ( AporiaGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/azure/__init__.py index 243c4ad408b..8eac1fe85a5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Union +from typing import TYPE_CHECKING from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -24,10 +24,9 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" raise ValueError("Azure Content Safety: guardrail_name is required") if azure_guardrail == "prompt_shield": - azure_content_safety_guardrail: Union[ - AzureContentSafetyPromptShieldGuardrail, - AzureContentSafetyTextModerationGuardrail, - ] = AzureContentSafetyPromptShieldGuardrail( + azure_content_safety_guardrail: ( + AzureContentSafetyPromptShieldGuardrail | AzureContentSafetyTextModerationGuardrail + ) = AzureContentSafetyPromptShieldGuardrail( guardrail_name=guardrail_name, **{ **litellm_params.model_dump(exclude_none=True), diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index b178efbda59..717ba3e70d4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -1,5 +1,5 @@ import re -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -41,7 +41,7 @@ class AzureGuardrailBase: self.api_base = api_base self.api_version: str = kwargs.get("api_version") or "2024-09-01" - async def _post_to_content_safety(self, endpoint_path: str, request_body: Dict[str, Any]) -> Dict[str, Any]: + async def _post_to_content_safety(self, endpoint_path: str, request_body: dict[str, Any]) -> dict[str, Any]: """POST to an Azure Content Safety endpoint with standard auth headers. Args: @@ -64,12 +64,12 @@ class AzureGuardrailBase: headers=headers, json=request_body, ) - response_json: Dict[str, Any] = response.json() + response_json: dict[str, Any] = response.json() verbose_proxy_logger.debug("Azure Content Safety response [%s]: %s", endpoint_path, response_json) return response_json @staticmethod - def split_text_by_words(text: str, max_length: int) -> List[str]: + def split_text_by_words(text: str, max_length: int) -> list[str]: """ Split text into chunks at word boundaries without breaking words. @@ -92,7 +92,7 @@ class AzureGuardrailBase: # within each chunk. tokens = re.findall(r"\S+|\s+", text) - chunks: List[str] = [] + chunks: list[str] = [] current_chunk = "" for token in tokens: @@ -117,7 +117,7 @@ class AzureGuardrailBase: return chunks - def get_user_prompt(self, messages: List["AllMessageValues"]) -> Optional[str]: + def get_user_prompt(self, messages: list["AllMessageValues"]) -> str | None: """ Get the last consecutive block of messages from the user. diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index befb1b7ae56..3df4f230dbf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -3,7 +3,7 @@ Azure Prompt Shield Native Guardrail Integrationfor LiteLLM """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, cast +from typing import TYPE_CHECKING, Any, cast from fastapi import HTTPException @@ -69,15 +69,16 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai chunk is analysed independently; an attack in *any* chunk raises an HTTPException immediately. """ - from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import ( AzurePromptShieldGuardrailRequestBody, AzurePromptShieldGuardrailResponse, ) + from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH + chunks = self.split_text_by_words(user_prompt, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH) - last_response: Optional[AzurePromptShieldGuardrailResponse] = None + last_response: AzurePromptShieldGuardrailResponse | None = None for chunk in chunks: request_body = AzurePromptShieldGuardrailRequestBody(documents=[], userPrompt=chunk) @@ -107,9 +108,9 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai self, user_api_key_dict: "UserAPIKeyAuth", cache: Any, - data: Dict[str, Any], + data: dict[str, Any], call_type: CallTypesLiteral, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """ Pre-call hook to scan user prompts before sending to LLM. @@ -119,7 +120,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai "Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s", call_type, ) - new_messages: Optional[List[AllMessageValues]] = data.get("messages") + new_messages: list[AllMessageValues] | None = data.get("messages") if new_messages is None: verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data") return data @@ -135,7 +136,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai return None @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Get the config model for the Azure Prompt Shield guardrail. """ @@ -146,7 +147,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai return AzurePromptShieldGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index 91f5df0e9b8..0b1faf99469 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -3,7 +3,7 @@ Azure Text Moderation Native Guardrail Integrationfor LiteLLM """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast +from typing import TYPE_CHECKING, Any, Literal, Union, cast from fastapi import HTTPException @@ -44,7 +44,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr default_severity_threshold: int = 2 @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -55,8 +55,8 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr guardrail_name: str, api_key: str, api_base: str, - severity_threshold: Optional[int] = None, - severity_threshold_by_category: Optional[Dict[str, int]] = None, + severity_threshold: int | None = None, + severity_threshold_by_category: dict[str, int] | None = None, **kwargs, ): """Initialize Azure Text Moderation guardrail handler.""" @@ -82,7 +82,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr "SelfHarm", "Violence", ], - "blocklistNames": cast(Optional[List[str]], kwargs.get("blocklistNames") or None), + "blocklistNames": cast(list[str] | None, kwargs.get("blocklistNames") or None), "haltOnBlocklistHit": kwargs.get("haltOnBlocklistHit") or False, "outputType": kwargs.get("outputType") or "FourSeverityLevels", } @@ -93,7 +93,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr verbose_proxy_logger.info(f"Initialized Azure Text Moderation Guardrail: {guardrail_name}") @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import ( AzureContentSafetyTextModerationConfigModel, ) @@ -109,15 +109,16 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr chunk is analysed independently; a severity-threshold violation in *any* chunk raises an HTTPException immediately. """ - from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import ( AzureTextModerationGuardrailRequestBody, AzureTextModerationGuardrailResponse, ) + from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH + chunks = self.split_text_by_words(text, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH) - last_response: Optional[AzureTextModerationGuardrailResponse] = None + last_response: AzureTextModerationGuardrailResponse | None = None for chunk in chunks: request_body = AzureTextModerationGuardrailRequestBody( @@ -203,9 +204,9 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr self, user_api_key_dict: "UserAPIKeyAuth", cache: Any, - data: Dict[str, Any], + data: dict[str, Any], call_type: CallTypesLiteral, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """ Pre-call hook to scan user prompts before sending to LLM. @@ -215,7 +216,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr "Azure Text Moderation: Running pre-call prompt scan, on call_type: %s", call_type, ) - new_messages: Optional[List[AllMessageValues]] = data.get("messages") + new_messages: list[AllMessageValues] | None = data.get("messages") if new_messages is None: verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data") return data diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index cf9c5b859c0..f8fedb22872 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -18,13 +18,9 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, - Dict, - List, Literal, NamedTuple, Optional, - Tuple, - Union, cast, ) @@ -102,7 +98,7 @@ _BEDROCK_CHECKS_PII_LOCATION_KEYS = ( # The model response is qualified as ``guard_content`` directly by the OUTPUT builder; # the existing ``guarded_text`` marker is intentionally left unmapped here so its # guardrail-hook payload is unchanged by this feature. -_CONTENT_TYPE_TO_QUALIFIER: Dict[str, BedrockGuardrailQualifier] = { +_CONTENT_TYPE_TO_QUALIFIER: dict[str, BedrockGuardrailQualifier] = { "grounding_source": "grounding_source", "query": "query", } @@ -118,22 +114,22 @@ class QualifiedTextBlock(NamedTuple): """A piece of message text paired with its Bedrock grounding qualifier (if any).""" text: str - qualifier: Optional[BedrockGuardrailQualifier] + qualifier: BedrockGuardrailQualifier | None class GuardrailMessageFilterResult(NamedTuple): - payload_messages: Optional[List[AllMessageValues]] - original_messages: Optional[List[AllMessageValues]] - target_indices: Optional[List[int]] + payload_messages: list[AllMessageValues] | None + original_messages: list[AllMessageValues] | None + target_indices: list[int] | None class ApplyGuardrailMessageSelection(NamedTuple): """Messages selected for an apply_guardrail scan + write-back metadata.""" - filtered_messages: Optional[list[AllMessageValues]] + filtered_messages: list[AllMessageValues] | None # Slice of the flat `texts` list actually scanned (offset, length), # used to write masked content back to the right positions. None = whole list. - scanned_slice: Optional[tuple[int, int]] + scanned_slice: tuple[int, int] | None # True when messages were selected by their original role. scanned_role_subset: bool # True when there is nothing to scan (e.g. no user-role message). @@ -151,7 +147,7 @@ def _redact_pii_matches(response_json: dict) -> dict: return redacted if isinstance(redacted, dict) else response_json -def _redact_assessment_match_fields(assessments: List[dict]) -> List[dict]: +def _redact_assessment_match_fields(assessments: list[dict]) -> list[dict]: """ Redact sensitive match-like fields from blocked assessment summaries. @@ -170,9 +166,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def __init__( self, - guardrailIdentifier: Optional[str] = None, - guardrailVersion: Optional[str] = None, - disable_exception_on_block: Optional[bool] = False, + guardrailIdentifier: str | None = None, + guardrailVersion: str | None = None, + disable_exception_on_block: bool | None = False, checks: BedrockChecksConfigModel | Mapping[str, object] | None = None, content_filter_threshold: float | None = 0.5, prompt_attack_threshold: float | None = 0.5, @@ -234,7 +230,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -271,12 +267,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) return cleaned or None - def _create_bedrock_input_content_request(self, messages: Optional[List[AllMessageValues]]) -> BedrockRequest: + def _create_bedrock_input_content_request(self, messages: list[AllMessageValues] | None) -> BedrockRequest: """ Create a bedrock request for the input content - the LLM request. """ bedrock_request: BedrockRequest = BedrockRequest(source="INPUT") - bedrock_request_content: List[BedrockContentItem] = [] + bedrock_request_content: list[BedrockContentItem] = [] if messages is None: return bedrock_request for message in messages: @@ -295,8 +291,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def _create_bedrock_output_content_request( self, - response: Union[Any, ModelResponse], - messages: Optional[List[AllMessageValues]] = None, + response: Any | ModelResponse, + messages: list[AllMessageValues] | None = None, ) -> BedrockRequest: """ Create a bedrock request for the output content - the LLM response. @@ -309,7 +305,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """ bedrock_request: BedrockRequest = BedrockRequest(source="OUTPUT") grounding_blocks = self._collect_grounding_blocks(messages) - bedrock_request_content: List[BedrockContentItem] = [ + bedrock_request_content: list[BedrockContentItem] = [ self._build_content_item(block) for block in grounding_blocks ] has_grounding = len(bedrock_request_content) > 0 @@ -320,12 +316,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return bedrock_request def _build_response_content_items( - self, response: Union[Any, ModelResponse], has_grounding: bool - ) -> List[BedrockContentItem]: + self, response: Any | ModelResponse, has_grounding: bool + ) -> list[BedrockContentItem]: """Build content item(s) from the model response. When the request supplied grounding, the response is qualified ``guard_content`` so Bedrock can score it. """ - items: List[BedrockContentItem] = [] + items: list[BedrockContentItem] = [] if not isinstance(response, litellm.ModelResponse): return items for choice in response.choices: @@ -344,8 +340,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def convert_to_bedrock_format( self, source: Literal["INPUT", "OUTPUT"], - messages: Optional[List[AllMessageValues]] = None, - response: Optional[Union[Any, ModelResponse]] = None, + messages: list[AllMessageValues] | None = None, + response: Any | ModelResponse | None = None, ) -> BedrockRequest: """ Convert the litellm messages/response to the bedrock request format. @@ -363,7 +359,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_request = self._create_bedrock_output_content_request(response=response, messages=messages) return bedrock_request - def get_content_items_for_message(self, message: AllMessageValues) -> Optional[List[QualifiedTextBlock]]: + def get_content_items_for_message(self, message: AllMessageValues) -> list[QualifiedTextBlock] | None: """ Flatten a message into text blocks, preserving any contextual-grounding qualifier carried by the content-block ``type`` (grounding_source / query). @@ -373,7 +369,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): content = message.get("content") if content is None: return None - blocks: List[QualifiedTextBlock] = [] + blocks: list[QualifiedTextBlock] = [] if isinstance(content, str): blocks.append(QualifiedTextBlock(text=content, qualifier=None)) elif isinstance(content, list): @@ -392,7 +388,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): text_content["qualifiers"] = [block.qualifier] return BedrockContentItem(text=text_content) - def _collect_grounding_blocks(self, messages: Optional[List[AllMessageValues]]) -> List[QualifiedTextBlock]: + def _collect_grounding_blocks(self, messages: list[AllMessageValues] | None) -> list[QualifiedTextBlock]: """Harvest grounding_source/query blocks from the request for an OUTPUT scan. ``grounding_source`` is honored only from app-authored roles (system / @@ -402,19 +398,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): contextual-grounding check to grade the response against. ``query`` is accepted from any role (it is the user's question). """ - grounding: List[QualifiedTextBlock] = [] + grounding: list[QualifiedTextBlock] = [] for message in messages or []: role = message.get("role") for block in self.get_content_items_for_message(message=message) or []: - if block.qualifier == "query": - grounding.append(block) - elif block.qualifier == "grounding_source" and role in _GROUNDING_SOURCE_TRUSTED_ROLES: + if ( + block.qualifier == "query" + or block.qualifier == "grounding_source" + and role in _GROUNDING_SOURCE_TRUSTED_ROLES + ): grounding.append(block) return grounding def _prepare_guardrail_messages_for_role( self, - messages: Optional[List[AllMessageValues]], + messages: list[AllMessageValues] | None, ) -> GuardrailMessageFilterResult: """Return payload + merge metadata for the latest user message.""" # NOTE: This logic probably belongs in CustomGuardrail once other guardrails adopt the feature. @@ -437,7 +435,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): target_indices=[latest_index], ) - def _find_latest_message_index(self, messages: List[AllMessageValues], target_role: str) -> Optional[int]: + def _find_latest_message_index(self, messages: list[AllMessageValues], target_role: str) -> int | None: for index in range(len(messages) - 1, -1, -1): if messages[index].get("role", None) == target_role: return index @@ -458,7 +456,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): structured_messages: list[AllMessageValues], target_index: int, texts: list[str], - ) -> Optional[tuple[int, int]]: + ) -> tuple[int, int] | None: """ Map one message's text segments to their (offset, length) slice in the flat `texts` list built by the guardrail translation handler. @@ -517,7 +515,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # None, and the write-back guard below safely skips masking rather than # corrupting positions. structured_messages = cast( - Optional[list[AllMessageValues]], + list[AllMessageValues] | None, inputs.get("structured_messages") or request_data.get("messages"), ) if input_type != "request" or not structured_messages: @@ -557,7 +555,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self, masked_texts: list, texts: list, - scanned_slice: Optional[tuple[int, int]], + scanned_slice: tuple[int, int] | None, scanned_role_subset: bool, ) -> list: """ @@ -591,10 +589,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def _merge_filtered_messages( self, - original_messages: Optional[List[AllMessageValues]], - updated_target_messages: List[AllMessageValues], - target_indices: Optional[List[int]], - ) -> List[AllMessageValues]: + original_messages: list[AllMessageValues] | None, + updated_target_messages: list[AllMessageValues], + target_indices: list[int] | None, + ) -> list[AllMessageValues]: if not target_indices: return updated_target_messages @@ -656,8 +654,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): data: dict, optional_params: dict, aws_region_name: str, - api_key: Optional[str] = None, - extra_headers: Optional[dict] = None, + api_key: str | None = None, + extra_headers: dict | None = None, request_path: str | None = None, ): headers = {"Content-Type": "application/json"} @@ -680,7 +678,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # first check api-key, if none, fall back to sigV4 if api_key is not None: - aws_bearer_token: Optional[str] = api_key + aws_bearer_token: str | None = api_key else: aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") @@ -764,7 +762,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self.convert_to_bedrock_format(source=source, messages=messages, response=response) ) bedrock_guardrail_response: BedrockGuardrailResponse = BedrockGuardrailResponse() - api_key: Optional[str] = None + api_key: str | None = None if request_data: dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data) bedrock_request_data.update( @@ -1158,7 +1156,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def _get_block_exception_for_checks( self, violations: list[BedrockChecksViolation], request_data: dict | None = None - ) -> Union[HTTPException, ModifyResponseException]: + ) -> HTTPException | ModifyResponseException: """Build the block exception for an over-threshold InvokeGuardrailChecks result. Mirrors ``_get_http_exception_for_blocked_guardrail``'s return-type branching. @@ -1238,7 +1236,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return "success" return "guardrail_failed_to_respond" - def _parse_bedrock_guardrail_error_response(self, response: httpx.Response) -> Tuple[int, str]: + def _parse_bedrock_guardrail_error_response(self, response: httpx.Response) -> tuple[int, str]: """ Parse AWS Bedrock guardrail error response body to extract status code and message. @@ -1283,7 +1281,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): tracing_detail["guardrail_action"] = bedrock_action return tracing_detail - def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> List[str]: + def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> list[str]: """ Flatten the BLOCKED assessments into a list of human-readable category names suitable for queryable OTEL / standard-logging attributes. @@ -1299,7 +1297,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): can still see the count in ``_extract_blocked_assessments`` which feeds the HTTP error detail. """ - names: List[str] = [] + names: list[str] = [] for block in self._extract_blocked_assessments(response): for match in block.get("matches", []) or []: # Allow-list non-sensitive labels only. Never fall back to @@ -1309,7 +1307,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): names.append(label) return names - def _extract_blocked_assessments(self, response: BedrockGuardrailResponse) -> List[dict]: + def _extract_blocked_assessments(self, response: BedrockGuardrailResponse) -> list[dict]: """ Walk the Bedrock guardrail response and emit a structured list of BLOCKED assessment entries describing exactly which policies fired. @@ -1320,7 +1318,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): matched term where available, so the client can render a precise explanation of the violation. """ - blocked: List[dict] = [] + blocked: list[dict] = [] assessments = response.get("assessments", []) or [] for assessment in assessments: @@ -1360,7 +1358,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # Word policy word_policy = assessment.get("wordPolicy") if word_policy: - word_matches: List[dict] = [] + word_matches: list[dict] = [] for w in word_policy.get("customWords") or []: if w.get("action") == "BLOCKED": word_matches.append( @@ -1386,7 +1384,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # Sensitive information policy (PII) sensitive_info = assessment.get("sensitiveInformationPolicy") if sensitive_info: - pii_matches: List[dict] = [] + pii_matches: list[dict] = [] for p in sensitive_info.get("piiEntities") or []: if p.get("action") == "BLOCKED": pii_matches.append( @@ -1441,13 +1439,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return blocked def _get_http_exception_for_blocked_guardrail( - self, response: BedrockGuardrailResponse, request_data: Optional[dict] = None - ) -> Union[HTTPException, ModifyResponseException]: + self, response: BedrockGuardrailResponse, request_data: dict | None = None + ) -> HTTPException | ModifyResponseException: """ Get the HTTP exception for a blocked guardrail. """ bedrock_guardrail_output_text: str = "" - outputs: Optional[List[BedrockGuardrailOutput]] = response.get("outputs", []) or [] + outputs: list[BedrockGuardrailOutput] | None = response.get("outputs", []) or [] if outputs: for output in outputs: if output.get("text"): @@ -1462,7 +1460,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): guardrail_name=self.guardrail_name, ) - detail: Dict[str, Any] = { + detail: dict[str, Any] = { "error": "Violated guardrail policy", "bedrock_guardrail_response": bedrock_guardrail_output_text, } @@ -1560,7 +1558,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Inside Bedrock Pre-Call Hook for call_type: %s", call_type) from litellm.proxy.common_utils.callback_utils import ( @@ -1704,7 +1702,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: return - new_messages: Optional[List[AllMessageValues]] = data.get("messages") + new_messages: list[AllMessageValues] | None = data.get("messages") if new_messages is None: verbose_proxy_logger.warning("Bedrock AI: not running guardrail. No messages in data") return @@ -1768,9 +1766,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ############################################################################## def _update_messages_with_updated_bedrock_guardrail_response( self, - messages: List[AllMessageValues], + messages: list[AllMessageValues], bedrock_guardrail_response: BedrockGuardrailResponse, - ) -> List[AllMessageValues]: + ) -> list[AllMessageValues]: """ Use the output from the bedrock guardrail to mask sensitive content in messages. @@ -1817,11 +1815,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): from litellm.types.utils import TextCompletionResponse # Collect all chunks to process them together - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] async for chunk in response: all_chunks.append(chunk) - assembled_model_response: Optional[Union[ModelResponse, TextCompletionResponse]] = stream_chunk_builder( + assembled_model_response: ModelResponse | TextCompletionResponse | None = stream_chunk_builder( chunks=all_chunks, ) if isinstance(assembled_model_response, ModelResponse): @@ -1890,7 +1888,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): for chunk in all_chunks: yield chunk - def _extract_masked_texts_from_response(self, bedrock_guardrail_response: BedrockGuardrailResponse) -> List[str]: + def _extract_masked_texts_from_response(self, bedrock_guardrail_response: BedrockGuardrailResponse) -> list[str]: """ Extract all masked text outputs from the guardrail response. @@ -1900,22 +1898,22 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): Returns: List of masked text strings """ - masked_output_text: List[str] = [] - masked_outputs: Optional[List[BedrockGuardrailOutput]] = bedrock_guardrail_response.get("outputs", []) or [] + masked_output_text: list[str] = [] + masked_outputs: list[BedrockGuardrailOutput] | None = bedrock_guardrail_response.get("outputs", []) or [] if not masked_outputs: verbose_proxy_logger.debug("No masked outputs found in guardrail response") return [] for output in masked_outputs: - text_content: Optional[str] = output.get("text") + text_content: str | None = output.get("text") if text_content is not None: masked_output_text.append(text_content) return masked_output_text def _apply_masking_to_messages( - self, messages: List[AllMessageValues], masked_texts: List[str] - ) -> List[AllMessageValues]: + self, messages: list[AllMessageValues], masked_texts: list[str] + ) -> list[AllMessageValues]: """ Apply masked texts to message content using index tracking. @@ -1956,8 +1954,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return updated_messages def _mask_content_list( - self, content_list: List[Any], masked_texts: List[str], masking_index: int - ) -> Tuple[List[Any], int]: + self, content_list: list[Any], masked_texts: list[str], masking_index: int + ) -> tuple[list[Any], int]: """ Apply masking to a list of content items. @@ -1969,7 +1967,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): Returns: Updated content list with masked items """ - new_content: List[Union[dict, str]] = [] + new_content: list[dict | str] = [] for item in content_list: if isinstance(item, dict) and "text" in item: new_item = item.copy() @@ -1988,7 +1986,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): def _apply_masking_to_response( self, - response: Union[ModelResponse, Any], + response: ModelResponse | Any, bedrock_guardrail_response: BedrockGuardrailResponse, ) -> None: """ @@ -2013,7 +2011,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): else: verbose_proxy_logger.warning("Unsupported response type for masking: %s", type(response)) - def _apply_masking_to_model_response(self, response: litellm.ModelResponse, masked_texts: List[str]) -> None: + def _apply_masking_to_model_response(self, response: litellm.ModelResponse, masked_texts: list[str]) -> None: """ Apply masked texts to a ModelResponse object. @@ -2268,4 +2266,4 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): raise except Exception as e: verbose_proxy_logger.error("Bedrock Guardrail: Failed to apply guardrail: %s", str(e)) - raise Exception(f"Bedrock guardrail failed: {str(e)}") + raise Exception(f"Bedrock guardrail failed: {e!s}") diff --git a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/__init__.py index 40ed634d39a..8640ba4419a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/__init__.py @@ -1,6 +1,6 @@ """Block Code Execution guardrail: blocks or masks fenced code blocks by language.""" -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Literal, cast from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations @@ -43,8 +43,8 @@ def initialize_guardrail( if not guardrail_name: raise ValueError("Block Code Execution guardrail requires a guardrail_name") - blocked_languages: Optional[List[str]] = cast( - Optional[List[str]], + blocked_languages: list[str] | None = cast( + list[str] | None, _get_param(litellm_params, guardrail, "blocked_languages"), ) action = cast( @@ -53,14 +53,14 @@ def initialize_guardrail( ) confidence_threshold = float( cast( - Union[int, float, str], + int | float | str, _get_param(litellm_params, guardrail, "confidence_threshold", 0.5), ) ) detect_execution_intent = bool(_get_param(litellm_params, guardrail, "detect_execution_intent", True)) mode = _get_param(litellm_params, guardrail, "mode") event_hook = cast( - Optional[Union[Literal["pre_call", "post_call", "during_call"], List[str]]], + Literal["pre_call", "post_call", "during_call"] | list[str] | None, mode if mode is not None else DEFAULT_EVENT_HOOKS, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py index cfb4a78fa6e..8dd43f88dea 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py +++ b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py @@ -11,12 +11,8 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Union, cast, ) @@ -43,7 +39,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj # Language tag aliases (normalize to canonical for comparison) -LANGUAGE_ALIASES: Dict[str, str] = { +LANGUAGE_ALIASES: dict[str, str] = { "js": "javascript", "py": "python", "sh": "bash", @@ -62,7 +58,7 @@ FENCED_BLOCK_RE = re.compile(r"```(\w*)\n(.*?)```", re.DOTALL) # NOTE: Since matching uses substring search (p in text), shorter phrases subsume longer ones. # e.g. "don't run" matches any text containing "don't run it", "but don't run", etc. # Keep only the minimal set; do not add entries subsumed by existing shorter phrases. -_NO_EXECUTION_PHRASES: Tuple[str, ...] = ( +_NO_EXECUTION_PHRASES: tuple[str, ...] = ( # Core negation phrases (short — each subsumes many longer variants) "don't run", "do not run", @@ -122,7 +118,7 @@ _NO_EXECUTION_PHRASES: Tuple[str, ...] = ( # NOTE: Since matching uses substring search (p in text), shorter phrases subsume longer ones. # e.g. "run `" matches any text containing "run `git", "run `docker", etc. # Keep only the minimal set; do not add entries subsumed by existing shorter phrases. -_EXECUTION_REQUEST_PHRASES: Tuple[str, ...] = ( +_EXECUTION_REQUEST_PHRASES: tuple[str, ...] = ( # Direct execution requests (short — each subsumes many longer variants) "run this ", "run these ", @@ -289,7 +285,7 @@ def _normalize_language(tag: str) -> str: def _is_blocked_language( tag: str, - blocked_languages: Optional[List[str]], + blocked_languages: list[str] | None, block_all: bool, ) -> bool: """True if this language tag should be considered blocked.""" @@ -333,17 +329,17 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): def __init__( self, - guardrail_name: Optional[str] = None, - blocked_languages: Optional[List[str]] = None, + guardrail_name: str | None = None, + blocked_languages: list[str] | None = None, action: Literal["block", "mask"] = "block", confidence_threshold: float = 0.5, detect_execution_intent: bool = True, - event_hook: Optional[Union[Literal["pre_call", "post_call", "during_call"], List[str]]] = None, + event_hook: Literal["pre_call", "post_call", "during_call"] | list[str] | None = None, default_on: bool = False, **kwargs: Any, ) -> None: # Normalize to type expected by CustomGuardrail - _event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = None + _event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None if event_hook is not None: if isinstance(event_hook, list): _event_hook = [GuardrailEventHooks(h) if isinstance(h, str) else h for h in event_hook] @@ -367,7 +363,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): self.detect_execution_intent = detect_execution_intent @staticmethod - def get_config_model() -> Optional[type[GuardrailConfigModel]]: + def get_config_model() -> type[GuardrailConfigModel] | None: from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import ( BlockCodeExecutionGuardrailConfigModel, ) @@ -375,19 +371,19 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): return BlockCodeExecutionGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, GuardrailEventHooks.during_call, ] - def _find_blocks(self, text: str) -> List[Tuple[int, int, str, str, float, CodeBlockActionTaken]]: + def _find_blocks(self, text: str) -> list[tuple[int, int, str, str, float, CodeBlockActionTaken]]: """ Find all fenced code blocks in text. Returns list of (start, end, language_tag, block_content, confidence, action_taken). """ - results: List[Tuple[int, int, str, str, float, CodeBlockActionTaken]] = [] + results: list[tuple[int, int, str, str, float, CodeBlockActionTaken]] = [] for m in FENCED_BLOCK_RE.finditer(text): tag = (m.group(1) or "").strip() body = m.group(2) @@ -408,9 +404,9 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): def _scan_text( self, text: str, - detections: Optional[List[CodeBlockDetection]] = None, + detections: list[CodeBlockDetection] | None = None, input_type: Literal["request", "response"] = "request", - ) -> Tuple[str, bool]: + ) -> tuple[str, bool]: """ Scan one text: find blocks, apply block/mask/allow by confidence. When detect_execution_intent is True and input_type is "request", only block if @@ -463,7 +459,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): should_raise = False last_end = 0 - parts: List[str] = [] + parts: list[str] = [] for start, end, tag, _body, confidence, action_taken in blocks: # For responses, always enforce the block action (no intent check needed). # For requests with detect_execution_intent, require execution intent. @@ -525,7 +521,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: start_time = datetime.now() - detections: List[CodeBlockDetection] = [] + detections: list[CodeBlockDetection] = [] status: GuardrailStatus = "success" exception_str = "" @@ -535,7 +531,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): return inputs is_output = input_type == "response" - processed: List[str] = [] + processed: list[str] = [] for text in texts: new_text, should_raise = self._scan_text(text, detections, input_type) processed.append(new_text) @@ -561,15 +557,15 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): exception_str = str(e) raise finally: - guardrail_response: Union[List[dict], str] = [dict(d) for d in detections] + guardrail_response: list[dict] | str = [dict(d) for d in detections] if status != "success" and not detections: guardrail_response = exception_str - max_confidence: Optional[float] = None + max_confidence: float | None = None for d in detections: c = d.get("confidence") if c is not None and (max_confidence is None or c > max_confidence): max_confidence = c - tracing_kw: Dict[str, Any] = { + tracing_kw: dict[str, Any] = { "guardrail_id": self.guardrail_name, "detection_method": "fenced_code_block", "match_details": guardrail_response, diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py index f1d08bfbd56..e0411de2db4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py @@ -10,7 +10,7 @@ import json import os import ssl from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, List, Optional, Type, Union +from typing import TYPE_CHECKING, Any from fastapi import HTTPException from pydantic import BaseModel @@ -52,14 +52,14 @@ class CatoNetworksGuardrailMissingSecrets(Exception): class CatoNetworksGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, GuardrailEventHooks.post_call, ] - def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None, **kwargs): + def __init__(self, api_key: str | None = None, api_base: str | None = None, **kwargs): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) ssl_verify = kwargs.pop("ssl_verify", None) self.async_handler = get_async_httpx_client( @@ -80,7 +80,7 @@ class CatoNetworksGuardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def _build_ws_ssl_kwargs(ssl_verify: Optional[Union[bool, str]], ws_api_base: str) -> dict: + def _build_ws_ssl_kwargs(ssl_verify: bool | str | None, ws_api_base: str) -> dict: """Resolve the ``ssl`` argument for ``websockets.connect``. Mirrors the ``ssl_verify`` handling applied to the HTTP handler so a custom Cato instance behind TLS honours the same verification settings for streaming.""" @@ -94,7 +94,7 @@ class CatoNetworksGuardrail(CustomGuardrail): return {"ssl": ssl_config} @staticmethod - def _resolve_cato_user_email(user_api_key_dict: UserAPIKeyAuth) -> Optional[str]: + def _resolve_cato_user_email(user_api_key_dict: UserAPIKeyAuth) -> str | None: """Only the key/JWT-bound user email is trusted. ``end_user_id`` is derived from caller-supplied request fields (OpenAI ``user``, headers, metadata) and is spoofable, so it must never be forwarded as the Cato user identity.""" @@ -112,7 +112,7 @@ class CatoNetworksGuardrail(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Inside Cato Pre-Call Hook") return await self.call_cato_guardrail( data, @@ -126,7 +126,7 @@ class CatoNetworksGuardrail(CustomGuardrail): data: dict, user_api_key_dict: UserAPIKeyAuth, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Inside Cato Moderation Hook") return await self.call_cato_guardrail( data, @@ -235,8 +235,8 @@ class CatoNetworksGuardrail(CustomGuardrail): self, data: dict, hook: str, - key_alias: Optional[str], - user_email: Optional[str] = None, + key_alias: str | None, + user_email: str | None = None, ) -> dict: call_id = data.get("litellm_call_id") headers = self._build_cato_headers( @@ -346,9 +346,9 @@ class CatoNetworksGuardrail(CustomGuardrail): request_data: dict, output: str, hook: str, - key_alias: Optional[str], - user_email: Optional[str] = None, - ) -> Optional[dict]: + key_alias: str | None, + user_email: str | None = None, + ) -> dict | None: call_id = request_data.get("litellm_call_id") inspection_messages = self._inspection_messages(request_data) assistant_index = len(inspection_messages) @@ -392,9 +392,9 @@ class CatoNetworksGuardrail(CustomGuardrail): self, *, hook: str, - key_alias: Optional[str], - user_email: Optional[str], - litellm_call_id: Optional[str], + key_alias: str | None, + user_email: str | None, + litellm_call_id: str | None, ): """ A helper function to build the http headers that are required by Cato guardrails. @@ -485,8 +485,8 @@ class CatoNetworksGuardrail(CustomGuardrail): data: dict, text: str, user_api_key_dict: UserAPIKeyAuth, - user_email: Optional[str], - ) -> Optional[str]: + user_email: str | None, + ) -> str | None: """Run the Cato output guardrail on a single assistant text fragment. Raises on a block action and returns the redacted replacement, or ``None`` when the fragment must be left unchanged.""" @@ -505,7 +505,7 @@ class CatoNetworksGuardrail(CustomGuardrail): self, data: dict, user_api_key_dict: UserAPIKeyAuth, - response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse], + response: Any | ModelResponse | EmbeddingResponse | ImageResponse, ) -> Any: user_email = self._resolve_cato_user_email(user_api_key_dict) if isinstance(response, ModelResponse) and response.choices: @@ -589,7 +589,7 @@ class CatoNetworksGuardrail(CustomGuardrail): await websocket.send(json.dumps({"done": True})) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.cato_networks import ( CatoNetworksGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py index 4e191c3db52..c28ccadb2b1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py @@ -84,7 +84,7 @@ __all__ = [ "CiscoAIDefenseGuardrail", "CiscoAIDefenseGuardrailAPIError", "CiscoAIDefenseGuardrailMissingSecrets", - "initialize_guardrail", - "guardrail_initializer_registry", "guardrail_class_registry", + "guardrail_initializer_registry", + "initialize_guardrail", ] diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index 10f2def7d77..a1068739791 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -25,13 +25,7 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, - Type, - Union, ) import httpx @@ -75,13 +69,13 @@ CISCO_MCP_INSPECT_PATH = "/api/v1/inspect/mcp" CISCO_API_KEY_HEADER = "X-Cisco-AI-Defense-API-Key" DEFAULT_TIMEOUT_SECONDS = 10.0 -SUPPORTED_INSPECTION_TYPES: Tuple[str, ...] = ("chat", "mcp") +SUPPORTED_INSPECTION_TYPES: tuple[str, ...] = ("chat", "mcp") DEFAULT_INSPECTION_TYPE = "chat" # LiteLLM marks MCP guardrail calls with these call_type values; the proxy # routes pre_mcp_call / during_mcp_call events through async_pre_call_hook / # async_moderation_hook with the call_type set accordingly. -_MCP_CALL_TYPES: Tuple[str, ...] = ("mcp_call", "call_mcp_tool") +_MCP_CALL_TYPES: tuple[str, ...] = ("mcp_call", "call_mcp_tool") # Action vocabulary Cisco AI Defense can return. _ACTION_BLOCK = "block" @@ -101,16 +95,16 @@ class _ScanContext: class _CiscoVerdict: """Parsed Cisco AI Defense decision plus any sanitized rewrites it carries.""" - is_safe: Optional[bool] - classifications: List[str] - severity: Optional[str] - rules: List[Dict[str, Any]] - explanation: Optional[str] - event_id: Optional[str] - action: Optional[str] = None - sanitized_text: Optional[str] = None - sanitized_messages: Optional[List[Dict[str, Any]]] = None - sanitized_mcp_arguments: Optional[Dict[str, Any]] = None + is_safe: bool | None + classifications: list[str] + severity: str | None + rules: list[dict[str, Any]] + explanation: str | None + event_id: str | None + action: str | None = None + sanitized_text: str | None = None + sanitized_messages: list[dict[str, Any]] | None = None + sanitized_mcp_arguments: dict[str, Any] | None = None class CiscoAIDefenseGuardrailMissingSecrets(Exception): @@ -132,28 +126,28 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): ``cisco_ai_defense_mcp.py``. """ - SUPPORTED_ON_FLAGGED_ACTIONS: Tuple[str, ...] = ("block", "monitor") + SUPPORTED_ON_FLAGGED_ACTIONS: tuple[str, ...] = ("block", "monitor") DEFAULT_ON_FLAGGED_ACTION: str = "block" - SUPPORTED_FALLBACK_ACTIONS: Tuple[str, ...] = ("allow", "block") + SUPPORTED_FALLBACK_ACTIONS: tuple[str, ...] = ("allow", "block") DEFAULT_FALLBACK_ON_ERROR: str = "block" _PROVIDER_NAME = "cisco_ai_defense" def __init__( self, - guardrail_name: Optional[str] = "cisco-ai-defense", - api_key: Optional[str] = None, - api_base: Optional[str] = None, - inspection_type: Optional[str] = None, - inspect_path: Optional[str] = None, - enabled_rules: Optional[List[Dict[str, Any]]] = None, - integration_profile_id: Optional[str] = None, - integration_profile_version: Optional[str] = None, - integration_tenant_id: Optional[str] = None, - integration_type: Optional[str] = None, - on_flagged_action: Optional[str] = None, - fallback_on_error: Optional[str] = None, - timeout: Optional[float] = None, + guardrail_name: str | None = "cisco-ai-defense", + api_key: str | None = None, + api_base: str | None = None, + inspection_type: str | None = None, + inspect_path: str | None = None, + enabled_rules: list[dict[str, Any]] | None = None, + integration_profile_id: str | None = None, + integration_profile_version: str | None = None, + integration_tenant_id: str | None = None, + integration_type: str | None = None, + on_flagged_action: str | None = None, + fallback_on_error: str | None = None, + timeout: float | None = None, **kwargs: Any, ) -> None: resolved_api_key = api_key or os.environ.get("CISCO_AI_DEFENSE_API_KEY") @@ -213,7 +207,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): setting_name="fallback_on_error", ) - resolved_timeout: Optional[float] + resolved_timeout: float | None if timeout is not None: resolved_timeout = self._coerce_timeout(timeout) else: @@ -251,9 +245,9 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _resolve_choice( - value: Optional[str], + value: str | None, env_var: str, - allowed: Tuple[str, ...], + allowed: tuple[str, ...], default: str, setting_name: str, ) -> str: @@ -272,7 +266,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return default @staticmethod - def _coerce_timeout(value: Union[str, float]) -> Optional[float]: + def _coerce_timeout(value: str | float) -> float | None: try: parsed = float(value) except (TypeError, ValueError): @@ -289,7 +283,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return parsed @staticmethod - def _is_mcp_call_type(call_type: Optional[str]) -> bool: + def _is_mcp_call_type(call_type: str | None) -> bool: return bool(call_type) and call_type in _MCP_CALL_TYPES # ------------------------------------------------------------------ @@ -314,7 +308,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: # Trust proxy call_type, not caller-controlled request shape. is_mcp = self._is_mcp_call_type(call_type) @@ -363,7 +357,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: is_mcp = self._is_mcp_call_type(call_type) if not self._surface_matches(is_mcp): @@ -447,7 +441,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): self.guardrail_name, ) - all_chunks: List[Any] = [] + all_chunks: list[Any] = [] try: async for chunk in response: all_chunks.append(chunk) @@ -507,7 +501,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): response_obj=assembled, ) except HTTPException as exc: - error_obj: Dict[str, Any] = self._http_exception_to_error_obj(exc) + error_obj: dict[str, Any] = self._http_exception_to_error_obj(exc) verbose_proxy_logger.warning( "Cisco AI Defense guardrail (%s): streaming response " "blocked — emitting SSE error event instead of " @@ -541,7 +535,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): for chunk in all_chunks: yield chunk - def _build_block_payload(self, context: _ScanContext, verdict: _CiscoVerdict) -> Dict[str, Any]: + def _build_block_payload(self, context: _ScanContext, verdict: _CiscoVerdict) -> dict[str, Any]: """Canonical block payload used across all four block paths. Same dict is the ``HTTPException.detail`` for chat / MCP request @@ -565,28 +559,28 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): "event_id": verdict.event_id, } - def _http_exception_to_error_obj(self, exc: HTTPException) -> Dict[str, Any]: + def _http_exception_to_error_obj(self, exc: HTTPException) -> dict[str, Any]: """Wrap an ``HTTPException`` detail into the SSE ``error`` payload. For Cisco's own blocks the detail is already the canonical block payload, so this is a near-passthrough that just adds ``code`` / ``guardrail`` defaults for non-Cisco / unstructured details. """ - error_obj: Dict[str, Any] = dict(exc.detail) if isinstance(exc.detail, dict) else {"message": str(exc.detail)} + error_obj: dict[str, Any] = dict(exc.detail) if isinstance(exc.detail, dict) else {"message": str(exc.detail)} error_obj.setdefault("message", error_obj.get("error", "Guardrail block")) error_obj.setdefault("code", exc.status_code) error_obj.setdefault("guardrail", self.guardrail_name) return error_obj @classmethod - def _streaming_content_was_modified(cls, original_chunks: List[Any], assembled: ModelResponse) -> bool: + def _streaming_content_was_modified(cls, original_chunks: list[Any], assembled: ModelResponse) -> bool: """Decide whether redact changed content or tool/function arguments.""" original_text = cls._extract_streaming_chunk_scan_text(original_chunks) assembled_text = " ".join(m.get("content", "") for m in cls._extract_response_messages(assembled)) return original_text != assembled_text @classmethod - def _extract_streaming_chunk_scan_text(cls, chunks: List[Any]) -> str: + def _extract_streaming_chunk_scan_text(cls, chunks: list[Any]) -> str: original_text = "" argument_text = "" for chunk in chunks: @@ -629,7 +623,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _normalize_event_hooks(event_hook: object) -> set: """Coerce a ``mode`` arg (str, enum, or list of either) to a set of values.""" - def _norm(hook: object) -> Optional[str]: + def _norm(hook: object) -> str | None: value = getattr(hook, "value", None) if isinstance(value, str): return value @@ -683,7 +677,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): allow, WARNING for intervened/redacted, ERROR is left for upstream API failures. """ - fields: Dict[str, Any] = { + fields: dict[str, Any] = { "guardrail": self.guardrail_name, "surface": context.surface, "direction": context.direction, @@ -757,12 +751,12 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): async def _inspect_chat( self, - messages: List[Dict[str, str]], + messages: list[dict[str, str]], request_data: dict, user_api_key_dict: UserAPIKeyAuth, direction: str = "input", response_obj: object = None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: url = f"{self.api_base}{self.inspect_path}" payload = self._build_chat_payload(messages, request_data, user_api_key_dict) start_time = datetime.now() @@ -791,10 +785,10 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _build_chat_payload( self, - messages: List[Dict[str, str]], + messages: list[dict[str, str]], request_data: dict, user_api_key_dict: UserAPIKeyAuth, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: return { "messages": messages, "metadata": self._build_metadata(request_data, user_api_key_dict), @@ -808,9 +802,9 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): async def _post_inspection( self, url: str, - payload: Dict[str, Any], + payload: dict[str, Any], surface: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: headers = self._build_headers() verbose_proxy_logger.debug( "Cisco AI Defense guardrail: posting %s inspection to %s", @@ -854,7 +848,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): f"Cisco AI Defense {surface} API returned a non-JSON response" ) from exc - def _build_headers(self) -> Dict[str, str]: + def _build_headers(self) -> dict[str, str]: return { CISCO_API_KEY_HEADER: self.api_key, "Content-Type": "application/json", @@ -866,8 +860,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): self, request_data: dict, user_api_key_dict: UserAPIKeyAuth, - ) -> Dict[str, Any]: - metadata: Dict[str, Any] = {} + ) -> dict[str, Any]: + metadata: dict[str, Any] = {} user = request_data.get("user") or getattr(user_api_key_dict, "user_id", None) if user: @@ -894,8 +888,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return metadata - def _build_config(self) -> Dict[str, Any]: - config: Dict[str, Any] = {} + def _build_config(self) -> dict[str, Any]: + config: dict[str, Any] = {} if self.enabled_rules: config["enabled_rules"] = self.enabled_rules if self.integration_profile_id: @@ -909,7 +903,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return config @staticmethod - def _normalize_rule(rule: object) -> Dict[str, Any]: + def _normalize_rule(rule: object) -> dict[str, Any]: """Coerce a user-supplied rule into the wire-shape dict Cisco expects. Accepts ``str``, ``dict``, and Pydantic model inputs. @@ -932,7 +926,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): rule = dumped if isinstance(rule, dict): - normalized: Dict[str, Any] = {} + normalized: dict[str, Any] = {} rule_name = rule.get("rule_name") if rule_name: normalized["rule_name"] = rule_name @@ -955,12 +949,12 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _finalize_inspection( self, - inspect_response: Dict[str, Any], + inspect_response: dict[str, Any], request_data: dict, context: _ScanContext, start_time: datetime, response_obj: object = None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Parse, log, and (optionally) raise/redact on the Cisco verdict. ``context.direction`` is ``"input"`` for request scans and ``"output"`` @@ -1129,10 +1123,10 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @classmethod def _sanitize_response_for_logging( cls, - inspect_response: Dict[str, Any], + inspect_response: dict[str, Any], surface: str, - action: Optional[str] = None, - ) -> Dict[str, Any]: + action: str | None = None, + ) -> dict[str, Any]: """Drop bulky / privacy-sensitive fields, recursing into nested dicts. MCP verdicts are commonly nested under ``result``, so a @@ -1148,9 +1142,9 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return sanitized @classmethod - def _strip_sensitive_keys(cls, d: Dict[str, Any]) -> Dict[str, Any]: + def _strip_sensitive_keys(cls, d: dict[str, Any]) -> dict[str, Any]: """Recursively strip privacy-sensitive keys from a verdict dict.""" - out: Dict[str, Any] = {} + out: dict[str, Any] = {} for key, value in d.items(): if key.startswith("_") or key in cls._REDACTED_LOG_KEYS: continue @@ -1164,7 +1158,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): # Verdict extraction helpers (sanitized content + JSON-RPC errors) # ------------------------------------------------------------------ - _DECISION_FIELDS: Tuple[str, ...] = ( + _DECISION_FIELDS: tuple[str, ...] = ( "action", "allowed", "blocked", @@ -1195,7 +1189,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return any(key in payload for key in cls._DECISION_FIELDS) @classmethod - def _unwrap_verdict_envelope(cls, inspect_response: Dict[str, Any]) -> Dict[str, Any]: + def _unwrap_verdict_envelope(cls, inspect_response: dict[str, Any]) -> dict[str, Any]: """Return the dict that actually holds is_safe / action / rules. Cisco AI Defense returns the verdict at different nesting depths @@ -1232,8 +1226,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _extract_jsonrpc_error( - inspect_response: Dict[str, Any], - ) -> Optional[Dict[str, Any]]: + inspect_response: dict[str, Any], + ) -> dict[str, Any] | None: """Detect a JSON-RPC error envelope inside an HTTP 200 response. The Cisco Inspect API can return ``{"error": {...}}`` (or nest one @@ -1280,8 +1274,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _extract_sanitized_text( - inspect_response: Dict[str, Any], - ) -> Optional[str]: + inspect_response: dict[str, Any], + ) -> str | None: """Pull ``sanitized_text`` (or camelCase variant) off the verdict.""" for key in ("sanitized_text", "sanitizedText"): value = inspect_response.get(key) @@ -1297,8 +1291,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _extract_sanitized_messages( - inspect_response: Dict[str, Any], - ) -> Optional[List[Dict[str, Any]]]: + inspect_response: dict[str, Any], + ) -> list[dict[str, Any]] | None: """Pull a sanitized OpenAI-format messages array off the verdict. Cisco can return the rewrite under several keys; we accept any of @@ -1363,8 +1357,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _redact_mcp_input( request_data: dict, - sanitized_text: Optional[str], - sanitized_mcp_arguments: Optional[Dict[str, Any]], + sanitized_text: str | None, + sanitized_mcp_arguments: dict[str, Any] | None, ) -> bool: """Rewrite MCP request arguments in all locations the proxy reads.""" if sanitized_mcp_arguments is not None: @@ -1397,8 +1391,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _redact_chat_input( self, request_data: dict, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: """Rewrite chat request input (``messages`` or ``input``).""" if sanitized_messages and self._extract_tool_definition_text(request_data): @@ -1453,8 +1447,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _redact_responses_instructions( cls, request_data: dict, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: if sanitized_messages: instruction_text = cls._instruction_text_from_messages(sanitized_messages) @@ -1467,7 +1461,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return False @classmethod - def _instruction_text_from_messages(cls, messages: List[Dict[str, Any]]) -> Optional[str]: + def _instruction_text_from_messages(cls, messages: list[dict[str, Any]]) -> str | None: for message in messages: if not isinstance(message, dict): continue @@ -1478,7 +1472,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return None @classmethod - def _non_instruction_messages(cls, messages: Optional[List[Dict[str, Any]]]) -> Optional[List[Dict[str, Any]]]: + def _non_instruction_messages(cls, messages: list[dict[str, Any]] | None) -> list[dict[str, Any]] | None: if messages is None: return None return [ @@ -1508,8 +1502,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _redact_chat_output( self, response_obj: object, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: """Rewrite chat response (``ModelResponse`` or ``ResponsesAPIResponse``).""" if response_obj is None: @@ -1535,8 +1529,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _redact_model_response_choices( choices: list, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: """Redact every returned choice, including tool-call/reasoning fields.""" if sanitized_messages: @@ -1579,8 +1573,8 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _redact_text_completion_choices( choices: list, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: """Rewrite ``/v1/completions`` text choices after Cisco redaction.""" replacement = sanitized_text @@ -1647,10 +1641,10 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): def _redact_responses_api_output( self, output_items: list, - sanitized_text: Optional[str], - sanitized_messages: Optional[List[Dict[str, Any]]], + sanitized_text: str | None, + sanitized_messages: list[dict[str, Any]] | None, ) -> bool: - replacement_text: Optional[str] = sanitized_text + replacement_text: str | None = sanitized_text if not replacement_text and sanitized_messages: replacement_text = " ".join( self._normalize_message_content(m.get("content")) for m in sanitized_messages if isinstance(m, dict) @@ -1682,14 +1676,14 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _sanitized_messages_to_responses_input( - sanitized_messages: List[Dict[str, Any]], - ) -> Optional[List[Dict[str, Any]]]: + sanitized_messages: list[dict[str, Any]], + ) -> list[dict[str, Any]] | None: """Convert chat-shape sanitized_messages to Responses API ``input``. Returns ``None`` if nothing usable could be converted, so the caller falls back to ``on_flagged_action``. """ - out: List[Dict[str, Any]] = [] + out: list[dict[str, Any]] = [] for m in sanitized_messages: if not isinstance(m, dict): continue @@ -1703,7 +1697,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return out or None @staticmethod - def _rewrite_responses_input_text(original_input: object, sanitized_text: str) -> Optional[object]: + def _rewrite_responses_input_text(original_input: object, sanitized_text: str) -> object | None: """Apply ``sanitized_text`` to a Responses API ``input`` value. Handles plain string, list of message items (rewrites the last @@ -1746,12 +1740,12 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _extract_masked_entity_count( - rules: List[Dict[str, Any]], - ) -> Optional[Dict[str, int]]: + rules: list[dict[str, Any]], + ) -> dict[str, int] | None: """Count entity-type detections per Cisco rule for the logging payload.""" if not rules: return None - counts: Dict[str, int] = {} + counts: dict[str, int] = {} for rule in rules: if not isinstance(rule, dict): continue @@ -1770,11 +1764,11 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): self, error: Exception, *, - request_data: Optional[dict] = None, - start_time: Optional[datetime] = None, + request_data: dict | None = None, + start_time: datetime | None = None, surface: str = "chat", direction: str = "input", - ) -> Dict[str, Any]: + ) -> dict[str, Any]: verbose_proxy_logger.error( "Cisco AI Defense guardrail (%s): API communication failed: %s", surface, @@ -1839,9 +1833,9 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): @staticmethod def _extract_inspect_messages_from_request( data: dict, - ) -> List[Dict[str, str]]: + ) -> list[dict[str, str]]: """Build {role, content} messages for the Cisco AI Defense chat API.""" - messages: List[Dict[str, str]] = [] + messages: list[dict[str, str]] = [] instructions_text = CiscoAIDefenseGuardrail._normalize_message_content(data.get("instructions")) if instructions_text: @@ -1854,7 +1848,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): role = message.get("role") if not role: continue - parts: List[str] = [] + parts: list[str] = [] text = CiscoAIDefenseGuardrail._normalize_message_content(message.get("content")) if text: parts.append(text) @@ -1889,13 +1883,13 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): be inspected too; otherwise it bypasses the guardrail by hiding in ``tools[].function.description`` and similar metadata. """ - parts: List[str] = [] + parts: list[str] = [] for key in ("tools", "functions"): CiscoAIDefenseGuardrail._collect_strings(data.get(key), parts) return " ".join(parts) @staticmethod - def _collect_strings(value: object, out: List[str]) -> None: + def _collect_strings(value: object, out: list[str]) -> None: if isinstance(value, str): if value: out.append(value) @@ -1907,7 +1901,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): CiscoAIDefenseGuardrail._collect_strings(item, out) @staticmethod - def _flatten_responses_input(input_value: object) -> List[Dict[str, str]]: + def _flatten_responses_input(input_value: object) -> list[dict[str, str]]: """Flatten the OpenAI Responses API ``input`` into chat-message form. Recognized shapes: @@ -1930,7 +1924,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return [{"role": "user", "content": text}] if text else [] if any(isinstance(item, dict) and "role" in item for item in input_value): - result: List[Dict[str, str]] = [] + result: list[dict[str, str]] = [] for item in input_value: if not isinstance(item, dict): continue @@ -1963,7 +1957,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): if isinstance(content, str): return content if isinstance(content, list): - parts: List[str] = [] + parts: list[str] = [] for part in content: if not isinstance(part, dict): continue @@ -1984,7 +1978,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return str(content) @staticmethod - def _extract_response_messages(response: object) -> List[Dict[str, str]]: + def _extract_response_messages(response: object) -> list[dict[str, str]]: """Extract scannable assistant text from a chat response. Handles both ``ModelResponse`` (Chat Completions) and @@ -1994,11 +1988,11 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): scan by placing content there. """ if isinstance(response, ModelResponse): - result: List[Dict[str, str]] = [] + result: list[dict[str, str]] = [] for choice in getattr(response, "choices", None) or []: if not isinstance(choice, Choices): continue - parts: List[str] = [] + parts: list[str] = [] content = CiscoAIDefenseGuardrail._normalize_message_content(getattr(choice.message, "content", None)) if content: parts.append(content) @@ -2009,7 +2003,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return result if isinstance(response, TextCompletionResponse): - text_parts: List[str] = [] + text_parts: list[str] = [] for choice in getattr(response, "choices", None) or []: text = getattr(choice, "text", None) if isinstance(text, str) and text: @@ -2020,7 +2014,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): output_items = getattr(response, "output", None) if not isinstance(output_items, list): return [] - output_parts: List[str] = [] + output_parts: list[str] = [] for item in output_items: get = item.get if isinstance(item, dict) else (lambda k: getattr(item, k, None)) for part in get("content") or []: @@ -2039,9 +2033,9 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return [{"role": "assistant", "content": joined}] if joined else [] @classmethod - def _extract_message_reasoning_parts(cls, message: object) -> List[str]: + def _extract_message_reasoning_parts(cls, message: object) -> list[str]: """Extract inspectable reasoning fields from a message/delta object.""" - parts: List[str] = [] + parts: list[str] = [] reasoning_content = cls._field(message, "reasoning_content") if isinstance(reasoning_content, str) and reasoning_content: parts.append(reasoning_content) @@ -2070,13 +2064,13 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return getattr(obj, key, None) @classmethod - def _field_list(cls, obj: object, key: str) -> List[Any]: + def _field_list(cls, obj: object, key: str) -> list[Any]: value = cls._field(obj, key) return value if isinstance(value, list) else [] @classmethod - def _extract_message_tool_argument_parts(cls, message: object) -> List[str]: - parts: List[str] = [] + def _extract_message_tool_argument_parts(cls, message: object) -> list[str]: + parts: list[str] = [] tool_calls = message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None) for tool_call in tool_calls or []: args = cls._extract_tool_call_arguments(tool_call) @@ -2092,7 +2086,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return parts @staticmethod - def _extract_tool_call_arguments(tool_call: object) -> Optional[str]: + def _extract_tool_call_arguments(tool_call: object) -> str | None: """Pull ``function.arguments`` off a tool_calls entry (dict or model).""" if tool_call is None: return None @@ -2100,7 +2094,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return CiscoAIDefenseGuardrail._extract_function_call_arguments(function) @staticmethod - def _extract_function_call_arguments(function_call: object) -> Optional[str]: + def _extract_function_call_arguments(function_call: object) -> str | None: """Pull ``arguments`` off a function_call entry (dict or model).""" if function_call is None: return None @@ -2118,7 +2112,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): # ------------------------------------------------------------------ @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import ( CiscoAIDefenseGuardrailConfigModel, ) @@ -2126,7 +2120,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): return CiscoAIDefenseGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py index 2b53d71e8be..0161ba296b9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py @@ -6,7 +6,7 @@ while preserving the existing public import path. """ from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Optional from fastapi import HTTPException @@ -23,7 +23,7 @@ if TYPE_CHECKING: from .cisco_ai_defense import _ScanContext -def _serialize_mcp_content_item(item: object) -> Dict[str, Any]: +def _serialize_mcp_content_item(item: object) -> dict[str, Any]: """Serialize an MCP content item to a JSON-friendly dict. Handles raw dicts, MCP SDK Pydantic models, and simple ``.text`` objects. @@ -53,30 +53,30 @@ class _CiscoAIDefenseMcpMixin: inspect_path: str inspection_type: str _PROVIDER_NAME: str - guardrail_name: Optional[str] + guardrail_name: str | None def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: ... - async def _post_inspection(self, url: str, payload: Dict[str, Any], surface: str) -> Dict[str, Any]: ... + async def _post_inspection(self, url: str, payload: dict[str, Any], surface: str) -> dict[str, Any]: ... def _handle_api_error( self, error: Exception, *, - request_data: Optional[dict] = ..., - start_time: Optional[datetime] = ..., + request_data: dict | None = ..., + start_time: datetime | None = ..., surface: str = ..., direction: str = ..., - ) -> Dict[str, Any]: ... + ) -> dict[str, Any]: ... def _finalize_inspection( self, - inspect_response: Dict[str, Any], + inspect_response: dict[str, Any], request_data: dict, context: "_ScanContext", start_time: datetime, response_obj: object = ..., - ) -> Dict[str, Any]: ... + ) -> dict[str, Any]: ... # ------------------------------------------------------------------ # MCP post-tool hook (dispatcher contract) @@ -95,7 +95,7 @@ class _CiscoAIDefenseMcpMixin: if self.inspection_type != "mcp": return None - request_data: Dict[str, Any] = {} + request_data: dict[str, Any] = {} for key in ( "name", "litellm_call_id", @@ -168,9 +168,10 @@ class _CiscoAIDefenseMcpMixin: """Build a synthetic MCPPostCallResponseObject for blocked output.""" import json as _json + from mcp.types import TextContent + from litellm.types.llms.base import HiddenParams from litellm.types.mcp import MCPPostCallResponseObject - from mcp.types import TextContent if isinstance(detail, dict): payload = detail @@ -252,7 +253,7 @@ class _CiscoAIDefenseMcpMixin: @staticmethod def _replacement_structured_content( replacement: object, - ) -> Optional[Dict[str, str]]: + ) -> dict[str, str] | None: if not isinstance(replacement, list) or not replacement: return None first = replacement[0] @@ -275,7 +276,7 @@ class _CiscoAIDefenseMcpMixin: self, data: dict, user_api_key_dict: UserAPIKeyAuth, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: del user_api_key_dict # carried via logging metadata, not the wire payload url = f"{self.api_base}{self.inspect_path}" payload = self._build_mcp_request_payload(data=data) @@ -309,9 +310,9 @@ class _CiscoAIDefenseMcpMixin: self, request_data: dict, response: object, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, redact_response_obj: object = None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: del user_api_key_dict # carried via logging metadata, not the wire payload url = f"{self.api_base}{self.inspect_path}" payload = self._build_mcp_response_payload( @@ -348,7 +349,7 @@ class _CiscoAIDefenseMcpMixin: def _build_mcp_request_payload( self, data: dict, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """Build the JSON-RPC ``tools/call`` envelope sent to ``/inspect/mcp``. The Cisco AI Defense MCP inspect endpoint expects the JSON-RPC @@ -389,7 +390,7 @@ class _CiscoAIDefenseMcpMixin: self, request_data: dict, response: object, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """Build the MCP response-inspection body sent to ``/inspect/mcp``.""" request_payload = self._build_mcp_request_payload(data=request_data) if request_payload is None: @@ -414,7 +415,7 @@ class _CiscoAIDefenseMcpMixin: return payload @staticmethod - def _hydrate_mcp_tool_context(request_data: Dict[str, Any]) -> None: + def _hydrate_mcp_tool_context(request_data: dict[str, Any]) -> None: metadata = request_data.get("mcp_tool_call_metadata") if metadata is None: nested = request_data.get("metadata") or request_data.get("litellm_metadata") @@ -439,7 +440,7 @@ class _CiscoAIDefenseMcpMixin: request_data.setdefault("server_name", server_name) @staticmethod - def _normalize_mcp_response(response: object) -> Optional[Dict[str, Any]]: + def _normalize_mcp_response(response: object) -> dict[str, Any] | None: """Normalize an MCP tool response into a JSON-RPC envelope. Handles JSON-RPC dicts, raw content lists, MCP SDK models, and @@ -501,10 +502,10 @@ class _CiscoAIDefenseMcpMixin: @staticmethod def _build_mcp_result( - content: List[Any], + content: list[Any], source: object = None, - ) -> Dict[str, Any]: - result: Dict[str, Any] = {"content": [_serialize_mcp_content_item(item) for item in content]} + ) -> dict[str, Any]: + result: dict[str, Any] = {"content": [_serialize_mcp_content_item(item) for item in content]} for key in ("structuredContent", "isError"): value = source.get(key) if isinstance(source, dict) else getattr(source, key, None) if value is not None and (key != "isError" or isinstance(value, bool)): @@ -558,7 +559,7 @@ class _CiscoAIDefenseMcpMixin: pass elif isinstance(response_obj, dict): result = response_obj.get("result") - target: Dict[Any, Any] = result if isinstance(result, dict) else response_obj + target: dict[Any, Any] = result if isinstance(result, dict) else response_obj if "structuredContent" in target: target["structuredContent"] = replacement replaced = True @@ -566,7 +567,7 @@ class _CiscoAIDefenseMcpMixin: return replaced @staticmethod - def _coerce_to_content_list(response_obj: object) -> Optional[List[Any]]: + def _coerce_to_content_list(response_obj: object) -> list[Any] | None: """Find the MCP content list inside supported response shapes.""" if response_obj is None: return None @@ -593,8 +594,8 @@ class _CiscoAIDefenseMcpMixin: @staticmethod def _extract_sanitized_mcp_arguments( - inspect_response: Dict[str, Any], - ) -> Optional[Dict[str, Any]]: + inspect_response: dict[str, Any], + ) -> dict[str, Any] | None: """Pull sanitized MCP tool-call arguments off the verdict. Cisco can return them at the top level (``params.arguments``) or diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index fa71e7fc301..600da6ecfc4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -4,11 +4,9 @@ from collections.abc import Mapping, Sequence from typing import ( TYPE_CHECKING, Annotated, - List, Literal, NamedTuple, Optional, - Union, cast, ) @@ -44,8 +42,6 @@ if TYPE_CHECKING: class CrowdStrikeAIDRGuardrailMissingSecrets(Exception): """Custom exception for missing CrowdStrike AIDR secrets.""" - pass - class _TextContentPart(BaseModel): model_config = ConfigDict(extra="forbid") @@ -65,32 +61,32 @@ class _ImageUrlContentPart(BaseModel): image_url: _ImageUrl -_ContentPart = Annotated[Union[_TextContentPart, _ImageUrlContentPart], Field(discriminator="type")] +_ContentPart = Annotated[_TextContentPart | _ImageUrlContentPart, Field(discriminator="type")] class _Message(BaseModel): role: str - content: Optional[Union[str, list[_ContentPart]]] = None + content: str | list[_ContentPart] | None = None class _GuardInput(BaseModel): messages: list[_Message] - tools: Optional[Sequence[OpenAIChatCompletionToolParam]] = None + tools: Sequence[OpenAIChatCompletionToolParam] | None = None class _GuardChatCompletionsResult(BaseModel): - guard_output: Optional[_GuardInput] = None + guard_output: _GuardInput | None = None """Updated structured prompt.""" - blocked: Optional[bool] = None + blocked: bool | None = None """Whether or not the prompt triggered a block detection.""" - transformed: Optional[bool] = None + transformed: bool | None = None """Whether or not the original input was transformed.""" - detectors: Optional[dict[str, Any]] = None + detectors: dict[str, Any] | None = None """Result of the policy analyzing and input prompt.""" class _GuardChatCompletionsResponse(BaseModel): - result: Optional[_GuardChatCompletionsResult] = None + result: _GuardChatCompletionsResult | None = None class _FilteredMessages(NamedTuple): @@ -239,7 +235,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): """ @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py index 3b031999f40..fe19ece6237 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py @@ -63,8 +63,8 @@ guardrail_class_registry = { } __all__ = [ - "CustomCodeGuardrail", "DEFAULT_REJECTION_PHRASES", "RESPONSE_REJECTION_GUARDRAIL_CODE", + "CustomCodeGuardrail", "initialize_guardrail", ] diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index 245b9806e71..c7a036562f7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -36,7 +36,7 @@ Example: block when response rejects the user (input_type response only): import asyncio import threading -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, cast from fastapi import HTTPException @@ -59,7 +59,7 @@ if TYPE_CHECKING: class CustomCodeGuardrailError(Exception): """Raised when custom code guardrail execution fails.""" - def __init__(self, message: str, details: Optional[Dict[str, Any]] = None) -> None: + def __init__(self, message: str, details: dict[str, Any] | None = None) -> None: super().__init__(message) self.details = details or {} @@ -105,7 +105,7 @@ class CustomCodeGuardrail(CustomGuardrail): def __init__( self, custom_code: str, - guardrail_name: Optional[str] = "custom_code", + guardrail_name: str | None = "custom_code", **kwargs: Any, ) -> None: """ @@ -117,9 +117,9 @@ class CustomCodeGuardrail(CustomGuardrail): **kwargs: Additional arguments passed to CustomGuardrail """ self.custom_code = custom_code - self._compiled_function: Optional[Any] = None + self._compiled_function: Any | None = None self._compile_lock = threading.Lock() - self._compile_error: Optional[str] = None + self._compile_error: str | None = None super().__init__( guardrail_name=guardrail_name, @@ -131,12 +131,12 @@ class CustomCodeGuardrail(CustomGuardrail): self._compile_custom_code() @staticmethod - def get_config_model() -> Optional[Type[GuardrailConfigModel]]: + def get_config_model() -> type[GuardrailConfigModel] | None: """Returns the config model for the UI.""" return CustomCodeGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -263,7 +263,7 @@ class CustomCodeGuardrail(CustomGuardrail): }, ) from e - def _prepare_safe_request_data(self, request_data: dict) -> Dict[str, Any]: + def _prepare_safe_request_data(self, request_data: dict) -> dict[str, Any]: """ Prepare a safe subset of request_data for code execution. diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py index e60b900428c..c019d445fd4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -7,7 +7,7 @@ and provide safe, sandboxed functionality for common guardrail operations. import json import re -from typing import Any, Dict, List, Optional, Tuple, Type, Union +from typing import Any from urllib.parse import urlparse import httpx @@ -21,7 +21,7 @@ from litellm.types.llms.custom_http import httpxSpecialProvider # ============================================================================= -def allow() -> Dict[str, Any]: +def allow() -> dict[str, Any]: """ Allow the request/response to proceed unchanged. @@ -31,7 +31,7 @@ def allow() -> Dict[str, Any]: return {"action": "allow"} -def block(reason: str, detection_info: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: +def block(reason: str, detection_info: dict[str, Any] | None = None) -> dict[str, Any]: """ Block the request/response with a reason. @@ -42,17 +42,17 @@ def block(reason: str, detection_info: Optional[Dict[str, Any]] = None) -> Dict[ Returns: Dict indicating the request should be blocked """ - result: Dict[str, Any] = {"action": "block", "reason": reason} + result: dict[str, Any] = {"action": "block", "reason": reason} if detection_info: result["detection_info"] = detection_info return result def modify( - texts: Optional[List[str]] = None, - images: Optional[List[Any]] = None, - tool_calls: Optional[List[Any]] = None, -) -> Dict[str, Any]: + texts: list[str] | None = None, + images: list[Any] | None = None, + tool_calls: list[Any] | None = None, +) -> dict[str, Any]: """ Modify the request/response content. @@ -64,7 +64,7 @@ def modify( Returns: Dict indicating the content should be modified """ - result: Dict[str, Any] = {"action": "modify"} + result: dict[str, Any] = {"action": "modify"} if texts is not None: result["texts"] = texts if images is not None: @@ -137,7 +137,7 @@ def regex_replace(text: str, pattern: str, replacement: str, flags: int = 0) -> return text -def regex_find_all(text: str, pattern: str, flags: int = 0) -> List[str]: +def regex_find_all(text: str, pattern: str, flags: int = 0) -> list[str]: """ Find all occurrences of a pattern in text. @@ -161,7 +161,7 @@ def regex_find_all(text: str, pattern: str, flags: int = 0) -> List[str]: # ============================================================================= -def json_parse(text: str) -> Optional[Any]: +def json_parse(text: str) -> Any | None: """ Parse a JSON string into a Python object. @@ -195,7 +195,7 @@ def json_stringify(obj: Any) -> str: return "" -def json_schema_valid(obj: Any, schema: Dict[str, Any]) -> bool: +def json_schema_valid(obj: Any, schema: dict[str, Any]) -> bool: """ Validate an object against a JSON schema. @@ -226,7 +226,7 @@ def json_schema_valid(obj: Any, schema: Dict[str, Any]) -> bool: return False -def _basic_json_schema_validate(obj: Any, schema: Dict[str, Any], max_depth: int = 50) -> bool: +def _basic_json_schema_validate(obj: Any, schema: dict[str, Any], max_depth: int = 50) -> bool: """ Basic JSON schema validation without external library. Handles: type, required, properties @@ -234,7 +234,7 @@ def _basic_json_schema_validate(obj: Any, schema: Dict[str, Any], max_depth: int Uses an iterative approach with a stack to avoid recursion limits. max_depth limits nesting to prevent infinite loops from circular schemas. """ - type_map: Dict[str, Union[Type, Tuple[Type, ...]]] = { + type_map: dict[str, type | tuple[type, ...]] = { "object": dict, "array": list, "string": str, @@ -245,7 +245,7 @@ def _basic_json_schema_validate(obj: Any, schema: Dict[str, Any], max_depth: int } # Stack of (obj, schema, depth) tuples to process - stack: List[Tuple[Any, Dict[str, Any], int]] = [(obj, schema, 0)] + stack: list[tuple[Any, dict[str, Any], int]] = [(obj, schema, 0)] while stack: current_obj, current_schema, depth = stack.pop() @@ -286,7 +286,7 @@ def _basic_json_schema_validate(obj: Any, schema: Dict[str, Any], max_depth: int _URL_PATTERN = re.compile(r"https?://(?:[-\w.]|(?:%[\da-fA-F]{2}))+[^\s]*", re.IGNORECASE) -def extract_urls(text: str) -> List[str]: +def extract_urls(text: str) -> list[str]: """ Extract all URLs from text. @@ -330,7 +330,7 @@ def all_urls_valid(text: str) -> bool: return all(is_valid_url(url) for url in urls) -def get_url_domain(url: str) -> Optional[str]: +def get_url_domain(url: str) -> str | None: """ Extract the domain from a URL. @@ -358,7 +358,7 @@ _HTTP_DEFAULT_TIMEOUT = 30.0 _HTTP_MAX_TIMEOUT = 60.0 -def _http_error_response(error: str) -> Dict[str, Any]: +def _http_error_response(error: str) -> dict[str, Any]: """Create a standardized error response for HTTP requests.""" return { "status_code": 0, @@ -369,7 +369,7 @@ def _http_error_response(error: str) -> Dict[str, Any]: } -def _http_success_response(response: httpx.Response) -> Dict[str, Any]: +def _http_success_response(response: httpx.Response) -> dict[str, Any]: """Create a standardized success response from an httpx Response.""" parsed_body: Any try: @@ -387,8 +387,8 @@ def _http_success_response(response: httpx.Response) -> Dict[str, Any]: def _prepare_http_body( - body: Optional[Any], -) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: + body: Any | None, +) -> tuple[dict[str, Any] | None, str | None]: """Prepare body arguments for HTTP request - returns (json_body, data_body).""" if body is None: return None, None @@ -404,10 +404,10 @@ def _prepare_http_body( async def http_request( url: str, method: str = "GET", - headers: Optional[Dict[str, str]] = None, - body: Optional[Any] = None, - timeout: Optional[float] = None, -) -> Dict[str, Any]: + headers: dict[str, str] | None = None, + body: Any | None = None, + timeout: float | None = None, +) -> dict[str, Any]: """ Make an async HTTP request to an external service. @@ -480,18 +480,18 @@ async def http_request( return _http_success_response(e.response) except httpx.RequestError as e: verbose_proxy_logger.warning(f"Custom code http_request error: {e}") - return _http_error_response(f"Request failed: {str(e)}") + return _http_error_response(f"Request failed: {e!s}") except Exception as e: verbose_proxy_logger.warning(f"Custom code http_request unexpected error: {e}") - return _http_error_response(f"Unexpected error: {str(e)}") + return _http_error_response(f"Unexpected error: {e!s}") async def _execute_http_request( client: Any, method: str, url: str, - headers: Optional[Dict[str, str]], - body: Optional[Any], + headers: dict[str, str] | None, + body: Any | None, timeout: float, ) -> httpx.Response: """Execute the HTTP request using the appropriate client method.""" @@ -513,9 +513,9 @@ async def _execute_http_request( async def http_get( url: str, - headers: Optional[Dict[str, str]] = None, - timeout: Optional[float] = None, -) -> Dict[str, Any]: + headers: dict[str, str] | None = None, + timeout: float | None = None, +) -> dict[str, Any]: """ Make an async HTTP GET request. @@ -534,10 +534,10 @@ async def http_get( async def http_post( url: str, - body: Optional[Any] = None, - headers: Optional[Dict[str, str]] = None, - timeout: Optional[float] = None, -) -> Dict[str, Any]: + body: Any | None = None, + headers: dict[str, str] | None = None, + timeout: float | None = None, +) -> dict[str, Any]: """ Make an async HTTP POST request. @@ -625,7 +625,7 @@ def detect_code(text: str) -> bool: return len(detect_code_languages(text)) > 0 -def detect_code_languages(text: str) -> List[str]: +def detect_code_languages(text: str) -> list[str]: """ Detect which programming languages are present in text. @@ -647,7 +647,7 @@ def detect_code_languages(text: str) -> List[str]: return detected -def contains_code_language(text: str, languages: List[str]) -> bool: +def contains_code_language(text: str, languages: list[str]) -> bool: """ Check if text contains code from specific languages. @@ -681,7 +681,7 @@ def contains(text: str, substring: str) -> bool: return substring in text -def contains_any(text: str, substrings: List[str]) -> bool: +def contains_any(text: str, substrings: list[str]) -> bool: """ Check if text contains any of the given substrings. @@ -695,7 +695,7 @@ def contains_any(text: str, substrings: List[str]) -> bool: return any(s in text for s in substrings) -def contains_all(text: str, substrings: List[str]) -> bool: +def contains_all(text: str, substrings: list[str]) -> bool: """ Check if text contains all of the given substrings. @@ -755,7 +755,7 @@ def trim(text: str) -> str: # ============================================================================= -def get_custom_code_primitives() -> Dict[str, Any]: +def get_custom_code_primitives() -> dict[str, Any]: """ Get all primitives to inject into the custom code environment. diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py index 2250daaf899..c1d6a52903f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py @@ -15,7 +15,7 @@ restriction intact. """ import operator -from typing import Any, Dict +from typing import Any from RestrictedPython import ( RestrictingNodeTransformer, @@ -58,7 +58,7 @@ class AsyncAwareTransformer(RestrictingNodeTransformer): return self.node_contents_visit(node) -_INPLACE_OPS: Dict[str, Any] = { +_INPLACE_OPS: dict[str, Any] = { "+=": operator.iadd, "-=": operator.isub, "*=": operator.imul, @@ -86,7 +86,7 @@ def _inplacevar_(op: str, x: Any, y: Any) -> Any: return fn(x, y) -def _build_sandbox_builtins() -> Dict[str, Any]: +def _build_sandbox_builtins() -> dict[str, Any]: # ``limited_builtins`` overrides ``list``/``tuple``/``range`` from # ``safe_builtins`` with bounds-checking variants (e.g. ``limited_range`` # rejects ``range(10**18)``). ``utility_builtins`` adds ``set``, @@ -98,14 +98,14 @@ def _build_sandbox_builtins() -> Dict[str, Any]: } -def build_sandbox_globals() -> Dict[str, Any]: +def build_sandbox_globals() -> dict[str, Any]: """Assemble the globals dict for executing guardrail code. Includes the LiteLLM-provided primitives (``regex_match``, ``http_get``, ``allow``/``block``/``modify``, etc.) plus the RestrictedPython guards that the compiled bytecode expects to find by name. """ - sandbox: Dict[str, Any] = get_custom_code_primitives().copy() + sandbox: dict[str, Any] = get_custom_code_primitives().copy() sandbox["__builtins__"] = _build_sandbox_builtins() sandbox["_getattr_"] = safer_getattr sandbox["_getitem_"] = default_guarded_getitem diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py index 51277936069..e4c21229f6a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py @@ -1,4 +1,4 @@ -from typing import Literal, Optional, Union +from typing import Literal import litellm from litellm._logging import verbose_proxy_logger @@ -38,7 +38,7 @@ class myCustomGuardrail(CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call Runs on only Input diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index cef359d5c21..c7ed4028218 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -37,14 +37,10 @@ _DEEPKEEP_GUARDRAIL_ENDPOINT = "/v3/openai/beta/litellm_basic_guardrail_api" class DeepKeepGuardrailMissingSecrets(Exception): """Exception raised when DeepKeep API key or firewall_id is missing.""" - pass - class DeepKeepGuardrailAPIError(Exception): """Exception raised when there's an error calling the DeepKeep API.""" - pass - class DeepKeepGuardrail(CustomGuardrail): """ @@ -232,7 +228,7 @@ class DeepKeepGuardrail(CustomGuardrail): **({"http_status_code": http_status_code} if http_status_code else {}), ) verbose_proxy_logger.error("DeepKeep guardrail API error: %s", str(error)) - raise DeepKeepGuardrailAPIError(f"DeepKeep guardrail API failed: {str(error)}") + raise DeepKeepGuardrailAPIError(f"DeepKeep guardrail API failed: {error!s}") @staticmethod def _build_return_inputs( diff --git a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py index d684a9c58a9..4cb55df2968 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py @@ -8,7 +8,7 @@ import os from collections.abc import AsyncGenerator from datetime import datetime -from typing import Any, Dict, List, Optional, Type, Union +from typing import Any import httpx @@ -43,10 +43,10 @@ class DynamoAIGuardrails(CustomGuardrail): def __init__( self, guardrail_name: str = "litellm_test", - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, model_id: str = "", - policy_ids: List[str] = [], + policy_ids: list[str] = [], **kwargs, ): self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -86,10 +86,10 @@ class DynamoAIGuardrails(CustomGuardrail): async def _call_dynamoai_guardrails( self, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], event_type: GuardrailEventHooks, text_type: str = "input", - request_data: Optional[dict] = None, + request_data: dict | None = None, ) -> DynamoAIResponse: """ Call DynamoAI Guardrails API to analyze messages for policy violations. @@ -187,8 +187,8 @@ class DynamoAIGuardrails(CustomGuardrail): final_action = response.get("finalAction", "NONE") applied_policies = response.get("appliedPolicies", []) - violations_detected: List[str] = [] - violation_details: Dict[str, Any] = {} + violations_detected: list[str] = [] + violation_details: dict[str, Any] = {} # For now, only handle BLOCK action if final_action == "BLOCK": @@ -291,7 +291,7 @@ class DynamoAIGuardrails(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call Runs on only Input @@ -404,7 +404,7 @@ class DynamoAIGuardrails(CustomGuardrail): # to avoid sending empty content to DynamoAI (e.g., during tool calls) if isinstance(response, litellm.ModelResponse): has_text_content = False - dynamoai_messages: List[Dict[str, Any]] = [] + dynamoai_messages: list[dict[str, Any]] = [] for choice in response.choices: if isinstance(choice, litellm.Choices): @@ -460,7 +460,7 @@ class DynamoAIGuardrails(CustomGuardrail): yield item @staticmethod - def get_config_model() -> Optional[Type[GuardrailConfigModel]]: + def get_config_model() -> type[GuardrailConfigModel] | None: from litellm.types.proxy.guardrails.guardrail_hooks.dynamoai import ( DynamoAIGuardrailConfigModel, ) @@ -468,7 +468,7 @@ class DynamoAIGuardrails(CustomGuardrail): return DynamoAIGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py index f3f19334473..056297b90aa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py @@ -11,11 +11,8 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Union, ) import httpx @@ -54,9 +51,9 @@ class EnkryptAIGuardrails(CustomGuardrail): def __init__( self, guardrail_name: str = "litellm_test", - api_key: Optional[str] = None, - api_base: Optional[str] = None, - policy_name: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + policy_name: str | None = None, **kwargs, ): self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -93,7 +90,7 @@ class EnkryptAIGuardrails(CustomGuardrail): async def _call_enkryptai_guardrails( self, prompt: str, - request_data: Optional[dict] = None, + request_data: dict | None = None, ) -> EnkryptAIResponse: """ Call Enkrypt AI Guardrails API to detect potential issues in the given prompt. @@ -192,8 +189,8 @@ class EnkryptAIGuardrails(CustomGuardrail): summary = response.get("summary", {}) details = response.get("details", {}) - detected_attacks: List[str] = [] - attack_details: Dict[str, Any] = {} + detected_attacks: list[str] = [] + attack_details: dict[str, Any] = {} for key, value in summary.items(): # Check if attack is detected @@ -280,7 +277,7 @@ class EnkryptAIGuardrails(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call Runs on only Input @@ -497,7 +494,7 @@ class EnkryptAIGuardrails(CustomGuardrail): return EnkryptAIGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index a3ac1b94004..fef7ef17dd4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -8,7 +8,7 @@ if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def _get_config_value(litellm_params: Any, optional_params: Any, attribute_name: str) -> Optional[Any]: +def _get_config_value(litellm_params: Any, optional_params: Any, attribute_name: str) -> Any | None: if optional_params is not None: value = ( optional_params.get(attribute_name) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 6b1c9a8b2b3..a694ef897ab 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -7,7 +7,7 @@ import fnmatch import os -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Set +from typing import TYPE_CHECKING, Any, Literal, Optional import httpx @@ -58,7 +58,7 @@ _HEADER_PRESENT_PLACEHOLDER = "[present]" def _header_value_allowed( header_name: str, - extra_allowlist: Set[str] | None = None, + extra_allowlist: set[str] | None = None, ) -> bool: """Return True if this header's value may be forwarded (allowlist, including globs and extra_headers).""" lower = header_name.lower() @@ -74,8 +74,8 @@ def _header_value_allowed( def _sanitize_inbound_headers( headers: Any, - extra_allowlist: Set[str] | None = None, -) -> Dict[str, str] | None: + extra_allowlist: set[str] | None = None, +) -> dict[str, str] | None: """ Sanitize inbound headers before passing them to a 3rd party guardrail service. @@ -86,7 +86,7 @@ def _sanitize_inbound_headers( if not headers or not isinstance(headers, dict): return None - sanitized: Dict[str, str] = {} + sanitized: dict[str, str] = {} for k, v in headers.items(): if k is None: continue @@ -105,8 +105,8 @@ def _sanitize_inbound_headers( def _extract_inbound_headers( request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"], - extra_allowlist: Set[str] | None = None, -) -> Dict[str, str] | None: + extra_allowlist: set[str] | None = None, +) -> dict[str, str] | None: """ Extract inbound headers from available request context. @@ -172,10 +172,10 @@ class GenericGuardrailAPI(CustomGuardrail): def __init__( self, - headers: Dict[str, Any] | None = None, + headers: dict[str, Any] | None = None, api_base: str | None = None, api_key: str | None = None, - additional_provider_specific_params: Dict[str, Any] | None = None, + additional_provider_specific_params: dict[str, Any] | None = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", fail_on_error: bool | None = True, extra_headers: list | None = None, @@ -357,7 +357,7 @@ class GenericGuardrailAPI(CustomGuardrail): **({"http_status_code": http_status_code} if http_status_code else {}), ) verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error)) - raise Exception(f"Generic Guardrail API failed: {str(error)}") + raise Exception(f"Generic Guardrail API failed: {error!s}") @log_guardrail_information async def apply_guardrail( @@ -497,7 +497,7 @@ class GenericGuardrailAPI(CustomGuardrail): return GenericGuardrailAPIConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index 72409a61c30..feb7f9a55ad 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -2,7 +2,7 @@ import os import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Optional from fastapi import HTTPException @@ -34,7 +34,7 @@ class GraySwanGuardrailMissingSecrets(Exception): class GraySwanGuardrailAPIError(Exception): """Raised when the Gray Swan API returns an error.""" - def __init__(self, message: str, status_code: Optional[int] = None) -> None: + def __init__(self, message: str, status_code: int | None = None) -> None: super().__init__(message) self.status_code = status_code @@ -63,18 +63,18 @@ class GraySwanGuardrail(CustomGuardrail): def __init__( self, - guardrail_name: Optional[str] = "grayswan", - api_key: Optional[str] = None, - api_base: Optional[str] = None, - on_flagged_action: Optional[str] = None, - violation_threshold: Optional[float] = None, - reasoning_mode: Optional[str] = None, - categories: Optional[Dict[str, str]] = None, - policy_id: Optional[str] = None, + guardrail_name: str | None = "grayswan", + api_key: str | None = None, + api_base: str | None = None, + on_flagged_action: str | None = None, + violation_threshold: float | None = None, + reasoning_mode: str | None = None, + categories: dict[str, str] | None = None, + policy_id: str | None = None, streaming_end_of_stream_only: bool = False, streaming_sampling_rate: int = 5, - fail_open: Optional[bool] = True, - guardrail_timeout: Optional[float] = 30.0, + fail_open: bool | None = True, + guardrail_timeout: float | None = 30.0, **kwargs: Any, ) -> None: self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -126,7 +126,7 @@ class GraySwanGuardrail(CustomGuardrail): ) @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -266,7 +266,7 @@ class GraySwanGuardrail(CustomGuardrail): # Legacy Test Interface (for backward compatibility) # ------------------------------------------------------------------ - async def run_grayswan_guardrail(self, payload: dict) -> Dict[str, Any]: + async def run_grayswan_guardrail(self, payload: dict) -> dict[str, Any]: """ Run the GraySwan guardrail on a payload. @@ -286,8 +286,8 @@ class GraySwanGuardrail(CustomGuardrail): def _process_grayswan_response( self, response_json: dict, - data: Optional[dict] = None, - hook_type: Optional[GuardrailEventHooks] = None, + data: dict | None = None, + hook_type: GuardrailEventHooks | None = None, ) -> None: """ Legacy method for processing GraySwan API responses. @@ -385,7 +385,7 @@ class GraySwanGuardrail(CustomGuardrail): # Core GraySwan API interaction # ------------------------------------------------------------------ - async def _call_grayswan_api(self, payload: dict) -> Dict[str, Any]: + async def _call_grayswan_api(self, payload: dict) -> dict[str, Any]: """Call the GraySwan monitoring API.""" headers = self._prepare_headers() @@ -406,7 +406,7 @@ class GraySwanGuardrail(CustomGuardrail): def _process_response_internal( self, - response_json: Dict[str, Any], + response_json: dict[str, Any], request_data: dict, inputs: GenericGuardrailAPIInputs, is_output: bool, @@ -497,7 +497,7 @@ class GraySwanGuardrail(CustomGuardrail): # Helpers # ------------------------------------------------------------------ - def _prepare_headers(self) -> Dict[str, str]: + def _prepare_headers(self) -> dict[str, str]: return { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", @@ -508,7 +508,7 @@ class GraySwanGuardrail(CustomGuardrail): self, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"] = None, - ) -> Optional[dict[str, str]]: + ) -> dict[str, str] | None: headers = (request_data.get("proxy_server_request") or {}).get("headers") if not headers: headers = request_data.get("headers") @@ -534,7 +534,7 @@ class GraySwanGuardrail(CustomGuardrail): dynamic_body: dict, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"] = None, - ) -> Optional[dict[str, Any]]: + ) -> dict[str, Any] | None: payload: dict[str, Any] = {"messages": messages} categories = dynamic_body.get("categories") or self.categories @@ -613,9 +613,9 @@ class GraySwanGuardrail(CustomGuardrail): return "\n".join(message_parts) - def _format_violated_rules(self, violated_rules: List) -> str: + def _format_violated_rules(self, violated_rules: list) -> str: """Format violated rules list into a readable string.""" - formatted: List[str] = [] + formatted: list[str] = [] for rule in violated_rules: if isinstance(rule, dict): # New format: {'rule': 6, 'name': 'Illegal Activities...', 'description': '...'} @@ -637,7 +637,7 @@ class GraySwanGuardrail(CustomGuardrail): return ", ".join(formatted) - def _resolve_threshold(self, value: Optional[float]) -> float: + def _resolve_threshold(self, value: float | None) -> float: if value is not None: return float(value) env_val = os.getenv("GRAYSWAN_VIOLATION_THRESHOLD") @@ -648,7 +648,7 @@ class GraySwanGuardrail(CustomGuardrail): pass return 0.5 - def _resolve_reasoning_mode(self, value: Optional[str]) -> Optional[str]: + def _resolve_reasoning_mode(self, value: str | None) -> str | None: if value and value.lower() in self.SUPPORTED_REASONING_MODES: return value.lower() env_val = os.getenv("GRAYSWAN_REASONING_MODE") @@ -662,7 +662,7 @@ class GraySwanGuardrail(CustomGuardrail): request_data: dict, start_time: float, end_time: float, - status_code: Optional[int] = None, + status_code: int | None = None, ) -> None: """Log guardrail failure and attach standard logging metadata.""" try: diff --git a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py index a7c94a4742d..49d09ec3bab 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py @@ -10,13 +10,8 @@ import os from typing import ( TYPE_CHECKING, Any, - List, Literal, - Optional, - Tuple, - Type, TypedDict, - Union, ) from fastapi import HTTPException @@ -46,22 +41,22 @@ class GuardrailsAIResponse(TypedDict): class InferenceData(TypedDict): name: str - shape: List[int] - data: List + shape: list[int] + data: list datatype: str class GuardrailsAIResponsePreCall(TypedDict): modelname: str modelversion: str - outputs: List[InferenceData] + outputs: list[InferenceData] class GuardrailsAI(CustomGuardrail): def __init__( self, guard_name: str, - api_base: Optional[str] = None, + api_base: str | None = None, guardrails_ai_api_input_format: Literal["inputs", "llmOutput"] = "llmOutput", **kwargs, ): @@ -186,12 +181,12 @@ class GuardrailsAI(CustomGuardrail): "rerank", "mcp_call", ], - ) -> Optional[ - Union[Exception, str, dict] - ]: # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm + ) -> ( + Exception | str | dict | None + ): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm return await self.process_input(data=data, call_type=call_type) - async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: if call_type == "acompletion" or call_type == "completion": kwargs = await self.process_input(data=kwargs, call_type=call_type) @@ -229,7 +224,7 @@ class GuardrailsAI(CustomGuardrail): return @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.guardrails_ai import ( GuardrailsAIGuardrailConfigModel, ) @@ -237,7 +232,7 @@ class GuardrailsAI(CustomGuardrail): return GuardrailsAIGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.post_call, GuardrailEventHooks.pre_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 99b83e10a4f..04ed5868e58 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -5,7 +5,7 @@ import re import time import uuid from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, ClassVar, List, Literal, Optional, TypeGuard +from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypeGuard import httpx from fastapi import HTTPException @@ -190,7 +190,7 @@ def _build_headroom_retrieve_tool() -> dict[str, object]: } -def _resolve_call_id(logging_obj: object, request_state: dict[str, object]) -> Optional[str]: +def _resolve_call_id(logging_obj: object, request_state: dict[str, object]) -> str | None: """Resolve the litellm_call_id shared by a request's pre-call hook and its agentic-loop hooks, so CCR hash validation can be scoped per call instead of trusting any hash-shaped string that shows up in message text.""" @@ -301,7 +301,7 @@ def _build_responses_followup_items( model wrote alongside the tool call is preserved. """ text = assistant_text_from_response(response) - items: List[dict[str, object]] = [{"role": "assistant", "content": text}] if text else [] + items: list[dict[str, object]] = [{"role": "assistant", "content": text}] if text else [] for tool_call, content in retrieved: call_id = tool_call.get("id") items.append( @@ -320,7 +320,7 @@ class HeadroomGuardrail(CustomGuardrail): records_own_guardrail_information: ClassVar[bool] = True @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -696,7 +696,7 @@ class HeadroomGuardrail(CustomGuardrail): response: Any, model: str, messages: list[dict], - tools: Optional[list[dict]], + tools: list[dict] | None, stream: bool, custom_llm_provider: str, kwargs: dict, @@ -763,7 +763,7 @@ class HeadroomGuardrail(CustomGuardrail): ] follow_up_messages = list(messages) + [assistant_message] + tool_results - max_tokens: Optional[int] = anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get( + max_tokens: int | None = anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get( "max_tokens" ) optional_params_without_max_tokens = { diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py index 1566c90ac0c..73306ce1154 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py @@ -1,11 +1,11 @@ from __future__ import annotations -from uuid import uuid4 -import httpx import os -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal from urllib.parse import urlparse +from uuid import uuid4 +import httpx import requests from fastapi import HTTPException from httpx import HTTPStatusError @@ -65,7 +65,7 @@ class HiddenlayerGuardrail(CustomGuardrail): """Custom guardrail wrapper for HiddenLayer's safety checks.""" @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -73,10 +73,10 @@ class HiddenlayerGuardrail(CustomGuardrail): def __init__( self, - api_id: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - auth_url: Optional[str] = None, + api_id: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + auth_url: str | None = None, **kwargs: Any, ) -> None: kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -114,7 +114,7 @@ class HiddenlayerGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"], - logging_obj: Optional["LiteLLMLoggingObj"] = None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: """Validate (and optionally redact) text via HiddenLayer before/after LLM calls.""" @@ -265,7 +265,7 @@ class HiddenlayerGuardrail(CustomGuardrail): return result @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type[GuardrailConfigModel] | None: from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( HiddenlayerGuardrailConfigModel, ) @@ -278,10 +278,10 @@ class HiddenlayerGuardrailV2(CustomGuardrail): def __init__( self, - api_id: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - auth_url: Optional[str] = None, + api_id: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + auth_url: str | None = None, **kwargs: Any, ) -> None: self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID") @@ -318,7 +318,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"], - logging_obj: Optional["LiteLLMLoggingObj"] = None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: """Validate (and optionally redact) text via HiddenLayer before/after LLM calls.""" @@ -461,7 +461,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail): return response @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type[GuardrailConfigModel] | None: from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( HiddenlayerGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py index 622f21e531e..562ae002634 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py @@ -8,7 +8,7 @@ import os from collections.abc import AsyncGenerator from datetime import datetime -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -36,13 +36,13 @@ class IBMGuardrailDetector(CustomGuardrail): def __init__( self, guardrail_name: str = "ibm_detector", - auth_token: Optional[str] = None, - base_url: Optional[str] = None, - detector_id: Optional[str] = None, + auth_token: str | None = None, + base_url: str | None = None, + detector_id: str | None = None, is_detector_server: bool = True, - detector_params: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, str]] = None, - score_threshold: Optional[float] = None, + detector_params: dict[str, Any] | None = None, + extra_headers: dict[str, str] | None = None, + score_threshold: float | None = None, block_on_detection: bool = True, verify_ssl: bool = True, **kwargs, @@ -100,10 +100,10 @@ class IBMGuardrailDetector(CustomGuardrail): async def _call_detector_server( self, - contents: List[str], + contents: list[str], event_type: GuardrailEventHooks, - request_data: Optional[dict] = None, - ) -> List[List[IBMDetectorDetection]]: + request_data: dict | None = None, + ) -> list[list[IBMDetectorDetection]]: """ Call IBM Detector Server directly. @@ -142,7 +142,7 @@ class IBMGuardrailDetector(CustomGuardrail): headers=headers, ) response.raise_for_status() - response_json: List[List[IBMDetectorDetection]] = response.json() + response_json: list[list[IBMDetectorDetection]] = response.json() end_time = datetime.now() duration = (end_time - start_time).total_seconds() @@ -192,8 +192,8 @@ class IBMGuardrailDetector(CustomGuardrail): self, content: str, event_type: GuardrailEventHooks, - request_data: Optional[dict] = None, - ) -> List[IBMDetectorDetection]: + request_data: dict | None = None, + ) -> list[IBMDetectorDetection]: """ Call IBM FMS Guardrails Orchestrator. @@ -275,7 +275,7 @@ class IBMGuardrailDetector(CustomGuardrail): raise - def _filter_detections_by_threshold(self, detections: List[IBMDetectorDetection]) -> List[IBMDetectorDetection]: + def _filter_detections_by_threshold(self, detections: list[IBMDetectorDetection]) -> list[IBMDetectorDetection]: """ Filter detections based on score threshold. @@ -291,7 +291,7 @@ class IBMGuardrailDetector(CustomGuardrail): return [detection for detection in detections if detection.get("score", 0.0) >= self.score_threshold] def _determine_guardrail_status_detector_server( - self, response_json: List[List[IBMDetectorDetection]] + self, response_json: list[list[IBMDetectorDetection]] ) -> GuardrailStatus: """ Determine the guardrail status based on IBM Detector Server response. @@ -352,7 +352,7 @@ class IBMGuardrailDetector(CustomGuardrail): verbose_proxy_logger.error("Error determining IBM Orchestrator guardrail status: %s", str(e)) return "guardrail_failed_to_respond" - def _create_error_message_detector_server(self, detections_list: List[List[IBMDetectorDetection]]) -> str: + def _create_error_message_detector_server(self, detections_list: list[list[IBMDetectorDetection]]) -> str: """ Create a detailed error message from detector server response. @@ -383,7 +383,7 @@ class IBMGuardrailDetector(CustomGuardrail): error_message = f"IBM Guardrail Detector failed: {total_detections} violation(s) detected\n\n" + error_message return error_message.strip() - def _create_error_message_orchestrator(self, detections: List[IBMDetectorDetection]) -> str: + def _create_error_message_orchestrator(self, detections: list[IBMDetectorDetection]) -> str: """ Create a detailed error message from orchestrator response. @@ -414,7 +414,7 @@ class IBMGuardrailDetector(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call Runs on only Input @@ -431,7 +431,7 @@ class IBMGuardrailDetector(CustomGuardrail): return data # Covers multimodal list content + Responses-API input. - contents_to_check: List[str] = list(iter_message_text(data)) + contents_to_check: list[str] = list(iter_message_text(data)) if contents_to_check: if self.is_detector_server: # Call detector server with all contents at once @@ -500,7 +500,7 @@ class IBMGuardrailDetector(CustomGuardrail): return # Covers multimodal list content + Responses-API input. - contents_to_check: List[str] = list(iter_message_text(data)) + contents_to_check: list[str] = list(iter_message_text(data)) if contents_to_check: if self.is_detector_server: # Call detector server with all contents at once @@ -587,7 +587,7 @@ class IBMGuardrailDetector(CustomGuardrail): ) return - contents_to_check: List[str] = [] + contents_to_check: list[str] = [] for choice in response.choices: if isinstance(choice, litellm.Choices): verbose_proxy_logger.debug("async_post_call_success_hook choice: %s", choice) @@ -667,7 +667,7 @@ class IBMGuardrailDetector(CustomGuardrail): return IBMDetectorGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py index 4da3e75d03e..d645aeb7576 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py +++ b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py @@ -1,5 +1,5 @@ from datetime import datetime -from typing import TYPE_CHECKING, Dict, List, Optional, Type, Union +from typing import TYPE_CHECKING from fastapi import HTTPException @@ -26,22 +26,22 @@ if TYPE_CHECKING: class JavelinGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, ] def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, default_on: bool = True, guardrail_name: str = "trustsafety", - javelin_guard_name: Optional[str] = None, + javelin_guard_name: str | None = None, api_version: str = "v1", - metadata: Optional[Dict] = None, - config: Optional[Dict] = None, - application: Optional[str] = None, + metadata: dict | None = None, + config: dict | None = None, + application: str | None = None, **kwargs, ): f""" @@ -100,7 +100,7 @@ class JavelinGuardrail(CustomGuardrail): headers["x-javelin-application"] = self.application status: GuardrailStatus = "guardrail_failed_to_respond" - javelin_response: Optional[JavelinGuardResponse] = None + javelin_response: JavelinGuardResponse | None = None exception_str = "" try: @@ -129,7 +129,7 @@ class JavelinGuardrail(CustomGuardrail): #################################################### # Create Guardrail Trace for logging on Langfuse, Datadog, etc. #################################################### - guardrail_json_response: Union[Exception, str, dict, List[dict]] = {} + guardrail_json_response: Exception | str | dict | list[dict] = {} if status == "success" and javelin_response is not None: guardrail_json_response = dict(javelin_response) else: @@ -163,9 +163,9 @@ class JavelinGuardrail(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: litellm.DualCache, - data: Dict, + data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, Dict]]: + ) -> Exception | str | dict | None: """ Pre-call hook for the Javelin guardrail. """ @@ -270,7 +270,7 @@ class JavelinGuardrail(CustomGuardrail): return data @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Get the config model for the Javelin guardrail. """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index d3360bbe641..4eb68327f4b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -11,7 +11,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import json import sys -from typing import Dict, List, Literal, Optional, Union +from typing import Literal import httpx from fastapi import HTTPException @@ -48,7 +48,7 @@ INPUT_POSITIONING_MAP = { class lakeraAI_Moderation(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -57,9 +57,9 @@ class lakeraAI_Moderation(CustomGuardrail): def __init__( self, moderation_check: Literal["pre_call", "in_parallel"] = "in_parallel", - category_thresholds: Optional[LakeraCategoryThresholds] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + category_thresholds: LakeraCategoryThresholds | None = None, + api_base: str | None = None, + api_key: str | None = None, **kwargs, ): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -77,7 +77,7 @@ class lakeraAI_Moderation(CustomGuardrail): return flagged = _results[0].get("flagged", False) - category_scores: Optional[dict] = _results[0].get("category_scores", None) + category_scores: dict | None = _results[0].get("category_scores", None) if self.category_thresholds is not None: if category_scores is not None: @@ -110,7 +110,7 @@ class lakeraAI_Moderation(CustomGuardrail): }, ) - return None + return async def _check( self, @@ -141,7 +141,7 @@ class lakeraAI_Moderation(CustomGuardrail): text = "" _json_data: str = "" if "messages" in data and isinstance(data["messages"], list): - prompt_injection_obj: Optional[GuardrailItem] = litellm.guardrail_name_config_map.get("prompt_injection") + prompt_injection_obj: GuardrailItem | None = litellm.guardrail_name_config_map.get("prompt_injection") if prompt_injection_obj is not None: enabled_roles = prompt_injection_obj.enabled_roles else: @@ -150,16 +150,16 @@ class lakeraAI_Moderation(CustomGuardrail): if enabled_roles is None: enabled_roles = default_roles - stringified_roles: List[str] = [] + stringified_roles: list[str] = [] if enabled_roles is not None: # convert to list of str for role in enabled_roles: if isinstance(role, Role): stringified_roles.append(role.value) elif isinstance(role, str): stringified_roles.append(role) - lakera_input_dict: Dict = {role: None for role in INPUT_POSITIONING_MAP.keys()} + lakera_input_dict: dict = {role: None for role in INPUT_POSITIONING_MAP} system_message = None - tool_call_messages: List = [] + tool_call_messages: list = [] for message in data["messages"]: role = message.get("role") if role in stringified_roles: @@ -288,7 +288,7 @@ class lakeraAI_Moderation(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: litellm.DualCache, - data: Dict, + data: dict, call_type: Literal[ "completion", "text_completion", @@ -301,7 +301,7 @@ class lakeraAI_Moderation(CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, Dict]]: + ) -> Exception | str | dict | None: from litellm.types.guardrails import GuardrailEventHooks if self.event_hook is None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 76603579d6c..f778e113691 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -1,7 +1,6 @@ import copy import os from datetime import datetime -from typing import Dict, List, Optional, Tuple, Union from fastapi import HTTPException @@ -30,7 +29,7 @@ from litellm.types.utils import CallTypesLiteral, GuardrailStatus, ModelResponse class LakeraAIGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -39,14 +38,14 @@ class LakeraAIGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - project_id: Optional[str] = None, - payload: Optional[bool] = True, - breakdown: Optional[bool] = True, - metadata: Optional[Dict] = None, - dev_info: Optional[bool] = True, - on_flagged: Optional[str] = "block", + api_key: str | None = None, + api_base: str | None = None, + project_id: str | None = None, + payload: bool | None = True, + breakdown: bool | None = True, + metadata: dict | None = None, + dev_info: bool | None = True, + on_flagged: str | None = "block", **kwargs, ): """ @@ -71,29 +70,29 @@ class LakeraAIGuardrail(CustomGuardrail): self.lakera_api_key = api_key or os.environ.get("LAKERA_API_KEY") or "" self.project_id = project_id self.api_base = api_base or get_secret_str("LAKERA_API_BASE") or "https://api.lakera.ai" - self.payload: Optional[bool] = payload - self.breakdown: Optional[bool] = breakdown - self.metadata: Optional[Dict] = metadata - self.dev_info: Optional[bool] = dev_info + self.payload: bool | None = payload + self.breakdown: bool | None = breakdown + self.metadata: dict | None = metadata + self.dev_info: bool | None = dev_info self.on_flagged = on_flagged or "block" kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) super().__init__(**kwargs) async def call_v2_guard( self, - messages: List[AllMessageValues], - request_data: Dict, + messages: list[AllMessageValues], + request_data: dict, event_type: GuardrailEventHooks, - ) -> Tuple[LakeraAIResponse, Dict]: + ) -> tuple[LakeraAIResponse, dict]: """ Call the Lakera AI v2 guard API. """ status: GuardrailStatus = "success" exception_str: str = "" start_time: datetime = datetime.now() - lakera_response: Optional[LakeraAIResponse] = None - request: Dict = {} - masked_entity_count: Dict = {} + lakera_response: LakeraAIResponse | None = None + request: dict = {} + masked_entity_count: dict = {} try: request = dict( LakeraAIRequest( @@ -122,7 +121,7 @@ class LakeraAIGuardrail(CustomGuardrail): #################################################### # Create Guardrail Trace for logging on Langfuse, Datadog, etc. #################################################### - guardrail_json_response: Union[Exception, str, dict, List[dict]] = {} + guardrail_json_response: Exception | str | dict | list[dict] = {} if status == "success": copy_lakera_response_dict = dict(copy.deepcopy(lakera_response)) if lakera_response else {} # payload contains PII, we don't want to log it @@ -143,10 +142,10 @@ class LakeraAIGuardrail(CustomGuardrail): def _mask_pii_in_messages( self, - messages: List[AllMessageValues], - lakera_response: Optional[LakeraAIResponse], - masked_entity_count: Dict, - ) -> List[AllMessageValues]: + messages: list[AllMessageValues], + lakera_response: LakeraAIResponse | None, + masked_entity_count: dict, + ) -> list[AllMessageValues]: """ Return a copy of messages with any detected PII replaced by “[MASKED ]” tokens. @@ -204,9 +203,9 @@ class LakeraAIGuardrail(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: litellm.DualCache, - data: Dict, + data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, Dict]]: + ) -> Exception | str | dict | None: from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, ) @@ -355,15 +354,15 @@ class LakeraAIGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return response - original_messages: Optional[List[AllMessageValues]] = data.get("messages", []) + original_messages: list[AllMessageValues] | None = data.get("messages", []) if original_messages is None: original_messages = [] # Extract assistant messages from the response, keeping only role/content. # Track choice indices so we write masked content back to the correct choice # when some choices have null content (e.g. tool-call-only). - response_messages: List[AllMessageValues] = [] - choice_indices: List[int] = [] + response_messages: list[AllMessageValues] = [] + choice_indices: list[int] = [] response_dict = response.model_dump() if hasattr(response, "model_dump") else {} for i, choice in enumerate(response_dict.get("choices", [])): msg = choice.get("message") @@ -389,7 +388,7 @@ class LakeraAIGuardrail(CustomGuardrail): if lakera_guardrail_response.get("flagged") is True: # If only PII violations exist, mask the PII in the response and allow if self._is_only_pii_violation(lakera_guardrail_response): - masked_entity_count: Dict[str, int] = {} + masked_entity_count: dict[str, int] = {} masked_messages = self._mask_pii_in_messages( messages=post_call_messages, lakera_response=lakera_guardrail_response, @@ -414,7 +413,7 @@ class LakeraAIGuardrail(CustomGuardrail): return response - def _is_only_pii_violation(self, lakera_response: Optional[LakeraAIResponse]) -> bool: + def _is_only_pii_violation(self, lakera_response: LakeraAIResponse | None) -> bool: """ Returns True if there are only PII violations in the response. """ @@ -437,7 +436,7 @@ class LakeraAIGuardrail(CustomGuardrail): # Return True only if there are violations and they are all PII return has_violations - def _get_http_exception_for_blocked_guardrail(self, lakera_response: Optional[LakeraAIResponse]) -> HTTPException: + def _get_http_exception_for_blocked_guardrail(self, lakera_response: LakeraAIResponse | None) -> HTTPException: """ Get the HTTP exception for a blocked guardrail, similar to Bedrock's implementation. """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 31ce0cc2214..dba82eeb32c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -11,13 +11,7 @@ import uuid from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, - Optional, - Tuple, - Type, - Union, TypedDict, ) @@ -39,6 +33,7 @@ except ImportError: from fastapi import HTTPException +import litellm from litellm import DualCache from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -46,7 +41,6 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.integrations.custom_guardrail import dc as global_cache - from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -57,16 +51,15 @@ from litellm.proxy.guardrails._content_utils import ( has_non_string_content, ) from litellm.types.guardrails import GuardrailEventHooks -import litellm class LassoResponse(TypedDict): """Type definition for Lasso API response.""" violations_detected: bool - deputies: Dict[str, bool] - findings: Dict[str, List[Dict[str, Any]]] - messages: Optional[List[Dict[str, str]]] + deputies: dict[str, bool] + findings: dict[str, list[dict[str, Any]]] + messages: list[dict[str, str]] | None if TYPE_CHECKING: @@ -76,14 +69,10 @@ if TYPE_CHECKING: class LassoGuardrailMissingSecrets(Exception): """Exception raised when Lasso API key is missing.""" - pass - class LassoGuardrailAPIError(Exception): """Exception raised when there's an error calling the Lasso API.""" - pass - class LassoGuardrail(CustomGuardrail): """ @@ -94,7 +83,7 @@ class LassoGuardrail(CustomGuardrail): """ @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -103,12 +92,12 @@ class LassoGuardrail(CustomGuardrail): def __init__( self, - lasso_api_key: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - user_id: Optional[str] = None, - conversation_id: Optional[str] = None, - mask: Optional[bool] = False, + lasso_api_key: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + user_id: str | None = None, + conversation_id: str | None = None, + mask: bool | None = False, **kwargs, ): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -143,7 +132,7 @@ class LassoGuardrail(CustomGuardrail): @staticmethod def _extract_tool_call_fields( call: Any, - ) -> Tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]: + ) -> tuple[str | None, str | None, dict[str, Any] | None]: """Extract (call_id, name, parsed_input) from a tool call. Handles both dict-style and Pydantic object-style tool_calls. @@ -156,7 +145,7 @@ class LassoGuardrail(CustomGuardrail): return call_id, None, None name = get(func, "name") args_str = get(func, "arguments") - input_data: Optional[Dict[str, Any]] = None + input_data: dict[str, Any] | None = None if args_str: try: parsed = json.loads(args_str) @@ -200,7 +189,7 @@ class LassoGuardrail(CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call to validate and potentially modify input. Uses 'PROMPT' messageType as this is input to the model. @@ -262,7 +251,7 @@ class LassoGuardrail(CustomGuardrail): # Extract messages from the response for validation if isinstance(response, litellm.ModelResponse): - response_messages: List[Dict[str, Any]] = [] + response_messages: list[dict[str, Any]] = [] for choice in response.choices: if not hasattr(choice, "message"): continue @@ -310,8 +299,8 @@ class LassoGuardrail(CustomGuardrail): except Exception as e: if isinstance(e, HTTPException): raise e - verbose_proxy_logger.error(f"Error in post-call Lasso masking: {str(e)}") - raise LassoGuardrailAPIError(f"Failed to apply post-call masking: {str(e)}") + verbose_proxy_logger.error(f"Error in post-call Lasso masking: {e!s}") + raise LassoGuardrailAPIError(f"Failed to apply post-call masking: {e!s}") else: # Use the same data for conversation_id consistency (no cache access needed) await self._run_lasso_guardrail(response_data, cache=global_cache, message_type="COMPLETION") @@ -406,8 +395,8 @@ class LassoGuardrail(CustomGuardrail): LassoGuardrailAPIError: If the Lasso API call fails HTTPException: If blocking violations are detected """ - raw_messages: List[Dict[str, Any]] = data.get("messages") or [] - messages: List[Dict[str, Any]] = self._expand_messages_for_classification(raw_messages) if raw_messages else [] + raw_messages: list[dict[str, Any]] = data.get("messages") or [] + messages: list[dict[str, Any]] = self._expand_messages_for_classification(raw_messages) if raw_messages else [] messages_count = len(messages) if data.get("input") is not None: # Responses-API payloads carry text in data["input"]. Inspect it @@ -431,7 +420,7 @@ class LassoGuardrail(CustomGuardrail): data: dict, cache: DualCache, message_type: Literal["PROMPT", "COMPLETION"], - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], ) -> dict: """Handle classification without masking.""" try: @@ -449,7 +438,7 @@ class LassoGuardrail(CustomGuardrail): data: dict, cache: DualCache, message_type: Literal["PROMPT", "COMPLETION"], - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], messages_count: int, ) -> dict: """Handle masking with classifix endpoint. @@ -489,9 +478,9 @@ class LassoGuardrail(CustomGuardrail): def _map_masked_messages_back( self, - original_messages: List[Dict[str, Any]], - masked_messages: List[Dict[str, Any]], - ) -> List[Dict[str, Any]]: + original_messages: list[dict[str, Any]], + masked_messages: list[dict[str, Any]], + ) -> list[dict[str, Any]]: """Map Lasso-format masked messages back onto the original OpenAI-format messages. Lasso receives expanded messages (tool_use / tool_result blocks) and returns them @@ -501,9 +490,9 @@ class LassoGuardrail(CustomGuardrail): while preserving the original structure. """ # Index masked content by type so we can look up by id without caring about order. - masked_tool_use: Dict[str, Dict[str, Any]] = {} - masked_tool_result: Dict[str, str] = {} - masked_text: List[str] = [] + masked_tool_use: dict[str, dict[str, Any]] = {} + masked_tool_result: dict[str, str] = {} + masked_text: list[str] = [] for msg in masked_messages: content = msg.get("content") @@ -538,7 +527,7 @@ class LassoGuardrail(CustomGuardrail): }, ) - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] text_cursor = 0 for orig_msg in original_messages: @@ -577,9 +566,9 @@ class LassoGuardrail(CustomGuardrail): def _update_tool_calls_from_masked( self, - tool_calls: List[Any], - masked_tool_use: Dict[str, Dict[str, Any]], - ) -> List[Any]: + tool_calls: list[Any], + masked_tool_use: dict[str, dict[str, Any]], + ) -> list[Any]: """Replace tool_call arguments with masked values returned by Lasso.""" updated = [] for call in tool_calls: @@ -610,7 +599,7 @@ class LassoGuardrail(CustomGuardrail): # Log error with context verbose_proxy_logger.error( - f"Error calling Lasso API: {str(error)}", + f"Error calling Lasso API: {error!s}", extra={ "guardrail_name": getattr(self, "guardrail_name", "unknown"), "message_type": message_type, @@ -631,12 +620,12 @@ class LassoGuardrail(CustomGuardrail): raise LassoGuardrailAPIError(f"API error: {error.response.status_code}") # Generic error handling - raise LassoGuardrailAPIError(f"Failed to verify request safety with Lasso API: {str(error)}") + raise LassoGuardrailAPIError(f"Failed to verify request safety with Lasso API: {error!s}") def _log_masking_applied( self, message_type: Literal["PROMPT", "COMPLETION"], - response: Dict[str, Any], + response: dict[str, Any], ) -> None: """Log masking application with structured context.""" conversation_id = getattr(self, "conversation_id", "unknown") @@ -651,7 +640,7 @@ class LassoGuardrail(CustomGuardrail): }, ) - def _expand_messages_for_classification(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + def _expand_messages_for_classification(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]: """ Convert raw OpenAI-format messages to Lasso API format with content blocks. @@ -659,7 +648,7 @@ class LassoGuardrail(CustomGuardrail): - role=tool messages → developer role + tool_result block - plain text messages pass through unchanged """ - expanded: List[Dict[str, Any]] = [] + expanded: list[dict[str, Any]] = [] for msg in messages: role = msg.get("role", "") content = msg.get("content") @@ -732,7 +721,7 @@ class LassoGuardrail(CustomGuardrail): return expanded - def _prepare_headers(self, data: dict, cache: DualCache) -> Dict[str, str]: + def _prepare_headers(self, data: dict, cache: DualCache) -> dict[str, str]: """Prepare headers for the Lasso API request.""" if not self.lasso_api_key: raise LassoGuardrailMissingSecrets( @@ -740,7 +729,7 @@ class LassoGuardrail(CustomGuardrail): "pass it as a parameter to the guardrail in the config file" ) - headers: Dict[str, str] = { + headers: dict[str, str] = { "lasso-api-key": self.lasso_api_key, "Content-Type": "application/json", } @@ -758,11 +747,11 @@ class LassoGuardrail(CustomGuardrail): def _prepare_payload( self, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], data: dict, cache: DualCache, message_type: Literal["PROMPT", "COMPLETION"] = "PROMPT", - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Prepare the payload for the Lasso API request. @@ -772,7 +761,7 @@ class LassoGuardrail(CustomGuardrail): data: Request data (used for conversation_id generation and tools extraction) cache: Cache instance for storing conversation_id (optional for post-call) """ - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "messages": messages, "messageType": message_type, # Drives the "Used By" badge on Lasso Application API Keys: every call from this @@ -789,7 +778,7 @@ class LassoGuardrail(CustomGuardrail): payload["sessionId"] = conversation_id # Map OpenAI ChatCompletionToolParam array → ToolDefinition array - tools_data: List[Dict[str, Any]] = data.get("tools") or [] + tools_data: list[dict[str, Any]] = data.get("tools") or [] if tools_data: get = self._get_field tool_definitions = [] @@ -800,7 +789,7 @@ class LassoGuardrail(CustomGuardrail): name = get(func, "name") if not name: continue - td: Dict[str, Any] = {"name": name} + td: dict[str, Any] = {"name": name} description = get(func, "description") if description: td["description"] = description @@ -815,9 +804,9 @@ class LassoGuardrail(CustomGuardrail): async def _call_lasso_api( self, - headers: Dict[str, str], - payload: Dict[str, Any], - api_url: Optional[str] = None, + headers: dict[str, str], + payload: dict[str, Any], + api_url: str | None = None, ) -> LassoResponse: """Call the Lasso API and return the response.""" url = api_url or f"{self.api_base}/classify" @@ -880,7 +869,7 @@ class LassoGuardrail(CustomGuardrail): f"Non-blocking Lasso violations detected, continuing with warning: {violated_deputies}" ) - def _check_for_blocking_actions(self, response: LassoResponse) -> List[str]: + def _check_for_blocking_actions(self, response: LassoResponse) -> list[str]: """ Check findings for actions that should block the request/response. @@ -918,7 +907,7 @@ class LassoGuardrail(CustomGuardrail): return blocking_violations - def _parse_violated_deputies(self, response: LassoResponse) -> List[str]: + def _parse_violated_deputies(self, response: LassoResponse) -> list[str]: """Parse the response to extract violated deputies.""" violated_deputies = [] if "deputies" in response: @@ -930,12 +919,12 @@ class LassoGuardrail(CustomGuardrail): def _apply_masking_to_model_response( self, model_response: litellm.ModelResponse, - masked_messages: List[Dict[str, Any]], + masked_messages: list[dict[str, Any]], ) -> None: """Apply masking to the actual model response when mask=True and masked content is available.""" # Index masked tool_use blocks by id for O(1) lookup. - masked_tool_use: Dict[str, Dict[str, Any]] = {} - masked_text: List[str] = [] + masked_tool_use: dict[str, dict[str, Any]] = {} + masked_text: list[str] = [] for masked_msg in masked_messages: content = masked_msg.get("content") if isinstance(content, dict) and content.get("type") == "tool_use": @@ -984,7 +973,7 @@ class LassoGuardrail(CustomGuardrail): verbose_proxy_logger.debug(f"Applied masked tool_call arguments for call_id={call_id}") @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.lasso import ( LassoGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/__init__.py index d373fc24816..077a9e2bb00 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/__init__.py @@ -14,8 +14,8 @@ from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_ ) __all__ = [ - "BaseCompetitorIntentChecker", "AirlineCompetitorIntentChecker", + "BaseCompetitorIntentChecker", "normalize", "text_for_entity_matching", ] diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py index 04d8459b821..b4fc682f288 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py @@ -11,7 +11,7 @@ brand_self so all other major airlines are treated as competitors. import json from pathlib import Path -from typing import Any, Dict, List, Set, Tuple +from typing import Any from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( BaseCompetitorIntentChecker, @@ -113,7 +113,7 @@ AIRLINE_EXPLICIT_OTHER_MEANING_MARKER = ( _MAJOR_AIRLINES_PATH = Path(__file__).resolve().parent / "major_airlines.json" -def _load_competitors_excluding_brand(brand_self: List[str]) -> List[str]: +def _load_competitors_excluding_brand(brand_self: list[str]) -> list[str]: """ Load competitor tokens from major_airlines.json (harm_toxic_abuse-style format). Exclude any airline whose id or match variants overlap with brand_self. @@ -127,13 +127,13 @@ def _load_competitors_excluding_brand(brand_self: List[str]) -> List[str]: airlines = json.load(f) except (json.JSONDecodeError, OSError): return [] - result: List[str] = [] + result: list[str] = [] for entry in airlines: if not isinstance(entry, dict): continue match_str = entry.get("match") or "" variants = [v.strip().lower() for v in match_str.split("|") if v.strip()] - words_in_match: Set[str] = set() + words_in_match: set[str] = set() for v in variants: words_in_match.update(v.split()) if brand_set & words_in_match or any(v in brand_set for v in variants): @@ -149,8 +149,8 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): with other_meaning/competitor signals and explicit markers. """ - def __init__(self, config: Dict[str, Any]) -> None: - merged: Dict[str, Any] = dict(config) + def __init__(self, config: dict[str, Any]) -> None: + merged: dict[str, Any] = dict(config) if not merged.get("other_meaning_signals"): merged["other_meaning_signals"] = AIRLINE_OTHER_MEANING_SIGNALS if not merged.get("competitor_signals"): @@ -173,7 +173,7 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): self._explicit_competitor_marker = _compile_marker(merged.get("explicit_competitor_marker")) self._explicit_other_meaning_marker = _compile_marker(merged.get("explicit_other_meaning_marker")) - def _classify_ambiguous(self, text: str, token: str) -> Tuple[str, float]: + def _classify_ambiguous(self, text: str, token: str) -> tuple[str, float]: """Other meaning vs competitor using airline signals and explicit markers.""" text_lower = text.lower() if ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py index 3e91afd2d6f..e1e700d0f65 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py @@ -5,7 +5,7 @@ Generic competitor intent checker: two entity sets and overridable disambiguatio import re import unicodedata from re import Pattern -from typing import Any, Dict, List, Optional, Set, Tuple, cast +from typing import Any, cast from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( CompetitorActionHint, @@ -36,12 +36,12 @@ def _word_boundary_match(text: str, token: str) -> bool: return bool(re.search(r"\b" + re.escape(token) + r"\b", text)) -def _count_signals(text: str, patterns: List[str]) -> int: +def _count_signals(text: str, patterns: list[str]) -> int: """Count how many of the patterns appear in text.""" return sum(1 for p in patterns if re.search(p, text, re.IGNORECASE)) -def _compile_marker(pattern: Optional[str]) -> Optional[Pattern[str]]: +def _compile_marker(pattern: str | None) -> Pattern[str] | None: """Compile optional regex string to a pattern.""" if not pattern or not pattern.strip(): return None @@ -64,12 +64,12 @@ class BaseCompetitorIntentChecker: _classify_ambiguous(). Base implementation: treat as non-competitor. """ - def __init__(self, config: Dict[str, Any]) -> None: - self.brand_self: List[str] = [s.lower().strip() for s in (config.get("brand_self") or []) if s] - competitors: List[str] = [s.lower().strip() for s in (config.get("competitors") or []) if s] - aliases_map: Dict[str, List[str]] = config.get("competitor_aliases") or {} - self.competitor_canonical: Dict[str, str] = {} - self._competitor_tokens: Set[str] = set() + def __init__(self, config: dict[str, Any]) -> None: + self.brand_self: list[str] = [s.lower().strip() for s in (config.get("brand_self") or []) if s] + competitors: list[str] = [s.lower().strip() for s in (config.get("competitors") or []) if s] + aliases_map: dict[str, list[str]] = config.get("competitor_aliases") or {} + self.competitor_canonical: dict[str, str] = {} + self._competitor_tokens: set[str] = set() for c in competitors: self._competitor_tokens.add(c) self.competitor_canonical[c] = c @@ -79,17 +79,17 @@ class BaseCompetitorIntentChecker: self._competitor_tokens.add(a) self.competitor_canonical[a] = c - other: List[str] = [s.lower().strip() for s in (config.get("locations") or []) if s] - self._other_meaning_tokens: Set[str] = set(other) - self._ambiguous: Set[str] = self._competitor_tokens & self._other_meaning_tokens + other: list[str] = [s.lower().strip() for s in (config.get("locations") or []) if s] + self._other_meaning_tokens: set[str] = set(other) + self._ambiguous: set[str] = self._competitor_tokens & self._other_meaning_tokens - self.policy: Dict[str, str] = config.get("policy") or {} + self.policy: dict[str, str] = config.get("policy") or {} self.threshold_high = float(config.get("threshold_high", 0.70)) self.threshold_medium = float(config.get("threshold_medium", 0.45)) self.threshold_low = float(config.get("threshold_low", 0.30)) - self.reframe_message_template: Optional[str] = config.get("reframe_message_template") - self.refuse_message_template: Optional[str] = config.get("refuse_message_template") - self._comparison_words: List[str] = list( + self.reframe_message_template: str | None = config.get("reframe_message_template") + self.refuse_message_template: str | None = config.get("refuse_message_template") + self._comparison_words: list[str] = list( config.get("comparison_words") or [ "better", @@ -103,19 +103,19 @@ class BaseCompetitorIntentChecker: "ranked", ] ) - self._domain_words: List[str] = [s.lower().strip() for s in (config.get("domain_words") or []) if s] + self._domain_words: list[str] = [s.lower().strip() for s in (config.get("domain_words") or []) if s] - def _classify_ambiguous(self, text: str, token: str) -> Tuple[str, float]: + def _classify_ambiguous(self, text: str, token: str) -> tuple[str, float]: """ Override in subclasses for industry-specific logic. Base: treat as non-competitor. """ return "OTHER_MEANING", 0.5 - def _find_matches(self, text: str) -> List[Tuple[str, str, bool]]: + def _find_matches(self, text: str) -> list[tuple[str, str, bool]]: """Find competitor matches; mark ambiguous (also in other-meaning set).""" normalized = normalize(text) - found: List[Tuple[str, str, bool]] = [] - seen: Set[Tuple[str, str]] = set() + found: list[tuple[str, str, bool]] = [] + seen: set[tuple[str, str]] = set() for token in self._competitor_tokens: if not _word_boundary_match(normalized, token): continue @@ -131,8 +131,8 @@ class BaseCompetitorIntentChecker: def run(self, text: str) -> CompetitorIntentResult: """Classify competitor intent; non-competitor when ambiguous or low confidence.""" normalized = normalize(text) - evidence: List[CompetitorIntentEvidenceEntry] = [] - entities: Dict[str, List[str]] = { + evidence: list[CompetitorIntentEvidenceEntry] = [] + entities: dict[str, list[str]] = { "brand_self": [], "competitors": [], "category": [], @@ -178,7 +178,7 @@ class BaseCompetitorIntentChecker: "evidence": evidence, } - competitor_resolved: List[str] = [] + competitor_resolved: list[str] = [] for token, canonical, _ in matches: label, conf = self._classify_ambiguous(normalized, token) if label == "OTHER_MEANING": diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 0b990719577..5650a6b07ac 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -15,12 +15,8 @@ from re import Pattern from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Union, cast, ) @@ -101,13 +97,13 @@ class CategoryConfig: category_name: str, description: str, default_action: ContentFilterAction, - keywords: List[Dict[str, str]], - exceptions: List[str], - identifier_words: Optional[List[str]] = None, - always_block_keywords: Optional[List[Dict[str, str]]] = None, - inherit_from: Optional[str] = None, - additional_block_words: Optional[List[str]] = None, - phrase_patterns: Optional[List[str]] = None, + keywords: list[dict[str, str]], + exceptions: list[str], + identifier_words: list[str] | None = None, + always_block_keywords: list[dict[str, str]] | None = None, + inherit_from: str | None = None, + additional_block_words: list[str] | None = None, + phrase_patterns: list[str] | None = None, ): self.category_name = category_name self.description = description @@ -120,7 +116,7 @@ class CategoryConfig: self.inherit_from = inherit_from self.additional_block_words = [w.lower() for w in additional_block_words] if additional_block_words else [] # Phrase patterns: regex patterns for catching paraphrases - self.phrase_patterns: List[Tuple[str, Pattern]] = [] + self.phrase_patterns: list[tuple[str, Pattern]] = [] for p in phrase_patterns or []: try: self.phrase_patterns.append((p, re.compile(p, re.IGNORECASE))) @@ -146,21 +142,21 @@ class ContentFilterGuardrail(CustomGuardrail): def __init__( self, - guardrail_name: Optional[str] = None, - guardrail_id: Optional[str] = None, - policy_template: Optional[str] = None, - patterns: Optional[List[ContentFilterPattern]] = None, - blocked_words: Optional[List[BlockedWord]] = None, - blocked_words_file: Optional[str] = None, - event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None, + guardrail_name: str | None = None, + guardrail_id: str | None = None, + policy_template: str | None = None, + patterns: list[ContentFilterPattern] | None = None, + blocked_words: list[BlockedWord] | None = None, + blocked_words_file: str | None = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, - pattern_redaction_format: Optional[str] = None, - keyword_redaction_tag: Optional[str] = None, - categories: Optional[List[ContentFilterCategoryConfig]] = None, + pattern_redaction_format: str | None = None, + keyword_redaction_tag: str | None = None, + categories: list[ContentFilterCategoryConfig] | None = None, severity_threshold: str = "medium", - llm_router: Optional[Router] = None, - image_model: Optional[str] = None, - competitor_intent_config: Optional[Dict[str, Any]] = None, + llm_router: Router | None = None, + image_model: str | None = None, + competitor_intent_config: dict[str, Any] | None = None, **kwargs, ): """ @@ -196,19 +192,19 @@ class ContentFilterGuardrail(CustomGuardrail): self.llm_router = llm_router self.image_model = image_model # Store loaded categories - self.loaded_categories: Dict[str, CategoryConfig] = {} - self.category_keywords: Dict[ - str, Tuple[str, str, ContentFilterAction] + self.loaded_categories: dict[str, CategoryConfig] = {} + self.category_keywords: dict[ + str, tuple[str, str, ContentFilterAction] ] = {} # keyword -> (category, severity, action) # Always-block keywords are checked after exceptions (exceptions take precedence) - self.always_block_category_keywords: Dict[str, Tuple[str, str, ContentFilterAction]] = {} + self.always_block_category_keywords: dict[str, tuple[str, str, ContentFilterAction]] = {} # Store conditional categories (identifier_words + block_words) - self.conditional_categories: Dict[ - str, Dict[str, Any] + self.conditional_categories: dict[ + str, dict[str, Any] ] = {} # category_name -> {identifier_words, block_words, action, severity} # Competitor intent checker (optional; airline uses major_airlines.json, generic requires competitors) - self._competitor_intent_checker: Optional[BaseCompetitorIntentChecker] = None + self._competitor_intent_checker: BaseCompetitorIntentChecker | None = None if competitor_intent_config and isinstance(competitor_intent_config, dict): self._init_competitor_intent_checker(competitor_intent_config) @@ -221,7 +217,7 @@ class ContentFilterGuardrail(CustomGuardrail): normalized_blocked_words = self._normalize_blocked_words(blocked_words) # Compile regex patterns - self.compiled_patterns: List[Dict[str, Any]] = [] + self.compiled_patterns: list[dict[str, Any]] = [] for pattern_config in normalized_patterns: self._add_pattern(pattern_config) @@ -235,7 +231,7 @@ class ContentFilterGuardrail(CustomGuardrail): ) # Load blocked words - always initialize as dict - self.blocked_words: Dict[str, Tuple[ContentFilterAction, Optional[str]]] = {} + self.blocked_words: dict[str, tuple[ContentFilterAction, str | None]] = {} for word in normalized_blocked_words: self.blocked_words[word.keyword.lower()] = (word.action, word.description) @@ -258,7 +254,7 @@ class ContentFilterGuardrail(CustomGuardrail): f"Loaded {len(self.loaded_categories)} categories with {len(self.category_keywords)} keywords" ) - def _init_competitor_intent_checker(self, competitor_intent_config: Dict[str, Any]) -> None: + def _init_competitor_intent_checker(self, competitor_intent_config: dict[str, Any]) -> None: try: competitor_intent_type = competitor_intent_config.get("competitor_intent_type", "airline") if competitor_intent_type == "generic": @@ -277,9 +273,9 @@ class ContentFilterGuardrail(CustomGuardrail): @staticmethod def _normalize_patterns( - patterns: Optional[List[ContentFilterPattern]], - ) -> List[ContentFilterPattern]: - result: List[ContentFilterPattern] = [] + patterns: list[ContentFilterPattern] | None, + ) -> list[ContentFilterPattern]: + result: list[ContentFilterPattern] = [] if patterns: for pattern_config in patterns: if isinstance(pattern_config, dict): @@ -290,9 +286,9 @@ class ContentFilterGuardrail(CustomGuardrail): @staticmethod def _normalize_blocked_words( - blocked_words: Optional[List[BlockedWord]], - ) -> List[BlockedWord]: - result: List[BlockedWord] = [] + blocked_words: list[BlockedWord] | None, + ) -> list[BlockedWord]: + result: list[BlockedWord] = [] if blocked_words: for word in blocked_words: if isinstance(word, dict): @@ -388,7 +384,7 @@ class ContentFilterGuardrail(CustomGuardrail): self._assert_within_categories_dir(os.path.join(module_dir, file_path), module_dir) return file_path - def _load_categories(self, categories: List[ContentFilterCategoryConfig]) -> None: + def _load_categories(self, categories: list[ContentFilterCategoryConfig]) -> None: """ Load content categories from configuration. @@ -640,7 +636,7 @@ class ContentFilterGuardrail(CustomGuardrail): # Derive category name from filename (e.g. harm_toxic_abuse.json -> harm_toxic_abuse) category_name = os.path.splitext(os.path.basename(file_path))[0] severity_map = {4: "high", 3: "high", 2: "medium", 1: "low"} - keywords: List[Dict[str, str]] = [] + keywords: list[dict[str, str]] = [] seen = set() for item in entries: if not isinstance(item, dict): @@ -684,7 +680,7 @@ class ContentFilterGuardrail(CustomGuardrail): pattern_config: ContentFilterPattern configuration """ try: - extra_config: Dict[str, Any] = {} + extra_config: dict[str, Any] = {} if pattern_config.pattern_type == "prebuilt": if not pattern_config.pattern_name: raise ValueError("pattern_name is required for prebuilt patterns") @@ -699,7 +695,7 @@ class ContentFilterGuardrail(CustomGuardrail): else: raise ValueError(f"Unknown pattern_type: {pattern_config.pattern_type}") - keyword_regex: Optional[Pattern] = None + keyword_regex: Pattern | None = None if extra_config.get("keyword_pattern"): keyword_regex = re.compile(extra_config["keyword_pattern"], re.IGNORECASE) @@ -754,22 +750,22 @@ class ContentFilterGuardrail(CustomGuardrail): except FileNotFoundError: raise FileNotFoundError(f"Blocked words file not found: {file_path}") except Exception as e: - raise Exception(f"Error loading blocked words file {file_path}: {str(e)}") + raise Exception(f"Error loading blocked words file {file_path}: {e!s}") - def _find_pattern_spans(self, text: str, pattern_entry: Dict[str, Any]) -> List[Tuple[int, int]]: + def _find_pattern_spans(self, text: str, pattern_entry: dict[str, Any]) -> list[tuple[int, int]]: """Return all match spans for a pattern, applying contextual rules if required.""" regex: Pattern = pattern_entry["regex"] - keyword_regex: Optional[Pattern] = pattern_entry.get("keyword_regex") + keyword_regex: Pattern | None = pattern_entry.get("keyword_regex") allow_word_numbers: bool = pattern_entry.get("allow_word_numbers", False) - keyword_matches: Optional[List[re.Match]] = None + keyword_matches: list[re.Match] | None = None if keyword_regex is not None: keyword_matches = list(keyword_regex.finditer(text)) if not keyword_matches: return [] - match_spans: List[Tuple[int, int]] = [] + match_spans: list[tuple[int, int]] = [] for match in regex.finditer(text): if keyword_matches is not None and not self._match_near_keyword( @@ -797,7 +793,7 @@ class ContentFilterGuardrail(CustomGuardrail): self, value_start: int, value_end: int, - keyword_matches: List[re.Match], + keyword_matches: list[re.Match], text: str, ) -> bool: """Check if a value is separated from a keyword by an allowed gap.""" @@ -828,14 +824,14 @@ class ContentFilterGuardrail(CustomGuardrail): words = GAP_WORD_TOKENIZER.findall(gap_text) return len(words) <= MAX_KEYWORD_VALUE_GAP_WORDS - def _merge_spans(self, spans: List[Tuple[int, int]]) -> List[Tuple[int, int]]: + def _merge_spans(self, spans: list[tuple[int, int]]) -> list[tuple[int, int]]: """Merge overlapping spans to avoid double-masking.""" if not spans: return [] spans.sort(key=lambda item: item[0]) - merged: List[Tuple[int, int]] = [spans[0]] + merged: list[tuple[int, int]] = [spans[0]] for start, end in spans[1:]: last_start, last_end = merged[-1] @@ -845,13 +841,13 @@ class ContentFilterGuardrail(CustomGuardrail): merged.append((start, end)) return merged - def _mask_spans(self, text: str, spans: List[Tuple[int, int]], redaction: str) -> str: + def _mask_spans(self, text: str, spans: list[tuple[int, int]], redaction: str) -> str: """Apply masking for the provided spans using the given redaction tag.""" if not spans: return text - result_parts: List[str] = [] + result_parts: list[str] = [] previous_end = 0 for start, end in spans: result_parts.append(text[previous_end:start]) @@ -860,14 +856,14 @@ class ContentFilterGuardrail(CustomGuardrail): result_parts.append(text[previous_end:]) return "".join(result_parts) - def _convert_word_number_sequence(self, sequence: str) -> Optional[str]: + def _convert_word_number_sequence(self, sequence: str) -> str | None: """Convert a spelled-out digit sequence (e.g., 'One-Two') into digits.""" tokens = WORD_NUMBER_TOKEN_FINDER.findall(sequence) if not tokens: return None - digits: List[str] = [] + digits: list[str] = [] for token in tokens: digit = WORD_NUMBER_MAP.get(token.lower()) if digit is None: @@ -876,7 +872,7 @@ class ContentFilterGuardrail(CustomGuardrail): return "".join(digits) if digits else None - def _check_patterns(self, text: str) -> Optional[Tuple[str, str, ContentFilterAction]]: + def _check_patterns(self, text: str) -> tuple[str, str, ContentFilterAction] | None: """ Check text against all compiled regex patterns. @@ -898,8 +894,8 @@ class ContentFilterGuardrail(CustomGuardrail): return None def _check_conditional_categories( - self, text: str, exceptions: List[str] - ) -> Optional[Tuple[str, str, str, ContentFilterAction]]: + self, text: str, exceptions: list[str] + ) -> tuple[str, str, str, ContentFilterAction] | None: """ Check text for conditional category matches (identifier + block word in same sentence). @@ -986,8 +982,8 @@ class ContentFilterGuardrail(CustomGuardrail): return None def _check_phrase_patterns( - self, text: str, exceptions: List[str] - ) -> Optional[Tuple[str, str, str, ContentFilterAction]]: + self, text: str, exceptions: list[str] + ) -> tuple[str, str, str, ContentFilterAction] | None: """ Check text against phrase patterns from loaded categories. @@ -1035,8 +1031,8 @@ class ContentFilterGuardrail(CustomGuardrail): return None def _check_category_keywords( - self, text: str, exceptions: List[str] - ) -> Optional[Tuple[str, str, str, ContentFilterAction]]: + self, text: str, exceptions: list[str] + ) -> tuple[str, str, str, ContentFilterAction] | None: """ Check text for category keywords. @@ -1112,7 +1108,7 @@ class ContentFilterGuardrail(CustomGuardrail): return (keyword, category, severity, action) return None - def _check_blocked_words(self, text: str) -> Optional[Tuple[str, ContentFilterAction, Optional[str]]]: + def _check_blocked_words(self, text: str) -> tuple[str, ContentFilterAction, str | None] | None: """ Check text for blocked keywords. @@ -1129,7 +1125,7 @@ class ContentFilterGuardrail(CustomGuardrail): "This suggests an old guardrail instance is still in use. Please restart the server." ) # Convert list to dict on-the-fly - temp_dict: Dict[str, Tuple[ContentFilterAction, Optional[str]]] = {} + temp_dict: dict[str, tuple[ContentFilterAction, str | None]] = {} for word in self.blocked_words: if isinstance(word, dict): temp_dict[word.get("keyword", "").lower()] = ( @@ -1154,7 +1150,7 @@ class ContentFilterGuardrail(CustomGuardrail): category_name: str, severity: str, action: ContentFilterAction, - detections: Optional[List[ContentFilterDetection]], + detections: list[ContentFilterDetection] | None, ) -> None: """Handle conditional category match detection and action.""" if detections is not None: @@ -1193,7 +1189,7 @@ class ContentFilterGuardrail(CustomGuardrail): severity: str, action: ContentFilterAction, text: str, - detections: Optional[List[ContentFilterDetection]], + detections: list[ContentFilterDetection] | None, ) -> str: """Handle category keyword match detection and action.""" if detections is not None: @@ -1237,8 +1233,8 @@ class ContentFilterGuardrail(CustomGuardrail): pattern_name: str, action: ContentFilterAction, text: str, - spans: List[Tuple[int, int]], - detections: Optional[List[ContentFilterDetection]], + spans: list[tuple[int, int]], + detections: list[ContentFilterDetection] | None, ) -> str: """Handle regex pattern match detection and action.""" if detections is not None: @@ -1267,9 +1263,9 @@ class ContentFilterGuardrail(CustomGuardrail): self, keyword: str, action: ContentFilterAction, - description: Optional[str], + description: str | None, text: str, - detections: Optional[List[ContentFilterDetection]], + detections: list[ContentFilterDetection] | None, ) -> str: """Handle blocked word match detection and action.""" verbose_proxy_logger.debug(f"Blocked word '{keyword}' found with action {action}") @@ -1308,7 +1304,7 @@ class ContentFilterGuardrail(CustomGuardrail): return text - def _filter_single_text(self, text: str, detections: Optional[List[ContentFilterDetection]] = None) -> str: + def _filter_single_text(self, text: str, detections: list[ContentFilterDetection] | None = None) -> str: """ Apply all content filtering checks to a single text. @@ -1393,7 +1389,7 @@ class ContentFilterGuardrail(CustomGuardrail): redaction_tag = self.pattern_redaction_format.format(pattern_name=pattern_name.upper()) return redaction_tag - async def _process_images(self, images: List[str], detections: List[ContentFilterDetection]) -> None: + async def _process_images(self, images: list[str], detections: list[ContentFilterDetection]) -> None: """ Process images by describing them and applying content filtering. @@ -1447,7 +1443,7 @@ class ContentFilterGuardrail(CustomGuardrail): except HTTPException as e: # e.detail can be a string or dict if isinstance(e.detail, dict) and "error" in e.detail: - detail_dict = cast(Dict[str, Any], e.detail) + detail_dict = cast(dict[str, Any], e.detail) detail_dict["error"] = detail_dict["error"] + " (Image description): " + description elif isinstance(e.detail, str): e.detail = e.detail + " (Image description): " + description @@ -1457,8 +1453,8 @@ class ContentFilterGuardrail(CustomGuardrail): def _count_masked_entities( self, - detections: List[ContentFilterDetection], - masked_entity_count: Dict[str, int], + detections: list[ContentFilterDetection], + masked_entity_count: dict[str, int], ) -> None: """ Count masked entities by type from detections. @@ -1484,9 +1480,9 @@ class ContentFilterGuardrail(CustomGuardrail): category = category_detection["category"] masked_entity_count[category] = masked_entity_count.get(category, 0) + 1 - def _build_match_details(self, detections: List[ContentFilterDetection]) -> List[dict]: + def _build_match_details(self, detections: list[ContentFilterDetection]) -> list[dict]: """Build match_details list from content filter detections.""" - match_details: List[dict] = [] + match_details: list[dict] = [] for detection in detections: action_taken = detection.get("action", detection.get("action_hint", "")) detail: dict = {"type": detection["type"], "action_taken": action_taken} @@ -1508,7 +1504,7 @@ class ContentFilterGuardrail(CustomGuardrail): match_details.append(detail) return match_details - def _get_detection_methods(self, detections: List[ContentFilterDetection]) -> str: + def _get_detection_methods(self, detections: list[ContentFilterDetection]) -> str: """Get comma-separated detection methods used.""" methods: set = set() for detection in detections: @@ -1529,7 +1525,7 @@ class ContentFilterGuardrail(CustomGuardrail): + len(self.always_block_category_keywords) ) - def _get_policy_templates(self) -> Optional[str]: + def _get_policy_templates(self) -> str | None: """Get comma-separated policy template names from loaded categories.""" if not self.loaded_categories: return None @@ -1538,8 +1534,8 @@ class ContentFilterGuardrail(CustomGuardrail): def _compute_risk_score( self, - detections: List[ContentFilterDetection], - masked_entity_count: Dict[str, int], + detections: list[ContentFilterDetection], + masked_entity_count: dict[str, int], status: "GuardrailStatus", ) -> float: """ @@ -1575,7 +1571,7 @@ class ContentFilterGuardrail(CustomGuardrail): self, intent_result: CompetitorIntentResult, request_data: dict, - detections: List[ContentFilterDetection], + detections: list[ContentFilterDetection], ) -> None: """ Apply policy for competitor intent result: refuse (raise), reframe (passthrough), or log_only/allow (return). @@ -1636,10 +1632,10 @@ class ContentFilterGuardrail(CustomGuardrail): def _log_guardrail_information( self, request_data: dict, - detections: List[ContentFilterDetection], + detections: list[ContentFilterDetection], status: "GuardrailStatus", start_time: datetime, - masked_entity_count: Dict[str, int], + masked_entity_count: dict[str, int], exception_str: str, ) -> None: """ @@ -1654,12 +1650,12 @@ class ContentFilterGuardrail(CustomGuardrail): exception_str: Exception string if guardrail failed """ # Convert TypedDict detections to regular dicts for JSON serialization - guardrail_json_response: Union[Exception, str, dict, List[dict]] = [dict(detection) for detection in detections] + guardrail_json_response: Exception | str | dict | list[dict] = [dict(detection) for detection in detections] if status != "success": guardrail_json_response = exception_str if exception_str else [dict(detection) for detection in detections] # Competitor intent: add confidence and classification to tracing if present - tracing_kw: Dict[str, Any] = { + tracing_kw: dict[str, Any] = { "guardrail_id": self.config_guardrail_id or self.guardrail_name, "policy_template": self.config_policy_template or self._get_policy_templates(), "detection_method": (self._get_detection_methods(detections) if detections else None), @@ -1771,8 +1767,8 @@ class ContentFilterGuardrail(CustomGuardrail): HTTPException: If sensitive content is detected and action is BLOCK """ start_time = datetime.now() - detections: List[ContentFilterDetection] = [] - masked_entity_count: Dict[str, int] = {} + detections: list[ContentFilterDetection] = [] + masked_entity_count: dict[str, int] = {} status: GuardrailStatus = "success" exception_str: str = "" @@ -1844,14 +1840,14 @@ class ContentFilterGuardrail(CustomGuardrail): and the UI Request Lifecycle panel. Mirrors apply_guardrail's finally-block contract. """ - accumulated_text_by_choice: Dict[int, str] = {} - yielded_masked_text_len_by_choice: Dict[int, int] = {} - latest_detections_by_choice: Dict[int, List[ContentFilterDetection]] = {} + accumulated_text_by_choice: dict[int, str] = {} + yielded_masked_text_len_by_choice: dict[int, int] = {} + latest_detections_by_choice: dict[int, list[ContentFilterDetection]] = {} buffer_size = 50 # Increased buffer to catch patterns split across many chunks start_time = datetime.now() - detections: List[ContentFilterDetection] = [] - masked_entity_count: Dict[str, int] = {} + detections: list[ContentFilterDetection] = [] + masked_entity_count: dict[str, int] = {} status: GuardrailStatus = "success" exception_str: str = "" @@ -1885,7 +1881,7 @@ class ContentFilterGuardrail(CustomGuardrail): # Add a space at the end if it's the final chunk to trigger word boundaries (\b) text_to_scan = text_to_check + (" " if is_final else "") - choice_detections: List[ContentFilterDetection] = [] + choice_detections: list[ContentFilterDetection] = [] try: # _filter_single_text scans the whole accumulated @@ -1964,7 +1960,7 @@ class ContentFilterGuardrail(CustomGuardrail): return LitellmContentFilterGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py index 3e20ada1cfa..6a5bcbad757 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/test_eval.py @@ -21,7 +21,7 @@ import os import re import time from datetime import datetime, timezone -from typing import Any, Dict, List +from typing import Any import pytest from fastapi import HTTPException @@ -33,7 +33,7 @@ RESULTS_DIR = os.path.join(os.path.dirname(__file__), "results") # ── Helpers ─────────────────────────────────────────────────────── -def _load_jsonl(filename: str) -> List[dict]: +def _load_jsonl(filename: str) -> list[dict]: """Load eval cases from a JSONL file. One JSON object per line.""" cases = [] path = os.path.join(EVAL_DIR, filename) @@ -60,7 +60,7 @@ def _run(checker, text: str) -> dict: return {"decision": "ALLOW", "score": 0.0, "matched_topic": None} except HTTPException as e: if e.status_code == 400: - detail: Dict[str, Any] = e.detail if isinstance(e.detail, dict) else {} + detail: dict[str, Any] = e.detail if isinstance(e.detail, dict) else {} return { "decision": "BLOCK", "score": detail.get("score", 1.0), @@ -149,7 +149,7 @@ def _save_confusion_results(label: str, metrics: dict, wrong: list, rows: list) return result -def _confusion_matrix(checker, cases: List[dict], label: str): +def _confusion_matrix(checker, cases: list[dict], label: str): """Run all cases, print confusion matrix, save results JSON.""" tp = fp = tn = fn = 0 wrong = [] diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py index 3a3431aa0ee..d30de723443 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py @@ -10,10 +10,10 @@ import os import re from enum import Enum from re import Pattern -from typing import Any, Dict, List +from typing import Any -def _load_patterns_from_json() -> Dict: +def _load_patterns_from_json() -> dict: """Load pattern definitions from patterns.json file""" json_path = os.path.join(os.path.dirname(__file__), "patterns.json") with open(json_path, "r") as f: @@ -27,8 +27,6 @@ _PATTERNS_DATA = _load_patterns_from_json() class PrebuiltPatternName(str, Enum): """Enum for prebuilt pattern names - dynamically generated from JSON""" - pass - # Dynamically create enum values from JSON for pattern_data in _PATTERNS_DATA["patterns"]: @@ -36,7 +34,7 @@ for pattern_data in _PATTERNS_DATA["patterns"]: # Build lookup dictionaries from JSON -PREBUILT_PATTERNS: Dict[str, str] = { +PREBUILT_PATTERNS: dict[str, str] = { pattern_data["name"]: pattern_data["pattern"] for pattern_data in _PATTERNS_DATA["patterns"] } @@ -51,7 +49,7 @@ KNOWN_PATTERN_KEYS = { "description", } -PATTERN_EXTRA_CONFIG: Dict[str, Dict[str, Any]] = {} +PATTERN_EXTRA_CONFIG: dict[str, dict[str, Any]] = {} for pattern_data in _PATTERNS_DATA["patterns"]: extra_config = {key: value for key, value in pattern_data.items() if key not in KNOWN_PATTERN_KEYS} PATTERN_EXTRA_CONFIG[pattern_data["name"]] = extra_config @@ -77,7 +75,7 @@ def get_compiled_pattern(pattern_name: str) -> Pattern: return re.compile(PREBUILT_PATTERNS[pattern_name], re.IGNORECASE) -def get_all_pattern_names() -> List[str]: +def get_all_pattern_names() -> list[str]: """ Get a list of all available prebuilt pattern names. @@ -88,7 +86,7 @@ def get_all_pattern_names() -> List[str]: # Build category mapping from JSON -PATTERN_CATEGORIES: Dict[str, List[str]] = {} +PATTERN_CATEGORIES: dict[str, list[str]] = {} for pattern_data in _PATTERNS_DATA["patterns"]: category = pattern_data["category"] if category not in PATTERN_CATEGORIES: @@ -97,18 +95,18 @@ for pattern_data in _PATTERNS_DATA["patterns"]: # Build display names mapping from JSON -PATTERN_DISPLAY_NAMES: Dict[str, str] = { +PATTERN_DISPLAY_NAMES: dict[str, str] = { pattern_data["name"]: pattern_data["display_name"] for pattern_data in _PATTERNS_DATA["patterns"] } # Build descriptions mapping from JSON -PATTERN_DESCRIPTIONS: Dict[str, str] = { +PATTERN_DESCRIPTIONS: dict[str, str] = { pattern_data["name"]: pattern_data["description"] for pattern_data in _PATTERNS_DATA["patterns"] } -def get_pattern_metadata() -> List[Dict[str, str]]: +def get_pattern_metadata() -> list[dict[str, str]]: """ Return pattern metadata for UI display. @@ -126,7 +124,7 @@ def get_pattern_metadata() -> List[Dict[str, str]]: ] -def get_available_content_categories() -> List[Dict[str, str]]: +def get_available_content_categories() -> list[dict[str, str]]: """ Return available content categories for UI display. @@ -170,7 +168,7 @@ def get_available_content_categories() -> List[Dict[str, str]]: # Skip files that can't be loaded but log the error for debugging from litellm._logging import verbose_proxy_logger - verbose_proxy_logger.warning(f"Failed to load category file {filename}: {str(e)}") + verbose_proxy_logger.warning(f"Failed to load category file {filename}: {e!s}") continue elif filename.endswith(".json"): # JSON category files (e.g. harm_toxic_abuse.json) - no YAML header, use filename diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py index 193793fc50a..afbc67f2abb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py @@ -4,7 +4,7 @@ import json import re from collections.abc import Callable from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, cast from fastapi import HTTPException @@ -45,7 +45,7 @@ def _default_router_provider() -> "Router | None": _JSON_FENCE_RE = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL | re.IGNORECASE) -def _parse_judge_verdict(raw: str) -> Dict[str, Any]: +def _parse_judge_verdict(raw: str) -> dict[str, Any]: """Parse the judge's JSON verdict, tolerating markdown fences and surrounding prose.""" text = raw.strip() fenced = _JSON_FENCE_RE.search(text) @@ -62,7 +62,7 @@ def _parse_judge_verdict(raw: str) -> Dict[str, Any]: parsed = json.loads(text[start : end + 1]) if not isinstance(parsed, dict): raise ValueError("judge response is not a JSON object") - return cast(Dict[str, Any], parsed) # cast-ok: narrowed to dict by the isinstance guard above + return cast(dict[str, Any], parsed) # cast-ok: narrowed to dict by the isinstance guard above def _extract_text_from_content(content: Any) -> str: @@ -98,8 +98,8 @@ def _get_litellm_param( def _build_judge_prompt( - criteria: List[Dict[str, Any]], - messages: List[Dict[str, Any]], + criteria: list[dict[str, Any]], + messages: list[dict[str, Any]], response_text: str, ) -> str: criteria_block = "\n".join( @@ -124,15 +124,15 @@ class LLMAsAJudgeGuardrail(CustomGuardrail): self, guardrail_name: str, judge_model: str, - criteria: List[Dict[str, Any]], + criteria: list[dict[str, Any]], overall_threshold: float = 80.0, on_failure: Literal["block", "log"] = "block", - event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None, default_on: bool = False, router_provider: "Callable[[], Router | None] | None" = None, **kwargs: Any, ) -> None: - _event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = None + _event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None if event_hook is not None: if isinstance(event_hook, list): _event_hook = [GuardrailEventHooks(h) if isinstance(h, str) else h for h in event_hook] @@ -153,14 +153,14 @@ class LLMAsAJudgeGuardrail(CustomGuardrail): self._router_provider = router_provider or _default_router_provider @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [GuardrailEventHooks.post_call] async def _run_judge( self, - messages: List[Dict[str, Any]], + messages: list[dict[str, Any]], response_text: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: judge_messages = [ {"role": "system", "content": JUDGE_SYSTEM_PROMPT}, { @@ -208,10 +208,10 @@ class LLMAsAJudgeGuardrail(CustomGuardrail): start_time = datetime.now() status: GuardrailStatus = "success" - judge_result: Dict[str, Any] = {} + judge_result: dict[str, Any] = {} try: - messages: List[Dict[str, Any]] = request_data.get("messages") or [] + messages: list[dict[str, Any]] = request_data.get("messages") or [] try: judge_result = await self._run_judge(messages, response_text) @@ -228,7 +228,7 @@ class LLMAsAJudgeGuardrail(CustomGuardrail): passed = overall_score >= self.overall_threshold - eval_info: "StandardLoggingEvalInformation" = { + eval_info: StandardLoggingEvalInformation = { "eval_name": self.guardrail_name or "", "overall_score": overall_score, "passed": passed, @@ -304,7 +304,7 @@ def initialize_guardrail( overall_threshold = float(_get_litellm_param(litellm_params, guardrail, "overall_threshold", 80.0)) mode = _get_litellm_param(litellm_params, guardrail, "mode") - event_hook: Optional[GuardrailEventHooks] = None + event_hook: GuardrailEventHooks | None = None if isinstance(mode, str) and mode in {e.value for e in GuardrailEventHooks}: event_hook = GuardrailEventHooks(mode) diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py index 237364f9714..2294703004f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, cast +from typing import TYPE_CHECKING, Any, cast from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -14,7 +14,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" # Default to always-on. Only disable if the user explicitly sets default_on: false. # We check the raw guardrail dict because LitellmParams normalizes None → False, # making it impossible to distinguish "not set" from "explicitly false" via litellm_params. - _raw_default_on = cast(Dict[str, Any], guardrail).get("litellm_params", {}).get("default_on") + _raw_default_on = cast(dict[str, Any], guardrail).get("litellm_params", {}).get("default_on") _default_on = False if _raw_default_on is False else True _callback = MCPEndUserPermissionGuardrail( diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py index d084a7e088b..1760d01e247 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py @@ -10,7 +10,7 @@ Permission logic: - end_user_id + mcp_servers → allow only those servers """ -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -54,7 +54,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"] = "request", - logging_obj: Optional[Any] = None, + logging_obj: Any | None = None, ) -> GenericGuardrailAPIInputs: """ Filters MCP tools the end user cannot access based on their @@ -70,7 +70,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): async def _check_request_tools( self, inputs: GenericGuardrailAPIInputs, - object_permission: Optional[LiteLLM_ObjectPermissionTable], + object_permission: LiteLLM_ObjectPermissionTable | None, ) -> GenericGuardrailAPIInputs: tools = inputs.get("tools") if not tools: @@ -116,7 +116,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): @staticmethod async def _resolve_end_user_object_permission( request_data: dict, - ) -> Optional[LiteLLM_ObjectPermissionTable]: + ) -> LiteLLM_ObjectPermissionTable | None: """ Resolve the end user's object_permission via the cached auth lookup. @@ -131,7 +131,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): return end_user_object.object_permission if end_user_object is not None else None @staticmethod - def _get_end_user_id_from_request_data(request_data: dict) -> Optional[str]: + def _get_end_user_id_from_request_data(request_data: dict) -> str | None: return request_data.get("user_api_key_end_user_id") or request_data.get("litellm_metadata", {}).get( "user_api_key_end_user_id" ) @@ -171,8 +171,8 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): @staticmethod async def _get_allowed_mcp_servers_from_object_permission( - object_permission: Optional[LiteLLM_ObjectPermissionTable], - ) -> Optional[List[str]]: + object_permission: LiteLLM_ObjectPermissionTable | None, + ) -> list[str] | None: """ Returns: None — no restrictions configured, allow all MCP servers @@ -200,7 +200,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): # ------------------------------------------------------------------ @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.mcp_end_user_permission import ( MCPEndUserPermissionGuardrailConfigModel, ) @@ -208,7 +208,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): return MCPEndUserPermissionGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, ] @@ -218,7 +218,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): # ------------------------------------------------------------------ @staticmethod - def _extract_mcp_server_name(tool_name: str) -> Optional[str]: + def _extract_mcp_server_name(tool_name: str) -> str | None: """ Split "github-create_issue" → "github". Returns None if the tool name has no '-' prefix (not an MCP tool). @@ -228,7 +228,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): return tool_name.split("-", 1)[0] @staticmethod - def _get_tool_name_from_definition(tool: Any) -> Optional[str]: + def _get_tool_name_from_definition(tool: Any) -> str | None: """ Extract tool name from a definition dict. diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py index 434a136dc91..5bcad0c9bb9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py @@ -77,6 +77,6 @@ guardrail_class_registry = { __all__ = [ "MCPJWTSigner", - "initialize_guardrail", "get_mcp_jwt_signer", + "initialize_guardrail", ] diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index e70a4e1d8e7..d5457ef4421 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -73,7 +73,7 @@ import hashlib import os import re import time -from typing import Any, Dict, List, Optional, Union +from typing import Any, Optional import jwt from cryptography.hazmat.primitives import serialization @@ -96,7 +96,7 @@ _mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None _MCP_JWT_CALL_TYPES = frozenset({"call_mcp_tool", "list_mcp_tools"}) # Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at). -_jwks_cache: Dict[str, tuple] = {} +_jwks_cache: dict[str, tuple] = {} _JWKS_CACHE_TTL = 3600 # 1 hour @@ -142,7 +142,7 @@ def _compute_kid(public_key: Any) -> str: return hashlib.sha256(der_bytes).hexdigest()[:16] -async def _fetch_jwks(jwks_uri: str) -> List[Dict[str, Any]]: +async def _fetch_jwks(jwks_uri: str) -> list[dict[str, Any]]: """ Fetch and cache a JWKS from the given URI. @@ -168,7 +168,7 @@ async def _fetch_jwks(jwks_uri: str) -> List[Dict[str, Any]]: return keys # type: ignore[return-value] -async def _fetch_oidc_discovery(discovery_uri: str) -> Dict[str, Any]: +async def _fetch_oidc_discovery(discovery_uri: str) -> dict[str, Any]: """Fetch an OIDC discovery document and return its parsed JSON.""" from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -213,36 +213,36 @@ class MCPJWTSigner(CustomGuardrail): SIGNING_KEY_ENV = "MCP_JWT_SIGNING_KEY" @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [GuardrailEventHooks.pre_mcp_call] def __init__( self, # Core signing config - issuer: Optional[str] = None, - audience: Optional[str] = None, - ttl_seconds: Optional[int] = None, + issuer: str | None = None, + audience: str | None = None, + ttl_seconds: int | None = None, # FR-5: Verify + re-sign - access_token_discovery_uri: Optional[str] = None, - token_introspection_endpoint: Optional[str] = None, - verify_issuer: Optional[str] = None, - verify_audience: Optional[str] = None, + access_token_discovery_uri: str | None = None, + token_introspection_endpoint: str | None = None, + verify_issuer: str | None = None, + verify_audience: str | None = None, # FR-12: End-user identity mapping - end_user_claim_sources: Optional[List[str]] = None, + end_user_claim_sources: list[str] | None = None, # FR-13: Claim operations - add_claims: Optional[Dict[str, Any]] = None, - set_claims: Optional[Dict[str, Any]] = None, - remove_claims: Optional[List[str]] = None, + add_claims: dict[str, Any] | None = None, + set_claims: dict[str, Any] | None = None, + remove_claims: list[str] | None = None, # FR-14: Two-token model - channel_token_audience: Optional[str] = None, - channel_token_ttl: Optional[int] = None, + channel_token_audience: str | None = None, + channel_token_ttl: int | None = None, # FR-15: Incoming claim validation - required_claims: Optional[List[str]] = None, - optional_claims: Optional[List[str]] = None, + required_claims: list[str] | None = None, + optional_claims: list[str] | None = None, # FR-9: Debug headers debug_headers: bool = False, # FR-10: Configurable scopes - allowed_scopes: Optional[List[str]] = None, + allowed_scopes: list[str] | None = None, **kwargs: Any, ) -> None: kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -278,39 +278,39 @@ class MCPJWTSigner(CustomGuardrail): self.ttl_seconds: int = resolved_ttl # --- FR-5: Verify + re-sign --- - self.access_token_discovery_uri: Optional[str] = access_token_discovery_uri - self.token_introspection_endpoint: Optional[str] = token_introspection_endpoint - self.verify_issuer: Optional[str] = verify_issuer - self.verify_audience: Optional[str] = verify_audience + self.access_token_discovery_uri: str | None = access_token_discovery_uri + self.token_introspection_endpoint: str | None = token_introspection_endpoint + self.verify_issuer: str | None = verify_issuer + self.verify_audience: str | None = verify_audience # Cached OIDC discovery document (fetched lazily, TTL = 24 h) - self._oidc_discovery_doc: Optional[Dict[str, Any]] = None + self._oidc_discovery_doc: dict[str, Any] | None = None self._oidc_discovery_fetched_at: float = 0.0 # --- FR-12: End-user identity mapping --- # Default chain: try incoming JWT sub, fall back to litellm user_id - self.end_user_claim_sources: List[str] = end_user_claim_sources or [ + self.end_user_claim_sources: list[str] = end_user_claim_sources or [ "token:sub", "litellm:user_id", ] # --- FR-13: Claim operations --- - self.add_claims: Dict[str, Any] = add_claims or {} - self.set_claims: Dict[str, Any] = set_claims or {} - self.remove_claims: List[str] = remove_claims or [] + self.add_claims: dict[str, Any] = add_claims or {} + self.set_claims: dict[str, Any] = set_claims or {} + self.remove_claims: list[str] = remove_claims or [] # --- FR-14: Two-token model --- - self.channel_token_audience: Optional[str] = channel_token_audience + self.channel_token_audience: str | None = channel_token_audience self.channel_token_ttl: int = channel_token_ttl if channel_token_ttl is not None else self.ttl_seconds # --- FR-15: Incoming claim validation --- - self.required_claims: List[str] = required_claims or [] - self.optional_claims: List[str] = optional_claims or [] + self.required_claims: list[str] = required_claims or [] + self.optional_claims: list[str] = optional_claims or [] # --- FR-9: Debug headers --- self.debug_headers: bool = debug_headers # --- FR-10: Configurable scopes --- - self.allowed_scopes: Optional[List[str]] = allowed_scopes + self.allowed_scopes: list[str] | None = allowed_scopes # Register singleton for JWKS/OIDC discovery endpoints. global _mcp_jwt_signer_instance @@ -347,7 +347,7 @@ class MCPJWTSigner(CustomGuardrail): """ return 3600 if self._persistent_key else 300 - def get_jwks(self) -> Dict[str, Any]: + def get_jwks(self) -> dict[str, Any]: """ Return the JWKS for the RSA public key. Used by GET /.well-known/jwks.json so MCP servers can verify tokens. @@ -374,7 +374,7 @@ class MCPJWTSigner(CustomGuardrail): # the IdP, short enough to pick up jwks_uri changes after key rotation. _OIDC_DISCOVERY_TTL = 86400 - async def _get_oidc_discovery(self) -> Dict[str, Any]: + async def _get_oidc_discovery(self) -> dict[str, Any]: """Fetch and cache the OIDC discovery document with a 24-hour TTL. Only caches when the doc contains a 'jwks_uri' so that a transient or @@ -391,7 +391,7 @@ class MCPJWTSigner(CustomGuardrail): return doc return self._oidc_discovery_doc or {} - async def _verify_incoming_jwt(self, raw_token: str) -> Dict[str, Any]: + async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, Any]: """ Verify an incoming Bearer JWT against the configured IdP's JWKS. @@ -442,8 +442,8 @@ class MCPJWTSigner(CustomGuardrail): # it infers from the key type (RSAPublicKey → RS256). alg = getattr(signing_jwk, "algorithm_name", None) or "RS256" - decode_options: Dict[str, Any] = {"verify_exp": True} - decode_kwargs: Dict[str, Any] = { + decode_options: dict[str, Any] = {"verify_exp": True} + decode_kwargs: dict[str, Any] = { "algorithms": [alg], "options": decode_options, } @@ -455,10 +455,10 @@ class MCPJWTSigner(CustomGuardrail): if self.verify_issuer: decode_kwargs["issuer"] = self.verify_issuer - payload: Dict[str, Any] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs) + payload: dict[str, Any] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs) return payload - async def _introspect_opaque_token(self, token: str) -> Dict[str, Any]: + async def _introspect_opaque_token(self, token: str) -> dict[str, Any]: """ Perform RFC 7662 token introspection for opaque (non-JWT) tokens. @@ -483,7 +483,7 @@ class MCPJWTSigner(CustomGuardrail): headers={"Accept": "application/json"}, ) resp.raise_for_status() - result: Dict[str, Any] = resp.json() + result: dict[str, Any] = resp.json() if not result.get("active", False): raise jwt.exceptions.ExpiredSignatureError( # type: ignore[attr-defined] "MCPJWTSigner: incoming token is inactive (introspection returned active=false)" @@ -496,7 +496,7 @@ class MCPJWTSigner(CustomGuardrail): def _validate_required_claims( self, - jwt_claims: Optional[Dict[str, Any]], + jwt_claims: dict[str, Any] | None, ) -> None: """ Raise HTTP 403 if any required_claims are absent from the verified @@ -526,7 +526,7 @@ class MCPJWTSigner(CustomGuardrail): def _resolve_end_user_identity( self, user_api_key_dict: UserAPIKeyAuth, - jwt_claims: Optional[Dict[str, Any]], + jwt_claims: dict[str, Any] | None, ) -> str: """ Resolve the outbound JWT 'sub' using the ordered end_user_claim_sources list. @@ -541,7 +541,7 @@ class MCPJWTSigner(CustomGuardrail): Falls back to a stable hash of the API token for service-account callers. """ for source in self.end_user_claim_sources: - value: Optional[str] = None + value: str | None = None if source.startswith("token:"): claim_name = source[len("token:") :] @@ -584,7 +584,7 @@ class MCPJWTSigner(CustomGuardrail): def _build_scope( self, raw_tool_name: str, - call_type: Optional[CallTypesLiteral] = None, + call_type: CallTypesLiteral | None = None, ) -> str: """ Build the JWT scope string. @@ -619,7 +619,7 @@ class MCPJWTSigner(CustomGuardrail): # FR-13: Claim operations # ------------------------------------------------------------------ - def _apply_claim_operations(self, claims: Dict[str, Any]) -> Dict[str, Any]: + def _apply_claim_operations(self, claims: dict[str, Any]) -> dict[str, Any]: """Apply add_claims, set_claims, and remove_claims to the claim dict.""" # add_claims: insert only when key is absent for k, v in self.add_claims.items(): @@ -641,9 +641,9 @@ class MCPJWTSigner(CustomGuardrail): def _passthrough_optional_claims( self, - claims: Dict[str, Any], - jwt_claims: Optional[Dict[str, Any]], - ) -> Dict[str, Any]: + claims: dict[str, Any], + jwt_claims: dict[str, Any] | None, + ) -> dict[str, Any]: """Forward optional_claims from verified incoming token into the outbound JWT.""" if not self.optional_claims or not jwt_claims: return claims @@ -660,9 +660,9 @@ class MCPJWTSigner(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, data: dict, - jwt_claims: Optional[Dict[str, Any]] = None, - call_type: Optional[CallTypesLiteral] = None, - ) -> Dict[str, Any]: + jwt_claims: dict[str, Any] | None = None, + call_type: CallTypesLiteral | None = None, + ) -> dict[str, Any]: """ Build JWT claims for the outbound MCP access token. @@ -673,7 +673,7 @@ class MCPJWTSigner(CustomGuardrail): jwt_claims if available. None for pure API-key requests. """ now = int(time.time()) - claims: Dict[str, Any] = { + claims: dict[str, Any] = { "iss": self.issuer, "aud": self.audience, "iat": now, @@ -714,8 +714,8 @@ class MCPJWTSigner(CustomGuardrail): def _build_channel_token_claims( self, - base_claims: Dict[str, Any], - ) -> Dict[str, Any]: + base_claims: dict[str, Any], + ) -> dict[str, Any]: """ Build claims for the channel token (FR-14 two-token model). @@ -737,7 +737,7 @@ class MCPJWTSigner(CustomGuardrail): # ------------------------------------------------------------------ @staticmethod - def _build_debug_header(claims: Dict[str, Any], kid: str) -> str: + def _build_debug_header(claims: dict[str, Any], kid: str) -> str: """ Build the x-litellm-mcp-debug header value. @@ -763,7 +763,7 @@ class MCPJWTSigner(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Verifies the incoming token (when configured), validates required claims, then signs an outbound JWT and injects it as the Authorization header. @@ -780,8 +780,8 @@ class MCPJWTSigner(CustomGuardrail): # ------------------------------------------------------------------ # FR-5: Verify incoming token before re-signing # ------------------------------------------------------------------ - jwt_claims: Optional[Dict[str, Any]] = None - raw_token: Optional[str] = hook_data.get("incoming_bearer_token") + jwt_claims: dict[str, Any] | None = None + raw_token: str | None = hook_data.get("incoming_bearer_token") if self.access_token_discovery_uri and raw_token: # Three-dot pattern → JWT; otherwise opaque. @@ -835,8 +835,8 @@ class MCPJWTSigner(CustomGuardrail): # Merge into existing extra_headers — a prior guardrail in the chain may # have already injected tracing headers or correlation IDs. - existing_headers: Dict[str, str] = hook_data.get("extra_headers") or {} - new_headers: Dict[str, str] = { + existing_headers: dict[str, str] = hook_data.get("extra_headers") or {} + new_headers: dict[str, str] = { **existing_headers, "Authorization": f"Bearer {signed_token}", } @@ -877,13 +877,13 @@ class MCPJWTSigner(CustomGuardrail): async def inject_mcp_jwt_headers_for_upstream( - user_api_key_dict: Optional[UserAPIKeyAuth], - extra_headers: Optional[Dict[str, str]] = None, - raw_headers: Optional[Dict[str, str]] = None, + user_api_key_dict: UserAPIKeyAuth | None, + extra_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, *, for_list_tools: bool = False, mcp_tool_name: str = "", -) -> Dict[str, str]: +) -> dict[str, str]: """ Sign outbound MCP headers when MCPJWTSigner is configured. @@ -895,12 +895,12 @@ async def inject_mcp_jwt_headers_for_upstream( return merged normalized_raw = {k.lower(): v for k, v in (raw_headers or {}).items()} - incoming_bearer_token: Optional[str] = None + incoming_bearer_token: str | None = None auth_hdr = normalized_raw.get("authorization", "") if auth_hdr.lower().startswith("bearer "): incoming_bearer_token = auth_hdr[len("bearer ") :] - hook_data: Dict[str, Any] = { + hook_data: dict[str, Any] = { "mcp_tool_name": "" if for_list_tools else mcp_tool_name, "incoming_bearer_token": incoming_bearer_token, "extra_headers": merged, diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py index 9b5c221b9c4..4328bb2484e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py @@ -5,7 +5,7 @@ Validates that MCP servers referenced in request tools are registered on the LiteLLM gateway. Blocks or alerts when unregistered servers are found. """ -from typing import Any, List, Literal, Optional, Set, Union +from typing import Any, Literal from fastapi import HTTPException @@ -32,7 +32,7 @@ class MCPSecurityGuardrail(CustomGuardrail): self.on_violation = on_violation @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [GuardrailEventHooks.pre_call] @log_guardrail_information @@ -42,7 +42,7 @@ class MCPSecurityGuardrail(CustomGuardrail): cache: Any, data: dict, call_type: str, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: return data @@ -72,9 +72,9 @@ class MCPSecurityGuardrail(CustomGuardrail): return data @staticmethod - def _extract_mcp_server_names_from_tools(tools: List[dict]) -> Set[str]: + def _extract_mcp_server_names_from_tools(tools: list[dict]) -> set[str]: """Extract MCP server names from tools with type=mcp and litellm_proxy server_url.""" - server_names: Set[str] = set() + server_names: set[str] = set() for tool in tools: if not isinstance(tool, dict): continue @@ -90,7 +90,7 @@ class MCPSecurityGuardrail(CustomGuardrail): return server_names @staticmethod - def _find_unregistered_mcp_servers(data: dict) -> Set[str]: + def _find_unregistered_mcp_servers(data: dict) -> set[str]: """Check tools in data against the MCP server registry. Returns set of unregistered server names.""" tools = data.get("tools") if not tools or not isinstance(tools, list): diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index 830056db0f4..8a984e6eaae 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -2,13 +2,13 @@ import threading import time import uuid from collections import OrderedDict -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger -from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, ) +from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -55,11 +55,11 @@ class PurviewGuardrailBase: self.user_id_field = user_id_field # Token cache: (access_token, expires_at_epoch) - self._token_cache: Optional[Tuple[str, float]] = None + self._token_cache: tuple[str, float] | None = None # Protection scope cache: user_id -> (etag, scope_response, fetched_at) # Capped at 1000 entries (LRU eviction) to avoid unbounded growth. - self._scope_cache: OrderedDict[str, Tuple[str, Dict[str, Any], float]] = OrderedDict() + self._scope_cache: OrderedDict[str, tuple[str, dict[str, Any], float]] = OrderedDict() self._scope_cache_maxsize = 1000 # Use a threading.Lock (not asyncio.Lock) because this lock is acquired # from both the proxy's main asyncio event loop and from short-lived @@ -117,9 +117,9 @@ class PurviewGuardrailBase: async def _graph_post( self, url: str, - json_body: Dict[str, Any], - extra_headers: Optional[Dict[str, str]] = None, - ) -> Tuple[Dict[str, Any], Dict[str, str]]: + json_body: dict[str, Any], + extra_headers: dict[str, str] | None = None, + ) -> tuple[dict[str, Any], dict[str, str]]: """POST to Graph API with bearer auth. Returns: @@ -136,7 +136,7 @@ class PurviewGuardrailBase: verbose_proxy_logger.debug("Purview Graph POST %s", url) response = await self.async_handler.post(url=url, headers=headers, json=json_body) response.raise_for_status() - response_json: Dict[str, Any] = response.json() + response_json: dict[str, Any] = response.json() response_headers = dict(response.headers) verbose_proxy_logger.debug("Purview Graph response: %s", response_json) return response_json, response_headers @@ -145,7 +145,7 @@ class PurviewGuardrailBase: # Protection scopes # ------------------------------------------------------------------ - async def _compute_protection_scopes(self, user_id: str) -> Tuple[str, Dict[str, Any]]: + async def _compute_protection_scopes(self, user_id: str) -> tuple[str, dict[str, Any]]: """Call protectionScopes/compute and cache with ETag. Returns: @@ -161,7 +161,7 @@ class PurviewGuardrailBase: return cached[0], cached[1] url = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/protectionScopes/compute" - body: Dict[str, Any] = { + body: dict[str, Any] = { "activities": "uploadText,downloadText", "locations": [ { @@ -198,8 +198,8 @@ class PurviewGuardrailBase: text: str, activity: str, etag: str, - correlation_id: Optional[str] = None, - ) -> Dict[str, Any]: + correlation_id: str | None = None, + ) -> dict[str, Any]: """Call processContent for DLP policy evaluation. Args: @@ -211,7 +211,7 @@ class PurviewGuardrailBase: """ encoded_user_id = self._encode_graph_user_id(user_id) url = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/processContent" - body: Dict[str, Any] = { + body: dict[str, Any] = { "contentToProcess": { "contentEntries": [ { @@ -244,7 +244,7 @@ class PurviewGuardrailBase: } } - extra_headers: Dict[str, str] = {} + extra_headers: dict[str, str] = {} if etag: extra_headers["If-None-Match"] = etag @@ -261,7 +261,7 @@ class PurviewGuardrailBase: # User ID resolution # ------------------------------------------------------------------ - def _resolve_user_id(self, data: Dict[str, Any], user_api_key_dict: Any) -> Optional[str]: + def _resolve_user_id(self, data: dict[str, Any], user_api_key_dict: Any) -> str | None: """Resolve the Entra user object ID from request data or auth context. Returns the strongest available identity walking down four sources, in @@ -296,7 +296,7 @@ class PurviewGuardrailBase: return None @staticmethod - def _logging_kwargs_metadata(kwargs: Dict[str, Any]) -> Dict[str, Any]: + def _logging_kwargs_metadata(kwargs: dict[str, Any]) -> dict[str, Any]: """Metadata dict from ``model_call_details`` / logging kwargs.""" litellm_params = kwargs.get("litellm_params") or {} if not isinstance(litellm_params, dict): @@ -304,7 +304,7 @@ class PurviewGuardrailBase: md = litellm_params.get("metadata") return md if isinstance(md, dict) else {} - def _resolve_trusted_user_id(self, data: Dict[str, Any], user_api_key_dict: Any) -> Optional[str]: + def _resolve_trusted_user_id(self, data: dict[str, Any], user_api_key_dict: Any) -> str | None: """Resolve user ID from API-key/JWT-bound identity for blocking DLP. Uses only ``UserAPIKeyAuth.user_id`` (bound on the LiteLLM key or JWT). @@ -325,7 +325,7 @@ class PurviewGuardrailBase: return None - def _resolve_user_id_from_logging_kwargs(self, kwargs: Dict[str, Any]) -> Optional[str]: + def _resolve_user_id_from_logging_kwargs(self, kwargs: dict[str, Any]) -> str | None: """Trusted-identity-only resolver for logging-only hooks. Uses only the proxy-injected ``user_api_key_user_id`` (populated from @@ -348,7 +348,7 @@ class PurviewGuardrailBase: # ------------------------------------------------------------------ @staticmethod - def _should_block(response: Dict[str, Any]) -> bool: + def _should_block(response: dict[str, Any]) -> bool: """Return True if any policyAction requires blocking.""" for action in response.get("policyActions", []): odata_type = action.get("@odata.type", "") @@ -383,7 +383,7 @@ class PurviewGuardrailBase: return False @staticmethod - def completion_prompt_to_str(prompt: Any) -> Optional[str]: + def completion_prompt_to_str(prompt: Any) -> str | None: """Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP. Supports string prompts and list-of-string prompts. List-of-token-id prompts @@ -408,7 +408,7 @@ class PurviewGuardrailBase: return None @staticmethod - def _extract_tool_call_args_from_message(message: Any) -> List[str]: + def _extract_tool_call_args_from_message(message: Any) -> list[str]: """Return plaintext arguments strings from tool_calls and function_call fields. Covers both the request path (assistant messages in chat histories that @@ -416,7 +416,7 @@ class PurviewGuardrailBase: tool calls returned in a ModelResponse). Both dict-style and object-style representations are handled. """ - args: List[str] = [] + args: list[str] = [] # tool_calls: [{"function": {"arguments": "..."}}] tool_calls = message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None) @@ -444,7 +444,7 @@ class PurviewGuardrailBase: return args - def get_prompt_text_for_dlp(self, messages: List["AllMessageValues"]) -> Optional[str]: + def get_prompt_text_for_dlp(self, messages: list["AllMessageValues"]) -> str | None: """Concatenate text from every chat message (all roles) for pre-call DLP. Evaluates the same payload the model receives, not only the trailing user @@ -459,9 +459,9 @@ class PurviewGuardrailBase: """ if not messages: return None - parts: List[str] = [] + parts: list[str] = [] for msg in messages: - segments: List[str] = [] + segments: list[str] = [] content = convert_content_list_to_str(message=msg).strip() if content: segments.append(content) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py index db07d00b37b..646c102dbb1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -15,11 +15,6 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Type, Union, cast, ) @@ -92,11 +87,11 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: return None # Config model can be added later for UI support @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -112,9 +107,9 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): user_id: str, text: str, activity: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], block_on_violation: bool = True, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Evaluate content against Purview DLP policies. Args: @@ -129,7 +124,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): """ start_time = datetime.now() status: GuardrailStatus = "success" - response: Dict[str, Any] = {} + response: dict[str, Any] = {} try: etag, _ = await self._compute_protection_scopes(user_id) @@ -158,7 +153,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): if block_on_violation: upstream_status = exc.response.status_code client_status = 502 if upstream_status in (401, 403) else upstream_status - headers: Optional[Dict[str, str]] = None + headers: dict[str, str] | None = None retry_after = exc.response.headers.get("retry-after") if retry_after: headers = {"Retry-After": retry_after} @@ -215,7 +210,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): return response @staticmethod - def _extract_responses_api_function_call_args(result: Any) -> List[str]: + def _extract_responses_api_function_call_args(result: Any) -> list[str]: """Return tool-call argument strings from a ``ResponsesAPIResponse.output``. ``ResponsesAPIResponse.output_text`` only aggregates ``output_text`` @@ -224,7 +219,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): extract them explicitly to keep DLP coverage consistent with the chat (``ModelResponse``) path. """ - args: List[str] = [] + args: list[str] = [] output = getattr(result, "output", None) if not output: return args @@ -240,14 +235,14 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): args.append(arguments) return args - def _completion_response_text_parts(self, result: Any) -> List[str]: + def _completion_response_text_parts(self, result: Any) -> list[str]: """Collect non-empty text segments from chat, text completions, or responses API. Includes assistant message content *and* model-generated tool-call arguments so that sensitive data returned inside function calls is not missed by the DLP scan. """ - parts: List[str] = [] + parts: list[str] = [] if isinstance(result, TextCompletionResponse) and result.choices: for text_choice in result.choices: if not isinstance(text_choice, TextChoices): @@ -276,7 +271,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): parts.extend(self._extract_tool_call_args_from_message(msg)) return parts - def _assemble_responses_api_from_chunks(self, chunks: List[Any]) -> Tuple[bool, Optional[ResponsesAPIResponse]]: + def _assemble_responses_api_from_chunks(self, chunks: list[Any]) -> tuple[bool, ResponsesAPIResponse | None]: """Extract the final ``ResponsesAPIResponse`` from a buffered Responses API stream. Returns a ``(is_responses_api_stream, assembled)`` tuple so the caller @@ -288,7 +283,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ``response.failed`` / ``response.incomplete`` as fallbacks). """ looks_like_responses_api = False - final: Optional[ResponsesAPIResponse] = None + final: ResponsesAPIResponse | None = None for chunk in chunks: event_type = getattr(chunk, "type", None) if isinstance(event_type, str) and event_type.startswith("response."): @@ -298,7 +293,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): final = candidate return looks_like_responses_api, final - def _responses_api_input_to_str(self, data: Dict[str, Any], raise_on_failure: bool = False) -> Optional[str]: + def _responses_api_input_to_str(self, data: dict[str, Any], raise_on_failure: bool = False) -> str | None: """Extract DLP-scannable text from a Responses API request ``input`` field. ``input`` may be a plain string or a list of input items (messages). In @@ -324,7 +319,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): input=input_data if input_data is not None else "", responses_api_request=data, ) - return self.get_prompt_text_for_dlp(cast(List[Any], messages)) + return self.get_prompt_text_for_dlp(cast(list[Any], messages)) except Exception: verbose_proxy_logger.warning( "Purview DLP: failed to transform responses API input", @@ -348,7 +343,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): def _resolve_user_id_for_blocking( self, - data: Dict[str, Any], + data: dict[str, Any], user_api_key_dict: Any, ) -> str: """Resolve user ID for blocking (pre_call / post_call) DLP hooks. @@ -397,13 +392,13 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): self, user_api_key_dict: "UserAPIKeyAuth", cache: Any, - data: Dict[str, Any], + data: dict[str, Any], call_type: "CallTypesLiteral", - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """Check user prompt against Purview DLP policies before LLM call.""" user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict) - prompt_text: Optional[str] = None + prompt_text: str | None = None if call_type in ("responses", "aresponses"): # Route Responses API calls to the responses-specific extractor # before the generic ``messages`` branch. This mirrors @@ -431,9 +426,9 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) prompt_text = self.completion_prompt_to_str(raw_prompt) else: - messages: Optional[List] = data.get("messages") + messages: list | None = data.get("messages") if messages: - prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages)) + prompt_text = self.get_prompt_text_for_dlp(cast(list[Any], messages)) if not prompt_text: return data @@ -503,7 +498,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): user_id = self._resolve_user_id_for_blocking(request_data, user_api_key_dict) # Buffer the entire stream before any DLP scan. - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] async for chunk in response: all_chunks.append(chunk) @@ -602,7 +597,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): # Logging-only hook — audit without blocking # ------------------------------------------------------------------ - def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """Fire-and-forget async audit logging; returns original (kwargs, result) immediately. In the proxy's async success path, litellm independently calls both @@ -650,7 +645,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): return kwargs, result - async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """Send both prompt and response to Purview for audit logging. Errors are logged but never raised — this mode is non-blocking. @@ -664,7 +659,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): # Log prompt (uploadText) try: - prompt_text: Optional[str] = None + prompt_text: str | None = None if call_type in ("responses", "aresponses"): # Responses API: route to the responses-specific extractor # before the generic ``messages`` branch. litellm's logging @@ -680,7 +675,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): else: messages = kwargs.get("messages") if messages: - prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages)) + prompt_text = self.get_prompt_text_for_dlp(cast(list[Any], messages)) if prompt_text: await self._check_content( diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 13081e68e9a..da90961f87a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -2,10 +2,7 @@ from collections.abc import AsyncGenerator, Mapping, Sequence from typing import ( TYPE_CHECKING, Any, - List, Literal, - Optional, - Type, Union, ) @@ -91,7 +88,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): """ @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -102,11 +99,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): def __init__( self, - template_id: Optional[str] = None, - project_id: Optional[str] = None, - location: Optional[str] = None, - credentials: Optional[Any] = None, - api_endpoint: Optional[str] = None, + template_id: str | None = None, + project_id: str | None = None, + location: str | None = None, + credentials: Any | None = None, + api_endpoint: str | None = None, sanitize_error_detail: "bool | None" = True, **kwargs, ): @@ -155,7 +152,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): else: return {"modelResponseData": {"text": content}} - def _extract_content_from_response(self, response: Union[Any, ModelResponse]) -> str: + def _extract_content_from_response(self, response: Any | ModelResponse) -> str: """ Extract text content from model response. @@ -237,11 +234,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): async def make_model_armor_request( self, - content: Optional[str] = None, + content: str | None = None, source: Literal["user_prompt", "model_response"] = "user_prompt", - request_data: Optional[dict] = None, - file_bytes: Optional[bytes] = None, - file_type: Optional[str] = None, + request_data: dict | None = None, + file_bytes: bytes | None = None, + file_type: str | None = None, ) -> dict: """ Make request to Model Armor API. Supports both text and file prompt sanitization. @@ -366,7 +363,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Fallback dict code removed; all cases handled above return False - def _get_sanitized_content(self, armor_response: dict) -> Optional[str]: + def _get_sanitized_content(self, armor_response: dict) -> str | None: """ Get the sanitized content from a Model Armor response, if available. Looks for sanitized text in deidentifyResult, and falls back to root-level fields if not found. @@ -422,13 +419,13 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): def _process_response( self, - response: Optional[dict], + response: dict | None, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, - original_inputs: Optional[dict] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, + original_inputs: dict | None = None, ): """ Override to store only the Model Armor API response, not the entire data dict. @@ -567,7 +564,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """Pre-call hook to sanitize user prompts.""" verbose_proxy_logger.debug("Inside Model Armor Pre-Call Hook") @@ -666,7 +663,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): data: dict, user_api_key_dict: UserAPIKeyAuth, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """During-call hook to sanitize user prompts in parallel with LLM call.""" verbose_proxy_logger.debug("Inside Model Armor Moderation Hook") @@ -851,7 +848,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): from litellm.main import stream_chunk_builder # Collect all chunks - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] async for chunk in response: all_chunks.append(chunk) @@ -945,7 +942,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): yield chunk @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Get the config model for the Model Armor guardrail. """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py index 77c9f1cad65..642c1dcbca8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py @@ -14,12 +14,8 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, Final, - List, Literal, - Optional, - Type, Union, ) from urllib.parse import urljoin @@ -52,8 +48,8 @@ from litellm.types.utils import ( ) # Constants -USER_ROLE: Final[Literal["user"]] = "user" -ASSISTANT_ROLE: Final[Literal["assistant"]] = "assistant" +USER_ROLE: Final = "user" +ASSISTANT_ROLE: Final = "assistant" SENSITIVE_DATA_DETECTOR_KEYS: Final[list[str]] = ["sensitiveData", "dataDetector"] # Type aliases @@ -77,7 +73,7 @@ class NomaBlockedMessage(HTTPException): }, ) - def _is_result_true(self, result_obj: Optional[Dict[str, Any]]) -> bool: + def _is_result_true(self, result_obj: dict[str, Any] | None) -> bool: """ Check if a result object has a "result" field that is True. @@ -105,7 +101,7 @@ class NomaGuardrail(CustomGuardrail): _AIDR_ENDPOINT = "/ai-dr/v2/prompt/scan" @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -115,12 +111,12 @@ class NomaGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - application_id: Optional[str] = None, - monitor_mode: Optional[bool] = None, - block_failures: Optional[bool] = None, - anonymize_input: Optional[bool] = None, + api_key: str | None = None, + api_base: str | None = None, + application_id: str | None = None, + monitor_mode: bool | None = None, + block_failures: bool | None = None, + anonymize_input: bool | None = None, **kwargs, ): global _LEGACY_NOMA_DEPRECATION_WARNED @@ -167,14 +163,14 @@ class NomaGuardrail(CustomGuardrail): try: asyncio.create_task(coro) except Exception as e: - verbose_proxy_logger.error(f"Failed to create background Noma task: {str(e)}") + verbose_proxy_logger.error(f"Failed to create background Noma task: {e!s}") async def _process_user_message_check( self, request_data: dict, user_auth: UserAPIKeyAuth, - event_type: Optional[GuardrailEventHooks] = None, - ) -> Optional[str]: + event_type: GuardrailEventHooks | None = None, + ) -> str | None: """Shared logic for processing user message checks""" start_time = datetime.now() extra_data = self.get_guardrail_dynamic_request_body_params(request_data) @@ -248,8 +244,8 @@ class NomaGuardrail(CustomGuardrail): request_data: dict, response: LLMResponse, user_auth: UserAPIKeyAuth, - event_type: Optional[GuardrailEventHooks] = None, - ) -> Optional[str]: + event_type: GuardrailEventHooks | None = None, + ) -> str | None: """Shared logic for processing LLM response checks""" start_time = datetime.now() @@ -352,7 +348,7 @@ class NomaGuardrail(CustomGuardrail): return "guardrail_failed_to_respond" except Exception as e: - verbose_proxy_logger.error(f"Error determining NOMA guardrail status: {str(e)}") + verbose_proxy_logger.error(f"Error determining NOMA guardrail status: {e!s}") return "guardrail_failed_to_respond" def _should_only_sensitive_data_failed(self, classification_obj: dict) -> bool: @@ -394,7 +390,7 @@ class NomaGuardrail(CustomGuardrail): # Return True only if sensitive data was detected AND no other detectors have result=true return sensitive_data_detected and len(failed_detectors) == 0 - def _extract_anonymized_content(self, response_json: dict, message_type: MessageRole) -> Optional[str]: + def _extract_anonymized_content(self, response_json: dict, message_type: MessageRole) -> str | None: """ Extract anonymized content from Noma API response. @@ -459,7 +455,7 @@ class NomaGuardrail(CustomGuardrail): return False - def _is_result_true(self, result_obj: Optional[Dict[str, Any]]) -> bool: + def _is_result_true(self, result_obj: dict[str, Any] | None) -> bool: """ Check if a result object has a "result" field that is True. @@ -517,7 +513,7 @@ class NomaGuardrail(CustomGuardrail): try: await self._process_user_message_check(request_data, user_auth) except Exception as e: - verbose_proxy_logger.error(f"Noma background user message check failed: {str(e)}") + verbose_proxy_logger.error(f"Noma background user message check failed: {e!s}") async def _check_llm_response_background( self, @@ -529,7 +525,7 @@ class NomaGuardrail(CustomGuardrail): try: await self._process_llm_response_check(request_data, response, user_auth) except Exception as e: - verbose_proxy_logger.error(f"Noma background response check failed: {str(e)}") + verbose_proxy_logger.error(f"Noma background response check failed: {e!s}") async def _handle_verdict_background( self, @@ -551,7 +547,7 @@ class NomaGuardrail(CustomGuardrail): msg = f"Noma guardrail allowed {type} message: {message}" verbose_proxy_logger.info(msg) except Exception as e: - verbose_proxy_logger.error(f"Noma background verdict handling failed: {str(e)}") + verbose_proxy_logger.error(f"Noma background verdict handling failed: {e!s}") async def async_pre_call_hook( self, @@ -559,7 +555,7 @@ class NomaGuardrail(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: verbose_proxy_logger.debug("Running Noma pre-call hook") event_type = GuardrailEventHooks.pre_call @@ -574,7 +570,7 @@ class NomaGuardrail(CustomGuardrail): try: self._create_background_noma_check(self._check_user_message_background(data, user_api_key_dict)) except Exception as e: - verbose_proxy_logger.error(f"Failed to start background Noma pre-call check: {str(e)}") + verbose_proxy_logger.error(f"Failed to start background Noma pre-call check: {e!s}") return data try: @@ -598,7 +594,7 @@ class NomaGuardrail(CustomGuardrail): event_type=GuardrailEventHooks.pre_call, ) - verbose_proxy_logger.error(f"Noma pre-call hook failed: {str(e)}") + verbose_proxy_logger.error(f"Noma pre-call hook failed: {e!s}") if self.block_failures: raise @@ -609,7 +605,7 @@ class NomaGuardrail(CustomGuardrail): data: dict, user_api_key_dict: UserAPIKeyAuth, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: event_type: GuardrailEventHooks = GuardrailEventHooks.during_call if call_type == CallTypes.call_mcp_tool.value: event_type = GuardrailEventHooks.pre_mcp_call @@ -622,7 +618,7 @@ class NomaGuardrail(CustomGuardrail): try: self._create_background_noma_check(self._check_user_message_background(data, user_api_key_dict)) except Exception as e: - verbose_proxy_logger.error(f"Failed to start background Noma moderation check: {str(e)}") + verbose_proxy_logger.error(f"Failed to start background Noma moderation check: {e!s}") return data try: @@ -646,7 +642,7 @@ class NomaGuardrail(CustomGuardrail): event_type=GuardrailEventHooks.during_call, ) - verbose_proxy_logger.error(f"Noma moderation hook failed: {str(e)}") + verbose_proxy_logger.error(f"Noma moderation hook failed: {e!s}") if self.block_failures: raise @@ -669,7 +665,7 @@ class NomaGuardrail(CustomGuardrail): self._check_llm_response_background(data, response, user_api_key_dict) ) except Exception as e: - verbose_proxy_logger.error(f"Failed to start background Noma post-call check: {str(e)}") + verbose_proxy_logger.error(f"Failed to start background Noma post-call check: {e!s}") return response try: @@ -693,7 +689,7 @@ class NomaGuardrail(CustomGuardrail): event_type=GuardrailEventHooks.post_call, ) - verbose_proxy_logger.error(f"Noma post-call hook failed: {str(e)}") + verbose_proxy_logger.error(f"Noma post-call hook failed: {e!s}") if self.block_failures: raise return response @@ -702,8 +698,8 @@ class NomaGuardrail(CustomGuardrail): self, request_data: dict, user_auth: UserAPIKeyAuth, - event_type: Optional[GuardrailEventHooks] = None, - ) -> Union[Exception, str, dict, None]: + event_type: GuardrailEventHooks | None = None, + ) -> Exception | str | dict | None: """Check user message for policy violations""" user_message = await self._process_user_message_check(request_data, user_auth, event_type) if not user_message: @@ -716,7 +712,7 @@ class NomaGuardrail(CustomGuardrail): request_data: dict, response: LLMResponse, user_auth: UserAPIKeyAuth, - event_type: Optional[GuardrailEventHooks] = None, + event_type: GuardrailEventHooks | None = None, ) -> Any: """Check LLM response for policy violations""" content = await self._process_llm_response_check(request_data, response, user_auth, event_type) @@ -728,7 +724,7 @@ class NomaGuardrail(CustomGuardrail): async def _call_noma_api( self, payload: dict, - llm_request_id: Optional[str], + llm_request_id: str | None, request_data: dict, user_auth: UserAPIKeyAuth, extra_data: dict, @@ -793,7 +789,7 @@ class NomaGuardrail(CustomGuardrail): verbose_proxy_logger.debug(msg) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.noma import ( NomaGuardrailConfigModel, ) @@ -808,14 +804,14 @@ class NomaGuardrail(CustomGuardrail): ) -> AsyncGenerator[ModelResponseStream, None]: """Process streaming response chunks with Noma guardrail.""" - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] async for chunk in response: all_chunks.append(chunk) if not all_chunks: return - assembled_model_response: Optional[Union[ModelResponse, TextCompletionResponse]] = stream_chunk_builder( + assembled_model_response: ModelResponse | TextCompletionResponse | None = stream_chunk_builder( chunks=all_chunks ) @@ -832,7 +828,7 @@ class NomaGuardrail(CustomGuardrail): except Exception as e: if self.block_failures: raise - verbose_proxy_logger.error(f"Noma streaming post-call hook failed: {str(e)}") + verbose_proxy_logger.error(f"Noma streaming post-call hook failed: {e!s}") for chunk in all_chunks: yield chunk return diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py index 9cf0986c122..e0dd9bca6ed 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py @@ -8,7 +8,7 @@ import enum import json import os from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, cast from urllib.parse import urlparse from litellm._logging import verbose_proxy_logger @@ -46,11 +46,11 @@ class _Action(str, enum.Enum): class NomaV2Guardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - application_id: Optional[str] = None, - monitor_mode: Optional[bool] = None, - block_failures: Optional[bool] = None, + api_key: str | None = None, + api_base: str | None = None, + application_id: str | None = None, + monitor_mode: bool | None = None, + block_failures: bool | None = None, **kwargs: Any, ) -> None: self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -76,7 +76,7 @@ class NomaV2Guardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.noma import ( NomaV2GuardrailConfigModel, ) @@ -84,7 +84,7 @@ class NomaV2Guardrail(CustomGuardrail): return NomaV2GuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -104,7 +104,7 @@ class NomaV2Guardrail(CustomGuardrail): return parsed.hostname == _DEFAULT_API_BASE_HOSTNAME @staticmethod - def _get_non_empty_str(value: Any) -> Optional[str]: + def _get_non_empty_str(value: Any) -> str | None: if not isinstance(value, str): return None stripped = value.strip() @@ -129,7 +129,7 @@ class NomaV2Guardrail(CustomGuardrail): request_data: dict, input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], - application_id: Optional[str], + application_id: str | None, ) -> dict: payload_request_data = self._sanitize_payload_for_transport(request_data) if logging_obj is not None: @@ -256,7 +256,7 @@ class NomaV2Guardrail(CustomGuardrail): dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data) if not isinstance(dynamic_params, dict): dynamic_params = {} - response_json: Optional[dict] = None + response_json: dict | None = None # Per-request dynamic params can override configured application context. application_id = self._get_non_empty_str(dynamic_params.get("application_id")) diff --git a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py index 606f0a4587d..b86273f754a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py +++ b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py @@ -6,7 +6,7 @@ # +-------------------------------------------------------------+ import os import uuid -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal, Optional import httpx from fastapi import HTTPException @@ -30,7 +30,7 @@ if TYPE_CHECKING: class OnyxGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -39,9 +39,9 @@ class OnyxGuardrail(CustomGuardrail): def __init__( self, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Optional[float] = 10.0, + api_base: str | None = None, + api_key: str | None = None, + timeout: float | None = 10.0, **kwargs, ): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -118,7 +118,7 @@ class OnyxGuardrail(CustomGuardrail): payload = parsed.get("response", {}) except Exception as e: verbose_proxy_logger.error( - f"Error in converting request_data to ModelResponse: {str(e)}", + f"Error in converting request_data to ModelResponse: {e!s}", extra={ "conversation_id": conversation_id, "input_type": input_type, @@ -133,13 +133,13 @@ class OnyxGuardrail(CustomGuardrail): raise e except Exception as e: verbose_proxy_logger.error( - f"Error in apply_guardrail guard: {str(e)}", + f"Error in apply_guardrail guard: {e!s}", extra={"conversation_id": conversation_id, "input_type": input_type}, ) return inputs @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.onyx import ( OnyxGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/base.py b/litellm/proxy/guardrails/guardrail_hooks/openai/base.py index 281afacd5c4..908bdfbd78a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/base.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_last_user_message, @@ -13,7 +13,7 @@ class OpenAIGuardrailBase: Base class for OpenAI guardrails. """ - def get_user_prompt(self, messages: List["AllMessageValues"]) -> Optional[str]: + def get_user_prompt(self, messages: list["AllMessageValues"]) -> str | None: """ Get the last consecutive block of messages from the user. diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index 44016bfd3dc..5ff04864e75 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -5,12 +5,8 @@ OpenAI Moderation Guardrail Integration for LiteLLM from typing import ( TYPE_CHECKING, - Dict, - List, Literal, Optional, - Type, - Union, ) from fastapi import HTTPException @@ -57,11 +53,11 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): def __init__( self, guardrail_name: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - model: Optional[Literal["omni-moderation-latest", "text-moderation-latest"]] = None, - streaming_end_of_stream_only: Optional[bool] = None, - streaming_sampling_rate: Optional[int] = None, + api_key: str | None = None, + api_base: str | None = None, + model: Literal["omni-moderation-latest", "text-moderation-latest"] | None = None, + streaming_end_of_stream_only: bool | None = None, + streaming_sampling_rate: int | None = None, **kwargs, ): """Initialize OpenAI Moderation guardrail handler.""" @@ -94,7 +90,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): f"Initialized OpenAI Moderation Guardrail: {guardrail_name} with model: {self.model}" ) - def _get_api_key(self) -> Optional[str]: + def _get_api_key(self) -> str | None: """Get API key from environment variables or litellm configuration""" import os @@ -201,7 +197,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): HTTPException: If content violates moderation policy """ # Extract text to moderate from inputs - text_to_moderate: Optional[str] = None + text_to_moderate: str | None = None # Prefer structured_messages if available (has role context) if structured_messages := inputs.get("structured_messages"): @@ -235,13 +231,13 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): def _process_response( self, - response: Optional[Dict], + response: dict | None, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, - original_inputs: Optional[Dict] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, + original_inputs: dict | None = None, ): """ Override to log the full OpenAI Moderation API response instead of @@ -276,10 +272,10 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): self, e: Exception, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, ): """ Override to log the full OpenAI Moderation API response on error @@ -296,7 +292,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): metadata = {} # Use the stashed moderation response if available, fall back to exception - guardrail_response: Union[dict, Exception, str] = metadata.pop("_openai_moderation_response", e) + guardrail_response: dict | Exception | str = metadata.pop("_openai_moderation_response", e) self.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=guardrail_response, @@ -312,8 +308,8 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): @staticmethod def _build_tracing_detail( - guardrail_response: Union[dict, str, Exception], - ) -> Optional[GuardrailTracingDetail]: + guardrail_response: dict | str | Exception, + ) -> GuardrailTracingDetail | None: """ Pull the flagged category names out of the moderation response so trace backends can index a short, queryable ``guardrail_violation_categories`` @@ -337,7 +333,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): return GuardrailTracingDetail(violation_categories=violation_categories) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Get the config model for the OpenAI Moderation guardrail. """ @@ -348,7 +344,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): return OpenAIModerationGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index e92ac37ca77..947a81c1b79 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -7,7 +7,7 @@ post_call (model output) checkpoints with optional correction/blocking. import datetime import hashlib import os -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal import httpx @@ -35,8 +35,6 @@ BLOCKED_ACTION_TYPE = "block" class OvalixGuardrailMissingSecrets(Exception): """Raised when required Ovalix config (API base, key, application/checkpoint IDs) is missing.""" - pass - class OvalixGuardrailBlockedException(GuardrailRaisedException): """ @@ -48,7 +46,7 @@ class OvalixGuardrailBlockedException(GuardrailRaisedException): def __init__( self, - guardrail_name: Optional[str] = None, + guardrail_name: str | None = None, message: str = "", should_wrap_with_default_message: bool = True, ): @@ -67,7 +65,7 @@ class OvalixGuardrail(CustomGuardrail): """ @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -75,11 +73,11 @@ class OvalixGuardrail(CustomGuardrail): def __init__( self, - tracker_api_base: Optional[str] = None, - tracker_api_key: Optional[str] = None, - application_id: Optional[str] = None, - pre_checkpoint_id: Optional[str] = None, - post_checkpoint_id: Optional[str] = None, + tracker_api_base: str | None = None, + tracker_api_key: str | None = None, + application_id: str | None = None, + pre_checkpoint_id: str | None = None, + post_checkpoint_id: str | None = None, **kwargs: Any, ): self._tracker_api_base = tracker_api_base or os.environ.get("OVALIX_TRACKER_API_BASE") @@ -112,9 +110,9 @@ class OvalixGuardrail(CustomGuardrail): self._post_checkpoint_id, ) - def _validate_config(self, supported_event_hooks: List[GuardrailEventHooks]) -> None: + def _validate_config(self, supported_event_hooks: list[GuardrailEventHooks]) -> None: """Ensure required secrets and checkpoint IDs are set; auto-add hooks when IDs are present.""" - errors: List[str] = [] + errors: list[str] = [] if not self._tracker_api_base: errors.append("Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base") @@ -171,7 +169,7 @@ class OvalixGuardrail(CustomGuardrail): checkpoint_id: str, actor: str, session_id: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Call the Ovalix Tracker checkpoint API and return the JSON response.""" application_id = self._application_id if not application_id or not checkpoint_id: @@ -197,7 +195,7 @@ class OvalixGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"], - logging_obj: Optional[Any] = None, + logging_obj: Any | None = None, ) -> GenericGuardrailAPIInputs: """ Apply Ovalix guardrail to the given inputs (request or response text). @@ -239,10 +237,10 @@ class OvalixGuardrail(CustomGuardrail): return inputs async def _generate_post_guardrail_llm_texts( - self, texts: List[str], actor: str, session_id: str, checkpoint_id: str - ) -> List[str]: + self, texts: list[str], actor: str, session_id: str, checkpoint_id: str + ) -> list[str]: """Generate post-guardrail LLM responses for the given LLM responses.""" - post_guardrail_texts: List[str] = [] + post_guardrail_texts: list[str] = [] is_first_response = True for llm_response in reversed(texts): @@ -276,7 +274,7 @@ class OvalixGuardrail(CustomGuardrail): should_wrap_with_default_message=False, ) - def _get_trackers_corrected_message(self, resp: dict) -> Optional[str]: + def _get_trackers_corrected_message(self, resp: dict) -> str | None: """Extract corrected/blocking message content from Tracker checkpoint response.""" modified = resp.get("modified_data") if isinstance(modified, dict) and "content" in modified: @@ -284,7 +282,7 @@ class OvalixGuardrail(CustomGuardrail): return None @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( OvalixGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py index d02d4b448e8..5c707153873 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py @@ -1,6 +1,6 @@ # litellm/proxy/guardrails/guardrail_hooks/pangea.py import os -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, cast +from typing import TYPE_CHECKING, Any, cast from fastapi import HTTPException @@ -33,8 +33,6 @@ if TYPE_CHECKING: class PangeaGuardrailMissingSecrets(Exception): """Custom exception for missing Pangea secrets.""" - pass - class _TextCompletionRequest: def __init__(self, body): @@ -61,10 +59,10 @@ class PangeaHandler(CustomGuardrail): def __init__( self, guardrail_name: str, - pangea_input_recipe: Optional[str] = None, - pangea_output_recipe: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + pangea_input_recipe: str | None = None, + pangea_output_recipe: str | None = None, + api_key: str | None = None, + api_base: str | None = None, **kwargs, ): """ @@ -229,7 +227,7 @@ class PangeaHandler(CustomGuardrail): messages = data.get("messages") if messages is None: return # No messages to check - input_messages = cast(List[Dict[Any, Any]], messages) + input_messages = cast(list[dict[Any, Any]], messages) else: return @@ -306,7 +304,7 @@ class PangeaHandler(CustomGuardrail): ) from e @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.pangea import ( PangeaGuardrailConfigModel, ) @@ -314,7 +312,7 @@ class PangeaHandler(CustomGuardrail): return PangeaGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 192c8c9bc77..5d134fd01c2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -9,17 +9,15 @@ import json import os import re from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Type +from typing import TYPE_CHECKING, Any, Literal, Optional from urllib.parse import urlparse import httpx - -from litellm._uuid import uuid -from litellm.caching import DualCache - from fastapi import HTTPException from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid +from litellm.caching import DualCache from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, @@ -29,10 +27,10 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth -from litellm.types.guardrails import GuardrailEventHooks from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ( CallTypes, CallTypesLiteral, @@ -70,17 +68,17 @@ class PanwPrismaAirsHandler(CustomGuardrail): def __init__( self, guardrail_name: str, - profile_name: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + profile_name: str | None = None, + api_key: str | None = None, + api_base: str | None = None, default_on: bool = True, mask_on_block: bool = False, mask_request_content: bool = False, mask_response_content: bool = False, - app_name: Optional[str] = None, + app_name: str | None = None, fallback_on_error: Literal["block", "allow"] = "block", timeout: float = 10.0, - violation_message_template: Optional[str] = None, + violation_message_template: str | None = None, **kwargs, ): """Initialize PANW Prisma AIRS guardrail handler.""" @@ -139,7 +137,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): self.timeout = float(timeout) if timeout is not None else 10.0 # Tri-state: None = not set (default-on for Anthropic), True = explicit on, False = explicit off - self.experimental_use_latest_role_message_only: Optional[bool] = kwargs.get( + self.experimental_use_latest_role_message_only: bool | None = kwargs.get( "experimental_use_latest_role_message_only" ) @@ -173,7 +171,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return True return False - def _extract_text_from_messages(self, messages: List[Dict[str, Any]]) -> str: + def _extract_text_from_messages(self, messages: list[dict[str, Any]]) -> str: """Extract text content from messages array.""" if not isinstance(messages, list) or not messages: return "" @@ -195,7 +193,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return "" - def _extract_text_from_content_list(self, content_list: List[Dict[str, Any]]) -> str: + def _extract_text_from_content_list(self, content_list: list[dict[str, Any]]) -> str: """Extract text from content list format.""" text_parts = [ part.get("text", "") @@ -233,17 +231,17 @@ class PanwPrismaAirsHandler(CustomGuardrail): return " ".join(text_parts) if text_parts else "" except (AttributeError, IndexError) as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS: Error extracting response text: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS: Error extracting response text: {e!s}") return "" async def _call_panw_api( self, content: str = "", is_response: bool = False, - metadata: Optional[Dict[str, Any]] = None, - call_id: Optional[str] = None, - tool_event: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + metadata: dict[str, Any] | None = None, + call_id: str | None = None, + tool_event: dict[str, Any] | None = None, + ) -> dict[str, Any]: """Call PANW Prisma AIRS API to scan content or a tool_event.""" if tool_event is None and not content.strip(): @@ -293,7 +291,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): panw_metadata["litellm_trace_id"] = metadata["litellm_trace_id"] # Build contents: tool_event takes priority, else prompt/response text - contents: List[Dict[str, Any]] + contents: list[dict[str, Any]] if tool_event is not None: contents = [{"tool_event": tool_event}] else: @@ -435,7 +433,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): } except httpx.TimeoutException as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS: Timeout error: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS: Timeout error: {e!s}") return { "action": "block", "category": "timeout_error", @@ -443,7 +441,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): } except httpx.RequestError as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS: Network/request error: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS: Network/request error: {e!s}") return { "action": "block", "category": "network_error", @@ -451,7 +449,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): } except Exception as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS: Unexpected error: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS: Unexpected error: {e!s}") return {"action": "block", "category": "api_error", "_is_transient": True} @staticmethod @@ -487,7 +485,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): ) return "unknown" - def _get_masked_text(self, scan_result: Dict[str, Any], is_response: bool = False) -> Optional[str]: + def _get_masked_text(self, scan_result: dict[str, Any], is_response: bool = False) -> str | None: """Extract masked text from PANW scan result.""" masked_key = "response_masked_data" if is_response else "prompt_masked_data" masked_data = scan_result.get(masked_key) @@ -496,7 +494,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return None @staticmethod - def _mask_content_list(content_list: List, masked_text: str) -> List: + def _mask_content_list(content_list: list, masked_text: str) -> list: """Replace text parts in a content list, preserving non-text parts (images, etc.).""" new_content = [] for part in content_list: @@ -570,7 +568,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): else: verbose_proxy_logger.info("PANW Prisma AIRS: MCP request allowed with PII masking applied") - def _apply_masking_to_messages(self, messages: List[Dict[str, Any]], masked_text: str) -> List[Dict[str, Any]]: + def _apply_masking_to_messages(self, messages: list[dict[str, Any]], masked_text: str) -> list[dict[str, Any]]: """Apply masked text to the last user message.""" if not messages: return messages @@ -622,7 +620,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): if hasattr(choice.message.function_call, "arguments"): choice.message.function_call.arguments = masked_text - def _build_error_detail(self, scan_result: Dict[str, Any], is_response: bool = False) -> Dict[str, Any]: + def _build_error_detail(self, scan_result: dict[str, Any], is_response: bool = False) -> dict[str, Any]: """Build enhanced error detail with scan information.""" action_type = "Response" if is_response else "Prompt" code_suffix = "_response_blocked" if is_response else "_blocked" @@ -672,12 +670,12 @@ class PanwPrismaAirsHandler(CustomGuardrail): def _handle_api_error_with_logging( self, - scan_result: Dict[str, Any], - data: Dict[str, Any], + scan_result: dict[str, Any], + data: dict[str, Any], start_time: datetime, event_type: GuardrailEventHooks, is_response: bool = False, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """Handle API errors with fail-open/fail-closed logic.""" end_time = datetime.now() duration = (end_time - start_time).total_seconds() @@ -736,7 +734,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): }, ) - def _prepare_metadata_from_request(self, data: Dict[str, Any]) -> Dict[str, Any]: + def _prepare_metadata_from_request(self, data: dict[str, Any]) -> dict[str, Any]: """ Extract and prepare metadata from request data for PANW API call. @@ -782,9 +780,9 @@ class PanwPrismaAirsHandler(CustomGuardrail): return metadata @staticmethod - def _extract_text_from_sse_bytes(chunks: List[bytes]) -> str: + def _extract_text_from_sse_bytes(chunks: list[bytes]) -> str: """Extract text from Anthropic SSE byte chunks (content_block_delta → text_delta).""" - texts: List[str] = [] + texts: list[str] = [] raw = b"".join(chunks).decode("utf-8", errors="replace") for line in raw.split("\n"): line = line.strip() @@ -812,7 +810,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): val = c.get(key) return val - parts: List[str] = [] + parts: list[str] = [] for chunk in chunks: if _attr(chunk, "type") == "response.output_text.delta": delta = _attr(chunk, "delta") @@ -957,9 +955,9 @@ class PanwPrismaAirsHandler(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, - data: Dict[str, Any], + data: dict[str, Any], call_type: CallTypesLiteral, - ) -> Optional[Dict[str, Any]]: + ) -> dict[str, Any] | None: """ Pre-call hook to scan user prompts before sending to LLM. @@ -1058,7 +1056,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {e!s}") raise HTTPException( status_code=500, detail={ @@ -1074,7 +1072,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): @log_guardrail_information async def async_post_call_success_hook( self, - data: Dict[str, Any], + data: dict[str, Any], user_api_key_dict: UserAPIKeyAuth, response: Any, ) -> Any: @@ -1172,7 +1170,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {e!s}") raise HTTPException( status_code=500, detail={ @@ -1190,7 +1188,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): assembled_model_response: ModelResponse, request_data: dict, start_time: datetime, - ) -> Tuple[bool, ModelResponse, Dict[str, Any]]: + ) -> tuple[bool, ModelResponse, dict[str, Any]]: """ Scan assembled streaming response and apply masking if needed. Returns (content_was_modified, response, scan_result). @@ -1364,18 +1362,18 @@ class PanwPrismaAirsHandler(CustomGuardrail): # returns a proper JSON error response with the correct status code. # (Raising from a generator hits create_response's generic except → 500.) detail = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} - error_obj: Dict[str, Any] = dict(detail.get("error", detail)) # type: ignore[arg-type] + error_obj: dict[str, Any] = dict(detail.get("error", detail)) # type: ignore[arg-type] error_obj["code"] = e.status_code yield f"data: {json.dumps({'error': error_obj})}\n\n" except Exception as e: - verbose_proxy_logger.error(f"PANW Prisma AIRS streaming error: {str(e)}") + verbose_proxy_logger.error(f"PANW Prisma AIRS streaming error: {e!s}") yield f"data: {json.dumps({'error': {'message': 'Security scan failed - streaming response blocked for safety', 'type': 'guardrail_scan_error', 'code': 500, 'guardrail': self.guardrail_name}})}\n\n" async def _scan_tool_calls_for_guardrail( self, tool_calls: list, is_response: bool, - metadata: Dict[str, Any], + metadata: dict[str, Any], call_id: str, request_data: dict, start_time: datetime, @@ -1400,8 +1398,8 @@ class PanwPrismaAirsHandler(CustomGuardrail): """ for tool_call in tool_calls: # --- extract tool_name and args_text -------------------------- - tool_name: Optional[str] = None - args_text: Optional[str] = None + tool_name: str | None = None + args_text: str | None = None if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"): args_text = tool_call.function.arguments @@ -1413,7 +1411,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): tool_name = func.get("name") # --- build tool_event payload (canonical PANW schema) ----------- - tool_event: Dict[str, Any] = { + tool_event: dict[str, Any] = { "metadata": { "ecosystem": "openai", "method": "tools/call", @@ -1512,9 +1510,9 @@ class PanwPrismaAirsHandler(CustomGuardrail): @staticmethod def _get_latest_user_text_indices( - texts: List[str], + texts: list[str], messages: list, - ) -> Optional[set]: + ) -> set | None: """Return text indices belonging to only the latest scannable human-authored (user or developer) message. Args: @@ -1525,7 +1523,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): Returns a set of scannable indices, or None on count mismatch or no user/developer message (safety fallback to existing role-filter behavior). """ - last_human_msg_idx: Optional[int] = None + last_human_msg_idx: int | None = None for idx in range(len(messages) - 1, -1, -1): msg = messages[idx] if isinstance(msg, dict) and msg.get("role") in ("user", "developer"): @@ -1563,9 +1561,9 @@ class PanwPrismaAirsHandler(CustomGuardrail): @staticmethod def _get_scannable_text_indices( - texts: List[str], + texts: list[str], structured_messages: list, - ) -> Optional[set]: + ) -> set | None: """Derive which ``texts`` indices originate from user/system messages. The unified guardrail framework flattens message content into ``texts`` @@ -1609,7 +1607,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return scannable @staticmethod - def _mcp_name_fallback(rd: dict) -> Optional[str]: + def _mcp_name_fallback(rd: dict) -> str | None: """Return rd['name'] only when 'arguments' or 'mcp_arguments' co-occurs (MCP shape). A bare 'name' key without 'arguments' is NOT an MCP request — it's a @@ -1695,11 +1693,11 @@ class PanwPrismaAirsHandler(CustomGuardrail): metadata = self._prepare_metadata_from_request(request_data) start_time = datetime.now() - new_texts: List[str] = [] + new_texts: list[str] = [] # On request side, determine which text indices correspond to scannable # messages so we can skip scanning assistant/tool history text. - scannable_indices: Optional[set] = None + scannable_indices: set | None = None if input_type == "request": structured_messages = inputs.get("structured_messages") if structured_messages: @@ -1783,7 +1781,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): # "mcp_tool_name"/"mcp_arguments". Check canonical first, then fallback. mcp_tool_name = request_data.get("mcp_tool_name") or self._mcp_name_fallback(request_data) if mcp_tool_name and input_type == "request": - mcp_tool_event: Dict[str, Any] = { + mcp_tool_event: dict[str, Any] = { "metadata": { "ecosystem": "mcp", "method": "tools/call", @@ -1841,7 +1839,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return inputs @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( PanwPrismaAirsGuardrailConfigModel, ) @@ -1849,7 +1847,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return PanwPrismaAirsGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py index 1d884352eec..f96fca11abc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py @@ -8,8 +8,8 @@ # Standard library imports import json import os +from typing import TYPE_CHECKING, Any, Literal from urllib.parse import quote -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Type, Union # Third-party imports from fastapi import HTTPException @@ -50,7 +50,7 @@ def _encode_json_for_header(data: Any) -> str: return quote(json_payload, safe="") -def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER_VALUE_BYTES) -> Tuple[Any, str, bool]: +def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER_VALUE_BYTES) -> tuple[Any, str, bool]: """ Truncate evidence payload so the encoded header value stays within max_bytes. @@ -66,7 +66,7 @@ def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER truncated_value = "[truncated]" return truncated_value, _encode_json_for_header(truncated_value), True - truncated: List[Any] = [] + truncated: list[Any] = [] encoded = _encode_json_for_header(truncated) truncated_flag = False @@ -105,11 +105,11 @@ def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER return truncated, encoded, truncated_flag -def build_pillar_response_headers(metadata_store: Dict[str, Any]) -> Dict[str, str]: +def build_pillar_response_headers(metadata_store: dict[str, Any]) -> dict[str, str]: """ Create URL-safe Pillar response headers and apply truncation metadata. """ - headers: Dict[str, str] = {} + headers: dict[str, str] = {} if "pillar_flagged" in metadata_store: headers["x-pillar-flagged"] = str(metadata_store["pillar_flagged"]).lower() @@ -140,14 +140,10 @@ def build_pillar_response_headers(metadata_store: Dict[str, Any]) -> Dict[str, s class PillarGuardrailMissingSecrets(Exception): """Exception raised when Pillar API key is missing.""" - pass - class PillarGuardrailAPIError(Exception): """Exception raised when there's an error calling the Pillar API.""" - pass - # Main guardrail class class PillarGuardrail(CustomGuardrail): @@ -167,16 +163,16 @@ class PillarGuardrail(CustomGuardrail): def __init__( self, - guardrail_name: Optional[str] = "pillar-security", - api_key: Optional[str] = None, - api_base: Optional[str] = None, - on_flagged_action: Optional[str] = None, - async_mode: Optional[bool] = None, - persist_session: Optional[bool] = None, - include_scanners: Optional[bool] = None, - include_evidence: Optional[bool] = None, - fallback_on_error: Optional[str] = None, - timeout: Optional[float] = None, + guardrail_name: str | None = "pillar-security", + api_key: str | None = None, + api_base: str | None = None, + on_flagged_action: str | None = None, + async_mode: bool | None = None, + persist_session: bool | None = None, + include_scanners: bool | None = None, + include_evidence: bool | None = None, + fallback_on_error: str | None = None, + timeout: float | None = None, **kwargs, ) -> None: """ @@ -297,7 +293,7 @@ class PillarGuardrail(CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Pre-call hook to scan the request for security threats before sending to LLM. @@ -341,7 +337,7 @@ class PillarGuardrail(CustomGuardrail): "mcp_call", "anthropic_messages", ], - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ During-call hook to scan the request in parallel with LLM processing. @@ -461,7 +457,7 @@ class PillarGuardrail(CustomGuardrail): raise e # Handle API communication errors based on fallback_on_error setting - verbose_proxy_logger.error(f"Pillar Guardrail: API communication failed - {str(e)}") + verbose_proxy_logger.error(f"Pillar Guardrail: API communication failed - {e!s}") return self._handle_api_error(e, data) @@ -501,7 +497,7 @@ class PillarGuardrail(CustomGuardrail): }, ) - def _prepare_headers(self, user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str]: + def _prepare_headers(self, user_api_key_dict: UserAPIKeyAuth) -> dict[str, str]: """ Prepare headers for the Pillar API request. @@ -518,7 +514,7 @@ class PillarGuardrail(CustomGuardrail): ) raise PillarGuardrailMissingSecrets(msg) - headers: Dict[str, str] = { + headers: dict[str, str] = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } @@ -545,7 +541,7 @@ class PillarGuardrail(CustomGuardrail): return headers - def _set_bool_header(self, headers: Dict[str, str], header_name: str, value: Optional[bool]) -> None: + def _set_bool_header(self, headers: dict[str, str], header_name: str, value: bool | None) -> None: """Apply a boolean value as a lowercase string HTTP header when provided.""" if value is None: @@ -554,11 +550,11 @@ class PillarGuardrail(CustomGuardrail): def _resolve_bool_config( self, - provided_value: Optional[Union[bool, str, int]], - env_var: Optional[str], - default: Optional[bool], + provided_value: bool | str | int | None, + env_var: str | None, + default: bool | None, setting_name: str, - ) -> Optional[bool]: + ) -> bool | None: """Resolve configuration precedence: explicit value -> environment -> default.""" if provided_value is not None: @@ -588,7 +584,7 @@ class PillarGuardrail(CustomGuardrail): return default @staticmethod - def _parse_bool_value(value: Union[bool, str, int]) -> bool: + def _parse_bool_value(value: bool | str | int) -> bool: """Normalise various truthy/falsey inputs to a strict boolean.""" if isinstance(value, bool): @@ -603,7 +599,7 @@ class PillarGuardrail(CustomGuardrail): return False raise ValueError(f"Unrecognised boolean value: {value}") - def _extract_model_and_provider(self, data: dict) -> Tuple[str, str]: + def _extract_model_and_provider(self, data: dict) -> tuple[str, str]: """ Extract the model and provider from the request data. @@ -633,7 +629,7 @@ class PillarGuardrail(CustomGuardrail): data.get("custom_llm_provider") or data.get("provider") or "unknown", ) - def _prepare_payload(self, data: dict) -> Dict[str, Any]: + def _prepare_payload(self, data: dict) -> dict[str, Any]: """ Prepare the payload for the Pillar API request following the /api/v1/protect contract. @@ -686,7 +682,7 @@ class PillarGuardrail(CustomGuardrail): ) return payload - async def _call_pillar_api(self, headers: Dict[str, str], payload: Dict[str, Any]) -> Dict[str, Any]: + async def _call_pillar_api(self, headers: dict[str, str], payload: dict[str, Any]) -> dict[str, Any]: """ Call the Pillar API and return the response. @@ -714,7 +710,7 @@ class PillarGuardrail(CustomGuardrail): verbose_proxy_logger.debug(f"Pillar Guardrail: Analysis complete - flagged={flagged}, session={session_id}") return res - def _process_pillar_response(self, pillar_response: Dict[str, Any], original_data: dict) -> None: + def _process_pillar_response(self, pillar_response: dict[str, Any], original_data: dict) -> None: """ Process the Pillar API response and handle detections based on configuration. @@ -774,7 +770,7 @@ class PillarGuardrail(CustomGuardrail): build_pillar_response_headers(metadata_store) - def _raise_pillar_detection_exception(self, pillar_response: Dict[str, Any]) -> None: + def _raise_pillar_detection_exception(self, pillar_response: dict[str, Any]) -> None: """ Raise an HTTPException for Pillar security detections. @@ -809,7 +805,7 @@ class PillarGuardrail(CustomGuardrail): # ========================================================================= @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Get the configuration model for this guardrail. @@ -823,7 +819,7 @@ class PillarGuardrail(CustomGuardrail): return PillarGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 86f5beeb3ee..9d0b8d2777c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -17,12 +17,8 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Union, cast, ) @@ -68,7 +64,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ad_hoc_recognizers = None @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -82,17 +78,17 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def __init__( self, mock_testing: bool = False, - mock_redacted_text: Optional[dict] = None, - presidio_analyzer_api_base: Optional[str] = None, - presidio_anonymizer_api_base: Optional[str] = None, - output_parse_pii: Optional[bool] = False, + mock_redacted_text: dict | None = None, + presidio_analyzer_api_base: str | None = None, + presidio_anonymizer_api_base: str | None = None, + output_parse_pii: bool | None = False, apply_to_output: bool = False, - presidio_ad_hoc_recognizers: Optional[str] = None, - logging_only: Optional[bool] = None, - pii_entities_config: Optional[Dict[Union[PiiEntityType, str], PiiAction]] = None, - presidio_language: Optional[str] = None, - presidio_score_thresholds: Optional[Dict[Union[PiiEntityType, str], float]] = None, - presidio_entities_deny_list: Optional[List[Union[PiiEntityType, str]]] = None, + presidio_ad_hoc_recognizers: str | None = None, + logging_only: bool | None = None, + pii_entities_config: dict[PiiEntityType | str, PiiAction] | None = None, + presidio_language: str | None = None, + presidio_score_thresholds: dict[PiiEntityType | str, float] | None = None, + presidio_entities_deny_list: list[PiiEntityType | str] | None = None, **kwargs, ): if logging_only is True: @@ -112,15 +108,15 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if (self.output_parse_pii or self.apply_to_output) and not logging_only: current_hook = self.event_hook if isinstance(current_hook, str) and current_hook != "post_call": - self.event_hook = cast(List[GuardrailEventHooks], [current_hook, "post_call"]) + self.event_hook = cast(list[GuardrailEventHooks], [current_hook, "post_call"]) elif isinstance(current_hook, list) and "post_call" not in current_hook: - self.event_hook = cast(List[GuardrailEventHooks], current_hook + ["post_call"]) - self.pii_entities_config: Dict[Union[PiiEntityType, str], PiiAction] = pii_entities_config or {} - self.presidio_score_thresholds: Dict[Union[PiiEntityType, str], float] = presidio_score_thresholds or {} - self.presidio_entities_deny_list: List[Union[PiiEntityType, str]] = presidio_entities_deny_list or [] + self.event_hook = cast(list[GuardrailEventHooks], current_hook + ["post_call"]) + self.pii_entities_config: dict[PiiEntityType | str, PiiAction] = pii_entities_config or {} + self.presidio_score_thresholds: dict[PiiEntityType | str, float] = presidio_score_thresholds or {} + self.presidio_entities_deny_list: list[PiiEntityType | str] = presidio_entities_deny_list or [] self.presidio_language = presidio_language or "en" # Shared HTTP session to prevent memory leaks (issue #14540) - self._http_session: Optional[aiohttp.ClientSession] = None + self._http_session: aiohttp.ClientSession | None = None # Lock to prevent race conditions when creating session under concurrent load # Note: asyncio.Lock() can be created without an event loop; it only needs one when awaited self._session_lock: asyncio.Lock = asyncio.Lock() @@ -130,7 +126,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self._main_thread_id = threading.get_ident() # Loop-bound session cache for background threads - self._loop_sessions: Dict[asyncio.AbstractEventLoop, aiohttp.ClientSession] = {} + self._loop_sessions: dict[asyncio.AbstractEventLoop, aiohttp.ClientSession] = {} if mock_testing is True: # for testing purposes only return @@ -143,9 +139,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): except FileNotFoundError: raise Exception(f"File not found. file_path={ad_hoc_recognizers}") except json.JSONDecodeError as e: - raise Exception(f"Error decoding JSON file: {str(e)}, file_path={ad_hoc_recognizers}") + raise Exception(f"Error decoding JSON file: {e!s}, file_path={ad_hoc_recognizers}") except Exception as e: - raise Exception(f"An error occurred: {str(e)}, file_path={ad_hoc_recognizers}") + raise Exception(f"An error occurred: {e!s}, file_path={ad_hoc_recognizers}") self.validate_environment( presidio_analyzer_api_base=presidio_analyzer_api_base, presidio_anonymizer_api_base=presidio_anonymizer_api_base, @@ -153,13 +149,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def validate_environment( self, - presidio_analyzer_api_base: Optional[str] = None, - presidio_anonymizer_api_base: Optional[str] = None, + presidio_analyzer_api_base: str | None = None, + presidio_anonymizer_api_base: str | None = None, ): - self.presidio_analyzer_api_base: Optional[str] = presidio_analyzer_api_base or get_secret( + self.presidio_analyzer_api_base: str | None = presidio_analyzer_api_base or get_secret( "PRESIDIO_ANALYZER_API_BASE", None ) # type: ignore - self.presidio_anonymizer_api_base: Optional[str] = presidio_anonymizer_api_base or litellm.get_secret( + self.presidio_anonymizer_api_base: str | None = presidio_anonymizer_api_base or litellm.get_secret( "PRESIDIO_ANONYMIZER_API_BASE", None ) # type: ignore @@ -227,7 +223,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def __del__(self): """Cleanup: we try to close, but doing async cleanup in __del__ is risky.""" - pass def _has_block_action(self) -> bool: """Return True if pii_entities_config has any BLOCK action (fail-closed on analyzer errors).""" @@ -238,7 +233,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def _get_presidio_analyze_request_payload( self, text: str, - presidio_config: Optional[PresidioPerRequestConfig], + presidio_config: PresidioPerRequestConfig | None, request_data: dict, ) -> PresidioAnalyzeRequest: """ @@ -274,9 +269,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async def analyze_text( self, text: str, - presidio_config: Optional[PresidioPerRequestConfig], + presidio_config: PresidioPerRequestConfig | None, request_data: dict, - ) -> Union[List[PresidioAnalyzeResponseItem], Dict]: + ) -> list[PresidioAnalyzeResponseItem] | dict: """ Send text to the Presidio analyzer endpoint and get analysis results """ @@ -309,7 +304,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def _fail_on_invalid_response( reason: str, - ) -> List[PresidioAnalyzeResponseItem]: + ) -> list[PresidioAnalyzeResponseItem]: should_fail_closed = bool(self.pii_entities_config) or self.output_parse_pii or self.apply_to_output if should_fail_closed: raise GuardrailRaisedException( @@ -428,8 +423,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def _finalize_presidio_anonymize_simple( self, - redacted_text: Dict[str, Any], - masked_entity_count: Dict[str, int], + redacted_text: dict[str, Any], + masked_entity_count: dict[str, int], ) -> str: # No need to build numbered tokens — just use Presidio's # already-anonymized text directly. The old code incorrectly @@ -445,8 +440,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, text: str, analyze_results: Any, - request_data: Optional[Dict], - masked_entity_count: Dict[str, int], + request_data: dict | None, + masked_entity_count: dict[str, int], ) -> str: # output_parse_pii is True — we need sequentially numbered # tokens and a pii_tokens mapping for later unmasking. @@ -496,8 +491,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): text: str, analyze_results: Any, output_parse_pii: bool, - masked_entity_count: Dict[str, int], - request_data: Optional[Dict] = None, + masked_entity_count: dict[str, int], + request_data: dict | None = None, ) -> str: """ Send analysis results to the Presidio anonymizer endpoint to get redacted text @@ -528,8 +523,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): raise Exception(f"Presidio PII anonymization failed: {type(e).__name__}") from e def filter_analyze_results_by_score( - self, analyze_results: Union[List[PresidioAnalyzeResponseItem], Dict] - ) -> Union[List[PresidioAnalyzeResponseItem], Dict]: + self, analyze_results: list[PresidioAnalyzeResponseItem] | dict + ) -> list[PresidioAnalyzeResponseItem] | dict: """ Drop detections that fall below configured per-entity score thresholds or match an entity type in the deny list. @@ -540,7 +535,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if not isinstance(analyze_results, list): return analyze_results - filtered_results: List[PresidioAnalyzeResponseItem] = [] + filtered_results: list[PresidioAnalyzeResponseItem] = [] deny_list_strings = [getattr(x, "value", str(x)) for x in self.presidio_entities_deny_list] for item in analyze_results: entity_type = item.get("entity_type") @@ -567,16 +562,14 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return filtered_results - def raise_exception_if_blocked_entities_detected( - self, analyze_results: Union[List[PresidioAnalyzeResponseItem], Dict] - ): + def raise_exception_if_blocked_entities_detected(self, analyze_results: list[PresidioAnalyzeResponseItem] | dict): """ Raise an exception if blocked entities are detected """ if self.pii_entities_config is None: return - if isinstance(analyze_results, Dict): + if isinstance(analyze_results, dict): # if mock testing is enabled, analyze_results is a dict # we don't need to raise an exception in this case return @@ -596,16 +589,16 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, text: str, output_parse_pii: bool, - presidio_config: Optional[PresidioPerRequestConfig], + presidio_config: PresidioPerRequestConfig | None, request_data: dict, ) -> str: """ Calls Presidio Analyze + Anonymize endpoints for PII Analysis + Masking """ start_time = datetime.now() - analyze_results: Optional[Union[List[PresidioAnalyzeResponseItem], Dict]] = None + analyze_results: list[PresidioAnalyzeResponseItem] | dict | None = None status: GuardrailStatus = "success" - masked_entity_count: Dict[str, int] = {} + masked_entity_count: dict[str, int] = {} exception_str: str = "" try: if self.mock_redacted_text is not None: @@ -646,9 +639,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): #################################################### # Create Guardrail Trace for logging on Langfuse, Datadog, etc. #################################################### - guardrail_json_response: Union[Exception, str, dict, List[dict]] = {} + guardrail_json_response: Exception | str | dict | list[dict] = {} if status == "success": - if isinstance(analyze_results, List): + if isinstance(analyze_results, list): guardrail_json_response = [dict(item) for item in analyze_results] else: guardrail_json_response = exception_str @@ -701,7 +694,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if messages is None: return data tasks = [] - task_mappings: List[Tuple[int, Optional[int]]] = [] # Track (message_index, content_index) for each task + task_mappings: list[tuple[int, int | None]] = [] # Track (message_index, content_index) for each task for msg_idx, m in enumerate(messages): content = m.get("content", None) @@ -738,7 +731,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): for task_idx, r in enumerate(responses): mapping = task_mappings[task_idx] msg_idx = cast(int, mapping[0]) - content_idx_optional = cast(Optional[int], mapping[1]) + content_idx_optional = cast(int | None, mapping[1]) content = messages[msg_idx].get("content", None) if content is None: continue @@ -753,7 +746,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): except Exception as e: raise e - def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: from concurrent.futures import ThreadPoolExecutor def run_in_new_loop(): @@ -781,14 +774,14 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): # No running event loop, we can safely run in this thread return run_in_new_loop() - async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> Tuple[dict, Any]: + async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """ Masks the input before logging to langfuse, datadog, etc. """ if call_type == "completion" or call_type == "acompletion": # /chat/completions requests - messages: Optional[List] = kwargs.get("messages", None) + messages: list | None = kwargs.get("messages", None) tasks = [] - task_mappings: List[Tuple[int, Optional[int]]] = [] # Track (message_index, content_index) for each task + task_mappings: list[tuple[int, int | None]] = [] # Track (message_index, content_index) for each task if messages is None: return kwargs, result @@ -830,7 +823,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): for task_idx, r in enumerate(responses): mapping = task_mappings[task_idx] msg_idx = cast(int, mapping[0]) - content_idx_optional = cast(Optional[int], mapping[1]) + content_idx_optional = cast(int | None, mapping[1]) content = messages[msg_idx].get("content", None) if content is None: continue @@ -848,7 +841,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, data: dict, user_api_key_dict: UserAPIKeyAuth, - response: Union[ModelResponse, EmbeddingResponse, ImageResponse], + response: ModelResponse | EmbeddingResponse | ImageResponse, ): """ Output parse the response object to replace the masked tokens with user sent values @@ -882,7 +875,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return response @staticmethod - def _unmask_pii_text(text: str, pii_tokens: Dict[str, str]) -> str: + def _unmask_pii_text(text: str, pii_tokens: dict[str, str]) -> str: """ Replace PII tokens in *text* with their original values. @@ -1040,7 +1033,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async def _mask_output_response( self, - response: Union[ModelResponse, EmbeddingResponse, ImageResponse], + response: ModelResponse | EmbeddingResponse | ImageResponse, request_data: dict, ): """ @@ -1064,7 +1057,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, response: Any, request_data: dict, - ) -> AsyncGenerator[Union[ModelResponseStream, bytes], None]: + ) -> AsyncGenerator[ModelResponseStream | bytes, None]: """Apply Presidio masking to streaming output (apply_to_output=True path).""" from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, @@ -1072,7 +1065,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): from litellm.main import stream_chunk_builder from litellm.types.utils import ModelResponse - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] passthrough_due_to_unknown_stream_shape = False try: async for chunk in response: @@ -1131,18 +1124,18 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield mock_response_stream except Exception as e: - verbose_proxy_logger.error(f"Error masking streaming PII output: {str(e)}") + verbose_proxy_logger.error(f"Error masking streaming PII output: {e!s}") for chunk in all_chunks: yield chunk @staticmethod - def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: Dict[str, str]) -> bytes: + def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes: try: text = chunk.decode("utf-8") except UnicodeDecodeError: return chunk - result_lines: List[str] = [] + result_lines: list[str] = [] for line in text.split("\n"): line = line.rstrip("\r") if line.startswith("data: ") and line != "data: [DONE]": @@ -1166,7 +1159,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return "\n".join(result_lines).encode("utf-8") - def _unmask_responses_api_completed_chunk(self, chunk: Any, pii_tokens: Dict[str, str]) -> None: + def _unmask_responses_api_completed_chunk(self, chunk: Any, pii_tokens: dict[str, str]) -> None: """ Unmask PII tokens in-place for a ``response.completed`` Responses API event. @@ -1193,7 +1186,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, response: Any, request_data: dict, - ) -> AsyncGenerator[Union[ModelResponseStream, bytes], None]: + ) -> AsyncGenerator[ModelResponseStream | bytes, None]: """Apply PII unmasking to streaming output (output_parse_pii=True path).""" from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, @@ -1202,9 +1195,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): from litellm.types.utils import ModelResponse metadata = (request_data.get("metadata") or {}) if request_data else {} - pii_tokens: Dict[str, str] = metadata.get("pii_tokens", {}) + pii_tokens: dict[str, str] = metadata.get("pii_tokens", {}) - remaining_chunks: List[ModelResponseStream] = [] + remaining_chunks: list[ModelResponseStream] = [] saw_non_chat_chunk = False try: async for chunk in response: @@ -1260,7 +1253,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield mock_response_stream except Exception as e: - verbose_proxy_logger.error(f"Error in PII streaming processing: {str(e)}") + verbose_proxy_logger.error(f"Error in PII streaming processing: {e!s}") for chunk in remaining_chunks: yield chunk @@ -1269,7 +1262,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): user_api_key_dict: UserAPIKeyAuth, response: Any, request_data: dict, - ) -> AsyncGenerator[Union[ModelResponseStream, bytes], None]: + ) -> AsyncGenerator[ModelResponseStream | bytes, None]: """ Process streaming response chunks to unmask PII tokens when needed. @@ -1297,7 +1290,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): @staticmethod def _preserve_usage_from_last_chunk( assembled_model_response: Any, - chunks: List[Any], + chunks: list[Any], ) -> None: """Copy usage metadata from the last chunk when stream_chunk_builder misses it.""" if not getattr(assembled_model_response, "usage", None) and chunks: @@ -1305,7 +1298,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if last_chunk_usage: setattr(assembled_model_response, "usage", last_chunk_usage) - def get_presidio_settings_from_request_data(self, data: dict) -> Optional[PresidioPerRequestConfig]: + def get_presidio_settings_from_request_data(self, data: dict) -> PresidioPerRequestConfig | None: if "metadata" in data: _metadata = data.get("metadata", None) if _metadata is None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index e60815dbcf7..9b55fcd8062 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -1,7 +1,7 @@ import asyncio import base64 import os -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type +from typing import TYPE_CHECKING, Any, Literal, Optional from fastapi import HTTPException @@ -28,7 +28,7 @@ class PromptSecurityGuardrailMissingSecrets(Exception): class PromptSecurityGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -37,11 +37,11 @@ class PromptSecurityGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - user: Optional[str] = None, - system_prompt: Optional[str] = None, - check_tool_results: Optional[bool] = None, + api_key: str | None = None, + api_base: str | None = None, + user: str | None = None, + system_prompt: str | None = None, + check_tool_results: bool | None = None, **kwargs, ): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -147,11 +147,11 @@ class PromptSecurityGuardrail(CustomGuardrail): async def _apply_guardrail_on_request( self, inputs: GenericGuardrailAPIInputs, - texts: List[str], - images: List[str], + texts: list[str], + images: list[str], structured_messages: list, request_data: dict, - user_api_key_alias: Optional[str], + user_api_key_alias: str | None, ) -> GenericGuardrailAPIInputs: """Handle request-side guardrail checks.""" # If we have structured messages, use them (they contain role information) @@ -228,8 +228,8 @@ class PromptSecurityGuardrail(CustomGuardrail): async def _apply_guardrail_on_response( self, inputs: GenericGuardrailAPIInputs, - texts: List[str], - user_api_key_alias: Optional[str], + texts: list[str], + user_api_key_alias: str | None, ) -> GenericGuardrailAPIInputs: """Handle response-side guardrail checks.""" if not texts: @@ -287,7 +287,7 @@ class PromptSecurityGuardrail(CustomGuardrail): return inputs - def _extract_texts_from_messages(self, messages: list) -> List[str]: + def _extract_texts_from_messages(self, messages: list) -> list[str]: """Extract text content from messages.""" texts = [] for message in messages: @@ -302,7 +302,7 @@ class PromptSecurityGuardrail(CustomGuardrail): texts.append(text) return texts - async def _process_standalone_images(self, images: List[str], user_api_key_alias: Optional[str]) -> None: + async def _process_standalone_images(self, images: list[str], user_api_key_alias: str | None) -> None: """Process standalone images from inputs (data URLs).""" for image_url in images: if image_url.startswith("data:"): @@ -326,10 +326,10 @@ class PromptSecurityGuardrail(CustomGuardrail): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error processing image: {str(e)}") + verbose_proxy_logger.error(f"Error processing image: {e!s}") @staticmethod - def _resolve_key_alias_from_request_data(request_data: dict) -> Optional[str]: + def _resolve_key_alias_from_request_data(request_data: dict) -> str | None: """Resolve user API key alias from request_data metadata.""" # Check litellm_metadata first (set by guardrail framework) litellm_metadata = request_data.get("litellm_metadata", {}) @@ -351,7 +351,7 @@ class PromptSecurityGuardrail(CustomGuardrail): self, file_data: bytes, filename: str, - user_api_key_alias: Optional[str] = None, + user_api_key_alias: str | None = None, ) -> dict: """ Sanitize file content using Prompt Security API. @@ -439,7 +439,7 @@ class PromptSecurityGuardrail(CustomGuardrail): raise HTTPException(status_code=408, detail="File sanitization timeout") - async def _process_image_url_item(self, item: dict, user_api_key_alias: Optional[str]) -> dict: + async def _process_image_url_item(self, item: dict, user_api_key_alias: str | None) -> dict: """Process and sanitize image_url items.""" image_url_data = item.get("image_url", {}) url = image_url_data.get("url", "") if isinstance(image_url_data, dict) else image_url_data @@ -481,10 +481,10 @@ class PromptSecurityGuardrail(CustomGuardrail): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error sanitizing image file: {str(e)}") - raise HTTPException(status_code=500, detail=f"File sanitization failed: {str(e)}") + verbose_proxy_logger.error(f"Error sanitizing image file: {e!s}") + raise HTTPException(status_code=500, detail=f"File sanitization failed: {e!s}") - async def _process_document_item(self, item: dict, user_api_key_alias: Optional[str]) -> dict: + async def _process_document_item(self, item: dict, user_api_key_alias: str | None) -> dict: """Process and sanitize document/file items.""" doc_data = item.get("document") or item.get("file") or item @@ -554,10 +554,10 @@ class PromptSecurityGuardrail(CustomGuardrail): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error sanitizing document: {str(e)}") - raise HTTPException(status_code=500, detail=f"Document sanitization failed: {str(e)}") + verbose_proxy_logger.error(f"Error sanitizing document: {e!s}") + raise HTTPException(status_code=500, detail=f"Document sanitization failed: {e!s}") - async def process_message_files(self, messages: list, user_api_key_alias: Optional[str] = None) -> list: + async def process_message_files(self, messages: list, user_api_key_alias: str | None = None) -> list: """Process messages and sanitize any file content (images, documents, PDFs, etc.).""" processed_messages = [] @@ -638,7 +638,7 @@ class PromptSecurityGuardrail(CustomGuardrail): return filtered_messages - def _build_headers(self, user_api_key_alias: Optional[str] = None) -> dict: + def _build_headers(self, user_api_key_alias: str | None = None) -> dict: headers = {"APP-ID": self.api_key, "Content-Type": "application/json"} if user_api_key_alias: headers["X-LiteLLM-Key-Alias"] = user_api_key_alias @@ -677,7 +677,7 @@ class PromptSecurityGuardrail(CustomGuardrail): ) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.prompt_security import ( PromptSecurityGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py b/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py index 6603183efef..11bc2725da3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py @@ -10,11 +10,8 @@ import os from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Type, ) from litellm._logging import verbose_proxy_logger @@ -49,9 +46,9 @@ class PromptGuardMissingCredentials(Exception): class PromptGuardGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - block_on_error: Optional[bool] = None, + api_key: str | None = None, + api_base: str | None = None, + block_on_error: bool | None = None, **kwargs: Any, ) -> None: self.api_key = api_key or os.environ.get( @@ -86,7 +83,7 @@ class PromptGuardGuardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import ( PromptGuardConfigModel, ) @@ -94,7 +91,7 @@ class PromptGuardGuardrail(CustomGuardrail): return PromptGuardConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -122,7 +119,7 @@ class PromptGuardGuardrail(CustomGuardrail): direction = "input" if input_type == "request" else "output" - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "messages": messages, "direction": direction, } @@ -194,14 +191,14 @@ class PromptGuardGuardrail(CustomGuardrail): return inputs @staticmethod - def _extract_texts_from_messages(messages: list) -> List[str]: + def _extract_texts_from_messages(messages: list) -> list[str]: """Extract text content from user-role messages only. Only user messages are extracted to avoid injecting system or assistant content into the ``texts`` list, which should mirror the original user-provided input. """ - texts: List[str] = [] + texts: list[str] = [] for message in messages: if message.get("role") != "user": continue diff --git a/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py b/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py index 6d1c14a934b..d2f8900dc38 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py @@ -3,7 +3,7 @@ Qostodian Nexus (by Qohash) — LiteLLM guardrail integration. """ import os -from typing import TYPE_CHECKING, Literal, Optional, Type +from typing import TYPE_CHECKING, Literal, Optional from litellm.integrations.custom_guardrail import log_guardrail_information from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( @@ -23,7 +23,7 @@ GUARDRAIL_NAME = "qostodian_nexus" class QostodianNexus(GenericGuardrailAPI): def __init__( self, - api_base: Optional[str] = None, + api_base: str | None = None, **kwargs, ): api_base = api_base or os.environ.get("QOSTODIAN_NEXUS_API_BASE", "http://nexus:8800") @@ -70,7 +70,7 @@ class QostodianNexus(GenericGuardrailAPI): ) @classmethod - def get_config_model(cls) -> Optional[Type[QostodianNexusConfigModel]]: + def get_config_model(cls) -> type[QostodianNexusConfigModel] | None: """ Returns the config model for Qostodian Nexus. """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index 9c62c2915ff..fe2cc40074f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -7,7 +7,7 @@ import json import os -from typing import Any, Dict, List, Literal, Optional, Type +from typing import Any, Literal from fastapi import HTTPException @@ -33,7 +33,7 @@ DEFAULT_QUALIFIRE_API_BASE = "https://proxy.qualifire.ai" class QualifireGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -42,17 +42,17 @@ class QualifireGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - evaluation_id: Optional[str] = None, - prompt_injections: Optional[bool] = None, - hallucinations_check: Optional[bool] = None, - grounding_check: Optional[bool] = None, - pii_check: Optional[bool] = None, - content_moderation_check: Optional[bool] = None, - tool_selection_quality_check: Optional[bool] = None, - assertions: Optional[List[str]] = None, - on_flagged: Optional[str] = "block", + api_key: str | None = None, + api_base: str | None = None, + evaluation_id: str | None = None, + prompt_injections: bool | None = None, + hallucinations_check: bool | None = None, + grounding_check: bool | None = None, + pii_check: bool | None = None, + content_moderation_check: bool | None = None, + tool_selection_quality_check: bool | None = None, + assertions: list[str] | None = None, + on_flagged: str | None = "block", **kwargs, ): """ @@ -112,7 +112,7 @@ class QualifireGuardrail(CustomGuardrail): ] ) - def _convert_messages_to_api_format(self, messages: List[AllMessageValues]) -> List[Dict[str, Any]]: + def _convert_messages_to_api_format(self, messages: list[AllMessageValues]) -> list[dict[str, Any]]: """ Convert LiteLLM messages to Qualifire API format. Supports tool calls for tool_selection_quality_check. @@ -140,7 +140,7 @@ class QualifireGuardrail(CustomGuardrail): text_parts.append(part) content = "\n".join(text_parts) - api_message: Dict[str, Any] = { + api_message: dict[str, Any] = { "role": role, "content": content if isinstance(content, str) else str(content), } @@ -178,7 +178,7 @@ class QualifireGuardrail(CustomGuardrail): return api_messages - def _convert_tools_to_api_format(self, tools: Optional[List[Any]]) -> Optional[List[Dict[str, Any]]]: + def _convert_tools_to_api_format(self, tools: list[Any] | None) -> list[dict[str, Any]] | None: """ Convert OpenAI-format tools to Qualifire API format. @@ -217,7 +217,7 @@ class QualifireGuardrail(CustomGuardrail): return api_tools if api_tools else None - def _check_if_flagged(self, result: Dict[str, Any]) -> bool: + def _check_if_flagged(self, result: dict[str, Any]) -> bool: """ Check if the Qualifire evaluation result indicates flagged content. @@ -237,13 +237,13 @@ class QualifireGuardrail(CustomGuardrail): def _build_evaluate_payload( self, - api_messages: List[Dict[str, Any]], - output: Optional[str], - assertions: Optional[List[str]], - available_tools: Optional[List[Dict[str, Any]]], - ) -> Dict[str, Any]: + api_messages: list[dict[str, Any]], + output: str | None, + assertions: list[str] | None, + available_tools: list[dict[str, Any]] | None, + ) -> dict[str, Any]: """Build payload dictionary for the /api/evaluation/evaluate endpoint.""" - payload: Dict[str, Any] = {"messages": api_messages} + payload: dict[str, Any] = {"messages": api_messages} if output is not None: payload["output"] = output @@ -275,10 +275,10 @@ class QualifireGuardrail(CustomGuardrail): async def _run_qualifire_check( self, - messages: List[AllMessageValues], - output: Optional[str], - dynamic_params: Dict[str, Any], - available_tools: Optional[List[Any]] = None, + messages: list[AllMessageValues], + output: str | None, + dynamic_params: dict[str, Any], + available_tools: list[Any] | None = None, ) -> None: """ Core Qualifire check logic - shared between hooks. @@ -398,7 +398,7 @@ class QualifireGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"], - logging_obj: Optional[LiteLLMLoggingObj] = None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: """ Apply Qualifire guardrail to the given inputs. @@ -425,13 +425,13 @@ class QualifireGuardrail(CustomGuardrail): dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data) # Extract messages from structured_messages or request_data - messages: Optional[List[AllMessageValues]] = inputs.get("structured_messages") + messages: list[AllMessageValues] | None = inputs.get("structured_messages") if not messages: messages = request_data.get("messages") # For response (post_call), messages may not be available in the inputs # We need to work with texts instead and construct messages if needed - output: Optional[str] = None + output: str | None = None texts = inputs.get("texts", []) if input_type == "response": @@ -465,7 +465,7 @@ class QualifireGuardrail(CustomGuardrail): return inputs @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: # type: ignore + def get_config_model() -> type["GuardrailConfigModel"] | None: # type: ignore from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import ( QualifireGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py index 9dc060ac9ae..e22996e945f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Union +from typing import TYPE_CHECKING from litellm.types.guardrails import ( GuardrailEventHooks, @@ -14,7 +14,7 @@ if TYPE_CHECKING: def _event_hook_from_mode( mode: str | list[str] | Mode, -) -> Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode]: +) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode: if isinstance(mode, Mode): return mode if isinstance(mode, list): diff --git a/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py b/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py index 7b648d4bf99..a91d06ff001 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import AsyncGenerator from datetime import datetime -from typing import List, Literal, TypeGuard +from typing import Literal, TypeGuard from fastapi import HTTPException from httpx import HTTPError @@ -64,7 +64,7 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin class RepelloAIGuardrail(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, diff --git a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py index ad05e7656c4..c30b2b0910c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py +++ b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py @@ -6,7 +6,7 @@ then builds a SemanticRouter for prompt matching. """ import os -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any import yaml @@ -27,7 +27,7 @@ class SemanticGuardRouteLoader: """Loads route definitions from YAML templates and custom configs, builds SemanticRouter.""" @staticmethod - def load_builtin_template(template_name: str) -> Dict[str, Any]: + def load_builtin_template(template_name: str) -> dict[str, Any]: """Load a built-in route template YAML by name.""" file_path = os.path.join(ROUTE_TEMPLATES_DIR, f"{template_name}.yaml") if not os.path.exists(file_path): @@ -39,7 +39,7 @@ class SemanticGuardRouteLoader: return yaml.safe_load(f) @staticmethod - def list_builtin_templates() -> List[str]: + def list_builtin_templates() -> list[str]: """List available built-in template names.""" templates = [] if os.path.isdir(ROUTE_TEMPLATES_DIR): @@ -49,7 +49,7 @@ class SemanticGuardRouteLoader: return sorted(templates) @staticmethod - def load_custom_routes_file(file_path: str) -> List[Dict[str, Any]]: + def load_custom_routes_file(file_path: str) -> list[dict[str, Any]]: """Load custom routes from a YAML file.""" if not os.path.exists(file_path): raise ValueError(f"SemanticGuard: custom routes file not found: {file_path}") @@ -64,15 +64,15 @@ class SemanticGuardRouteLoader: @classmethod def build_routes( cls, - route_templates: Optional[List[str]], - custom_routes_file: Optional[str], - custom_routes: Optional[List[Dict[str, Any]]], + route_templates: list[str] | None, + custom_routes_file: str | None, + custom_routes: list[dict[str, Any]] | None, global_threshold: float = DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD, - ) -> List["Route"]: + ) -> list["Route"]: """Build semantic-router Route objects from templates + custom config.""" from semantic_router.routers.base import Route - routes: List[Route] = [] + routes: list[Route] = [] if route_templates: for template_name in route_templates: @@ -118,7 +118,7 @@ class SemanticGuardRouteLoader: @classmethod def build_semantic_router( cls, - routes: List["Route"], + routes: list["Route"], litellm_router: "Router", embedding_model: str, global_threshold: float, diff --git a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/semantic_guard.py b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/semantic_guard.py index 3865251a48a..f57827d03c9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/semantic_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/semantic_guard.py @@ -6,7 +6,7 @@ via embedding similarity. Smarter than regex (understands intent), lighter than an LLM call (~20-50ms per request for embedding). """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.integrations.custom_guardrail import ( @@ -48,11 +48,11 @@ class SemanticGuardrail(CustomGuardrail): llm_router: "Router", embedding_model: str, similarity_threshold: float, - route_templates: Optional[List[str]] = None, - custom_routes_file: Optional[str] = None, - custom_routes: Optional[List[Dict[str, Any]]] = None, + route_templates: list[str] | None = None, + custom_routes_file: str | None = None, + custom_routes: list[dict[str, Any]] | None = None, on_flagged_action: str = "block", - event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, **kwargs, ): @@ -80,7 +80,7 @@ class SemanticGuardrail(CustomGuardrail): if not routes: raise ValueError("SemanticGuardrail: no routes configured. Provide route_templates or custom_routes.") - self.semantic_router: "SemanticRouter" = SemanticGuardRouteLoader.build_semantic_router( + self.semantic_router: SemanticRouter = SemanticGuardRouteLoader.build_semantic_router( routes=routes, litellm_router=llm_router, embedding_model=embedding_model, @@ -94,7 +94,7 @@ class SemanticGuardrail(CustomGuardrail): ) @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -111,11 +111,11 @@ class SemanticGuardrail(CustomGuardrail): """Check user messages against semantic routes before LLM call.""" messages = self.get_guardrails_messages_for_call_type(call_type=CallTypes(call_type), data=data) if not messages: - return None + return user_text = _extract_user_text(messages) if not user_text: - return None + return route_choice = _get_top_route_choice(self.semantic_router(text=user_text)) if route_choice is not None and route_choice.name: @@ -127,7 +127,7 @@ class SemanticGuardrail(CustomGuardrail): data=data, ) - return None + return @log_guardrail_information async def async_post_call_success_hook( @@ -166,7 +166,7 @@ def _get_top_route_choice(result: Any) -> Any: return result -def _extract_user_text(messages: List) -> str: +def _extract_user_text(messages: list) -> str: """Extract the latest user message text.""" for msg in reversed(messages): if isinstance(msg, dict) and msg.get("role") == "user": @@ -181,7 +181,7 @@ def _extract_user_text(messages: List) -> str: def _extract_response_text(response: Any) -> str: """Extract text from every LLM response choice.""" if hasattr(response, "choices") and response.choices: - text_parts: List[str] = [] + text_parts: list[str] = [] for choice in response.choices: if hasattr(choice, "message") and choice.message: text = _content_to_text(choice.message.content) @@ -205,7 +205,7 @@ def _content_to_text(content: Any) -> str: def _handle_match( guardrail: SemanticGuardrail, route_name: str, - similarity_score: Optional[float], + similarity_score: float | None, user_text: str, data: dict, ) -> None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 9f61a0670f3..b15e5b61243 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -1,7 +1,7 @@ import json import re from collections.abc import AsyncGenerator -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Literal from fastapi import HTTPException @@ -38,7 +38,7 @@ GUARDRAIL_NAME = "tool_permission" class ToolPermissionGuardrail(CustomGuardrail): def __init__( self, - rules: Optional[List[Dict]] = None, + rules: list[dict] | None = None, default_action: Literal["deny", "allow"] = "deny", on_disallowed_action: Literal["block", "rewrite"] = "block", **kwargs, @@ -71,7 +71,7 @@ class ToolPermissionGuardrail(CustomGuardrail): self.default_action, ) - def _load_rules(self, rules: Optional[List[Any]]) -> None: + def _load_rules(self, rules: list[Any] | None) -> None: """Parse ``rules`` and (re)build the compiled target/pattern lookups. ``self.rules`` plus ``_compiled_rule_targets`` / ``_compiled_rule_patterns`` @@ -80,14 +80,14 @@ class ToolPermissionGuardrail(CustomGuardrail): single source of truth, so an in-place update (PUT /guardrails, immediate sync) reflects rule changes instead of keeping the construction-time maps. """ - parsed_rules: List[ToolPermissionRule] = [] - compiled_targets: Dict[str, Dict[str, Optional[re.Pattern]]] = {} - compiled_patterns: Dict[str, Dict[str, re.Pattern]] = {} + parsed_rules: list[ToolPermissionRule] = [] + compiled_targets: dict[str, dict[str, re.Pattern | None]] = {} + compiled_patterns: dict[str, dict[str, re.Pattern]] = {} for rule_item in rules or []: rule = rule_item if isinstance(rule_item, ToolPermissionRule) else ToolPermissionRule(**rule_item) - target_patterns: Dict[str, Optional[re.Pattern]] = { + target_patterns: dict[str, re.Pattern | None] = { "tool_name": None, "tool_type": None, } @@ -102,7 +102,7 @@ class ToolPermissionGuardrail(CustomGuardrail): except re.error as exc: raise ValueError(f"Invalid regex for tool_type in rule '{rule.id}': {exc}") from exc - rule_patterns: Dict[str, re.Pattern] = {} + rule_patterns: dict[str, re.Pattern] = {} for path, pattern in (rule.allowed_param_patterns or {}).items(): try: rule_patterns[path] = re.compile(pattern) @@ -121,7 +121,7 @@ class ToolPermissionGuardrail(CustomGuardrail): self._compiled_rule_targets = compiled_targets self._compiled_rule_patterns = compiled_patterns - def update_in_memory_litellm_params(self, litellm_params: Union[LitellmParams, dict]) -> None: + def update_in_memory_litellm_params(self, litellm_params: LitellmParams | dict) -> None: """Apply updated params in place, rebuilding the compiled rule state. The base implementation only ``setattr``s raw fields, which would leave @@ -177,13 +177,13 @@ class ToolPermissionGuardrail(CustomGuardrail): return ToolPermissionGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, ] - def _matches_regex(self, pattern: Optional[re.Pattern], value: Optional[str]) -> bool: + def _matches_regex(self, pattern: re.Pattern | None, value: str | None) -> bool: if pattern is None: return True if value is None: @@ -194,8 +194,8 @@ class ToolPermissionGuardrail(CustomGuardrail): self, rule: ToolPermissionRule, *, - tool_name: Optional[str], - tool_type: Optional[str] = None, + tool_name: str | None, + tool_type: str | None = None, ) -> tuple[bool, bool]: target_patterns = self._compiled_rule_targets.get(rule.id, {}) name_pattern = target_patterns.get("tool_name") @@ -214,9 +214,9 @@ class ToolPermissionGuardrail(CustomGuardrail): def _check_tool_permission( self, - tool_name: Optional[str], - tool_type: Optional[str] = None, - ) -> tuple[bool, Optional[str], Optional[str]]: + tool_name: str | None, + tool_type: str | None = None, + ) -> tuple[bool, str | None, str | None]: """ Check if a tool is allowed based on the configured rules @@ -268,7 +268,7 @@ class ToolPermissionGuardrail(CustomGuardrail): def _parse_tool_call_arguments( self, tool_call: ChatCompletionMessageToolCall - ) -> tuple[Optional[Dict[str, Any]], Optional[str]]: + ) -> tuple[dict[str, Any] | None, str | None]: arguments = getattr(tool_call.function, "arguments", None) if not arguments: return None, "missing arguments" @@ -302,7 +302,7 @@ class ToolPermissionGuardrail(CustomGuardrail): self, value: Any, current_path: str, - collected: Dict[str, List[Any]], + collected: dict[str, list[Any]], depth: int = 0, ) -> None: from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -326,15 +326,15 @@ class ToolPermissionGuardrail(CustomGuardrail): def _patterns_match_for_rule( self, *, - arguments: Dict[str, Any], + arguments: dict[str, Any], rule: ToolPermissionRule, - tool_name: Optional[str], - ) -> tuple[bool, Optional[str]]: + tool_name: str | None, + ) -> tuple[bool, str | None]: compiled_patterns = self._compiled_rule_patterns.get(rule.id) if not compiled_patterns: return True, None - path_value_map: Dict[str, List[Any]] = {} + path_value_map: dict[str, list[Any]] = {} self._collect_argument_paths(arguments, "", path_value_map) for path, compiled_pattern in compiled_patterns.items(): @@ -356,7 +356,7 @@ class ToolPermissionGuardrail(CustomGuardrail): def _get_permission_for_tool_call( self, tool_call: ChatCompletionMessageToolCall - ) -> tuple[bool, Optional[str], Optional[str]]: + ) -> tuple[bool, str | None, str | None]: tool_name = tool_call.function.name if tool_call.function else None tool_type = getattr(tool_call, "type", None) if not tool_name and not tool_type: @@ -364,7 +364,7 @@ class ToolPermissionGuardrail(CustomGuardrail): tool_identifier = tool_name or tool_type or "unknown_tool" - last_pattern_failure_msg: Optional[str] = None + last_pattern_failure_msg: str | None = None for rule in self.rules: matches, should_check_params = self._rule_matches_tool( @@ -433,7 +433,7 @@ class ToolPermissionGuardrail(CustomGuardrail): def _legacy_function_call_to_tool_call( self, function_call: Any, choice_index: int - ) -> Optional[ChatCompletionMessageToolCall]: + ) -> ChatCompletionMessageToolCall | None: if function_call is None: return None @@ -448,7 +448,7 @@ class ToolPermissionGuardrail(CustomGuardrail): function={"name": function_name, "arguments": arguments}, ) - def _extract_tool_calls_from_response(self, response: ModelResponse) -> List[ChatCompletionMessageToolCall]: + def _extract_tool_calls_from_response(self, response: ModelResponse) -> list[ChatCompletionMessageToolCall]: """ Extract tool_calls from all choices in a model response. @@ -472,7 +472,7 @@ class ToolPermissionGuardrail(CustomGuardrail): return tool_calls - def _get_request_tool_name(self, tool: Any) -> tuple[Optional[str], Optional[str]]: + def _get_request_tool_name(self, tool: Any) -> tuple[str | None, str | None]: tool_type = self._get_mapping_value(tool, "type") if tool_type != "function": return None, tool_type @@ -481,10 +481,10 @@ class ToolPermissionGuardrail(CustomGuardrail): tool_name = self._get_mapping_value(function, "name") return tool_name, tool_type - def _get_legacy_function_name(self, function: Any) -> Optional[str]: + def _get_legacy_function_name(self, function: Any) -> str | None: return self._get_mapping_value(function, "name") - def _get_named_tool_choice(self, data: dict) -> Optional[str]: + def _get_named_tool_choice(self, data: dict) -> str | None: tool_choice = data.get("tool_choice") if not tool_choice or tool_choice in ("auto", "none", "required"): return None @@ -494,7 +494,7 @@ class ToolPermissionGuardrail(CustomGuardrail): return None return self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name") - def _get_named_function_call(self, data: dict) -> Optional[str]: + def _get_named_function_call(self, data: dict) -> str | None: function_call = data.get("function_call") if not function_call or function_call in ("auto", "none"): return None @@ -502,8 +502,8 @@ class ToolPermissionGuardrail(CustomGuardrail): return function_call return self._get_mapping_value(function_call, "name") - def _collect_request_tools(self, data: dict) -> List[tuple[str, Optional[str]]]: - request_tools: List[tuple[str, Optional[str]]] = [] + def _collect_request_tools(self, data: dict) -> list[tuple[str, str | None]]: + request_tools: list[tuple[str, str | None]] = [] for tool in data.get("tools") or []: tool_name, tool_type = self._get_request_tool_name(tool) @@ -527,7 +527,7 @@ class ToolPermissionGuardrail(CustomGuardrail): def _modify_request_with_permission_errors( self, data: dict, - denied_tool_names: List[str], + denied_tool_names: list[str], ): """ Modify the request to replace denied tool_calls blocks with error results @@ -546,7 +546,7 @@ class ToolPermissionGuardrail(CustomGuardrail): for tool_use in denied_tool_names: error_tool_names.add(tool_use) - tools: Optional[List[ChatCompletionToolParam]] = data.get("tools") + tools: list[ChatCompletionToolParam] | None = data.get("tools") if tools is not None: new_tools = [] for tool in tools: @@ -594,7 +594,7 @@ class ToolPermissionGuardrail(CustomGuardrail): def _modify_response_with_permission_errors( self, response: ModelResponse, - denied_tools: List[tuple[ChatCompletionMessageToolCall, PermissionError]], + denied_tools: list[tuple[ChatCompletionMessageToolCall, PermissionError]], ) -> None: """ Modify the response to replace denied tool_calls blocks with error results @@ -655,7 +655,7 @@ class ToolPermissionGuardrail(CustomGuardrail): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ """ verbose_proxy_logger.debug("Tool Permission Guardrail Pre-Call Hook") @@ -789,11 +789,11 @@ class ToolPermissionGuardrail(CustomGuardrail): from litellm.types.utils import TextCompletionResponse # Collect all chunks to process them together - all_chunks: List[ModelResponseStream] = [] + all_chunks: list[ModelResponseStream] = [] async for chunk in response: all_chunks.append(chunk) - assembled_model_response: Optional[Union[ModelResponse, TextCompletionResponse]] = stream_chunk_builder( + assembled_model_response: ModelResponse | TextCompletionResponse | None = stream_chunk_builder( chunks=all_chunks, ) if isinstance(assembled_model_response, ModelResponse): diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index 7b9e88fb6e5..3f139d1deb0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -20,7 +20,7 @@ Configuration in proxy config YAML: mode: post_call """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal, Optional from fastapi import HTTPException @@ -41,7 +41,7 @@ GUARDRAIL_NAME = "tool_policy" def _get_request_object_permission_ids( request_data: dict, -) -> Tuple[Optional[str], Optional[str]]: +) -> tuple[str | None, str | None]: """Extract object_permission_id and team_object_permission_id from request_data.""" if not request_data: return None, None @@ -68,7 +68,7 @@ def _get_request_object_permission_ids( return None, None -def _get_request_route_from_data(request_data: dict) -> Optional[str]: +def _get_request_route_from_data(request_data: dict) -> str | None: """Get request route from request_data (metadata or top-level).""" route = request_data.get("user_api_key_request_route") if route: @@ -77,12 +77,12 @@ def _get_request_route_from_data(request_data: dict) -> Optional[str]: return meta.get("user_api_key_request_route") -def _resolve_tool_names_from_messages(messages: List[dict]) -> Dict[str, str]: +def _resolve_tool_names_from_messages(messages: list[dict]) -> dict[str, str]: """ Build a map of tool_call_id -> tool_name from assistant messages' tool_calls. Used to resolve which tool produced each tool result in the conversation. """ - mapping: Dict[str, str] = {} + mapping: dict[str, str] = {} for msg in messages: if msg.get("role") != "assistant": continue @@ -192,7 +192,7 @@ class ToolPolicyGuardrail(CustomGuardrail): messages = request_data.get("messages") or [] tc_id_to_name = _resolve_tool_names_from_messages(messages) - untrusted_sources: List[str] = [] + untrusted_sources: list[str] = [] for msg in messages: if msg.get("role") != "tool": continue diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 8d94dd74180..6130eec3281 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -9,7 +9,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint import copy import json from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, List, Union +from typing import TYPE_CHECKING, Any from fastapi import HTTPException @@ -46,7 +46,7 @@ class _StreamTerminated(Exception): its terminal chunks (block message or in-stream error) and must stop.""" -def _get_a2a_request_id(responses_so_far: List[Any], request_data: dict) -> str | None: +def _get_a2a_request_id(responses_so_far: list[Any], request_data: dict) -> str | None: """Get JSON-RPC request id from first A2A chunk or request body for in-stream error reporting.""" for item in responses_so_far: if isinstance(item, dict) and "id" in item: @@ -100,7 +100,7 @@ class UnifiedLLMGuardrails(CustomLogger): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Union[Exception, str, dict, None]: + ) -> Exception | str | dict | None: """ Runs before the LLM API call Runs on only Input @@ -804,7 +804,7 @@ class UnifiedLLMGuardrails(CustomLogger): user_api_key_dict: UserAPIKeyAuth, response: Any, request_data: dict, - guardrail_to_apply: Union[CustomGuardrail, None] = None, + guardrail_to_apply: CustomGuardrail | None = None, buffer_until_moderated_default: bool = False, ) -> AsyncGenerator[Any, None]: """ @@ -926,7 +926,7 @@ class UnifiedLLMGuardrails(CustomLogger): # Infer call type from first chunk call_type = None chunk_counter = 0 - responses_so_far: List[Any] = [] + responses_so_far: list[Any] = [] responses_yielded: list[Any] = [] pending_end_of_stream_items: list[Any] = [] # Whether any real response chunk has been forwarded to the client. diff --git a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py index 7d72a233d3d..609f89514b1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py @@ -3,13 +3,9 @@ from json import JSONDecodeError from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, Protocol, - Tuple, - Type, cast, ) @@ -69,8 +65,8 @@ class _AsyncPostHandler(Protocol): self, *, url: str, - headers: Dict[str, str], - json: Dict[str, Any], + headers: dict[str, str], + json: dict[str, Any], timeout: httpx.Timeout, ) -> Awaitable[httpx.Response]: ... @@ -82,11 +78,11 @@ class VigilGuardMissingConfig(ValueError): class VigilGuardGuardrail(CustomGuardrail): def __init__( self, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - unreachable_fallback: Optional[str] = None, - timeout: Optional[float] = None, - async_handler: Optional[_AsyncPostHandler] = None, + api_base: str | None = None, + api_key: str | None = None, + unreachable_fallback: str | None = None, + timeout: float | None = None, + async_handler: _AsyncPostHandler | None = None, **kwargs: Any, ) -> None: resolved_base = api_base or get_secret_str("VIGIL_GUARD_URL") @@ -121,7 +117,7 @@ class VigilGuardGuardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import ( VigilGuardGuardrailConfigModel, ) @@ -129,7 +125,7 @@ class VigilGuardGuardrail(CustomGuardrail): return VigilGuardGuardrailConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -152,7 +148,7 @@ class VigilGuardGuardrail(CustomGuardrail): source = "user_input" if input_type == "request" else "model_output" metadata = self._collect_metadata(request_data, logging_obj) - result_texts: List[str] = [] + result_texts: list[str] = [] for index, text in enumerate(texts): if not isinstance(text, str) or not text.strip(): result_texts.append(text) @@ -255,7 +251,7 @@ class VigilGuardGuardrail(CustomGuardrail): exc: Exception, inputs: GenericGuardrailAPIInputs, source: str, - final_texts: List[Any], + final_texts: list[Any], final_tool_calls: Any, ) -> GenericGuardrailAPIInputs: if self.unreachable_fallback == "fail_open": @@ -282,7 +278,7 @@ class VigilGuardGuardrail(CustomGuardrail): @staticmethod def _build_output( inputs: GenericGuardrailAPIInputs, - final_texts: List[Any], + final_texts: list[Any], final_tool_calls: Any, ) -> GenericGuardrailAPIInputs: # When nothing was changed, return the input shape verbatim so the guardrail @@ -303,8 +299,8 @@ class VigilGuardGuardrail(CustomGuardrail): return guardrailed @staticmethod - def _tool_call_arguments(tool_calls: Any) -> List[Tuple[int, str]]: - pairs: List[Tuple[int, str]] = [] + def _tool_call_arguments(tool_calls: Any) -> list[tuple[int, str]]: + pairs: list[tuple[int, str]] = [] if isinstance(tool_calls, list): for index, tool_call in enumerate(tool_calls): function = tool_call.get("function") if isinstance(tool_call, dict) else None @@ -314,7 +310,7 @@ class VigilGuardGuardrail(CustomGuardrail): return pairs @staticmethod - def _set_tool_call_arguments(tool_calls: Any, index: int, arguments: str) -> List[Any]: + def _set_tool_call_arguments(tool_calls: Any, index: int, arguments: str) -> list[Any]: updated = list(tool_calls) tool_call = dict(updated[index]) function = dict(tool_call.get("function") or {}) @@ -323,7 +319,7 @@ class VigilGuardGuardrail(CustomGuardrail): updated[index] = tool_call return updated - async def _analyze(self, text: str, source: str, metadata: Dict[str, Any]) -> Dict[str, Any]: + async def _analyze(self, text: str, source: str, metadata: dict[str, Any]) -> dict[str, Any]: payload = { "text": text, "source": source, @@ -338,7 +334,7 @@ class VigilGuardGuardrail(CustomGuardrail): response = await self._post_with_retry(endpoint, headers, payload) return response.json() - async def _post_with_retry(self, endpoint: str, headers: Dict[str, str], payload: Dict[str, Any]) -> httpx.Response: + async def _post_with_retry(self, endpoint: str, headers: dict[str, str], payload: dict[str, Any]) -> httpx.Response: for attempt in range(2): try: response = await self.async_handler.post( @@ -375,7 +371,7 @@ class VigilGuardGuardrail(CustomGuardrail): ) @staticmethod - def _build_block_reason(analysis: Dict[str, Any]) -> str: + def _build_block_reason(analysis: dict[str, Any]) -> str: for key in ("blockMessage", "decisionReason"): value = analysis.get(key) if isinstance(value, str) and value.strip(): @@ -388,15 +384,15 @@ class VigilGuardGuardrail(CustomGuardrail): return "Blocked by policy" @staticmethod - def _resolve_sanitized_text(original: str, analysis: Dict[str, Any]) -> str: + def _resolve_sanitized_text(original: str, analysis: dict[str, Any]) -> str: for key in ("sanitizedText", "outputText"): value = analysis.get(key) if isinstance(value, str): return value return original - def _collect_metadata(self, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]) -> Dict[str, Any]: - sources: List[dict] = [] + def _collect_metadata(self, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]) -> dict[str, Any]: + sources: list[dict] = [] if isinstance(request_data, dict): sources.append(request_data) for nested_key in ("metadata", "litellm_metadata"): @@ -404,7 +400,7 @@ class VigilGuardGuardrail(CustomGuardrail): if isinstance(nested, dict): sources.append(nested) - collected: Dict[str, Any] = {} + collected: dict[str, Any] = {} for field in _METADATA_ALLOWLIST: for source in sources: if field in source and source[field] is not None: @@ -428,7 +424,7 @@ class VigilGuardGuardrail(CustomGuardrail): if isinstance(value, (int, float)): return value if isinstance(value, list): - clamped: List[Any] = [] + clamped: list[Any] = [] for item in value[:_METADATA_ARRAY_MAX_ITEMS]: if isinstance(item, bool): continue @@ -440,7 +436,7 @@ class VigilGuardGuardrail(CustomGuardrail): return None @staticmethod - def _extract_call_id(request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]) -> Optional[str]: + def _extract_call_id(request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]) -> str | None: if logging_obj is not None: call_id = getattr(logging_obj, "litellm_call_id", None) if isinstance(call_id, str) and call_id: diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index 7fe942bcb38..3f681d2e3d3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -24,19 +24,14 @@ Design notes (intentional divergences from the framework defaults): import asyncio import os +from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Type, ) -from datetime import datetime - from fastapi.exceptions import HTTPException from litellm._logging import verbose_proxy_logger @@ -94,12 +89,12 @@ class XecGuardMissingCredentials(Exception): class XecGuardGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - xecguard_model: Optional[str] = None, - policy_names: Optional[List[str]] = None, - block_on_error: Optional[bool] = None, - grounding_strictness: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + xecguard_model: str | None = None, + policy_names: list[str] | None = None, + block_on_error: bool | None = None, + grounding_strictness: str | None = None, **kwargs: Any, ) -> None: self.api_key = api_key or os.environ.get("XECGUARD_API_KEY") @@ -137,7 +132,7 @@ class XecGuardGuardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.xecguard import ( XecGuardConfigModel, ) @@ -145,7 +140,7 @@ class XecGuardGuardrail(CustomGuardrail): return XecGuardConfigModel @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, @@ -208,7 +203,7 @@ class XecGuardGuardrail(CustomGuardrail): kwargs: dict, result: Any, call_type: str, - ) -> Tuple[dict, Any]: + ) -> tuple[dict, Any]: """Observe-only scan for logging_only mode. Never blocks, never raises - all errors are swallowed. Records a @@ -287,7 +282,7 @@ class XecGuardGuardrail(CustomGuardrail): kwargs: dict, result: Any, call_type: str, - ) -> Tuple[dict, Any]: + ) -> tuple[dict, Any]: """Sync counterpart to ``async_logging_hook``. Runs the async version on an available loop, swallowing every @@ -316,11 +311,11 @@ class XecGuardGuardrail(CustomGuardrail): async def _call_scan( self, - messages: List[dict], + messages: list[dict], scan_type: str, suppress_errors: bool = False, - ) -> Optional[dict]: - payload: Dict[str, Any] = { + ) -> dict | None: + payload: dict[str, Any] = { "model": self.xecguard_model, "scan_type": scan_type, "messages": messages, @@ -334,9 +329,9 @@ class XecGuardGuardrail(CustomGuardrail): async def _call_grounding( self, - messages: List[dict], - documents: List[dict], - ) -> Optional[dict]: + messages: list[dict], + documents: list[dict], + ) -> dict | None: prompt = self._extract_last_text_by_role(messages, "user") response_text = self._extract_last_text_by_role(messages, "assistant") if prompt is None or response_text is None: @@ -355,7 +350,7 @@ class XecGuardGuardrail(CustomGuardrail): path: str, payload: dict, suppress_errors: bool = False, - ) -> Optional[dict]: + ) -> dict | None: endpoint = f"{self.api_base}{path}" verbose_proxy_logger.debug( "XecGuard: POST %s payload_keys=%s", @@ -397,7 +392,7 @@ class XecGuardGuardrail(CustomGuardrail): request_data: dict, inputs: Any, input_type: str, - ) -> List[dict]: + ) -> list[dict]: """Assemble the full message list that will be sent to XecGuard. Always reads from ``request_data['messages']`` so the framework's @@ -406,7 +401,7 @@ class XecGuardGuardrail(CustomGuardrail): the request data is incomplete. """ raw_messages = request_data.get("messages") or [] - messages: List[dict] = [self._normalize_message(m) for m in raw_messages if isinstance(m, dict)] + messages: list[dict] = [self._normalize_message(m) for m in raw_messages if isinstance(m, dict)] if input_type == "request": if not messages: @@ -433,7 +428,7 @@ class XecGuardGuardrail(CustomGuardrail): if isinstance(content, str): return {"role": role, "content": content} if isinstance(content, list): - parts: List[str] = [] + parts: list[str] = [] for item in content: if isinstance(item, dict) and item.get("type") == "text": text = item.get("text") @@ -443,7 +438,7 @@ class XecGuardGuardrail(CustomGuardrail): return {"role": role, "content": ""} @staticmethod - def _synthesize_user_from_inputs(inputs: Any) -> Optional[dict]: + def _synthesize_user_from_inputs(inputs: Any) -> dict | None: if not isinstance(inputs, dict): return None texts = inputs.get("texts") @@ -455,7 +450,7 @@ class XecGuardGuardrail(CustomGuardrail): return {"role": "user", "content": joined} @staticmethod - def _extract_last_text_by_role(messages: List[dict], role: str) -> Optional[str]: + def _extract_last_text_by_role(messages: list[dict], role: str) -> str | None: for message in reversed(messages): if message.get("role") == role: content = message.get("content") @@ -465,7 +460,7 @@ class XecGuardGuardrail(CustomGuardrail): return None @staticmethod - def _extract_assistant_text_from_response(response: Any) -> Optional[str]: + def _extract_assistant_text_from_response(response: Any) -> str | None: if response is None: return None choices = None @@ -475,7 +470,7 @@ class XecGuardGuardrail(CustomGuardrail): choices = response.get("choices") if not choices: return None - text_parts: List[str] = [] + text_parts: list[str] = [] for choice in choices: content = XecGuardGuardrail._extract_choice_content(choice) text = XecGuardGuardrail._content_to_text(content) @@ -500,7 +495,7 @@ class XecGuardGuardrail(CustomGuardrail): return None @staticmethod - def _content_to_text(content: Any) -> Optional[str]: + def _content_to_text(content: Any) -> str | None: if isinstance(content, str) and content: return content if isinstance(content, list): @@ -518,14 +513,14 @@ class XecGuardGuardrail(CustomGuardrail): # ------------------------------------------------------------------ @staticmethod - def _extract_grounding_documents(request_data: dict) -> List[dict]: + def _extract_grounding_documents(request_data: dict) -> list[dict]: metadata = request_data.get("metadata") or request_data.get("litellm_metadata") if not isinstance(metadata, dict): return [] raw_docs = metadata.get(_METADATA_GROUNDING_KEY) if not isinstance(raw_docs, list) or not raw_docs: return [] - valid_docs: List[dict] = [] + valid_docs: list[dict] = [] for doc in raw_docs: if ( isinstance(doc, dict) @@ -555,7 +550,7 @@ class XecGuardGuardrail(CustomGuardrail): violations = result.get("xecguard_result") if not isinstance(violations, list): violations = [] - seen: List[str] = [] + seen: list[str] = [] for v in violations: if not isinstance(v, dict): continue @@ -576,7 +571,7 @@ class XecGuardGuardrail(CustomGuardrail): def _format_grounding_block_message(result: dict) -> str: trace_id = result.get("trace_id", "") detail = result.get("xecguard_result") - rules: List[str] = [] + rules: list[str] = [] rationale = "" if isinstance(detail, dict): raw_rules = detail.get("violated_rules_list") diff --git a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py index 65338827e07..09e5ffff193 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py @@ -4,7 +4,7 @@ # # +-------------------------------------------------------------+ import os -from typing import TYPE_CHECKING, List, Literal, Optional +from typing import TYPE_CHECKING, Literal, Optional from fastapi import HTTPException @@ -29,7 +29,7 @@ GUARDRAIL_TIMEOUT = 5 class ZscalerAIGuard(CustomGuardrail): @classmethod - def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, @@ -37,12 +37,12 @@ class ZscalerAIGuard(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - policy_id: Optional[int] = None, - send_user_api_key_alias: Optional[bool] = None, - send_user_api_key_user_id: Optional[bool] = None, - send_user_api_key_team_id: Optional[bool] = None, + api_key: str | None = None, + api_base: str | None = None, + policy_id: int | None = None, + send_user_api_key_alias: bool | None = None, + send_user_api_key_user_id: bool | None = None, + send_user_api_key_team_id: bool | None = None, **kwargs, ): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -80,7 +80,7 @@ class ZscalerAIGuard(CustomGuardrail): verbose_proxy_logger.debug("ZscalerAIGuard Initializing ...") @staticmethod - def _resolve_metadata_value(request_data: Optional[dict], key: str) -> Optional[str]: + def _resolve_metadata_value(request_data: dict | None, key: str) -> str | None: """ Resolve metadata value from request_data, checking both metadata locations. @@ -351,11 +351,11 @@ class ZscalerAIGuard(CustomGuardrail): return self._handle_response(response, direction) except Exception as e: verbose_proxy_logger.error(f"{e}. Blocking request.") - user_facing_error = self._create_user_facing_error(f"{str(e)}") + user_facing_error = self._create_user_facing_error(f"{e!s}") raise HTTPException(status_code=500, detail=user_facing_error) @staticmethod - def get_config_model() -> Optional[type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: from litellm.types.proxy.guardrails.guardrail_hooks.zscaler_ai_guard import ( ZscalerAIGuardConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index e909c15382b..95ec74b32df 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -1,5 +1,5 @@ # litellm/proxy/guardrails/guardrail_initializers.py -from typing import Any, Dict, List, Optional +from typing import Any import litellm from litellm.proxy._types import CommonProxyErrors @@ -152,7 +152,7 @@ def initialize_tool_permission(litellm_params: LitellmParams, guardrail: Guardra ToolPermissionGuardrail, ) - rules: Optional[List[Dict[str, Any]]] = None + rules: list[dict[str, Any]] | None = None if litellm_params.rules: rules = [] for rule in litellm_params.rules: diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 82cc97df7f9..aaaef95f4a4 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -4,7 +4,7 @@ import importlib import os from datetime import datetime, timezone from itertools import chain, count -from typing import Any, Dict, List, Literal, Optional, Set, Type, cast +from typing import Any, Literal, Optional, cast from pydantic import ValidationError @@ -68,7 +68,7 @@ guardrail_initializer_registry = { CONFIG_GUARDRAIL_ID_NAMESPACE = uuid.UUID("625f63f4-935a-50e5-98b5-fbe77babc74a") -guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = { +guardrail_class_registry: dict[str, type[CustomGuardrail]] = { SupportedGuardrailIntegrations.BEDROCK.value: BedrockGuardrail, SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail, SupportedGuardrailIntegrations.LAKERA.value: lakeraAI_Moderation, @@ -239,7 +239,7 @@ class GuardrailRegistry: ########################################################### ########### In memory management helpers for guardrails ########### ############################################################ - def get_initialized_guardrail_callback(self, guardrail_name: str) -> Optional[CustomGuardrail]: + def get_initialized_guardrail_callback(self, guardrail_name: str) -> CustomGuardrail | None: """ Returns the initialized guardrail callback for a given guardrail name """ @@ -285,7 +285,7 @@ class GuardrailRegistry: return guardrail_dict except Exception as e: - raise Exception(f"Error adding guardrail to DB: {str(e)}") + raise Exception(f"Error adding guardrail to DB: {e!s}") async def delete_guardrail_from_db(self, guardrail_id: str, prisma_client: PrismaClient): """ @@ -297,7 +297,7 @@ class GuardrailRegistry: return {"message": f"Guardrail {guardrail_id} deleted successfully"} except Exception as e: - raise Exception(f"Error deleting guardrail from DB: {str(e)}") + raise Exception(f"Error deleting guardrail from DB: {e!s}") async def update_guardrail_in_db(self, guardrail_id: str, guardrail: Guardrail, prisma_client: PrismaClient): """ @@ -328,12 +328,12 @@ class GuardrailRegistry: # Convert to dict and return return dict(updated_guardrail) except Exception as e: - raise Exception(f"Error updating guardrail in DB: {str(e)}") + raise Exception(f"Error updating guardrail in DB: {e!s}") @staticmethod async def get_all_guardrails_from_db( prisma_client: PrismaClient, - ) -> List[Guardrail]: + ) -> list[Guardrail]: """ Get all active guardrails from the database. Only rows with status == "active" are returned (pending_review and rejected are excluded). @@ -344,15 +344,15 @@ class GuardrailRegistry: order={"created_at": "desc"}, ) - guardrails: List[Guardrail] = [] + guardrails: list[Guardrail] = [] for guardrail in guardrails_from_db: guardrails.append(Guardrail(**(dict(guardrail)))) # type: ignore return guardrails except Exception as e: - raise Exception(f"Error getting guardrails from DB: {str(e)}") + raise Exception(f"Error getting guardrails from DB: {e!s}") - async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Optional[Guardrail]: + async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Guardrail | None: """ Get a guardrail by its ID from the database """ @@ -366,11 +366,9 @@ class GuardrailRegistry: return Guardrail(**(dict(guardrail))) # type: ignore except Exception as e: - raise Exception(f"Error getting guardrail from DB: {str(e)}") + raise Exception(f"Error getting guardrail from DB: {e!s}") - async def get_guardrail_by_name_from_db( - self, guardrail_name: str, prisma_client: PrismaClient - ) -> Optional[Guardrail]: + async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None: """ Get a guardrail by its name from the database """ @@ -384,7 +382,7 @@ class GuardrailRegistry: return Guardrail(**(dict(guardrail))) # type: ignore except Exception as e: - raise Exception(f"Error getting guardrail from DB: {str(e)}") + raise Exception(f"Error getting guardrail from DB: {e!s}") class InMemoryGuardrailHandler: @@ -393,17 +391,17 @@ class InMemoryGuardrailHandler: """ def __init__(self): - self.IN_MEMORY_GUARDRAILS: Dict[str, Guardrail] = {} + self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {} """ Guardrail id to Guardrail object mapping """ - self.guardrail_id_to_custom_guardrail: Dict[str, Optional[CustomGuardrail]] = {} + self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {} """ Guardrail id to CustomGuardrail object mapping """ - self._sources: Dict[str, Literal["db", "config"]] = {} + self._sources: dict[str, Literal["db", "config"]] = {} """ Guardrail id to provenance marker. "db" entries are reconciled against the DB on each polling tick; "config" entries are owned by proxy_config.yaml @@ -418,10 +416,10 @@ class InMemoryGuardrailHandler: def initialize_guardrail( self, guardrail: Guardrail, - config_file_path: Optional[str] = None, + config_file_path: str | None = None, llm_router: Optional["Router"] = None, source: Literal["db", "config"] = "config", - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -437,7 +435,7 @@ class InMemoryGuardrailHandler: self._sources[guardrail_id] = source return self.IN_MEMORY_GUARDRAILS[guardrail_id] - custom_guardrail_callback: Optional[CustomGuardrail] = None + custom_guardrail_callback: CustomGuardrail | None = None litellm_params_data = guardrail["litellm_params"] verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) @@ -517,11 +515,11 @@ class InMemoryGuardrailHandler: def initialize_custom_guardrail( self, - guardrail: Dict, + guardrail: dict, guardrail_type: str, litellm_params: LitellmParams, - config_file_path: Optional[str] = None, - ) -> Optional[CustomGuardrail]: + config_file_path: str | None = None, + ) -> CustomGuardrail | None: """ Initialize a Custom Guardrail from a python file or module path @@ -606,25 +604,25 @@ class InMemoryGuardrailHandler: litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) - def list_in_memory_guardrails(self) -> List[Guardrail]: + def list_in_memory_guardrails(self) -> list[Guardrail]: """ List all guardrails in memory """ return list(self.IN_MEMORY_GUARDRAILS.values()) - def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]: + def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: """ Get a guardrail by its ID from memory """ return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) - def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]: + def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None: """ Return the provenance of an in-memory guardrail. """ return self._sources.get(guardrail_id) - def list_config_guardrails(self) -> List[Guardrail]: + def list_config_guardrails(self) -> list[Guardrail]: """ List in-memory guardrails owned by config.yaml. @@ -634,7 +632,7 @@ class InMemoryGuardrailHandler: """ return [g for gid, g in self.IN_MEMORY_GUARDRAILS.items() if self._sources.get(gid) == "config"] - def get_config_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]: + def get_config_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: """ Get a config-owned in-memory guardrail by its ID, or None. @@ -645,7 +643,7 @@ class InMemoryGuardrailHandler: return None return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) - def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]: + def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]: """ Drop in-memory entries that originated from the DB but are no longer present in db_guardrail_ids. Config-loaded guardrails are never touched. @@ -668,8 +666,8 @@ class InMemoryGuardrailHandler: @staticmethod def _normalize_litellm_params_for_comparison( - params: Optional[Any], - ) -> Optional[Dict[str, Any]]: + params: Any | None, + ) -> dict[str, Any] | None: """ Render litellm_params to a canonical dict so an in-memory LitellmParams and the raw dict loaded from the DB compare equal when they describe the same @@ -733,9 +731,9 @@ class InMemoryGuardrailHandler: def reinitialize_guardrail( self, guardrail: Guardrail, - config_file_path: Optional[str] = None, + config_file_path: str | None = None, source: Literal["db", "config"] = "config", - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Force re-initialization of a guardrail even if it exists in memory. Removes old callback from litellm.callbacks and creates fresh instance. @@ -752,9 +750,7 @@ class InMemoryGuardrailHandler: # Initialize fresh (will add new callback to litellm.callbacks) return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source) - def sync_guardrail_from_db( - self, guardrail: Guardrail, config_file_path: Optional[str] = None - ) -> Optional[Guardrail]: + def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None: """ Sync a guardrail from DB - initializes if new, re-initializes if changed. This is the method to call during DB polling. diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index 34961afb8ff..71ffc9d36ef 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast import litellm from litellm import Router @@ -8,7 +8,7 @@ from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_pr # v2 implementation from litellm.types.guardrails import Guardrail, GuardrailItem, GuardrailItemSpec -all_guardrails: List[GuardrailItem] = [] +all_guardrails: list[GuardrailItem] = [] """ Map guardrail_name: , , during_call @@ -17,13 +17,13 @@ Map guardrail_name: , , during_call def init_guardrails_v2( - all_guardrails: List[Dict], - config_file_path: Optional[str] = None, - llm_router: Optional[Router] = None, + all_guardrails: list[dict], + config_file_path: str | None = None, + llm_router: Router | None = None, ): from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER - guardrail_list: List[Guardrail] = [] + guardrail_list: list[Guardrail] = [] for guardrail in all_guardrails: initialized_guardrail = IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail( @@ -41,7 +41,7 @@ def init_guardrails_v2( _populate_router_guardrail_list(guardrail_list=guardrail_list) -def _populate_router_guardrail_list(guardrail_list: List[Guardrail]) -> None: +def _populate_router_guardrail_list(guardrail_list: list[Guardrail]) -> None: """ Populate the router's guardrail_list from initialized guardrails. @@ -56,7 +56,7 @@ def _populate_router_guardrail_list(guardrail_list: List[Guardrail]) -> None: verbose_proxy_logger.debug("Router not initialized yet, skipping guardrail_list population") return - router_guardrail_list: List[GuardrailTypedDict] = [] + router_guardrail_list: list[GuardrailTypedDict] = [] for guardrail in guardrail_list: guardrail_id = guardrail.get("guardrail_id") @@ -91,11 +91,11 @@ def _populate_router_guardrail_list(guardrail_list: List[Guardrail]) -> None: ### LEGACY IMPLEMENTATION ### def initialize_guardrails( - guardrails_config: List[Dict[str, GuardrailItemSpec]], + guardrails_config: list[dict[str, GuardrailItemSpec]], premium_user: bool, config_file_path: str, litellm_settings: dict, -) -> Dict[str, GuardrailItem]: +) -> dict[str, GuardrailItem]: try: verbose_proxy_logger.debug(f"validating guardrails passed {guardrails_config}") global all_guardrails @@ -141,5 +141,5 @@ def initialize_guardrails( return litellm.guardrail_name_config_map except Exception as e: - verbose_proxy_logger.exception("error initializing guardrails {}".format(str(e))) + verbose_proxy_logger.exception(f"error initializing guardrails {e!s}") raise e diff --git a/litellm/proxy/guardrails/tool_name_extraction.py b/litellm/proxy/guardrails/tool_name_extraction.py index 02f11e3c20a..f4afcc24232 100644 --- a/litellm/proxy/guardrails/tool_name_extraction.py +++ b/litellm/proxy/guardrails/tool_name_extraction.py @@ -6,19 +6,19 @@ knowledge lives in one place. Uses guardrail translation handlers where availabl with standalone extractors for generate_content and MCP. """ -from typing import Any, Dict, List +from typing import Any from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route from litellm.llms import load_guardrail_translation_mappings from litellm.types.utils import CallTypes # Call types that have no guardrail translation handler; we use standalone extractors -STANDALONE_EXTRACTORS: Dict[str, Any] = {} +STANDALONE_EXTRACTORS: dict[str, Any] = {} -def _extract_generate_content_tool_names(data: dict) -> List[str]: +def _extract_generate_content_tool_names(data: dict) -> list[str]: """Google generateContent: tools[].functionDeclarations[].name""" - names: List[str] = [] + names: list[str] = [] for tool in data.get("tools") or []: if not isinstance(tool, dict): continue @@ -28,9 +28,9 @@ def _extract_generate_content_tool_names(data: dict) -> List[str]: return names -def _extract_mcp_tool_names(data: dict) -> List[str]: +def _extract_mcp_tool_names(data: dict) -> list[str]: """MCP call_tool: name or mcp_tool_name in body""" - names: List[str] = [] + names: list[str] = [] name = data.get("name") or data.get("mcp_tool_name") if name: names.append(str(name)) @@ -60,7 +60,7 @@ TOOL_CAPABLE_CALL_TYPES = frozenset( ) -def extract_request_tool_names(route: str, data: dict) -> List[str]: +def extract_request_tool_names(route: str, data: dict) -> list[str]: """ Extract tool names from the request body for the given route. Uses guardrail translation handlers when available, else standalone extractors diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index eb1979074d4..bfa8806e89e 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -6,7 +6,7 @@ insert into SpendLogGuardrailIndex when spend logs are written. import json from collections import defaultdict from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.proxy.utils import PrismaClient @@ -16,7 +16,7 @@ from litellm.repositories.table_repositories import ( ) -def _guardrail_status_to_action(status: Optional[str]) -> str: +def _guardrail_status_to_action(status: str | None) -> str: """Map StandardLogging guardrail_status to blocked/passed/flagged.""" if not status: return "passed" @@ -28,7 +28,7 @@ def _guardrail_status_to_action(status: Optional[str]) -> str: return "passed" -def _parse_guardrail_info_from_payload(payload: Dict[str, Any]) -> List[Dict[str, Any]]: +def _parse_guardrail_info_from_payload(payload: dict[str, Any]) -> list[dict[str, Any]]: """Extract guardrail_information from spend log payload metadata.""" meta = payload.get("metadata") if not meta: @@ -55,7 +55,7 @@ def _date_str(dt: datetime) -> str: async def process_spend_logs_guardrail_usage( prisma_client: PrismaClient, - logs_to_process: List[Dict[str, Any]], + logs_to_process: list[dict[str, Any]], ) -> None: """ After spend logs are written: update DailyGuardrailMetrics and insert @@ -64,7 +64,7 @@ async def process_spend_logs_guardrail_usage( if not logs_to_process: return # Aggregate daily metrics by (guardrail_id, date). Latency/score metrics dropped. - daily_guardrail: Dict[tuple, Dict[str, Any]] = defaultdict( + daily_guardrail: dict[tuple, dict[str, Any]] = defaultdict( lambda: { "requests_evaluated": 0, "passed_count": 0, @@ -72,7 +72,7 @@ async def process_spend_logs_guardrail_usage( "flagged_count": 0, } ) - index_rows: List[Dict[str, Any]] = [] + index_rows: list[dict[str, Any]] = [] for payload in logs_to_process: request_id = payload.get("request_id") diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 29c023df1a6..f62bf27f774 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -7,7 +7,6 @@ import sys import threading import time from collections.abc import Mapping -from typing import List, Optional import litellm @@ -87,7 +86,7 @@ def _should_inject_health_check_max_tokens(model_info: Mapping[str, object], mod _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT = frozenset((None, "chat", "completion")) -def _get_process_rss_mb() -> Optional[float]: +def _get_process_rss_mb() -> float | None: """ Get process RSS memory in MB. On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes. @@ -119,7 +118,7 @@ def _get_random_llm_message(): return [{"role": "user", "content": random.choice(messages)}] -def _clean_endpoint_data(endpoint_data: dict, details: Optional[bool] = True): +def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True): """ Clean the endpoint data for display to users. """ @@ -132,7 +131,7 @@ def _clean_endpoint_data(endpoint_data: dict, details: Optional[bool] = True): def health_check_filter_kwargs_from_general_settings( - general_settings: Optional[dict], + general_settings: dict | None, ) -> dict: """ Build kwargs for ``perform_health_check`` from ``general_settings``. @@ -150,8 +149,8 @@ def health_check_filter_kwargs_from_general_settings( def filter_deployments_by_id( - model_list: List, -) -> List: + model_list: list, +) -> list: seen_ids = set() filtered_deployments = [] @@ -268,9 +267,9 @@ async def _run_health_checks_with_bounded_concurrency(models: list, concurrency_ async def _perform_health_check( model_list: list, - details: Optional[bool] = True, - max_concurrency: Optional[int] = None, - instrumentation_context: Optional[dict] = None, + details: bool | None = True, + max_concurrency: int | None = None, + instrumentation_context: dict | None = None, ): """ Perform a health check for each model in the list. @@ -396,7 +395,7 @@ def _health_check_deployment_is_wildcard(litellm_params: dict) -> bool: return "*" in _deployment_model_string_for_health_check(litellm_params) -def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) -> Optional[int]: +def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) -> int | None: """ Pick max_tokens for the health check request. @@ -494,8 +493,7 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di model = litellm_params["model"] # Strip only the bedrock/ prefix (preserve routes like converse/, invoke/) - if model.startswith("bedrock/"): - model = model[8:] # len("bedrock/") = 8 + model = model.removeprefix("bedrock/") # len("bedrock/") = 8 # Now check for region routing and strip it if present # Need to handle formats like: @@ -524,12 +522,12 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di async def perform_health_check( model_list: list, - model: Optional[str] = None, - cli_model: Optional[str] = None, - details: Optional[bool] = True, - model_id: Optional[str] = None, - max_concurrency: Optional[int] = None, - instrumentation_context: Optional[dict] = None, + model: str | None = None, + cli_model: str | None = None, + details: bool | None = True, + model_id: str | None = None, + max_concurrency: int | None = None, + instrumentation_context: dict | None = None, health_check_skip_disabled_background_models: bool = False, ): """ diff --git a/litellm/proxy/health_check_utils/shared_health_check_manager.py b/litellm/proxy/health_check_utils/shared_health_check_manager.py index eb6d27ca20b..89a0ddb6cbe 100644 --- a/litellm/proxy/health_check_utils/shared_health_check_manager.py +++ b/litellm/proxy/health_check_utils/shared_health_check_manager.py @@ -1,15 +1,15 @@ import asyncio import json import time -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from litellm._logging import verbose_proxy_logger -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.caching.redis_cache import RedisCache from litellm.constants import ( - DEFAULT_SHARED_HEALTH_CHECK_TTL, DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL, + DEFAULT_SHARED_HEALTH_CHECK_TTL, ) +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy.health_check import perform_health_check @@ -26,7 +26,7 @@ class SharedHealthCheckManager: def __init__( self, - redis_cache: Optional[RedisCache] = None, + redis_cache: RedisCache | None = None, health_check_ttl: int = DEFAULT_SHARED_HEALTH_CHECK_TTL, lock_ttl: int = DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL, ): @@ -100,7 +100,7 @@ class SharedHealthCheckManager: except Exception as e: verbose_proxy_logger.error("Error releasing health check lock: %s", str(e)) - async def get_cached_health_check_results(self) -> Optional[Dict[str, Any]]: + async def get_cached_health_check_results(self) -> dict[str, Any] | None: """ Get cached health check results from Redis. @@ -140,8 +140,8 @@ class SharedHealthCheckManager: async def cache_health_check_results( self, - healthy_endpoints: List[Dict[str, Any]], - unhealthy_endpoints: List[Dict[str, Any]], + healthy_endpoints: list[dict[str, Any]], + unhealthy_endpoints: list[dict[str, Any]], ) -> None: """ Cache health check results in Redis. @@ -181,11 +181,11 @@ class SharedHealthCheckManager: async def perform_shared_health_check( self, - model_list: List[Dict[str, Any]], + model_list: list[dict[str, Any]], details: bool = True, - max_concurrency: Optional[int] = None, + max_concurrency: int | None = None, health_check_skip_disabled_background_models: bool = False, - ) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]], Dict[str, Any]]: + ) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: """ Perform health check with shared state coordination. @@ -329,7 +329,7 @@ class SharedHealthCheckManager: verbose_proxy_logger.error("Error checking health check lock status: %s", str(e)) return False - async def get_health_check_status(self) -> Dict[str, Any]: + async def get_health_check_status(self) -> dict[str, Any]: """ Get the current status of health check coordination. diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 34841d2e692..03645c0b2fa 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -7,7 +7,7 @@ import time import traceback from collections.abc import Iterable from datetime import datetime, timedelta -from typing import Any, Dict, Literal, Optional, TypedDict, Union, cast +from typing import Any, Literal, TypedDict, Union, cast import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response, status @@ -237,7 +237,7 @@ async def health_services_endpoint( ) return { "status": "success", - "message": "Mock LLM request made - check {}.".format(service), + "message": f"Mock LLM request made - check {service}.", } elif service == "datadog": from litellm.integrations.datadog.datadog import DataDogLogger @@ -398,9 +398,7 @@ async def health_services_endpoint( else: raise HTTPException( status_code=422, - detail={ - "error": '"{}" not in proxy config: general_settings. Unable to test this.'.format(service) - }, + detail={"error": f'"{service}" not in proxy config: general_settings. Unable to test this.'}, ) if service == "email": webhook_event = WebhookEvent( @@ -427,13 +425,11 @@ async def health_services_endpoint( } except Exception as e: - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.health_services_endpoint(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.health_services_endpoint(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -484,8 +480,8 @@ async def _save_health_check_to_db( healthy_endpoints: list, unhealthy_endpoints: list, start_time: float, - user_id: Optional[str], - model_id: Optional[str] = None, + user_id: str | None, + model_id: str | None = None, ): """Helper function to save health check results to database""" try: @@ -608,7 +604,7 @@ async def _save_health_check_results_if_changed( model_results: dict, latest_checks_map: dict, start_time: float, - checked_by: Optional[str] = None, + checked_by: str | None = None, ): """ Save health check results to database, but only if status changed or >1 hour since last save. @@ -669,7 +665,7 @@ async def _save_background_health_checks_to_db( healthy_endpoints: list, unhealthy_endpoints: list, start_time: float, - checked_by: Optional[str] = None, + checked_by: str | None = None, ): """ Save background health check results to database for each model. @@ -760,7 +756,7 @@ def _strip_admin_only_fields_from_health_result(result: dict) -> dict: return out -def _resolve_targeted_model_ids(model_list: list, model: Optional[str], model_id: Optional[str]) -> Optional[set]: +def _resolve_targeted_model_ids(model_list: list, model: str | None, model_id: str | None) -> set | None: """ Resolve a ``/health`` ``model`` / ``model_id`` query param to the set of deployment IDs the response should be scoped to. @@ -871,10 +867,10 @@ async def _perform_health_check_and_save( def _health_endpoint_resolve_target_model_name( - model: Optional[str], - model_id: Optional[str], + model: str | None, + model_id: str | None, llm_router, -) -> Optional[str]: +) -> str | None: """Map ``model_id`` (without ``model``) to ``model_name`` for live health checks.""" if not model_id or model: return model @@ -903,8 +899,8 @@ def _health_endpoint_resolve_target_model_name( async def health_endpoint( response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - model: Optional[str] = fastapi.Query(None, description="Specify the model name (optional)"), - model_id: Optional[str] = fastapi.Query(None, description="Specify the model ID (optional)"), + model: str | None = fastapi.Query(None, description="Specify the model name (optional)"), + model_id: str | None = fastapi.Query(None, description="Specify the model ID (optional)"), ): """ 🚨 USE `/health/liveliness` to health check the proxy 🚨 @@ -1073,9 +1069,7 @@ async def health_endpoint( ) return _post_process(router_result) except Exception as e: - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.py::health_endpoint(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.py::health_endpoint(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) raise e @@ -1083,8 +1077,8 @@ async def health_endpoint( @router.get("/health/history", tags=["health"], dependencies=[Depends(user_api_key_auth)]) async def health_check_history_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - model: Optional[str] = fastapi.Query(None, description="Filter by specific model name"), - status_filter: Optional[str] = fastapi.Query(None, description="Filter by status (healthy/unhealthy)"), + model: str | None = fastapi.Query(None, description="Filter by specific model name"), + status_filter: str | None = fastapi.Query(None, description="Filter by status (healthy/unhealthy)"), limit: int = fastapi.Query(100, description="Number of records to return", ge=1, le=1000), offset: int = fastapi.Query(0, description="Number of records to skip", ge=0), ): @@ -1116,7 +1110,7 @@ async def health_check_history_endpoint( verbose_proxy_logger.error(f"Error getting health check history: {e}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to retrieve health check history: {str(e)}"}, + detail={"error": f"Failed to retrieve health check history: {e!s}"}, ) @@ -1148,7 +1142,7 @@ async def latest_health_checks_endpoint( verbose_proxy_logger.error(f"Error getting latest health checks: {e}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to retrieve latest health checks: {str(e)}"}, + detail={"error": f"Failed to retrieve latest health checks: {e!s}"}, ) @@ -1191,14 +1185,14 @@ async def shared_health_check_status_endpoint( verbose_proxy_logger.error(f"Error getting shared health check status: {e}") raise HTTPException( status_code=fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to retrieve shared health check status: {str(e)}"}, + detail={"error": f"Failed to retrieve shared health check status: {e!s}"}, ) -def _read_license_data() -> Optional[Dict[str, Any]]: +def _read_license_data() -> dict[str, Any] | None: from litellm.proxy.proxy_server import _license_check, premium_user_data - license_data: Optional[EnterpriseLicenseData] = premium_user_data or _license_check.airgapped_license_data + license_data: EnterpriseLicenseData | None = premium_user_data or _license_check.airgapped_license_data if ( license_data is None @@ -1217,10 +1211,10 @@ def _read_license_data() -> Optional[Dict[str, Any]]: if license_data is None: return None - return cast(Dict[str, Any], license_data) + return cast(dict[str, Any], license_data) -def _read_allowed_features(license_data: Dict[str, Any]) -> list: +def _read_allowed_features(license_data: dict[str, Any]) -> list: raw_allowed_features = license_data.get("allowed_features") if isinstance(raw_allowed_features, list): return list(raw_allowed_features) @@ -1408,8 +1402,8 @@ def callback_name(callback): async def _get_health_readiness_details( - response: Optional[Response] = None, -) -> Dict[str, Any]: + response: Response | None = None, +) -> dict[str, Any]: """ Detailed health payload for authenticated diagnostics. """ @@ -1479,7 +1473,7 @@ async def _get_health_readiness_details( "is_detailed_debug": is_detailed_debug, } except Exception as e: - raise HTTPException(status_code=503, detail=f"Service Unhealthy ({str(e)})") + raise HTTPException(status_code=503, detail=f"Service Unhealthy ({e!s})") def _allow_public_health_readiness_details() -> bool: @@ -1494,7 +1488,7 @@ def _drain_endpoint_enabled() -> bool: return general_settings.get("enable_drain_endpoint") is True -def _drain_endpoint_token() -> Optional[str]: +def _drain_endpoint_token() -> str | None: """ Shared secret required on the X-Drain-Token header to call /health/drain. @@ -1713,30 +1707,29 @@ async def health_liveliness_options(): ) async def test_model_connection( request: Request, - mode: Optional[ - Literal[ - "chat", - "completion", - "embedding", - "audio_speech", - "audio_transcription", - "image_generation", - "video_generation", - "batch", - "rerank", - "realtime", - "responses", - "ocr", - ] - ] = fastapi.Body( + mode: Literal[ + "chat", + "completion", + "embedding", + "audio_speech", + "audio_transcription", + "image_generation", + "video_generation", + "batch", + "rerank", + "realtime", + "responses", + "ocr", + ] + | None = fastapi.Body( None, description="The mode to test the model with. If not provided, auto-detected from model capabilities.", ), - litellm_params: Dict = fastapi.Body( + litellm_params: dict = fastapi.Body( None, description="Parameters for litellm.completion, litellm.embedding for the health check", ), - model_info: Dict = fastapi.Body( + model_info: dict = fastapi.Body( None, description="Model info for the health check", ), @@ -1813,7 +1806,7 @@ async def test_model_connection( # Look up model configuration from router if model name is provided # This gets the litellm_params from proxy config (with resolved env vars) config_litellm_params: dict = {} - loaded_model_info: Optional[dict] = None + loaded_model_info: dict | None = None if llm_router is not None: # Prefer disambiguation by deployment id (`model_info.id`) when # the caller supplies it. This is required when multiple @@ -1905,9 +1898,9 @@ async def test_model_connection( raise e except Exception as e: verbose_proxy_logger.debug( - f"litellm.proxy.health_endpoints.test_model_connection(): Exception occurred - {str(e)}" + f"litellm.proxy.health_endpoints.test_model_connection(): Exception occurred - {e!s}" ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to test connection: {str(e)}"}, + detail={"error": f"Failed to test connection: {e!s}"}, ) diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index d729e339a51..a4a27bb458f 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -1,5 +1,5 @@ import os -from typing import Literal, Union +from typing import Literal from . import * from .cache_control_check import _PROXY_CacheControlCheck @@ -33,15 +33,7 @@ if os.getenv("LEGACY_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": def get_proxy_hook( - hook_name: Union[ - Literal[ - "max_budget_limiter", - "managed_files", - "parallel_request_limiter", - "cache_control_check", - ], - str, - ], + hook_name: Literal["max_budget_limiter", "managed_files", "parallel_request_limiter", "cache_control_check"] | str, ): """ Factory method to get a proxy hook instance by name diff --git a/litellm/proxy/hooks/azure_content_safety.py b/litellm/proxy/hooks/azure_content_safety.py index c41effb4783..75ce0dff59c 100644 --- a/litellm/proxy/hooks/azure_content_safety.py +++ b/litellm/proxy/hooks/azure_content_safety.py @@ -1,5 +1,4 @@ import traceback -from typing import Optional from fastapi import HTTPException @@ -75,7 +74,7 @@ class _PROXY_AzureContentSafety( return result - async def test_violation(self, content: str, source: Optional[str] = None): + async def test_violation(self, content: str, source: str | None = None): verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content) # Construct a request @@ -124,9 +123,7 @@ class _PROXY_AzureContentSafety( raise e except Exception as e: verbose_proxy_logger.error( - "litellm.proxy.hooks.azure_content_safety.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.azure_content_safety.py::async_pre_call_hook(): Exception occured - {e!s}" ) verbose_proxy_logger.debug(traceback.format_exc()) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index d631fa7ee0c..651ede6f5bc 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -22,12 +22,8 @@ from collections.abc import Iterable from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, NoReturn, - Optional, - Tuple, Union, ) @@ -82,8 +78,8 @@ else: InternalUsageCache = Any Router = Any ParallelRequestLimiter = Any - RateLimitStatus = Dict[str, Any] - RateLimitDescriptor = Dict[str, Any] + RateLimitStatus = dict[str, Any] + RateLimitDescriptor = dict[str, Any] class BatchFileUsage(BaseModel): @@ -123,7 +119,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): self.parallel_request_limiter = parallel_request_limiter self._warned_unsupported_model_skip = False - def _get_file_bound_batch_model(self, data: Dict) -> Optional[str]: + def _get_file_bound_batch_model(self, data: dict) -> str | None: """Resolve the model bound to the batch input file ID. ``create_batch`` routes a file-bound id (model-embedded ``file-...`` or @@ -154,7 +150,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None - def _get_batch_routing_model(self, data: Dict) -> Optional[str]: + def _get_batch_routing_model(self, data: dict) -> str | None: """Resolve the deployment/model used for this batch from request data. Mirrors ``create_batch`` routing precedence: a model bound to the input @@ -173,7 +169,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None - def _resolve_batch_provider(self, batch_model: Optional[str]) -> Optional[str]: + def _resolve_batch_provider(self, batch_model: str | None) -> str | None: """Resolve the provider from the deployment that serves ``batch_model``. The provider is read from trusted router credentials rather than the @@ -206,8 +202,8 @@ class _PROXY_BatchRateLimiter(CustomLogger): def _create_batch_rate_limit_descriptors( self, user_api_key_dict: UserAPIKeyAuth, - data: Dict, - ) -> List["RateLimitDescriptor"]: + data: dict, + ) -> list["RateLimitDescriptor"]: return self.parallel_request_limiter._create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=data, @@ -218,9 +214,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): def _should_skip_batch_input_file_processing( self, - data: Dict, + data: dict, user_api_key_dict: UserAPIKeyAuth, - ) -> Tuple[bool, Optional[List["RateLimitDescriptor"]]]: + ) -> tuple[bool, list["RateLimitDescriptor"] | None]: """ Skip downloading batch input files when the operator disabled batch input-file rate limiting, when the batch runs entirely on a skip-listed @@ -273,7 +269,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): return False, descriptors - def _warn_if_unsupported_model_skip_configured(self, general_settings: Dict) -> None: + def _warn_if_unsupported_model_skip_configured(self, general_settings: dict) -> None: """Warn once that ``skip_batch_input_file_rate_limiting_for_models`` is a no-op. A per-model skip is intentionally not honored because the model a batch @@ -309,7 +305,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): @staticmethod def _has_applicable_batch_rate_limits( - descriptors: List["RateLimitDescriptor"], + descriptors: list["RateLimitDescriptor"], ) -> bool: for descriptor in descriptors: rate_limit = descriptor.get("rate_limit") or {} @@ -325,8 +321,8 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, file_id: str, custom_llm_provider: str, - data: Dict, - ) -> Tuple[str, Dict[str, Any]]: + data: dict, + ) -> tuple[str, dict[str, Any]]: """ Map proxy-facing file IDs to provider file IDs and credentials. @@ -341,7 +337,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): ) from litellm.proxy.proxy_server import llm_router - fetch_kwargs: Dict[str, Any] = { + fetch_kwargs: dict[str, Any] = { "custom_llm_provider": custom_llm_provider, } @@ -384,10 +380,10 @@ class _PROXY_BatchRateLimiter(CustomLogger): def _raise_rate_limit_error( self, status: "RateLimitStatus", - descriptors: List["RateLimitDescriptor"], + descriptors: list["RateLimitDescriptor"], batch_usage: BatchFileUsage, limit_type: str, - requested_model: Optional[str] = None, + requested_model: str | None = None, ) -> NoReturn: """Raise :class:`ProxyRateLimitError` (a 429) for batch rate limit exceeded.""" from datetime import datetime @@ -441,9 +437,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): async def _check_and_increment_batch_counters( self, user_api_key_dict: UserAPIKeyAuth, - data: Dict, + data: dict, batch_usage: BatchFileUsage, - descriptors: Optional[List["RateLimitDescriptor"]] = None, + descriptors: list["RateLimitDescriptor"] | None = None, ) -> None: """ Atomically check + increment rate-limit counters by the batch amounts. @@ -462,11 +458,11 @@ class _PROXY_BatchRateLimiter(CustomLogger): data=data, ) - increment: Dict[Literal["requests", "tokens"], int] = { + increment: dict[Literal["requests", "tokens"], int] = { "requests": batch_usage.request_count, "tokens": batch_usage.total_tokens, } - increments: List[Dict[Literal["requests", "tokens"], int]] = [increment for _ in descriptors] + increments: list[dict[Literal["requests", "tokens"], int]] = [increment for _ in descriptors] rate_limit_response = await self.parallel_request_limiter.atomic_check_and_increment_by_n( descriptors=descriptors, @@ -490,8 +486,8 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, file_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - data: Optional[Dict] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, + data: dict | None = None, ) -> BatchFileUsage: """ Count number of requests and tokens in a batch input file. @@ -604,14 +600,14 @@ class _PROXY_BatchRateLimiter(CustomLogger): ) raise except Exception as e: - verbose_proxy_logger.error(f"Error counting input file usage for {file_id}: {str(e)}") + verbose_proxy_logger.error(f"Error counting input file usage for {file_id}: {e!s}") raise async def _enforce_batch_file_model_access( self, user_api_key_dict: UserAPIKeyAuth, - models: Optional[Iterable[str]] = None, - target_model_names: Optional[List[str]] = None, + models: Iterable[str] | None = None, + target_model_names: list[str] | None = None, ) -> None: """Reject the batch if the caller is not authorized for the upload target. @@ -708,7 +704,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): detail={ "error": ( "Batch input file references a model the caller is " - f"not authorized to use: model={model_to_check}, reason={str(e)}" + f"not authorized to use: model={model_to_check}, reason={e!s}" ) }, ) @@ -738,8 +734,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): from litellm.proxy.proxy_server import llm_router, proxy_logging_obj except ImportError as e: raise ValueError( - f"Cannot import proxy_server dependencies: {str(e)}. " - "Managed files require proxy_server to be initialized." + f"Cannot import proxy_server dependencies: {e!s}. Managed files require proxy_server to be initialized." ) # Get the managed files hook @@ -770,9 +765,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, cache: Any, - data: Dict, + data: dict, call_type: str, - ) -> Union[Exception, str, Dict, None]: + ) -> Exception | str | dict | None: """ Pre-call hook for batch operations. @@ -851,6 +846,6 @@ class _PROXY_BatchRateLimiter(CustomLogger): # Re-raise HTTP exceptions (rate limit exceeded) raise except Exception as e: - verbose_proxy_logger.error(f"Error in batch rate limiting: {str(e)}", exc_info=True) + verbose_proxy_logger.error(f"Error in batch rate limiting: {e!s}", exc_info=True) # Don't block the request if rate limiting fails return data diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index 8b11e185fb9..effafdbcf35 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -4,7 +4,7 @@ ### [BETA] this is in Beta. And might change. import traceback -from typing import Literal, Optional +from typing import Literal from fastapi import HTTPException @@ -17,7 +17,7 @@ from litellm.proxy._types import UserAPIKeyAuth class _PROXY_BatchRedisRequests(CustomLogger): # Class variables or attributes - in_memory_cache: Optional[InMemoryCache] = None + in_memory_cache: InMemoryCache | None = None def __init__(self): if litellm.cache is not None: @@ -26,9 +26,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): ) # map the litellm 'get_cache' function to our custom function def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG"): - if debug_level == "DEBUG": - verbose_proxy_logger.debug(print_statement) - elif debug_level == "INFO": + if debug_level == "DEBUG" or debug_level == "INFO": verbose_proxy_logger.debug(print_statement) if litellm.set_verbose is True: print(print_statement) # noqa: T201 @@ -86,7 +84,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): raise e except Exception as e: verbose_proxy_logger.error( - "litellm.proxy.hooks.batch_redis_get.py::async_pre_call_hook(): Exception occured - {}".format(str(e)) + f"litellm.proxy.hooks.batch_redis_get.py::async_pre_call_hook(): Exception occured - {e!s}" ) verbose_proxy_logger.debug(traceback.format_exc()) @@ -100,7 +98,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): - return redis cache request """ try: # never block execution - cache_key: Optional[str] = None + cache_key: str | None = None if "cache_key" in kwargs: cache_key = kwargs["cache_key"] elif litellm.cache is not None: diff --git a/litellm/proxy/hooks/cache_control_check.py b/litellm/proxy/hooks/cache_control_check.py index 6e3fbf84fab..a5c26e0dad8 100644 --- a/litellm/proxy/hooks/cache_control_check.py +++ b/litellm/proxy/hooks/cache_control_check.py @@ -52,7 +52,5 @@ class _PROXY_CacheControlCheck(CustomLogger): raise e except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.cache_control_check.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.cache_control_check.py::async_pre_call_hook(): Exception occured - {e!s}" ) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index 766d02666be..d08c3488348 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -6,7 +6,6 @@ import asyncio import os from collections.abc import Callable from datetime import datetime -from typing import List, Optional, Tuple, Union import litellm from litellm import ModelResponse, Router @@ -37,17 +36,17 @@ class DynamicRateLimiterCache: self.ttl = 60 # 1 min ttl self.time_fn = time_fn - async def async_get_cache(self, model: str) -> Optional[int]: + async def async_get_cache(self, model: str) -> int | None: dt = self.time_fn() current_minute = dt.strftime("%H-%M") - key_name = "{}:{}".format(current_minute, model) + key_name = f"{current_minute}:{model}" _response = await self.cache.async_get_cache(key=key_name) - response: Optional[int] = None + response: int | None = None if _response is not None: response = len(_response) return response - async def async_set_cache_sadd(self, model: str, value: List): + async def async_set_cache_sadd(self, model: str, value: list): """ Add value to set. @@ -65,13 +64,11 @@ class DynamicRateLimiterCache: dt = self.time_fn() current_minute = dt.strftime("%H-%M") - key_name = "{}:{}".format(current_minute, model) + key_name = f"{current_minute}:{model}" await self.cache.async_set_cache_sadd(key=key_name, value=value, ttl=self.ttl) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.hooks.dynamic_rate_limiter.py::async_set_cache_sadd(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.dynamic_rate_limiter.py::async_set_cache_sadd(): Exception occured - {e!s}" ) raise e @@ -85,8 +82,8 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): self.llm_router = llm_router async def check_available_usage( - self, model: str, priority: Optional[str] = None - ) -> Tuple[Optional[int], Optional[int], Optional[int], Optional[int], Optional[int]]: + self, model: str, priority: str | None = None + ) -> tuple[int | None, int | None, int | None, int | None, int | None]: """ For a given model, get its available tpm @@ -104,14 +101,12 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): """ try: # Get model info first for conversion - model_group_info: Optional[ModelGroupInfo] = self.llm_router.get_model_group_info(model_group=model) + model_group_info: ModelGroupInfo | None = self.llm_router.get_model_group_info(model_group=model) weight: float = 1 if litellm.priority_reservation is None or priority not in litellm.priority_reservation: verbose_proxy_logger.error( - "Priority Reservation not set. priority={}, but litellm.priority_reservation is {}.".format( - priority, litellm.priority_reservation - ) + f"Priority Reservation not set. priority={priority}, but litellm.priority_reservation is {litellm.priority_reservation}." ) elif priority is not None and litellm.priority_reservation is not None: if os.getenv("LITELLM_LICENSE", None) is None: @@ -127,27 +122,27 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): current_model_tpm, current_model_rpm, ) = await self.llm_router.get_model_group_usage(model_group=model) - total_model_tpm: Optional[int] = None - total_model_rpm: Optional[int] = None + total_model_tpm: int | None = None + total_model_rpm: int | None = None if model_group_info is not None: if model_group_info.tpm is not None: total_model_tpm = model_group_info.tpm if model_group_info.rpm is not None: total_model_rpm = model_group_info.rpm - remaining_model_tpm: Optional[int] = None + remaining_model_tpm: int | None = None if total_model_tpm is not None and current_model_tpm is not None: remaining_model_tpm = total_model_tpm - current_model_tpm elif total_model_tpm is not None: remaining_model_tpm = total_model_tpm - remaining_model_rpm: Optional[int] = None + remaining_model_rpm: int | None = None if total_model_rpm is not None and current_model_rpm is not None: remaining_model_rpm = total_model_rpm - current_model_rpm elif total_model_rpm is not None: remaining_model_rpm = total_model_rpm - available_tpm: Optional[int] = None + available_tpm: int | None = None if remaining_model_tpm is not None: if active_projects is not None: @@ -158,7 +153,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): if available_tpm is not None and available_tpm < 0: available_tpm = 0 - available_rpm: Optional[int] = None + available_rpm: int | None = None if remaining_model_rpm is not None: if active_projects is not None: @@ -177,9 +172,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): ) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.hooks.dynamic_rate_limiter.py::check_available_usage: Exception occurred - {}".format( - str(e) - ) + f"litellm.proxy.hooks.dynamic_rate_limiter.py::check_available_usage: Exception occurred - {e!s}" ) return None, None, None, None, None @@ -189,16 +182,16 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Optional[ - Union[Exception, str, dict] - ]: # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm + ) -> ( + Exception | str | dict | None + ): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm """ - For a model group - Check if tpm/rpm available - Raise RateLimitError if no tpm/rpm available """ if "model" in data: - key_priority: Optional[str] = user_api_key_dict.metadata.get("priority", None) + key_priority: str | None = user_api_key_dict.metadata.get("priority", None) ( available_tpm, available_rpm, @@ -211,12 +204,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(data.get("model")) raise ProxyRateLimitError( detail={ - "error": "Key={} over available TPM={}. Model TPM={}, Active keys={}".format( - user_api_key_dict.api_key, - available_tpm, - model_tpm, - active_projects, - ) + "error": f"Key={user_api_key_dict.api_key} over available TPM={available_tpm}. Model TPM={model_tpm}, Active keys={active_projects}" }, rate_limit_type=RateLimitType.TOKENS, model=resolved_model, @@ -227,12 +215,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(data.get("model")) raise ProxyRateLimitError( detail={ - "error": "Key={} over available RPM={}. Model RPM={}, Active keys={}".format( - user_api_key_dict.api_key, - available_rpm, - model_rpm, - active_projects, - ) + "error": f"Key={user_api_key_dict.api_key} over available RPM={available_rpm}. Model RPM={model_rpm}, Active keys={active_projects}" }, rate_limit_type=RateLimitType.REQUESTS, model=resolved_model, @@ -255,7 +238,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): assert model_info is not None, "Model info for model with id={} is None".format( response._hidden_params["model_id"] ) - key_priority: Optional[str] = user_api_key_dict.metadata.get("priority", None) + key_priority: str | None = user_api_key_dict.metadata.get("priority", None) ( available_tpm, available_rpm, @@ -280,8 +263,6 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): ) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.hooks.dynamic_rate_limiter.py::async_post_call_success_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.dynamic_rate_limiter.py::async_post_call_success_hook(): Exception occured - {e!s}" ) return response diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 441187d5d82..cee11ff22ae 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -5,7 +5,7 @@ Dynamic rate limiter v3 - Saturation-aware priority-based rate limiting import os from collections.abc import Callable from datetime import datetime -from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Literal from fastapi import HTTPException @@ -78,7 +78,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): def __init__( self, internal_usage_cache: DualCache, - time_provider: Optional[Callable[[], datetime]] = None, + time_provider: Callable[[], datetime] | None = None, ): self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache) self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) @@ -93,7 +93,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): async def _get_saturation_value_from_cache( self, counter_key: str, - ) -> Optional[str]: + ) -> str | None: """ Get saturation value with configurable local cache TTL. @@ -115,7 +115,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): ttl=local_cache_ttl, ) - def _get_priority_weight(self, priority: Optional[str], model_info: Optional[ModelGroupInfo] = None) -> float: + def _get_priority_weight(self, priority: str | None, model_info: ModelGroupInfo | None = None) -> float: """Get the weight for a given priority from litellm.priority_reservation""" weight: float = _get_priority_settings().default_priority if litellm.priority_reservation is None or priority not in litellm.priority_reservation: @@ -130,7 +130,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): weight = convert_priority_to_percent(value, model_info) return weight - def _get_priority_from_user_api_key_dict(self, user_api_key_dict: UserAPIKeyAuth) -> Optional[str]: + def _get_priority_from_user_api_key_dict(self, user_api_key_dict: UserAPIKeyAuth) -> str | None: """ Get priority from user_api_key_dict. @@ -142,7 +142,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): Returns: Priority string if found, None otherwise """ - priority: Optional[str] = None + priority: str | None = None # Check team metadata first (takes precedence) if user_api_key_dict.team_metadata is not None: @@ -154,7 +154,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return priority - def _normalize_priority_weights(self, model_info: ModelGroupInfo) -> Dict[str, float]: + def _normalize_priority_weights(self, model_info: ModelGroupInfo) -> dict[str, float]: """ Normalize priority weights if they sum to > 1.0 @@ -165,7 +165,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return {} # Convert all values to percentages first - weights: Dict[str, float] = {} + weights: dict[str, float] = {} for k, v in litellm.priority_reservation.items(): weights[k] = convert_priority_to_percent(v, model_info) @@ -181,9 +181,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): def _get_priority_allocation( self, model: str, - priority: Optional[str], - normalized_weights: Dict[str, float], - model_info: Optional[ModelGroupInfo] = None, + priority: str | None, + normalized_weights: dict[str, float], + model_info: ModelGroupInfo | None = None, ) -> tuple[float, str]: """ Get priority weight and pool key for a given priority. @@ -282,7 +282,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return max_saturation except Exception as e: - verbose_proxy_logger.error(f"Error checking saturation for {model}: {str(e)}") + verbose_proxy_logger.error(f"Error checking saturation for {model}: {e!s}") # Fail open: assume not saturated on error return 0.0 @@ -290,8 +290,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): self, model: str, user_api_key_dict: UserAPIKeyAuth, - priority: Optional[str], - ) -> List[RateLimitDescriptor]: + priority: str | None, + ) -> list[RateLimitDescriptor]: """ Create rate limit descriptors with normalized priority weights. @@ -300,13 +300,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): For explicit priorities: each priority gets its own pool (e.g., prod gets 75%) For default priority: ALL keys without explicit priority share ONE pool (e.g., all share 25%) """ - descriptors: List[RateLimitDescriptor] = [] + descriptors: list[RateLimitDescriptor] = [] if litellm.priority_reservation is None: return descriptors # Get model group info - model_group_info: Optional[ModelGroupInfo] = self.llm_router.get_model_group_info(model_group=model) + model_group_info: ModelGroupInfo | None = self.llm_router.get_model_group_info(model_group=model) if model_group_info is None: return descriptors @@ -375,7 +375,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): model: str, model_group_info: ModelGroupInfo, user_api_key_dict: UserAPIKeyAuth, - priority: Optional[str], + priority: str | None, saturation: float, ) -> None: """ @@ -413,7 +413,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): should_enforce_priority = saturation >= saturation_threshold # Build ALL descriptors upfront - descriptors_to_check: List[RateLimitDescriptor] = [] + descriptors_to_check: list[RateLimitDescriptor] = [] # Model-wide descriptor (always enforce) model_wide_descriptor = self._create_model_tracking_descriptor( @@ -440,11 +440,11 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): # asyncio.Lock + in-memory fallback for single-process deployments. # All-or-nothing: if any enforced descriptor would exceed its limit, # no counter is modified and the response carries "OVER_LIMIT". - enforced_descriptors: List[RateLimitDescriptor] = [model_wide_descriptor] + enforced_descriptors: list[RateLimitDescriptor] = [model_wide_descriptor] if priority_descriptors and should_enforce_priority: enforced_descriptors.extend(priority_descriptors) - per_request_increment: Dict[Literal["requests", "tokens"], int] = { + per_request_increment: dict[Literal["requests", "tokens"], int] = { "requests": 1, "tokens": 0, } @@ -565,7 +565,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Saturation-aware pre-call hook for priority-based rate limiting. @@ -608,7 +608,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): priority = self._get_priority_from_user_api_key_dict(user_api_key_dict=user_api_key_dict) # Get model configuration - model_group_info: Optional[ModelGroupInfo] = self.llm_router.get_model_group_info(model_group=model) + model_group_info: ModelGroupInfo | None = self.llm_router.get_model_group_info(model_group=model) if model_group_info is None: verbose_proxy_logger.debug(f"No model group info for {model}, allowing request") return None @@ -640,7 +640,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error in dynamic rate limiter: {str(e)}, allowing request") + verbose_proxy_logger.error(f"Error in dynamic rate limiter: {e!s}, allowing request") # Fail open on unexpected errors return None @@ -676,7 +676,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return response except Exception as e: - verbose_proxy_logger.exception(f"Error in dynamic rate limiter v3 post-call hook: {str(e)}") + verbose_proxy_logger.exception(f"Error in dynamic rate limiter v3 post-call hook: {e!s}") return response async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -713,7 +713,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): # Get priority from user_api_key_auth_metadata in standard_logging_metadata # This is where user_api_key_dict.metadata is stored during pre-call user_api_key_auth_metadata = standard_logging_metadata.get("user_api_key_auth_metadata") or {} - key_priority: Optional[str] = user_api_key_auth_metadata.get("priority") + key_priority: str | None = user_api_key_auth_metadata.get("priority") # Get total tokens from response total_tokens = 0 @@ -733,7 +733,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return # Create pipeline operations for token increments - pipeline_operations: List[RedisPipelineIncrementOperation] = [] + pipeline_operations: list[RedisPipelineIncrementOperation] = [] # Model-wide token tracking (model_saturation_check) model_token_key = self.v3_limiter.create_rate_limit_keys( @@ -791,4 +791,4 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): ) except Exception as e: - verbose_proxy_logger.exception(f"Error in dynamic rate limiter success event: {str(e)}") + verbose_proxy_logger.exception(f"Error in dynamic rate limiter success event: {e!s}") diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index 8f2155a7fbc..7ef1341a168 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -1,7 +1,7 @@ import asyncio import json from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any import litellm from litellm._logging import verbose_proxy_logger @@ -30,7 +30,7 @@ class KeyManagementEventHooks: data: GenerateKeyRequest, response: GenerateKeyResponse, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ): """ Hook that runs after a successful /key/generate request @@ -92,7 +92,7 @@ class KeyManagementEventHooks: existing_key_row: Any, response: Any, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ): """ Post /key/update processing hook @@ -135,11 +135,11 @@ class KeyManagementEventHooks: @staticmethod async def async_key_rotated_hook( - data: Optional[RegenerateKeyRequest], + data: RegenerateKeyRequest | None, existing_key_row: LiteLLM_VerificationToken, response: GenerateKeyResponse, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ): from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, @@ -204,10 +204,10 @@ class KeyManagementEventHooks: @staticmethod async def async_key_deleted_hook( data: KeyRequest, - keys_being_deleted: List[LiteLLM_VerificationToken], + keys_being_deleted: list[LiteLLM_VerificationToken], response: dict, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ): """ Post /key/delete processing hook @@ -251,10 +251,9 @@ class KeyManagementEventHooks: ) # delete the keys from the secret manager await KeyManagementEventHooks._delete_virtual_keys_from_secret_manager(keys_being_deleted=keys_being_deleted) - pass @staticmethod - async def _store_virtual_key_in_secret_manager(secret_name: str, secret_token: str, team_id: Optional[str] = None): + async def _store_virtual_key_in_secret_manager(secret_name: str, secret_token: str, team_id: str | None = None): """ Store a virtual key in the secret manager @@ -290,7 +289,7 @@ class KeyManagementEventHooks: current_secret_name: str, new_secret_name: str, new_secret_value: str, - team_id: Optional[str] = None, + team_id: str | None = None, ): """ Update a virtual key in the secret manager @@ -326,7 +325,7 @@ class KeyManagementEventHooks: @staticmethod async def _delete_virtual_keys_from_secret_manager( - keys_being_deleted: List[LiteLLM_VerificationToken], + keys_being_deleted: list[LiteLLM_VerificationToken], ): """ Deletes virtual keys from the secret manager @@ -341,7 +340,7 @@ class KeyManagementEventHooks: ) if isinstance(litellm.secret_manager_client, BaseSecretManager): - team_settings_cache: Dict[Optional[str], Optional[dict]] = {} + team_settings_cache: dict[str | None, dict | None] = {} for key in keys_being_deleted: if key.key_alias is not None: team_id = getattr(key, "team_id", None) @@ -361,8 +360,8 @@ class KeyManagementEventHooks: @staticmethod async def _get_secret_manager_optional_params( - team_id: Optional[str], - ) -> Optional[dict]: + team_id: str | None, + ) -> dict | None: if team_id is None: return None @@ -511,7 +510,7 @@ class KeyManagementEventHooks: ) @staticmethod - async def _send_key_rotated_email(response: dict, existing_key_alias: Optional[str]): + async def _send_key_rotated_email(response: dict, existing_key_alias: str | None): """ Send key rotated email if email sending is enabled. diff --git a/litellm/proxy/hooks/litellm_skills/__init__.py b/litellm/proxy/hooks/litellm_skills/__init__.py index 1507b652ab4..751122ac51c 100644 --- a/litellm/proxy/hooks/litellm_skills/__init__.py +++ b/litellm/proxy/hooks/litellm_skills/__init__.py @@ -27,13 +27,13 @@ from litellm.proxy.hooks.litellm_skills.main import ( ) __all__ = [ - "SkillsInjectionHook", - "skills_injection_hook", + "LITELLM_CODE_EXECUTION_TOOL", "CodeExecutionHandler", "LiteLLMInternalTools", - "LITELLM_CODE_EXECUTION_TOOL", - "get_litellm_code_execution_tool", - "code_execution_handler", "SkillPromptInjectionHandler", + "SkillsInjectionHook", "SkillsSandboxExecutor", + "code_execution_handler", + "get_litellm_code_execution_tool", + "skills_injection_hook", ] diff --git a/litellm/proxy/hooks/litellm_skills/main.py b/litellm/proxy/hooks/litellm_skills/main.py index 12370ea1536..2fc9779e5fd 100644 --- a/litellm/proxy/hooks/litellm_skills/main.py +++ b/litellm/proxy/hooks/litellm_skills/main.py @@ -26,7 +26,7 @@ Usage: import base64 import json -from typing import Any, Dict, List, Optional, Union +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache @@ -74,7 +74,7 @@ class SkillsInjectionHook(CustomLogger): cache: DualCache, data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Process skills from container.skills before the LLM call. @@ -98,8 +98,8 @@ class SkillsInjectionHook(CustomLogger): verbose_proxy_logger.debug(f"SkillsInjectionHook: Processing {len(skills)} skills") - litellm_skills: List[LiteLLM_SkillsTable] = [] - anthropic_skills: List[Dict[str, Any]] = [] + litellm_skills: list[LiteLLM_SkillsTable] = [] + anthropic_skills: list[dict[str, Any]] = [] # Separate skills by prefix for skill in skills: @@ -137,7 +137,7 @@ class SkillsInjectionHook(CustomLogger): def _process_for_messages_api( self, data: dict, - litellm_skills: List[LiteLLM_SkillsTable], + litellm_skills: list[LiteLLM_SkillsTable], use_anthropic_format: bool = True, ) -> dict: """ @@ -153,9 +153,9 @@ class SkillsInjectionHook(CustomLogger): ) tools = data.get("tools", []) - skill_contents: List[str] = [] - all_skill_files: Dict[str, Dict[str, bytes]] = {} - all_module_paths: List[str] = [] + skill_contents: list[str] = [] + all_skill_files: dict[str, dict[str, bytes]] = {} + all_module_paths: list[str] = [] for skill in litellm_skills: # Convert skill to Anthropic-style tool @@ -208,7 +208,7 @@ class SkillsInjectionHook(CustomLogger): def _process_non_anthropic_model( self, data: dict, - litellm_skills: List[LiteLLM_SkillsTable], + litellm_skills: list[LiteLLM_SkillsTable], ) -> dict: """ Process skills for non-Anthropic models (OpenAI format tools). @@ -219,9 +219,9 @@ class SkillsInjectionHook(CustomLogger): - Stores skill files in metadata for sandbox execution """ tools = data.get("tools", []) - skill_contents: List[str] = [] - all_skill_files: Dict[str, Dict[str, bytes]] = {} - all_module_paths: List[str] = [] + skill_contents: list[str] = [] + all_skill_files: dict[str, dict[str, bytes]] = {} + all_module_paths: list[str] = [] for skill in litellm_skills: # Convert skill to OpenAI-style tool @@ -277,7 +277,7 @@ class SkillsInjectionHook(CustomLogger): self, skill_id: str, user_api_key_dict: UserAPIKeyAuth, - ) -> Optional[LiteLLM_SkillsTable]: + ) -> LiteLLM_SkillsTable | None: """ Fetch a skill from the LiteLLM database. @@ -323,8 +323,8 @@ class SkillsInjectionHook(CustomLogger): self, request_data: dict, response: Any, - call_type: Optional[CallTypes], - ) -> Optional[Any]: + call_type: CallTypes | None, + ) -> Any | None: """ Post-call hook to handle automatic code execution. @@ -354,7 +354,7 @@ class SkillsInjectionHook(CustomLogger): # Get skill files skill_files_by_id = litellm_metadata.get("_skill_files") or metadata.get("_skill_files", {}) - all_skill_files: Dict[str, bytes] = {} + all_skill_files: dict[str, bytes] = {} for files_dict in skill_files_by_id.values(): all_skill_files.update(files_dict) @@ -388,7 +388,7 @@ class SkillsInjectionHook(CustomLogger): skill_files=all_skill_files, ) - def _extract_tool_calls(self, response: Any) -> List[Dict[str, Any]]: + def _extract_tool_calls(self, response: Any) -> list[dict[str, Any]]: """Extract tool calls from response, handling both formats.""" tool_calls = [] @@ -438,7 +438,7 @@ class SkillsInjectionHook(CustomLogger): self, data: dict, response: Any, - skill_files: Dict[str, bytes], + skill_files: dict[str, bytes], ) -> Any: """ Execute the code execution loop for messages API (Anthropic format). @@ -464,7 +464,7 @@ class SkillsInjectionHook(CustomLogger): max_tokens = data.get("max_tokens", 4096) executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) - generated_files: List[Dict[str, Any]] = [] + generated_files: list[dict[str, Any]] = [] current_response = response for iteration in range(self.max_iterations): @@ -557,9 +557,9 @@ class SkillsInjectionHook(CustomLogger): async def _execute_code( self, code: str, - skill_files: Dict[str, bytes], + skill_files: dict[str, bytes], executor: Any, - generated_files: List[Dict[str, Any]], + generated_files: list[dict[str, Any]], ) -> str: """Execute code in sandbox and return result string.""" try: @@ -587,20 +587,20 @@ class SkillsInjectionHook(CustomLogger): return result or "Code executed successfully" except Exception as e: - return f"Code execution failed: {str(e)}" + return f"Code execution failed: {e!s}" async def _execute_skill_tool( self, tool_name: str, - tool_input: Dict[str, Any], - skill_files: Dict[str, bytes], + tool_input: dict[str, Any], + skill_files: dict[str, bytes], executor: Any, - generated_files: List[Dict[str, Any]], + generated_files: list[dict[str, Any]], ) -> str: """Execute a skill tool by generating and running code based on skill content.""" # Generate code based on available skill modules # Look for Python modules in the skill - python_modules = [p for p in skill_files.keys() if p.endswith(".py") and not p.endswith("__init__.py")] + python_modules = [p for p in skill_files if p.endswith(".py") and not p.endswith("__init__.py")] # Try to find the main builder/creator module main_module = None @@ -666,7 +666,7 @@ print('No executable skill module found') self, data: dict, response: Any, - skill_files: Dict[str, bytes], + skill_files: dict[str, bytes], ) -> Any: """ Execute the code execution loop until model gives final response. @@ -701,7 +701,7 @@ print('No executable skill module found') kwargs = {k: v for k, v in data.items() if k not in _EXCLUDED_ACOMPLETION_KEYS} executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) - generated_files: List[Dict[str, Any]] = [] + generated_files: list[dict[str, Any]] = [] current_response: Any = response for iteration in range(self.max_iterations): @@ -710,7 +710,7 @@ print('No executable skill module found') stop_reason = current_response.choices[0].finish_reason # type: ignore[union-attr] # Build assistant message for conversation history - assistant_msg_dict: Dict[str, Any] = { + assistant_msg_dict: dict[str, Any] = { "role": "assistant", "content": assistant_message.content, } @@ -776,9 +776,9 @@ print('No executable skill module found') async def _execute_code_tool( self, tool_call: Any, - skill_files: Dict[str, bytes], + skill_files: dict[str, bytes], executor: Any, - generated_files: List[Dict[str, Any]], + generated_files: list[dict[str, Any]], ) -> str: """Execute a litellm_code_execution tool call and return result string.""" try: @@ -821,12 +821,12 @@ print('No executable skill module found') except Exception as e: verbose_proxy_logger.error(f"SkillsInjectionHook: Code execution failed: {e}") - return f"Code execution failed: {str(e)}" + return f"Code execution failed: {e!s}" def _attach_files_to_response( self, response: Any, - generated_files: List[Dict[str, Any]], + generated_files: list[dict[str, Any]], ) -> Any: """ Attach generated files to the response object. diff --git a/litellm/proxy/hooks/max_budget_limiter.py b/litellm/proxy/hooks/max_budget_limiter.py index 1983675c5f3..0a1a09d0792 100644 --- a/litellm/proxy/hooks/max_budget_limiter.py +++ b/litellm/proxy/hooks/max_budget_limiter.py @@ -3,8 +3,8 @@ from fastapi import HTTPException from litellm import verbose_logger from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache -from litellm.integrations.custom_logger import CustomLogger from litellm.exceptions import RateLimitType +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -75,7 +75,5 @@ class _PROXY_MaxBudgetLimiter(CustomLogger): raise e except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.max_budget_limiter.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.max_budget_limiter.py::async_pre_call_hook(): Exception occured - {e!s}" ) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 53de9e0c5a3..6910c67af97 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -15,12 +15,12 @@ Follows the same pattern as max_iterations_limiter.py. """ import os -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any from litellm import DualCache from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_logger import CustomLogger from litellm.exceptions import RateLimitType +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -87,7 +87,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): cache: DualCache, data: dict, call_type: str, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Before each LLM call, check if max_budget_per_session is set and whether accumulated spend exceeds the budget (429 if so). @@ -171,7 +171,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): str(e), ) - def _get_session_id(self, data: dict) -> Optional[str]: + def _get_session_id(self, data: dict) -> str | None: """Extract session_id from request metadata.""" metadata = data.get("metadata") or {} session_id = metadata.get("session_id") @@ -185,7 +185,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return None - def _get_max_budget_per_session(self, user_api_key_dict: UserAPIKeyAuth) -> Optional[float]: + def _get_max_budget_per_session(self, user_api_key_dict: UserAPIKeyAuth) -> float | None: """Extract max_budget_per_session from agent litellm_params.""" agent_id = user_api_key_dict.agent_id if agent_id is None: diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 093351c7d6d..4cda50940d8 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -11,12 +11,12 @@ Follows the same pattern as parallel_request_limiter_v3.py. """ import os -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any from litellm import DualCache from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_logger import CustomLogger from litellm.exceptions import RateLimitType +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -86,7 +86,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): cache: DualCache, data: dict, call_type: str, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Check session iteration count before making the API call. @@ -133,7 +133,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return None - def _get_session_id(self, data: dict) -> Optional[str]: + def _get_session_id(self, data: dict) -> str | None: """Extract session_id from request metadata.""" metadata = data.get("metadata") or {} session_id = metadata.get("session_id") @@ -148,7 +148,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return None - def _get_max_iterations(self, user_api_key_dict: UserAPIKeyAuth) -> Optional[int]: + def _get_max_iterations(self, user_api_key_dict: UserAPIKeyAuth) -> int | None: """Extract max_iterations from agent litellm_params, with fallback to key metadata.""" # Try agent litellm_params first agent_id = user_api_key_dict.agent_id diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 5f1d061c7cb..10293bc5e5f 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -5,7 +5,7 @@ Pre-call hook that filters MCP tools semantically before LLM inference. Reduces context window size and improves tool selection accuracy. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Optional from fastapi import HTTPException @@ -66,7 +66,7 @@ class SemanticToolFilterHook(CustomLogger): f"enabled={semantic_filter.enabled}, top_k={semantic_filter.top_k}" ) - def _should_expand_mcp_tools(self, tools: List[Any]) -> bool: + def _should_expand_mcp_tools(self, tools: list[Any]) -> bool: """ Check if tools contain MCP references with server_url="litellm_proxy". @@ -80,9 +80,9 @@ class SemanticToolFilterHook(CustomLogger): async def _expand_mcp_tools( self, - tools: List[Any], + tools: list[Any], user_api_key_dict: "UserAPIKeyAuth", - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """ Expand MCP references to actual tool definitions. @@ -276,7 +276,7 @@ class SemanticToolFilterHook(CustomLogger): cache: "DualCache", data: dict, call_type: str, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Filter tools before LLM call based on user query. @@ -412,9 +412,9 @@ class SemanticToolFilterHook(CustomLogger): data: dict, user_api_key_dict: "UserAPIKeyAuth", response: Any, - request_headers: Optional[Dict[str, str]] = None, - litellm_call_info: Optional[Dict[str, Any]] = None, - ) -> Optional[Dict[str, str]]: + request_headers: dict[str, str] | None = None, + litellm_call_info: dict[str, Any] | None = None, + ) -> dict[str, str] | None: """Add semantic filter stats and tool names to response headers.""" from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH @@ -438,7 +438,7 @@ class SemanticToolFilterHook(CustomLogger): return headers - def _get_tool_names_csv(self, tools: List[Any]) -> str: + def _get_tool_names_csv(self, tools: list[Any]) -> str: """Extract tool names and return as CSV string.""" if not tools: return "" @@ -453,7 +453,7 @@ class SemanticToolFilterHook(CustomLogger): @staticmethod async def initialize_from_config( - config: Optional[Dict[str, Any]], + config: dict[str, Any] | None, llm_router: Optional["Router"], ) -> Optional["SemanticToolFilterHook"]: """ diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index e8cf5fbc718..2aeb505bd97 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -1,5 +1,4 @@ import json -from typing import List, Optional import litellm from litellm._logging import verbose_proxy_logger @@ -86,7 +85,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): self, user_api_key_dict: UserAPIKeyAuth, model: str, - ) -> Optional[str]: + ) -> str | None: budget_fallbacks: dict[str, list[str]] = user_api_key_dict.budget_fallbacks or {} for fallback_model in budget_fallbacks.get(model, []): try: @@ -153,7 +152,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): end_user_id: str, model: str, key_budget_config: BudgetConfig, - ) -> Optional[float]: + ) -> float | None: # 1. model: directly look up `model` end_user_model_spend_cache_key = ( f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{key_budget_config.budget_duration}" @@ -172,10 +171,10 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): async def _get_virtual_key_spend_for_model( self, - user_api_key_hash: Optional[str], + user_api_key_hash: str | None, model: str, key_budget_config: BudgetConfig, - ) -> Optional[float]: + ) -> float | None: """ Get the current spend for a virtual key for a model @@ -203,7 +202,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): def _get_request_model_budget_config( self, model: str, internal_model_max_budget: GenericBudgetConfigType - ) -> Optional[BudgetConfig]: + ) -> BudgetConfig | None: """ Get the budget config for the request model @@ -222,11 +221,11 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, # type: ignore - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, # type: ignore + ) -> list[dict]: return healthy_deployments async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -236,7 +235,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d """ verbose_proxy_logger.debug("in RouterBudgetLimiting.async_log_success_event") - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: verbose_proxy_logger.debug( "Skipping _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event: standard_logging_payload is None" @@ -245,8 +244,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): _litellm_params: dict = kwargs.get("litellm_params", {}) or {} _metadata: dict = _litellm_params.get("metadata", {}) or {} - user_api_key_model_max_budget: Optional[dict] = _metadata.get("user_api_key_model_max_budget", None) - user_api_key_end_user_model_max_budget: Optional[dict] = _metadata.get( + user_api_key_model_max_budget: dict | None = _metadata.get("user_api_key_model_max_budget", None) + user_api_key_end_user_model_max_budget: dict | None = _metadata.get( "user_api_key_end_user_model_max_budget", None ) if (user_api_key_model_max_budget is None or len(user_api_key_model_max_budget) == 0) and ( diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index ee6abb13d6b..b41fd960aac 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -1,7 +1,7 @@ import asyncio import sys from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Any, List, Literal, NoReturn, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Literal, NoReturn, Union from pydantic import BaseModel from typing_extensions import TypedDict @@ -9,10 +9,10 @@ from typing_extensions import TypedDict import litellm from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionResponse from litellm._logging import verbose_proxy_logger +from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth -from litellm.exceptions import RateLimitType from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, get_key_model_tpm_limit, @@ -34,12 +34,12 @@ else: class CacheObject(TypedDict): - current_global_requests: Optional[dict] - request_count_api_key: Optional[dict] - request_count_api_key_model: Optional[dict] - request_count_user_id: Optional[dict] - request_count_team_id: Optional[dict] - request_count_end_user_id: Optional[dict] + current_global_requests: dict | None + request_count_api_key: dict | None + request_count_api_key_model: dict | None + request_count_user_id: dict | None + request_count_team_id: dict | None + request_count_end_user_id: dict | None class _PROXY_MaxParallelRequestsHandler(CustomLogger): @@ -64,10 +64,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): max_parallel_requests: int, tpm_limit: int, rpm_limit: int, - current: Optional[dict], + current: dict | None, request_count_api_key: str, rate_limit_type: Literal["key", "model_per_key", "user", "customer", "team"], - values_to_update_in_cache: List[Tuple[Any, Any]], + values_to_update_in_cache: list[tuple[Any, Any]], ) -> dict: verbose_proxy_logger.info(f"Current Usage of {rate_limit_type} in this minute: {current}") if current is None: @@ -150,9 +150,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): def raise_rate_limit_error( self, - additional_details: Optional[str] = None, - rate_limit_type: Optional[RateLimitType] = None, - requested_model: Optional[str] = None, + additional_details: str | None = None, + rate_limit_type: RateLimitType | None = None, + requested_model: str | None = None, ) -> NoReturn: """ Raise a 429 with a retry-after header for litellm-proxy parallel-request limits. @@ -194,13 +194,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): async def get_all_cache_objects( self, - current_global_requests: Optional[str], - request_count_api_key: Optional[str], - request_count_api_key_model: Optional[str], - request_count_user_id: Optional[str], - request_count_team_id: Optional[str], - request_count_end_user_id: Optional[str], - parent_otel_span: Optional[Span] = None, + current_global_requests: str | None, + request_count_api_key: str | None, + request_count_api_key_model: str | None, + request_count_user_id: str | None, + request_count_team_id: str | None, + request_count_end_user_id: str | None, + parent_otel_span: Span | None = None, ) -> CacheObject: keys = [ current_global_requests, @@ -257,14 +257,14 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): if rpm_limit is None: rpm_limit = sys.maxsize - values_to_update_in_cache: List[ - Tuple[Any, Any] + values_to_update_in_cache: list[ + tuple[Any, Any] ] = [] # values that need to get updated in cache, will run a batch_set_cache after this function # ------------ # Setup values # ------------ - new_val: Optional[dict] = None + new_val: dict | None = None if global_max_parallel_requests is not None: # get value from cache @@ -480,14 +480,12 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) # don't block execution for cache updates ) - return - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) - litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(kwargs=kwargs) try: self.print_verbose("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") @@ -711,7 +709,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: self.print_verbose("Inside Max Parallel Request Failure Hook") - litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(kwargs=kwargs) _metadata = kwargs["litellm_params"].get("metadata", {}) or {} global_max_parallel_requests = _metadata.get("global_max_parallel_requests", None) user_api_key = _metadata.get("user_api_key", None) @@ -778,13 +776,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=litellm_parent_otel_span, ) # save in cache for up to 1 min. except Exception as e: - verbose_proxy_logger.exception("Inside Parallel Request Limiter: An exception occurred - {}".format(str(e))) + verbose_proxy_logger.exception(f"Inside Parallel Request Limiter: An exception occurred - {e!s}") async def get_internal_user_object( self, user_id: str, user_api_key_dict: UserAPIKeyAuth, - ) -> Optional[dict]: + ) -> dict | None: """ Helper to get the 'Internal User Object' @@ -824,15 +822,15 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): current_minute = datetime.now().strftime("%M") precise_minute = f"{current_date}-{current_hour}-{current_minute}" request_count_api_key = f"{api_key}::{precise_minute}::request_count" - current: Optional[CurrentItemRateLimit] = await self.internal_usage_cache.async_get_cache( + current: CurrentItemRateLimit | None = await self.internal_usage_cache.async_get_cache( key=request_count_api_key, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) - key_remaining_rpm_limit: Optional[int] = None - key_rpm_limit: Optional[int] = None - key_remaining_tpm_limit: Optional[int] = None - key_tpm_limit: Optional[int] = None + key_remaining_rpm_limit: int | None = None + key_rpm_limit: int | None = None + key_remaining_tpm_limit: int | None = None + key_tpm_limit: int | None = None if current is not None: if user_api_key_dict.rpm_limit is not None: key_remaining_rpm_limit = user_api_key_dict.rpm_limit - current["current_rpm"] diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 40127007c16..7253a684b3c 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -15,12 +15,7 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - FrozenSet, - List, Literal, - Optional, - Tuple, TypedDict, Union, cast, @@ -306,16 +301,16 @@ PARALLEL_REQUEST_SLOT_TTL_SECONDS = 3600 class RateLimitDescriptorRateLimitObject(TypedDict, total=False): - requests_per_unit: Optional[int] - tokens_per_unit: Optional[int] - max_parallel_requests: Optional[int] - window_size: Optional[int] + requests_per_unit: int | None + tokens_per_unit: int | None + max_parallel_requests: int | None + window_size: int | None class RateLimitDescriptor(TypedDict): key: str value: str - rate_limit: Optional[RateLimitDescriptorRateLimitObject] + rate_limit: RateLimitDescriptorRateLimitObject | None class ParallelRequestGauge(TypedDict): @@ -339,11 +334,11 @@ class RateLimitStatus(TypedDict): class RateLimitResponse(TypedDict): overall_code: str - statuses: List[RateLimitStatus] + statuses: list[RateLimitStatus] class RateLimitResponseWithDescriptors(TypedDict): - descriptors: List[RateLimitDescriptor] + descriptors: list[RateLimitDescriptor] response: RateLimitResponse @@ -370,21 +365,21 @@ class RequestRateLimiterStash: mint fresh ids and are ignored. """ - owner_litellm_call_id: Optional[str] = None - rate_limit_response: Optional[RateLimitResponse] = None - parallel_slot: Optional[ParallelSlotAcquisition] = None + owner_litellm_call_id: str | None = None + rate_limit_response: RateLimitResponse | None = None + parallel_slot: ParallelSlotAcquisition | None = None reserved_tokens: int = 0 - reserved_model: Optional[str] = None - reserved_scopes: FrozenSet[Tuple[str, str]] = field(default_factory=frozenset) + reserved_model: str | None = None + reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) reservation_released: bool = False -_request_stash: ContextVar[Optional[RequestRateLimiterStash]] = ContextVar( +_request_stash: ContextVar[RequestRateLimiterStash | None] = ContextVar( "litellm_v3_rate_limiter_request_stash", default=None ) -def get_request_stash() -> Optional[RequestRateLimiterStash]: +def get_request_stash() -> RequestRateLimiterStash | None: return _request_stash.get() @@ -404,7 +399,7 @@ def claim_request_stash_for_data(data: dict) -> RequestRateLimiterStash: return stash -def get_request_stash_for_call(litellm_call_id: Optional[str]) -> Optional[RequestRateLimiterStash]: +def get_request_stash_for_call(litellm_call_id: str | None) -> RequestRateLimiterStash | None: stash = _request_stash.get() if stash is None: return None @@ -413,7 +408,7 @@ def get_request_stash_for_call(litellm_call_id: Optional[str]) -> Optional[Reque return stash if litellm_call_id == stash.owner_litellm_call_id else None -def _call_id_from_callback_kwargs(kwargs: object) -> Optional[str]: +def _call_id_from_callback_kwargs(kwargs: object) -> str | None: if not isinstance(kwargs, dict): return None call_id = kwargs.get("litellm_call_id") @@ -424,7 +419,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def __init__( self, internal_usage_cache: InternalUsageCache, - time_provider: Optional[Callable[[], datetime]] = None, + time_provider: Callable[[], datetime] | None = None, ): self.internal_usage_cache = internal_usage_cache self._time_provider = time_provider or datetime.now @@ -464,7 +459,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.tpm_reservation_enabled = os.getenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "true").lower() == "true" # Batch rate limiter (lazy loaded) - self._batch_rate_limiter: Optional[Any] = None + self._batch_rate_limiter: Any | None = None # Serializes multi-phase check+increment sequences (batch + dynamic # limiters) within this process to close the TOCTOU window between @@ -482,7 +477,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # one round-trip. self._check_and_increment_lock = asyncio.Lock() - def _get_batch_rate_limiter(self) -> Optional[Any]: + def _get_batch_rate_limiter(self) -> Any | None: """Get or lazy-load the batch rate limiter.""" if self._batch_rate_limiter is None: try: @@ -495,7 +490,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parallel_request_limiter=self, ) except Exception as e: - verbose_proxy_logger.debug(f"Could not load batch rate limiter: {str(e)}") + verbose_proxy_logger.debug(f"Could not load batch rate limiter: {e!s}") return self._batch_rate_limiter def _get_current_time(self) -> datetime: @@ -504,7 +499,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): @staticmethod def _no_max_tokens_output_floor( - min_configured_tpm_limit: Optional[int], + min_configured_tpm_limit: int | None, ) -> int: """Output-budget floor used when the request omits max_tokens. @@ -520,8 +515,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _estimate_tokens_for_request( self, data: dict, - model: Optional[str] = None, - min_configured_tpm_limit: Optional[int] = None, + model: str | None = None, + min_configured_tpm_limit: int | None = None, ) -> int: """ Estimate total tokens this request will consume so we can reserve them @@ -606,15 +601,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def in_memory_cache_sliding_window( self, - keys: List[str], + keys: list[str], now_int: int, window_size: int, - ) -> List[Any]: + ) -> list[Any]: """ Implement sliding window rate limiting logic using in-memory cache operations. This follows the same logic as the Redis Lua script but uses async cache operations. """ - results: List[Any] = [] + results: list[Any] = [] # Process each window/counter pair for i in range(0, len(keys), 2): @@ -683,14 +678,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def is_cache_list_over_limit( self, - keys_to_fetch: List[str], - cache_values: List[Any], - key_metadata: Dict[str, Any], + keys_to_fetch: list[str], + cache_values: list[Any], + key_metadata: dict[str, Any], ) -> RateLimitResponse: """ Check if the cache values are over the limit. """ - statuses: List[RateLimitStatus] = [] + statuses: list[RateLimitStatus] = [] overall_code = "OK" for i in range(0, len(cache_values), 2): @@ -702,8 +697,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): tokens_limit = key_metadata[window_key]["tokens_limit"] # Determine which limit to use for current_limit and limit_remaining - current_limit: Optional[int] = None - rate_limit_type: Optional[Literal["requests", "tokens", "max_parallel_requests"]] = None + current_limit: int | None = None + rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"] | None = None if counter_key.endswith(":requests"): current_limit = requests_limit rate_limit_type = "requests" @@ -760,14 +755,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): crc = binascii.crc_hqx(key.encode("utf-8"), 0) return crc % REDIS_CLUSTER_SLOTS - def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]: + def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. For Redis clusters, uses slot calculation to group keys that belong to the same slot. For regular Redis, no grouping is needed - all keys can be processed together. """ - groups: Dict[str, List[str]] = {} + groups: dict[str, list[str]] = {} # Use slot calculation for Redis clusters only if self._is_redis_cluster(): @@ -786,9 +781,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _execute_redis_batch_rate_limiter_script( self, - keys_to_fetch: List[str], + keys_to_fetch: list[str], now_int: int, - ) -> List[Any]: + ) -> list[Any]: """ Execute Redis operations grouped by hash tag for cluster compatibility. @@ -813,7 +808,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) all_cache_values.extend(group_cache_values) except Exception as e: - verbose_proxy_logger.warning(f"Redis Lua script failed for hash tag {hash_tag}: {str(e)}") + verbose_proxy_logger.warning(f"Redis Lua script failed for hash tag {hash_tag}: {e!s}") # Fallback to in-memory cache for this group group_cache_values = await self.in_memory_cache_sliding_window( keys=group_keys, @@ -826,8 +821,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def should_rate_limit( self, - descriptors: List[RateLimitDescriptor], - parent_otel_span: Optional[Span] = None, + descriptors: list[RateLimitDescriptor], + parent_otel_span: Span | None = None, read_only: bool = False, skip_tpm_check: bool = False, parallel_slot_id: str | None = None, @@ -960,7 +955,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): list with its per-window metadata, and the concurrency gauges for descriptors carrying a max_parallel_requests limit. """ - keys_to_fetch: List[str] = [] + keys_to_fetch: list[str] = [] key_metadata: dict[str, dict[str, Any]] = {} gauges: list[ParallelRequestGauge] = [] for descriptor in descriptors: @@ -1060,7 +1055,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) counts = [max(0, int(value)) for value in raw_counts] except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror, never a 500 - verbose_proxy_logger.warning(f"parallel_count_script failed, using local mirror: {str(e)}") + verbose_proxy_logger.warning(f"parallel_count_script failed, using local mirror: {e!s}") counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span) else: counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span) @@ -1090,9 +1085,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ], ) except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to in-memory enforcement, never a 500 - verbose_proxy_logger.warning( - f"parallel_acquire_script failed, falling back to in-memory gauge: {str(e)}" - ) + verbose_proxy_logger.warning(f"parallel_acquire_script failed, falling back to in-memory gauge: {e!s}") async with self._check_and_increment_lock: return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span) if int(raw[0]) == 1: @@ -1175,9 +1168,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): statuses = [] for gauge, (registry, in_flight) in zip(gauges, states): - new_value: Union[dict[str, float], int] = ( - {**registry, slot_id: now} if registry is not None else in_flight + 1 - ) + new_value: dict[str, float] | int = {**registry, slot_id: now} if registry is not None else in_flight + 1 await self.internal_usage_cache.async_set_cache( key=gauge["counter_key"], value=new_value, @@ -1222,7 +1213,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500 verbose_proxy_logger.warning( - f"parallel_release_script failed, falling back to in-memory release: {str(e)}" + f"parallel_release_script failed, falling back to in-memory release: {e!s}" ) async with self._check_and_increment_lock: @@ -1235,9 +1226,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if isinstance(raw_value, dict): if slot_id not in raw_value: continue - new_value: Union[dict[str, float], int] = { - key: ts for key, ts in raw_value.items() if key != slot_id - } + new_value: dict[str, float] | int = {key: ts for key, ts in raw_value.items() if key != slot_id} elif raw_value is None: continue else: @@ -1252,9 +1241,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def atomic_check_and_increment_by_n( self, - descriptors: List[RateLimitDescriptor], - increments: List[Dict[Literal["requests", "tokens"], int]], - parent_otel_span: Optional[Span] = None, + descriptors: list[RateLimitDescriptor], + increments: list[dict[Literal["requests", "tokens"], int]], + parent_otel_span: Span | None = None, ) -> RateLimitResponse: """ Atomic check-and-increment-by-N across one or more descriptors. @@ -1288,7 +1277,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Build per-descriptor (keys, args, meta) groups. All keys within a # group share the descriptor's {key:value} hash tag, so a single Lua # call per group never triggers CROSSSLOT on Redis Cluster. - descriptor_groups: List[Tuple[List[str], List[Any], List[Dict[str, Any]]]] = [] + descriptor_groups: list[tuple[list[str], list[Any], list[dict[str, Any]]]] = [] for descriptor, increment_amounts in zip(descriptors, increments): keys, args, meta = self._build_descriptor_atomic_payload( descriptor=descriptor, @@ -1311,7 +1300,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parent_otel_span=parent_otel_span, ) - flat_meta: List[Dict[str, Any]] = [m for _keys, _args, group_meta in descriptor_groups for m in group_meta] + flat_meta: list[dict[str, Any]] = [m for _keys, _args, group_meta in descriptor_groups for m in group_meta] async with self._check_and_increment_lock: return await self._atomic_check_and_increment_in_memory( per_counter_meta=flat_meta, @@ -1321,8 +1310,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _build_descriptor_atomic_payload( self, descriptor: RateLimitDescriptor, - increment_amounts: Dict[Literal["requests", "tokens"], int], - ) -> Tuple[List[str], List[Any], List[Dict[str, Any]]]: + increment_amounts: dict[Literal["requests", "tokens"], int], + ) -> tuple[list[str], list[Any], list[dict[str, Any]]]: """ Build (KEYS, ARGV, per-counter meta) for a single descriptor's Lua call. All keys returned share the descriptor's {key:value} hash tag. @@ -1335,9 +1324,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): window_size = rate_limit.get("window_size") or self.window_size window_key = f"{{{descriptor_key}:{descriptor_value}}}:window" - keys: List[str] = [] - args: List[Any] = [] - meta: List[Dict[str, Any]] = [] + keys: list[str] = [] + args: list[Any] = [] + meta: list[dict[str, Any]] = [] for rate_limit_type in ("requests", "tokens"): rlt: Literal["requests", "tokens"] = cast(Literal["requests", "tokens"], rate_limit_type) @@ -1376,8 +1365,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _atomic_lua_per_descriptor( self, - descriptor_groups: List[Tuple[List[str], List[Any], List[Dict[str, Any]]]], - parent_otel_span: Optional[Span] = None, + descriptor_groups: list[tuple[list[str], list[Any], list[dict[str, Any]]]], + parent_otel_span: Span | None = None, ) -> RateLimitResponse: """ Run Lua check-and-increment one descriptor at a time so each call's @@ -1385,8 +1374,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptor i, refund descriptors 0..i-1's increments. On Lua failure mid-loop, refund applied increments and fall back to in-memory. """ - applied: List[List[Dict[str, Any]]] = [] - statuses: List[RateLimitStatus] = [] + applied: list[list[dict[str, Any]]] = [] + statuses: list[RateLimitStatus] = [] for _idx, (keys, args, meta) in enumerate(descriptor_groups): try: @@ -1408,7 +1397,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): f"{self.window_size}s)." ) await self._refund_applied_descriptor_groups(applied) - flat_meta: List[Dict[str, Any]] = [m for _k, _a, group_meta in descriptor_groups for m in group_meta] + flat_meta: list[dict[str, Any]] = [m for _k, _a, group_meta in descriptor_groups for m in group_meta] async with self._check_and_increment_lock: return await self._atomic_check_and_increment_in_memory( per_counter_meta=flat_meta, @@ -1426,7 +1415,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _refund_applied_descriptor_groups( self, - applied: List[List[Dict[str, Any]]], + applied: list[list[dict[str, Any]]], ) -> None: """ Decrement counters for descriptor groups already applied via Lua. @@ -1452,8 +1441,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _build_atomic_response( self, - raw: List[Any], - per_counter_meta: List[Dict[str, Any]], + raw: list[Any], + per_counter_meta: list[dict[str, Any]], ) -> RateLimitResponse: """Convert Lua script return value to RateLimitResponse. @@ -1489,7 +1478,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ], ) - statuses: List[RateLimitStatus] = [] + statuses: list[RateLimitStatus] = [] for meta, new_counter in zip(per_counter_meta, raw[1:]): statuses.append( RateLimitStatus( @@ -1504,8 +1493,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _atomic_check_and_increment_in_memory( self, - per_counter_meta: List[Dict[str, Any]], - parent_otel_span: Optional[Span] = None, + per_counter_meta: list[dict[str, Any]], + parent_otel_span: Span | None = None, ) -> RateLimitResponse: """In-memory all-or-nothing check-and-increment. Caller holds lock. @@ -1519,7 +1508,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): now_int = int(self._get_current_time().timestamp()) # Pass 1: read state, validate. - descriptor_state: List[Dict[str, Any]] = [] + descriptor_state: list[dict[str, Any]] = [] for meta in per_counter_meta: window_size = meta["window_size"] window_start = await self.internal_usage_cache.async_get_cache( @@ -1561,7 +1550,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptor_state.append({"window_expired": window_expired, "current": current_counter}) # Pass 2: apply increments. - statuses: List[RateLimitStatus] = [] + statuses: list[RateLimitStatus] = [] for meta, state in zip(per_counter_meta, descriptor_state): new_counter = meta["increment"] if state["window_expired"] else state["current"] + meta["increment"] if state["window_expired"]: @@ -1592,9 +1581,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def reserve_tpm_tokens( self, - descriptors: List[RateLimitDescriptor], + descriptors: list[RateLimitDescriptor], estimated_tokens: int, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> RateLimitResponse: """ Reserve ``estimated_tokens`` against every TPM-bearing descriptor @@ -1606,7 +1595,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): atomicity (Lua on Redis, asyncio-locked DualCache otherwise) to the shared primitive. """ - tpm_descriptors: List[RateLimitDescriptor] = [ + tpm_descriptors: list[RateLimitDescriptor] = [ RateLimitDescriptor( key=d["key"], value=d["value"], @@ -1621,7 +1610,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if not tpm_descriptors: return RateLimitResponse(overall_code="OK", statuses=[]) - increments: List[Dict[Literal["requests", "tokens"], int]] = [ + increments: list[dict[Literal["requests", "tokens"], int]] = [ {"tokens": estimated_tokens} for _ in tpm_descriptors ] return await self.atomic_check_and_increment_by_n( @@ -1631,9 +1620,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) def create_organization_rate_limit_descriptor( - self, user_api_key_dict: UserAPIKeyAuth, requested_model: Optional[str] = None - ) -> List[RateLimitDescriptor]: - descriptors: List[RateLimitDescriptor] = [] + self, user_api_key_dict: UserAPIKeyAuth, requested_model: str | None = None + ) -> list[RateLimitDescriptor]: + descriptors: list[RateLimitDescriptor] = [] # Global org rate limits if user_api_key_dict.org_id is not None and ( @@ -1666,9 +1655,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) should_check_rate_limit = False - if requested_model in _tpm_limit_for_team_model: - should_check_rate_limit = True - elif requested_model in _rpm_limit_for_team_model: + if requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model: should_check_rate_limit = True if should_check_rate_limit: @@ -1695,8 +1682,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _add_model_per_key_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, - requested_model: Optional[str], - descriptors: List[RateLimitDescriptor], + requested_model: str | None, + descriptors: list[RateLimitDescriptor], ) -> None: """ Add model-specific rate limit descriptor for API key if applicable. @@ -1732,8 +1719,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return # Get model-specific limits - model_specific_tpm_limit: Optional[int] = _tpm_limit_for_key_model.get(requested_model) - model_specific_rpm_limit: Optional[int] = _rpm_limit_for_key_model.get(requested_model) + model_specific_tpm_limit: int | None = _tpm_limit_for_key_model.get(requested_model) + model_specific_rpm_limit: int | None = _rpm_limit_for_key_model.get(requested_model) descriptors.append( RateLimitDescriptor( @@ -1787,8 +1774,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _add_mcp_per_key_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, - mcp_server_name: Optional[str], - descriptors: List[RateLimitDescriptor], + mcp_server_name: str | None, + descriptors: list[RateLimitDescriptor], ) -> None: """ Add a per-MCP-server rpm descriptor for the API key, if a limit is @@ -1825,8 +1812,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _add_mcp_per_team_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, - mcp_server_name: Optional[str], - descriptors: List[RateLimitDescriptor], + mcp_server_name: str | None, + descriptors: list[RateLimitDescriptor], ) -> None: """ Add a per-MCP-server rpm descriptor for the team, if a limit is @@ -1869,7 +1856,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _should_enforce_rate_limit( self, - limit_type: Optional[str], + limit_type: str | None, model_has_failures: bool, ) -> bool: """ @@ -1890,10 +1877,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _get_enforced_limit( self, - limit_value: Optional[int], - limit_type: Optional[str], + limit_value: int | None, + limit_type: str | None, model_has_failures: bool, - ) -> Optional[int]: + ) -> int | None: """ Get the rate limit value to enforce based on limit type and model health. @@ -1918,8 +1905,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _is_dynamic_rate_limiting_enabled( self, - rpm_limit_type: Optional[str], - tpm_limit_type: Optional[str], + rpm_limit_type: str | None, + tpm_limit_type: str | None, ) -> bool: """ Check if dynamic rate limiting is enabled for either RPM or TPM. @@ -1933,13 +1920,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ return rpm_limit_type == "dynamic" or tpm_limit_type == "dynamic" - def _get_agent_from_registry(self, agent_id: str) -> Optional[Any]: + def _get_agent_from_registry(self, agent_id: str) -> Any | None: """Look up an agent from the in-memory registry by ID.""" from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry return global_agent_registry.get_agent_by_id(agent_id=agent_id) - def _get_resolved_agent_id(self, user_api_key_dict: UserAPIKeyAuth, data: dict) -> Optional[str]: + def _get_resolved_agent_id(self, user_api_key_dict: UserAPIKeyAuth, data: dict) -> str | None: """ Resolve the agent_id from either the API key or request metadata. Key-level agent_id takes precedence over metadata/header-supplied agent_id. @@ -1950,7 +1937,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): metadata = data.get("metadata") or {} return metadata.get("agent_id") - def _get_session_id_from_data(self, data: dict) -> Optional[str]: + def _get_session_id_from_data(self, data: dict) -> str | None: """Extract session_id from request metadata or litellm_session_id.""" session_id = data.get("litellm_session_id") if session_id: @@ -1969,14 +1956,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self, agent_id: str, data: dict, - ) -> List[RateLimitDescriptor]: + ) -> list[RateLimitDescriptor]: """ Create rate limit descriptors for agent-level and session-level limits. Agent-level: caps total RPM/TPM across all sessions for a given agent. Session-level: caps RPM/TPM within a single session (identified by session_id). """ - descriptors: List[RateLimitDescriptor] = [] + descriptors: list[RateLimitDescriptor] = [] agent = self._get_agent_from_registry(agent_id) if agent is None: @@ -2020,11 +2007,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, data: dict, - rpm_limit_type: Optional[str], - tpm_limit_type: Optional[str], + rpm_limit_type: str | None, + tpm_limit_type: str | None, model_has_failures: bool, - call_type: Optional[str] = None, - ) -> List[RateLimitDescriptor]: + call_type: str | None = None, + ) -> list[RateLimitDescriptor]: """ Create all rate limit descriptors for the request. @@ -2168,9 +2155,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): _tpm_limit_for_team_model = get_team_model_tpm_limit(user_api_key_dict) or {} _rpm_limit_for_team_model = get_team_model_rpm_limit(user_api_key_dict) or {} should_check_rate_limit = False - if requested_model in _tpm_limit_for_team_model: - should_check_rate_limit = True - elif requested_model in _rpm_limit_for_team_model: + if requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model: should_check_rate_limit = True if should_check_rate_limit: @@ -2208,7 +2193,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _check_model_has_recent_failures( self, model: str, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> bool: """ Check if any deployment for this model has recent failures by using @@ -2255,11 +2240,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return False except Exception as e: - verbose_proxy_logger.debug(f"Error checking model failure status: {str(e)}, defaulting to enforce limits") + verbose_proxy_logger.debug(f"Error checking model failure status: {e!s}, defaulting to enforce limits") # Fail safe: enforce limits if we can't check return True - def get_rate_limiter_for_call_type(self, call_type: str) -> Optional[Any]: + def get_rate_limiter_for_call_type(self, call_type: str) -> Any | None: """Get the rate limiter for the call type.""" if call_type == "acreate_batch": batch_limiter = self._get_batch_rate_limiter() @@ -2269,8 +2254,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _add_team_model_rate_limit_descriptor_from_metadata( self, user_api_key_dict: UserAPIKeyAuth, - requested_model: Optional[str], - descriptors: List[RateLimitDescriptor], + requested_model: str | None, + descriptors: list[RateLimitDescriptor], ) -> None: """Add team model rate limit descriptor from team_metadata if applicable.""" if ( @@ -2305,8 +2290,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _add_project_model_rate_limit_descriptor_from_metadata( self, user_api_key_dict: UserAPIKeyAuth, - requested_model: Optional[str], - descriptors: List[RateLimitDescriptor], + requested_model: str | None, + descriptors: list[RateLimitDescriptor], ) -> None: """Add project model rate limit descriptor from project_metadata if applicable.""" if ( @@ -2341,8 +2326,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _handle_rate_limit_error( self, response: RateLimitResponse, - descriptors: List[RateLimitDescriptor], - requested_model: Optional[str] = None, + descriptors: list[RateLimitDescriptor], + requested_model: str | None = None, ) -> None: """Handle rate limit exceeded by raising :class:`ProxyRateLimitError` (a 429).""" for status in response["statuses"]: @@ -2598,11 +2583,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): value: str, rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"], total_tokens: int, - ) -> List["RedisPipelineIncrementOperation"]: + ) -> list["RedisPipelineIncrementOperation"]: """ Create pipeline operations for TPM increments """ - pipeline_operations: List[RedisPipelineIncrementOperation] = [] + pipeline_operations: list[RedisPipelineIncrementOperation] = [] counter_key = self.create_rate_limit_keys( key=key, value=value, @@ -2619,7 +2604,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations def _get_total_tokens_from_usage( - self, usage: Optional[Any], rate_limit_type: Literal["output", "input", "total"] + self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"] ) -> int: """ Get total tokens from response usage for rate limiting. @@ -2667,7 +2652,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return total_tokens @staticmethod - def _aggregate_only_total_tokens(usage: Union[Usage, dict, None]) -> int: + def _aggregate_only_total_tokens(usage: Usage | dict | None) -> int: """Total for usage that carries no input/output split, else 0. A source that can only report one number for the whole request (a @@ -2697,7 +2682,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _execute_token_increment_script( self, - pipeline_operations: List["RedisPipelineIncrementOperation"], + pipeline_operations: list["RedisPipelineIncrementOperation"], ) -> None: """ Execute token increment script grouped by hash tag for cluster compatibility. @@ -2734,8 +2719,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def async_increment_tokens_with_ttl_preservation( self, - pipeline_operations: List["RedisPipelineIncrementOperation"], - parent_otel_span: Optional[Span] = None, + pipeline_operations: list["RedisPipelineIncrementOperation"], + parent_otel_span: Span | None = None, ) -> None: """ Increment token counters using Lua script to preserve existing TTL. @@ -2761,7 +2746,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) except Exception as e: - verbose_proxy_logger.warning(f"TTL preservation failed, falling back to regular pipeline: {str(e)}") + verbose_proxy_logger.warning(f"TTL preservation failed, falling back to regular pipeline: {e!s}") # Fallback to regular pipeline on error await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( increment_list=pipeline_operations, @@ -2782,15 +2767,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): @staticmethod def _merge_ratelimit_statuses_into_additional_headers( - additional_headers: Dict[str, Any], - statuses: List[RateLimitStatus], - ) -> Dict[str, Any]: + additional_headers: dict[str, Any], + statuses: list[RateLimitStatus], + ) -> dict[str, Any]: """ Return ``additional_headers`` extended with ``x-ratelimit-{descriptor_key}-{remaining|limit}-{rate_limit_type}`` entries. Non-mutating so callers pick their own target dict. """ - merged: Dict[str, Any] = dict(additional_headers) + merged: dict[str, Any] = dict(additional_headers) for status in statuses: prefix = f"x-ratelimit-{status['descriptor_key']}" merged[f"{prefix}-remaining-{status['rate_limit_type']}"] = status["limit_remaining"] @@ -2799,10 +2784,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _collect_tpm_scope_targets( self, - standard_logging_metadata: Dict[str, Any], + standard_logging_metadata: dict[str, Any], kwargs: Any, - model_group: Optional[str], - ) -> List[Tuple[str, str]]: + model_group: str | None, + ) -> list[tuple[str, str]]: """ Enumerate every (scope_key, scope_value) pair that *might* carry a TPM counter for this request — independent of whether each scope had @@ -2821,7 +2806,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): agent_id = standard_logging_metadata.get("agent_id") session_id = standard_logging_metadata.get("session_id") or standard_logging_metadata.get("trace_id") - targets: List[Tuple[str, str]] = [] + targets: list[tuple[str, str]] = [] if user_api_key: targets.append(("api_key", user_api_key)) if user_api_key_user_id: @@ -2861,11 +2846,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _build_reservation_aware_tpm_ops( self, - targets: List[Tuple[str, str]], - reserved_scopes: FrozenSet[Tuple[str, str]], + targets: list[tuple[str, str]], + reserved_scopes: frozenset[tuple[str, str]], actual_tokens: int, reserved_tokens: int, - ) -> List[RedisPipelineIncrementOperation]: + ) -> list[RedisPipelineIncrementOperation]: """ Emit per-scope TPM increment ops with reservation awareness. @@ -2878,7 +2863,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): release, and failure refund — pass ``actual_tokens=0`` for the pure refund case (reserved scopes get -reserved, unreserved get 0/skip). """ - ops: List[RedisPipelineIncrementOperation] = [] + ops: list[RedisPipelineIncrementOperation] = [] for scope_key, scope_value in targets: if (scope_key, scope_value) in reserved_scopes: increment = actual_tokens - reserved_tokens @@ -2900,7 +2885,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): kwargs: Any, response_obj: Any, rate_limit_type: Literal["output", "input", "total"], - ) -> List[RedisPipelineIncrementOperation]: + ) -> list[RedisPipelineIncrementOperation]: """Build Redis pipeline increment ops for TPM / parallel-request counters.""" from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, @@ -2918,7 +2903,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # than parsed out of the body) carry their usage in # ``combined_usage_object`` instead, and would otherwise never charge # the TPM window. - _usage: Union[Usage, dict, None] = None + _usage: Usage | dict | None = None if isinstance( response_obj, ( @@ -2940,14 +2925,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) reserved_tokens = stash.reserved_tokens if stash is not None else 0 reserved_model = stash.reserved_model if stash is not None else None - reserved_scopes: FrozenSet[Tuple[str, str]] = stash.reserved_scopes if stash is not None else frozenset() + reserved_scopes: frozenset[tuple[str, str]] = stash.reserved_scopes if stash is not None else frozenset() # Reconciliation must target the same model-scoped counter that the # pre-call reservation incremented. If a reservation was made, # ``reserved_model`` is authoritative; otherwise fall back to the # router's ``model_group`` (covers the no-reservation charge path). reconcile_model = reserved_model or model_group - pipeline_operations: List[RedisPipelineIncrementOperation] = [] + pipeline_operations: list[RedisPipelineIncrementOperation] = [] # ---------------------------------------------------------------- # TPM reconciliation @@ -2992,7 +2977,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): rate_limit_type = self.get_rate_limit_type() - litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(kwargs) try: verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") @@ -3018,14 +3003,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) except Exception as e: - verbose_proxy_logger.exception(f"Error in rate limit success event: {str(e)}") + verbose_proxy_logger.exception(f"Error in rate limit success event: {e!s}") async def async_logging_hook( self, kwargs: dict, result: Any, call_type: str, - ) -> Tuple[dict, Any]: + ) -> tuple[dict, Any]: """ Mirror the pre-call rate-limit snapshot into the SLP so streaming success callbacks see the same ``x-ratelimit-*`` headers the @@ -3091,9 +3076,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) try: - litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(kwargs) - pipeline_operations: List[RedisPipelineIncrementOperation] = [] + pipeline_operations: list[RedisPipelineIncrementOperation] = [] stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) acquisition = stash.parallel_slot if stash is not None else None @@ -3135,7 +3120,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if stash is not None and reserved_tokens > 0: stash.reservation_released = True except Exception as e: - verbose_proxy_logger.exception(f"Error in rate limit failure event: {str(e)}") + verbose_proxy_logger.exception(f"Error in rate limit failure event: {e!s}") async def async_release_max_parallel_requests_on_disconnect( self, @@ -3200,14 +3185,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) except Exception as e: - verbose_proxy_logger.exception(f"Error in rate limit post-call hook: {str(e)}") + verbose_proxy_logger.exception(f"Error in rate limit post-call hook: {e!s}") async def async_post_call_failure_hook( self, request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, + traceback_str: str | None = None, ) -> None: """ Release the parallel-request slot and any TPM reservation when the @@ -3256,4 +3241,4 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): stash.reservation_released = True except Exception as e: verbose_proxy_logger.exception(f"Error releasing TPM reservation on post-call failure: {e}") - return None + return diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index 2d55a644bb2..3e8518d55dc 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -8,7 +8,7 @@ from difflib import SequenceMatcher -from typing import List, Literal, Optional +from typing import Literal from fastapi import HTTPException @@ -29,10 +29,10 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): # Class variables or attributes def __init__( self, - prompt_injection_params: Optional[LiteLLMPromptInjectionParams] = None, + prompt_injection_params: LiteLLMPromptInjectionParams | None = None, ): self.prompt_injection_params = prompt_injection_params - self.llm_router: Optional[Router] = None + self.llm_router: Router | None = None self.verbs = [ "Ignore", @@ -74,7 +74,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): if litellm.set_verbose is True: print(print_statement) # noqa: T201 - def update_environment(self, router: Optional[Router] = None): + def update_environment(self, router: Router | None = None): self.llm_router = router if self.prompt_injection_params is not None and self.prompt_injection_params.llm_api_check is True: @@ -94,7 +94,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): "PromptInjectionDetection: Invalid LLM API Name. LLM API Name must be a 'model_name' in 'model_list'." ) - def generate_injection_keywords(self) -> List[str]: + def generate_injection_keywords(self) -> list[str]: combinations = [] for verb in self.verbs: for adj in self.adjectives: @@ -197,9 +197,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): raise e except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {e!s}" ) async def async_moderation_hook( # type: ignore @@ -214,7 +212,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): "moderation", "audio_transcription", ], - ) -> Optional[bool]: + ) -> bool | None: self.print_verbose(f"IN ASYNC MODERATION HOOK - self.prompt_injection_params = {self.prompt_injection_params}") if self.prompt_injection_params is None: diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index b839426fcda..857429fa89f 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -1,568 +1,570 @@ -import asyncio -import traceback -from datetime import datetime -from typing import Any, List, Optional, Union, cast - -import litellm -from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, - get_litellm_metadata_from_kwargs, -) -from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - get_key_object, - get_team_object, - log_db_metrics, -) -from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup -from litellm.proxy.spend_tracking.spend_log_error_logger import ( - should_suppress_spend_log_tracebacks, - spend_log_error, -) -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - _sanitize_error_information_for_spend_logs, -) -from litellm.proxy.utils import ProxyUpdateSpend -from litellm.types.utils import ( - CallTypes, - StandardLoggingPayload, - StandardLoggingPayloadErrorInformation, -) -from litellm.utils import get_end_user_id_for_cost_tracking - -_PASS_THROUGH_CALL_TYPES: frozenset[str] = frozenset( - { - CallTypes.pass_through.value, - CallTypes.llm_passthrough_route.value, - CallTypes.allm_passthrough_route.value, - } -) - - -class _ProxyDBLogger(CustomLogger): - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) - - async def async_post_call_failure_hook( - self, - request_data: dict, - original_exception: Exception, - user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, - ): - try: - await _release_budget_reservation(budget_reservation=user_api_key_dict.budget_reservation) - except Exception: - verbose_proxy_logger.exception("Failed to release budget reservation during failure handling") - try: - await _invalidate_budget_reservation_counters(budget_reservation=user_api_key_dict.budget_reservation) - if user_api_key_dict.budget_reservation is not None: - user_api_key_dict.budget_reservation["finalized"] = True - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after failure release failed" - ) - - request_route = user_api_key_dict.request_route - if _ProxyDBLogger._should_track_errors_in_db() is False: - return - elif request_route is not None and not ( - RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) - ): - return - - from litellm.proxy.proxy_server import proxy_logging_obj - - _metadata = dict( - LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) - ) - _metadata["user_api_key"] = user_api_key_dict.api_key - _metadata["status"] = "failure" - _error_information = StandardLoggingPayloadSetup.get_error_information( - original_exception=original_exception, - traceback_str=traceback_str, - ) - if should_suppress_spend_log_tracebacks(): - # Drop the traceback key entirely so the per-row Metadata pane in - # the UI (which renders the JSON blob verbatim) doesn't show a - # noisy ``"traceback": ""`` line. Downstream consumers all use - # ``.get("traceback")`` / truthy checks, and the TypedDict marks - # the field as optional, so omitting is type-safe. - _error_information.pop("traceback", None) - # Strip echoed request input + apply DB-size cap before storing in - # the spend-log metadata column (LIT-2992). Result is never None - # here because the input above is constructed non-None. - _error_information = cast( - StandardLoggingPayloadErrorInformation, - _sanitize_error_information_for_spend_logs(_error_information), - ) - _metadata["error_information"] = _error_information - - _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( - metadata=_metadata, - ) - - existing_metadata: dict = request_data.get("metadata", None) or {} - existing_metadata.update(_metadata) - - if "litellm_params" not in request_data: - request_data["litellm_params"] = {} - - existing_litellm_params = request_data.get("litellm_params", {}) - existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {} - - # Preserve tags from existing metadata - if existing_litellm_metadata.get("tags"): - existing_metadata["tags"] = existing_litellm_metadata.get("tags") - - request_data["litellm_params"]["proxy_server_request"] = ( - request_data.get("proxy_server_request") or existing_litellm_params.get("proxy_server_request") or {} - ) - request_data["litellm_params"]["metadata"] = existing_metadata - - # Preserve model name and custom_llm_provider - if "model" not in request_data: - request_data["model"] = existing_litellm_params.get("model") or request_data.get("model", "") - if "custom_llm_provider" not in request_data: - request_data["custom_llm_provider"] = existing_litellm_params.get( - "custom_llm_provider" - ) or request_data.get("custom_llm_provider", "") - - # Propagate standard_logging_object and litellm_trace_id from the - # Logging instance so that _get_session_id_for_spend_log uses the same - # trace_id that Langfuse received (via async_failure_handler). - # Without this, the DB session_id would be a random UUID that doesn't - # match the Langfuse trace_id, making failed requests unsearchable. - _litellm_logging_obj = request_data.get("litellm_logging_obj") - if _litellm_logging_obj is not None: - if not request_data.get("standard_logging_object"): - request_data["standard_logging_object"] = getattr(_litellm_logging_obj, "model_call_details", {}).get( - "standard_logging_object" - ) - if request_data.get("litellm_trace_id") is None: - request_data["litellm_trace_id"] = getattr(_litellm_logging_obj, "litellm_trace_id", None) - - # Use the actual request start time from the logging object so that - # failed requests record the real duration instead of 0. - actual_start_time = datetime.now() - if _litellm_logging_obj is not None: - obj_start = getattr(_litellm_logging_obj, "start_time", None) - if obj_start is not None: - actual_start_time = obj_start - - # A stream that broke mid-flight still billed the provider for the - # chunks already delivered. ``post_call_failure_hook`` lifts that - # recovered cost onto request_data (the usage rides along in - # ``combined_usage_object`` for the token columns), so attribute the - # real partial spend to this failure row instead of zero. - recovered_response_cost = 0.0 - if isinstance(request_data.get("combined_usage_object"), litellm.Usage): - recovered_response_cost = max(float(request_data.get("response_cost") or 0.0), 0.0) - - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key_dict.api_key, - response_cost=recovered_response_cost, - user_id=user_api_key_dict.user_id, - end_user_id=user_api_key_dict.end_user_id, - team_id=user_api_key_dict.team_id, - kwargs=request_data, - completion_response=original_exception, - start_time=actual_start_time, - end_time=datetime.now(), - org_id=user_api_key_dict.org_id, - ) - - @log_db_metrics - async def _PROXY_track_cost_callback( - self, - kwargs, # kwargs to completion - completion_response: Optional[Union[litellm.ModelResponse, Any]], # response from completion - start_time=None, - end_time=None, # start/end time for completion - ): - from litellm.proxy.proxy_server import ( - increment_spend_counters, - proxy_logging_obj, - update_cache, - ) - - verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") - try: - verbose_proxy_logger.debug( - f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" - ) - parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) - litellm_params = kwargs.get("litellm_params", {}) or {} - end_user_id = get_end_user_id_for_cost_tracking(litellm_params) - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). - # Avoids a cache/DB lookup on every normal LLM request. - if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): - metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) - _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) - budget_reservation = _get_budget_reservation_from_metadata(metadata=metadata) - user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) - team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) - org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) - key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None)) - end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) - sl_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) - response_cost = ( - sl_object.get("response_cost", None) if sl_object is not None else kwargs.get("response_cost", None) - ) - tags = _get_request_tags_for_cost_tracking( - sl_object=sl_object, - metadata=metadata, - ) - - if response_cost is not None: - user_api_key = metadata.get("user_api_key", None) - if kwargs.get("cache_hit", False) is True: - response_cost = 0.0 - verbose_proxy_logger.debug(f"Cache Hit: response_cost {response_cost}, for user_id {user_id}") - - verbose_proxy_logger.debug( - f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" - ) - call_type: Optional[str] = kwargs.get("call_type") - if _should_track_cost_callback( - user_api_key=user_api_key, - user_id=user_id, - team_id=team_id, - end_user_id=end_user_id, - call_type=call_type, - ): - ## UPDATE DATABASE - await _update_database_and_spend_counters( - proxy_logging_obj=proxy_logging_obj, - increment_spend_counters=increment_spend_counters, - user_api_key=user_api_key, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - org_id=org_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - response_cost=response_cost, - budget_reservation=budget_reservation, - request_tags=tags, - ) - - # update cache (fire-and-forget for backward compat: - # cached object fields, soft budget alerts, etc.) - asyncio.create_task( - update_cache( - token=user_api_key, - user_id=user_id, - end_user_id=end_user_id, - response_cost=response_cost, - team_id=team_id, - parent_otel_span=parent_otel_span, - tags=tags, - ) - ) - - await proxy_logging_obj.slack_alerting_instance.customer_spend_alert( - token=user_api_key, - key_alias=key_alias, - end_user_id=end_user_id, - response_cost=response_cost, - max_budget=end_user_max_budget, - ) - elif budget_reservation is not None: - await _release_budget_reservation(budget_reservation=budget_reservation) - else: - await _release_budget_reservation(budget_reservation=budget_reservation) - # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. - # Use .get() for "stream" to avoid KeyError on health checks. - # WS session wrappers (_aresponses_websocket, _arealtime) also reach here with - # result=None; their per-turn costs are tracked on the inner aresponses/realtime calls. - if sl_object is None and ( - not kwargs.get("model") or kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime") - ): - verbose_proxy_logger.warning( - "Cost tracking - skipping, no standard_logging_object for call_type=%s", - kwargs.get("call_type", "unknown"), - ) - return - if kwargs.get("stream") is not True or ( - kwargs.get("stream") is True and "complete_streaming_response" in kwargs - ): - if sl_object is not None: - cost_tracking_failure_debug_info: Union[dict, str] = ( - sl_object["response_cost_failure_debug_info"] # type: ignore - or "response_cost_failure_debug_info is None in standard_logging_object" - ) - else: - cost_tracking_failure_debug_info = "standard_logging_object not found" - model = kwargs.get("model") - raise Exception( - f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing" - ) - except Exception as e: - error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}" - model = kwargs.get("model", "") - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - litellm_metadata = kwargs.get("litellm_params", {}).get("litellm_metadata", {}) - old_metadata = kwargs.get("litellm_params", {}).get("metadata", {}) - call_type = kwargs.get("call_type", "") - error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n" - asyncio.create_task( - proxy_logging_obj.failed_tracking_alert( - error_message=error_msg, - failing_model=model, - ) - ) - - spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) - - @staticmethod - async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: - """ - Enriches failure spend log metadata by looking up the key object (and team object) - from cache/DB when key fields are missing. - - This handles two scenarios: - 1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other - fields are null. We look up the full key object to fill in alias, user_id, - team_id, etc. - 2. Post-auth failures (provider errors, rate limits): key fields are populated - but team_alias is missing because LiteLLM_VerificationTokenView SQL view - doesn't include it. We look up the team object to fill in team_alias. - """ - api_key_hash = metadata.get("user_api_key") - if not api_key_hash: - return metadata - - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - # Step 1: If key fields are missing, look up the full key object - if metadata.get("user_api_key_alias") is None: - try: - key_obj = await get_key_object( - hashed_token=api_key_hash, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - if metadata.get("user_api_key_alias") is None: - metadata["user_api_key_alias"] = key_obj.key_alias - if metadata.get("user_api_key_user_id") is None: - metadata["user_api_key_user_id"] = key_obj.user_id - if metadata.get("user_api_key_team_id") is None: - metadata["user_api_key_team_id"] = key_obj.team_id - if metadata.get("user_api_key_org_id") is None: - metadata["user_api_key_org_id"] = key_obj.org_id - except Exception: - verbose_proxy_logger.debug( - "Failed to enrich failure metadata with key info for api_key=%s", - api_key_hash, - ) - - # Step 2: If team_id is known but team_alias is missing, look up the team object - team_id = metadata.get("user_api_key_team_id") - if team_id and metadata.get("user_api_key_team_alias") is None: - try: - team_obj = await get_team_object( - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - if team_obj.team_alias is not None: - metadata["user_api_key_team_alias"] = team_obj.team_alias - except Exception: - verbose_proxy_logger.debug( - "Failed to enrich failure metadata with team_alias for team_id=%s", - team_id, - ) - return metadata - - @staticmethod - def _should_track_errors_in_db(): - """ - Returns True if errors should be tracked in the database - - By default, errors are tracked in the database - - If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings - """ - from litellm.proxy.proxy_server import general_settings - - if general_settings.get("disable_error_logs") is True: - return False - return - - -def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: - patch = {k: v for k, v in metadata.items() if (k.startswith("user_api_key") or k == "tags") and v is not None} - if not patch: - return - - litellm_params = kwargs.setdefault("litellm_params", {}) - for bucket_name in ("litellm_metadata", "metadata"): - bucket = litellm_params.get(bucket_name) - if isinstance(bucket, dict): - for key, value in patch.items(): - if bucket.get(key) is None: - bucket[key] = value - - -def _should_track_cost_callback( - user_api_key: Optional[str], - user_id: Optional[str], - team_id: Optional[str], - end_user_id: Optional[str], - call_type: Optional[str] = None, -) -> bool: - """ - Determine if the cost callback should be tracked based on the kwargs - - Pass-through endpoints can be configured with ``auth=false``, which leaves - the request with no key/user/team/end-user to attribute spend to. Those - requests still forward real provider traffic that operators expect to see - in request/usage logs, so they are tracked even when unauthenticated. - """ - - # don't run track cost callback if user opted into disabling spend - if ProxyUpdateSpend.disable_spend_updates() is True: - return False - - if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None: - return True - return call_type in _PASS_THROUGH_CALL_TYPES - - -def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: - metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") - if isinstance(metadata_budget_reservation, dict): - return metadata_budget_reservation - - user_api_key_auth_obj = metadata.get("user_api_key_auth") - if user_api_key_auth_obj is None: - return None - if isinstance(user_api_key_auth_obj, dict): - budget_reservation = user_api_key_auth_obj.get("budget_reservation") - return budget_reservation if isinstance(budget_reservation, dict) else None - return getattr(user_api_key_auth_obj, "budget_reservation", None) - - -def _get_request_tags_for_cost_tracking( - sl_object: Optional[StandardLoggingPayload], - metadata: dict, -) -> Optional[List[str]]: - if sl_object is not None: - request_tags = sl_object.get("request_tags", None) - if isinstance(request_tags, list): - return request_tags - - metadata_tags = metadata.get("tags", None) - if isinstance(metadata_tags, list): - return metadata_tags - - return None - - -async def _update_database_and_spend_counters( - proxy_logging_obj: Any, - increment_spend_counters: Any, - user_api_key: Optional[str], - user_id: Optional[str], - end_user_id: Optional[str], - team_id: Optional[str], - org_id: Optional[str], - kwargs: dict, - completion_response: Optional[Union[litellm.ModelResponse, Any]], - start_time: Any, - end_time: Any, - response_cost: float, - budget_reservation: Optional[dict], - request_tags: Optional[List[str]] = None, -) -> None: - try: - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key, - response_cost=response_cost, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - org_id=org_id, - ) - except Exception: - if budget_reservation is not None: - try: - await _release_budget_reservation(budget_reservation=budget_reservation) - except Exception: - verbose_proxy_logger.exception("Failed to release budget reservation after database update failed") - try: - await _invalidate_budget_reservation_counters(budget_reservation=budget_reservation) - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after release failed" - ) - raise - - try: - await increment_spend_counters( - token=user_api_key, - team_id=team_id, - user_id=user_id, - response_cost=response_cost, - org_id=org_id, - budget_reservation=budget_reservation, - end_user_id=end_user_id, - tags=request_tags, - ) - except Exception: - if budget_reservation is not None: - try: - await _invalidate_budget_reservation_counters(budget_reservation=budget_reservation) - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after spend counter update failed" - ) - finally: - budget_reservation["finalized"] = True - raise - - -async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None: - if budget_reservation is None: - return - - from litellm.proxy.spend_tracking.budget_reservation import ( - release_budget_reservation, - ) - - await release_budget_reservation( - budget_reservation=budget_reservation, - ) - - -async def _invalidate_budget_reservation_counters( - budget_reservation: Optional[dict], -) -> None: - if budget_reservation is None: - return - - from litellm.proxy.spend_tracking.budget_reservation import ( - invalidate_budget_reservation_counters, - ) - - await invalidate_budget_reservation_counters( - budget_reservation=budget_reservation, - ) +import asyncio +import traceback +from datetime import datetime +from typing import Any, cast + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + _get_parent_otel_span_from_kwargs, + get_litellm_metadata_from_kwargs, +) +from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import ( + get_key_object, + get_team_object, + log_db_metrics, +) +from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_log_error_logger import ( + should_suppress_spend_log_tracebacks, + spend_log_error, +) +from litellm.proxy.spend_tracking.spend_tracking_utils import ( + _sanitize_error_information_for_spend_logs, +) +from litellm.proxy.utils import ProxyUpdateSpend +from litellm.types.utils import ( + CallTypes, + StandardLoggingPayload, + StandardLoggingPayloadErrorInformation, +) +from litellm.utils import get_end_user_id_for_cost_tracking + +_PASS_THROUGH_CALL_TYPES: frozenset[str] = frozenset( + { + CallTypes.pass_through.value, + CallTypes.llm_passthrough_route.value, + CallTypes.allm_passthrough_route.value, + } +) + + +class _ProxyDBLogger(CustomLogger): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) + + async def async_post_call_failure_hook( + self, + request_data: dict, + original_exception: Exception, + user_api_key_dict: UserAPIKeyAuth, + traceback_str: str | None = None, + ): + try: + await _release_budget_reservation(budget_reservation=user_api_key_dict.budget_reservation) + except Exception: + verbose_proxy_logger.exception("Failed to release budget reservation during failure handling") + try: + await _invalidate_budget_reservation_counters(budget_reservation=user_api_key_dict.budget_reservation) + if user_api_key_dict.budget_reservation is not None: + user_api_key_dict.budget_reservation["finalized"] = True + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after failure release failed" + ) + + request_route = user_api_key_dict.request_route + if ( + _ProxyDBLogger._should_track_errors_in_db() is False + or request_route is not None + and not ( + RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) + ) + ): + return + + from litellm.proxy.proxy_server import proxy_logging_obj + + _metadata = dict( + LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) + ) + _metadata["user_api_key"] = user_api_key_dict.api_key + _metadata["status"] = "failure" + _error_information = StandardLoggingPayloadSetup.get_error_information( + original_exception=original_exception, + traceback_str=traceback_str, + ) + if should_suppress_spend_log_tracebacks(): + # Drop the traceback key entirely so the per-row Metadata pane in + # the UI (which renders the JSON blob verbatim) doesn't show a + # noisy ``"traceback": ""`` line. Downstream consumers all use + # ``.get("traceback")`` / truthy checks, and the TypedDict marks + # the field as optional, so omitting is type-safe. + _error_information.pop("traceback", None) + # Strip echoed request input + apply DB-size cap before storing in + # the spend-log metadata column (LIT-2992). Result is never None + # here because the input above is constructed non-None. + _error_information = cast( + StandardLoggingPayloadErrorInformation, + _sanitize_error_information_for_spend_logs(_error_information), + ) + _metadata["error_information"] = _error_information + + _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata=_metadata, + ) + + existing_metadata: dict = request_data.get("metadata", None) or {} + existing_metadata.update(_metadata) + + if "litellm_params" not in request_data: + request_data["litellm_params"] = {} + + existing_litellm_params = request_data.get("litellm_params", {}) + existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {} + + # Preserve tags from existing metadata + if existing_litellm_metadata.get("tags"): + existing_metadata["tags"] = existing_litellm_metadata.get("tags") + + request_data["litellm_params"]["proxy_server_request"] = ( + request_data.get("proxy_server_request") or existing_litellm_params.get("proxy_server_request") or {} + ) + request_data["litellm_params"]["metadata"] = existing_metadata + + # Preserve model name and custom_llm_provider + if "model" not in request_data: + request_data["model"] = existing_litellm_params.get("model") or request_data.get("model", "") + if "custom_llm_provider" not in request_data: + request_data["custom_llm_provider"] = existing_litellm_params.get( + "custom_llm_provider" + ) or request_data.get("custom_llm_provider", "") + + # Propagate standard_logging_object and litellm_trace_id from the + # Logging instance so that _get_session_id_for_spend_log uses the same + # trace_id that Langfuse received (via async_failure_handler). + # Without this, the DB session_id would be a random UUID that doesn't + # match the Langfuse trace_id, making failed requests unsearchable. + _litellm_logging_obj = request_data.get("litellm_logging_obj") + if _litellm_logging_obj is not None: + if not request_data.get("standard_logging_object"): + request_data["standard_logging_object"] = getattr(_litellm_logging_obj, "model_call_details", {}).get( + "standard_logging_object" + ) + if request_data.get("litellm_trace_id") is None: + request_data["litellm_trace_id"] = getattr(_litellm_logging_obj, "litellm_trace_id", None) + + # Use the actual request start time from the logging object so that + # failed requests record the real duration instead of 0. + actual_start_time = datetime.now() + if _litellm_logging_obj is not None: + obj_start = getattr(_litellm_logging_obj, "start_time", None) + if obj_start is not None: + actual_start_time = obj_start + + # A stream that broke mid-flight still billed the provider for the + # chunks already delivered. ``post_call_failure_hook`` lifts that + # recovered cost onto request_data (the usage rides along in + # ``combined_usage_object`` for the token columns), so attribute the + # real partial spend to this failure row instead of zero. + recovered_response_cost = 0.0 + if isinstance(request_data.get("combined_usage_object"), litellm.Usage): + recovered_response_cost = max(float(request_data.get("response_cost") or 0.0), 0.0) + + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key_dict.api_key, + response_cost=recovered_response_cost, + user_id=user_api_key_dict.user_id, + end_user_id=user_api_key_dict.end_user_id, + team_id=user_api_key_dict.team_id, + kwargs=request_data, + completion_response=original_exception, + start_time=actual_start_time, + end_time=datetime.now(), + org_id=user_api_key_dict.org_id, + ) + + @log_db_metrics + async def _PROXY_track_cost_callback( + self, + kwargs, # kwargs to completion + completion_response: litellm.ModelResponse | Any | None, # response from completion + start_time=None, + end_time=None, # start/end time for completion + ): + from litellm.proxy.proxy_server import ( + increment_spend_counters, + proxy_logging_obj, + update_cache, + ) + + verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") + try: + verbose_proxy_logger.debug( + f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" + ) + parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_params = kwargs.get("litellm_params", {}) or {} + end_user_id = get_end_user_id_for_cost_tracking(litellm_params) + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). + # Avoids a cache/DB lookup on every normal LLM request. + if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): + metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) + budget_reservation = _get_budget_reservation_from_metadata(metadata=metadata) + user_id = cast(str | None, metadata.get("user_api_key_user_id", None)) + team_id = cast(str | None, metadata.get("user_api_key_team_id", None)) + org_id = cast(str | None, metadata.get("user_api_key_org_id", None)) + key_alias = cast(str | None, metadata.get("user_api_key_alias", None)) + end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) + sl_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + response_cost = ( + sl_object.get("response_cost", None) if sl_object is not None else kwargs.get("response_cost", None) + ) + tags = _get_request_tags_for_cost_tracking( + sl_object=sl_object, + metadata=metadata, + ) + + if response_cost is not None: + user_api_key = metadata.get("user_api_key", None) + if kwargs.get("cache_hit", False) is True: + response_cost = 0.0 + verbose_proxy_logger.debug(f"Cache Hit: response_cost {response_cost}, for user_id {user_id}") + + verbose_proxy_logger.debug( + f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" + ) + call_type: str | None = kwargs.get("call_type") + if _should_track_cost_callback( + user_api_key=user_api_key, + user_id=user_id, + team_id=team_id, + end_user_id=end_user_id, + call_type=call_type, + ): + ## UPDATE DATABASE + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + request_tags=tags, + ) + + # update cache (fire-and-forget for backward compat: + # cached object fields, soft budget alerts, etc.) + asyncio.create_task( + update_cache( + token=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + response_cost=response_cost, + team_id=team_id, + parent_otel_span=parent_otel_span, + tags=tags, + ) + ) + + await proxy_logging_obj.slack_alerting_instance.customer_spend_alert( + token=user_api_key, + key_alias=key_alias, + end_user_id=end_user_id, + response_cost=response_cost, + max_budget=end_user_max_budget, + ) + elif budget_reservation is not None: + await _release_budget_reservation(budget_reservation=budget_reservation) + else: + await _release_budget_reservation(budget_reservation=budget_reservation) + # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. + # Use .get() for "stream" to avoid KeyError on health checks. + # WS session wrappers (_aresponses_websocket, _arealtime) also reach here with + # result=None; their per-turn costs are tracked on the inner aresponses/realtime calls. + if sl_object is None and ( + not kwargs.get("model") or kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime") + ): + verbose_proxy_logger.warning( + "Cost tracking - skipping, no standard_logging_object for call_type=%s", + kwargs.get("call_type", "unknown"), + ) + return + if kwargs.get("stream") is not True or ( + kwargs.get("stream") is True and "complete_streaming_response" in kwargs + ): + if sl_object is not None: + cost_tracking_failure_debug_info: dict | str = ( + sl_object["response_cost_failure_debug_info"] # type: ignore + or "response_cost_failure_debug_info is None in standard_logging_object" + ) + else: + cost_tracking_failure_debug_info = "standard_logging_object not found" + model = kwargs.get("model") + raise Exception( + f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing" + ) + except Exception as e: + error_msg = f"Error in tracking cost callback - {e!s}\n Traceback:{traceback.format_exc()}" + model = kwargs.get("model", "") + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + litellm_metadata = kwargs.get("litellm_params", {}).get("litellm_metadata", {}) + old_metadata = kwargs.get("litellm_params", {}).get("metadata", {}) + call_type = kwargs.get("call_type", "") + error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n" + asyncio.create_task( + proxy_logging_obj.failed_tracking_alert( + error_message=error_msg, + failing_model=model, + ) + ) + + spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) + + @staticmethod + async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: + """ + Enriches failure spend log metadata by looking up the key object (and team object) + from cache/DB when key fields are missing. + + This handles two scenarios: + 1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other + fields are null. We look up the full key object to fill in alias, user_id, + team_id, etc. + 2. Post-auth failures (provider errors, rate limits): key fields are populated + but team_alias is missing because LiteLLM_VerificationTokenView SQL view + doesn't include it. We look up the team object to fill in team_alias. + """ + api_key_hash = metadata.get("user_api_key") + if not api_key_hash: + return metadata + + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + # Step 1: If key fields are missing, look up the full key object + if metadata.get("user_api_key_alias") is None: + try: + key_obj = await get_key_object( + hashed_token=api_key_hash, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if metadata.get("user_api_key_alias") is None: + metadata["user_api_key_alias"] = key_obj.key_alias + if metadata.get("user_api_key_user_id") is None: + metadata["user_api_key_user_id"] = key_obj.user_id + if metadata.get("user_api_key_team_id") is None: + metadata["user_api_key_team_id"] = key_obj.team_id + if metadata.get("user_api_key_org_id") is None: + metadata["user_api_key_org_id"] = key_obj.org_id + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with key info for api_key=%s", + api_key_hash, + ) + + # Step 2: If team_id is known but team_alias is missing, look up the team object + team_id = metadata.get("user_api_key_team_id") + if team_id and metadata.get("user_api_key_team_alias") is None: + try: + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if team_obj.team_alias is not None: + metadata["user_api_key_team_alias"] = team_obj.team_alias + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with team_alias for team_id=%s", + team_id, + ) + return metadata + + @staticmethod + def _should_track_errors_in_db(): + """ + Returns True if errors should be tracked in the database + + By default, errors are tracked in the database + + If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings + """ + from litellm.proxy.proxy_server import general_settings + + if general_settings.get("disable_error_logs") is True: + return False + return + + +def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: + patch = {k: v for k, v in metadata.items() if (k.startswith("user_api_key") or k == "tags") and v is not None} + if not patch: + return + + litellm_params = kwargs.setdefault("litellm_params", {}) + for bucket_name in ("litellm_metadata", "metadata"): + bucket = litellm_params.get(bucket_name) + if isinstance(bucket, dict): + for key, value in patch.items(): + if bucket.get(key) is None: + bucket[key] = value + + +def _should_track_cost_callback( + user_api_key: str | None, + user_id: str | None, + team_id: str | None, + end_user_id: str | None, + call_type: str | None = None, +) -> bool: + """ + Determine if the cost callback should be tracked based on the kwargs + + Pass-through endpoints can be configured with ``auth=false``, which leaves + the request with no key/user/team/end-user to attribute spend to. Those + requests still forward real provider traffic that operators expect to see + in request/usage logs, so they are tracked even when unauthenticated. + """ + + # don't run track cost callback if user opted into disabling spend + if ProxyUpdateSpend.disable_spend_updates() is True: + return False + + if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None: + return True + return call_type in _PASS_THROUGH_CALL_TYPES + + +def _get_budget_reservation_from_metadata(metadata: dict) -> dict | None: + metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") + if isinstance(metadata_budget_reservation, dict): + return metadata_budget_reservation + + user_api_key_auth_obj = metadata.get("user_api_key_auth") + if user_api_key_auth_obj is None: + return None + if isinstance(user_api_key_auth_obj, dict): + budget_reservation = user_api_key_auth_obj.get("budget_reservation") + return budget_reservation if isinstance(budget_reservation, dict) else None + return getattr(user_api_key_auth_obj, "budget_reservation", None) + + +def _get_request_tags_for_cost_tracking( + sl_object: StandardLoggingPayload | None, + metadata: dict, +) -> list[str] | None: + if sl_object is not None: + request_tags = sl_object.get("request_tags", None) + if isinstance(request_tags, list): + return request_tags + + metadata_tags = metadata.get("tags", None) + if isinstance(metadata_tags, list): + return metadata_tags + + return None + + +async def _update_database_and_spend_counters( + proxy_logging_obj: Any, + increment_spend_counters: Any, + user_api_key: str | None, + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + org_id: str | None, + kwargs: dict, + completion_response: litellm.ModelResponse | Any | None, + start_time: Any, + end_time: Any, + response_cost: float, + budget_reservation: dict | None, + request_tags: list[str] | None = None, +) -> None: + try: + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key, + response_cost=response_cost, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + org_id=org_id, + ) + except Exception: + if budget_reservation is not None: + try: + await _release_budget_reservation(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception("Failed to release budget reservation after database update failed") + try: + await _invalidate_budget_reservation_counters(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after release failed" + ) + raise + + try: + await increment_spend_counters( + token=user_api_key, + team_id=team_id, + user_id=user_id, + response_cost=response_cost, + org_id=org_id, + budget_reservation=budget_reservation, + end_user_id=end_user_id, + tags=request_tags, + ) + except Exception: + if budget_reservation is not None: + try: + await _invalidate_budget_reservation_counters(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) + finally: + budget_reservation["finalized"] = True + raise + + +async def _release_budget_reservation(budget_reservation: dict | None) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + release_budget_reservation, + ) + + await release_budget_reservation( + budget_reservation=budget_reservation, + ) + + +async def _invalidate_budget_reservation_counters( + budget_reservation: dict | None, +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + ) + + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/hooks/rate_limiter_utils.py b/litellm/proxy/hooks/rate_limiter_utils.py index 07440975476..e0b62ddfe74 100644 --- a/litellm/proxy/hooks/rate_limiter_utils.py +++ b/litellm/proxy/hooks/rate_limiter_utils.py @@ -2,8 +2,6 @@ Shared utility functions for rate limiter hooks. """ -from typing import Optional, Tuple, Union - import litellm from litellm._logging import verbose_proxy_logger from litellm.types.router import ModelGroupInfo @@ -13,8 +11,8 @@ PROXY_LLM_PROVIDER_FALLBACK = "litellm_proxy" def resolve_llm_provider_for_rate_limit( - model: Optional[str], -) -> Tuple[str, str]: + model: str | None, +) -> tuple[str, str]: """ Resolve ``(model, llm_provider)`` for a request being rejected by an internal proxy-side rate-limit hook. @@ -68,7 +66,7 @@ def resolve_llm_provider_for_rate_limit( def _resolve_provider_from_router_alias( model: str, -) -> Optional[Tuple[str, str]]: +) -> tuple[str, str] | None: """ Resolve a router ``model_name`` alias to ``(underlying_model, provider)`` by scanning the active router's ``model_list``. @@ -120,9 +118,7 @@ def _resolve_provider_from_router_alias( return None -def convert_priority_to_percent( - value: Union[float, PriorityReservationDict], model_info: Optional[ModelGroupInfo] -) -> float: +def convert_priority_to_percent(value: float | PriorityReservationDict, model_info: ModelGroupInfo | None) -> float: """ Convert priority reservation value to percentage (0.0-1.0). diff --git a/litellm/proxy/hooks/responses_id_security.py b/litellm/proxy/hooks/responses_id_security.py index b78688aa593..8b6bfd95892 100644 --- a/litellm/proxy/hooks/responses_id_security.py +++ b/litellm/proxy/hooks/responses_id_security.py @@ -6,7 +6,7 @@ instead of writing immediately on each request. """ from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, cast from fastapi import HTTPException @@ -38,7 +38,7 @@ class ResponsesIDSecurity(CustomLogger): cache: "DualCache", data: dict, call_type: CallTypesLiteral, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: # MAP all the responses api response ids to the encrypted response ids responses_api_call_types = { "aresponses", @@ -68,8 +68,8 @@ class ResponsesIDSecurity(CustomLogger): def check_user_access_to_response_id( self, - response_id_user_id: Optional[str], - response_id_team_id: Optional[str], + response_id_user_id: str | None, + response_id_team_id: str | None, user_api_key_dict: "UserAPIKeyAuth", ) -> bool: from litellm.proxy.proxy_server import general_settings @@ -119,7 +119,7 @@ class ResponsesIDSecurity(CustomLogger): return True return False - def _decrypt_response_id(self, response_id: str) -> Tuple[str, Optional[str], Optional[str]]: + def _decrypt_response_id(self, response_id: str) -> tuple[str, str | None, str | None]: """ Returns: - original_response_id: the original response id @@ -159,7 +159,7 @@ class ResponsesIDSecurity(CustomLogger): return response_id, None, None return response_id, None, None - def _get_signing_key(self) -> Optional[str]: + def _get_signing_key(self) -> str | None: """Get the signing key for encryption/decryption.""" import os @@ -174,7 +174,7 @@ class ResponsesIDSecurity(CustomLogger): self, response: BaseLiteLLMOpenAIResponseObject, user_api_key_dict: "UserAPIKeyAuth", - request_cache: Optional[dict[str, str]] = None, + request_cache: dict[str, str] | None = None, ) -> BaseLiteLLMOpenAIResponseObject: # encrypt the response id using the symmetric key # encrypt the response id, and encode the user id and response id in base64 diff --git a/litellm/proxy/hooks/sensitive_data_routing.py b/litellm/proxy/hooks/sensitive_data_routing.py index b4f44b5e41a..044f69c6768 100644 --- a/litellm/proxy/hooks/sensitive_data_routing.py +++ b/litellm/proxy/hooks/sensitive_data_routing.py @@ -11,7 +11,7 @@ Works across multiple proxy instances via DualCache (in-memory + Redis). """ import os -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache @@ -53,7 +53,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): return f"{{{SENSITIVE_ROUTING_CACHE_PREFIX}:{tenant}:{session_id}}}:model" @staticmethod - def _resolve_tenant(user_api_key_dict: Optional[UserAPIKeyAuth]) -> str: + def _resolve_tenant(user_api_key_dict: UserAPIKeyAuth | None) -> str: """ Identify the authenticated principal the routing override belongs to. @@ -77,7 +77,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): ] return "|".join(principal) if principal else "default" - async def _get_routed_model(self, session_id: str, user_api_key_dict: Optional[UserAPIKeyAuth]) -> Optional[str]: + async def _get_routed_model(self, session_id: str, user_api_key_dict: UserAPIKeyAuth | None) -> str | None: """Get the model this session should be routed to, if any.""" cache_key = self._make_cache_key(session_id, self._resolve_tenant(user_api_key_dict)) @@ -114,8 +114,8 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): self, session_id: str, model: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - guardrail_name: Optional[str] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, + guardrail_name: str | None = None, ) -> None: """ Store a routing override for a session. @@ -161,7 +161,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): cache: DualCache, data: dict, call_type: str, - ) -> Optional[Union[Exception, str, dict]]: + ) -> Exception | str | dict | None: """ Before each LLM call, check if this session has a routing override. If so, modify the request's model field. diff --git a/litellm/proxy/hooks/user_management_event_hooks.py b/litellm/proxy/hooks/user_management_event_hooks.py index e40db7f3f1f..444c39340a0 100644 --- a/litellm/proxy/hooks/user_management_event_hooks.py +++ b/litellm/proxy/hooks/user_management_event_hooks.py @@ -4,7 +4,6 @@ Hooks that are triggered when a litellm user event occurs import asyncio from datetime import datetime, timezone -from typing import Optional import litellm from litellm._logging import verbose_proxy_logger @@ -72,8 +71,7 @@ class UserManagementEventHooks: ) ) except Exception as e: - verbose_proxy_logger.warning("Unable to create audit log for user on `/user/new` - {}".format(str(e))) - pass + verbose_proxy_logger.warning(f"Unable to create audit log for user on `/user/new` - {e!s}") @staticmethod async def async_send_user_invitation_email( @@ -163,11 +161,11 @@ class UserManagementEventHooks: async def create_internal_user_audit_log( user_id: str, action: AUDIT_ACTIONS, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, user_api_key_dict: UserAPIKeyAuth, - litellm_proxy_admin_name: Optional[str], - before_value: Optional[str] = None, - after_value: Optional[str] = None, + litellm_proxy_admin_name: str | None, + before_value: str | None = None, + after_value: str | None = None, ): """ Create an audit log for an internal user. diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index 8178cad9038..7666ad0f065 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -1,6 +1,5 @@ import asyncio import traceback -from typing import List import orjson from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, status @@ -36,8 +35,8 @@ async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO: async def batch_to_bytesio( - uploads: Optional[List[UploadFile]], -) -> Optional[List[io.BytesIO]]: + uploads: list[UploadFile] | None, +) -> list[io.BytesIO] | None: """ Convert a list of UploadFiles to a list of BytesIO buffers, or None. """ @@ -68,7 +67,7 @@ async def image_generation( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - model: Optional[str] = None, + model: str | None = None, ): from litellm.proxy.proxy_server import ( add_litellm_data_to_request, @@ -186,9 +185,7 @@ async def image_generation( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.image_generation(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.image_generation(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -198,7 +195,7 @@ async def image_generation( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -228,11 +225,11 @@ async def image_edit_api( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - image: Optional[List[UploadFile]] = File(None), - image_array: Optional[List[UploadFile]] = File(None, alias="image[]"), - mask: Optional[List[UploadFile]] = File(None), - mask_array: Optional[List[UploadFile]] = File(None, alias="mask[]"), - model: Optional[str] = None, + image: list[UploadFile] | None = File(None), + image_array: list[UploadFile] | None = File(None, alias="image[]"), + mask: list[UploadFile] | None = File(None), + mask_array: list[UploadFile] | None = File(None, alias="mask[]"), + model: str | None = None, ): """ Follows the OpenAI Images API spec: https://platform.openai.com/docs/api-reference/images/create diff --git a/litellm/proxy/lambda.py b/litellm/proxy/lambda.py index 6b278c41188..f783f0f647e 100644 --- a/litellm/proxy/lambda.py +++ b/litellm/proxy/lambda.py @@ -1,4 +1,5 @@ from mangum import Mangum + from litellm.proxy.proxy_server import app handler = Mangum(app, lifespan="on") diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 1fad1954dc4..d58f953d2ef 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -4,7 +4,7 @@ import json import re import time from collections import OrderedDict -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any from fastapi import HTTPException, Request from pydantic import ValidationError as PydanticValidationError @@ -107,7 +107,7 @@ service_logger_obj = ServiceLogging() # used for tracking latency on OTEL _MAX_STALE_ALIAS_WARNING_KEYS = 10_000 _STALE_TEAM_ALIAS_WARNING_KEYS: OrderedDict[str, None] = OrderedDict() # Cache the stale alias bypass flag at module load to avoid hot-path secret lookups -_ENABLE_TEAM_STALE_ALIAS_BYPASS: Optional[bool] = None +_ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None if TYPE_CHECKING: @@ -242,7 +242,7 @@ _ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY = "allow_client_pricing_override" _URL_DESTINATION_REQUEST_FIELDS = ("model", "file_id") -def _reject_url_valued_destinations(data: Dict[str, Any]) -> None: +def _reject_url_valued_destinations(data: dict[str, Any]) -> None: """Reject URL-valued ``model``/``file_id`` unless admin-allowlisted. Some providers (HuggingFace, Oobabooga, Gemini files) accept a URL in the @@ -356,7 +356,7 @@ def _strip_client_message_redaction_opt_out(data: dict[str, Any]) -> None: ) -def _strip_client_pricing_overrides(data: Dict[str, Any]) -> None: +def _strip_client_pricing_overrides(data: dict[str, Any]) -> None: """Drop pricing overrides from the request body and any metadata variant. Skipped only when the calling key/team carries @@ -365,7 +365,7 @@ def _strip_client_pricing_overrides(data: Dict[str, Any]) -> None: trace why a client-supplied pricing override stopped being applied (otherwise the strip is invisible from the caller's perspective). """ - stripped: List[str] = [] + stripped: list[str] = [] for field in _CLIENT_PRICING_CONTROL_FIELDS: if field in data: stripped.append(field) @@ -409,8 +409,8 @@ def _get_metadata_variable_name(request: Request) -> str: def _extract_generic_session_id_from_headers( - normalized: Dict[str, str], -) -> Optional[str]: + normalized: dict[str, str], +) -> str | None: """ Scan a normalised (lower-cased keys) header dict for any header that looks like ``x--session-id`` and whose value is a plausible session/trace @@ -433,7 +433,7 @@ def _extract_generic_session_id_from_headers( return None -def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]: +def get_chain_id_from_headers(headers: dict[str, str] | None) -> str | None: """ Extract chain id for call chaining from request headers. @@ -524,7 +524,7 @@ def safe_add_api_version_from_query_params(data: dict, request: Request): def convert_key_logging_metadata_to_callback( - data: AddTeamCallback, team_callback_settings_obj: Optional[TeamCallbackMetadata] + data: AddTeamCallback, team_callback_settings_obj: TeamCallbackMetadata | None ) -> TeamCallbackMetadata: if team_callback_settings_obj is None: team_callback_settings_obj = TeamCallbackMetadata() @@ -567,7 +567,7 @@ def convert_key_logging_metadata_to_callback( return team_callback_settings_obj -def _get_validated_callback_metadata(item: dict, *, source: str) -> Optional[AddTeamCallback]: +def _get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallback | None: try: return AddTeamCallback(**item) except (PydanticValidationError, ValueError) as e: @@ -599,12 +599,12 @@ class KeyAndTeamLoggingSettings: def _get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig -) -> Optional[TeamCallbackMetadata]: - callback_settings_obj: Optional[TeamCallbackMetadata] = None - key_dynamic_logging_settings: Optional[dict] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings( +) -> TeamCallbackMetadata | None: + callback_settings_obj: TeamCallbackMetadata | None = None + key_dynamic_logging_settings: dict | None = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings( user_api_key_dict ) - team_dynamic_logging_settings: Optional[dict] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings( + team_dynamic_logging_settings: dict | None = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings( user_api_key_dict ) ######################################################################################### @@ -660,9 +660,9 @@ def _get_dynamic_logging_metadata( def clean_headers( headers: Headers, - litellm_key_header_name: Optional[str] = None, + litellm_key_header_name: str | None = None, forward_llm_provider_auth_headers: bool = False, - authenticated_with_header: Optional[str] = None, + authenticated_with_header: str | None = None, ) -> dict: """ Removes litellm api key from headers @@ -713,7 +713,7 @@ def clean_headers( class LiteLLMProxyRequestSetup: @staticmethod - def _get_timeout_from_request(headers: dict) -> Optional[float]: + def _get_timeout_from_request(headers: dict) -> float | None: """ Workaround for client request from Vercel's AI SDK. @@ -738,7 +738,7 @@ class LiteLLMProxyRequestSetup: return None @staticmethod - def _get_stream_timeout_from_request(headers: dict) -> Optional[float]: + def _get_stream_timeout_from_request(headers: dict) -> float | None: """ Get the `stream_timeout` from the request headers. """ @@ -748,7 +748,7 @@ class LiteLLMProxyRequestSetup: return None @staticmethod - def _get_num_retries_from_request(headers: dict) -> Optional[int]: + def _get_num_retries_from_request(headers: dict) -> int | None: """ Workaround for client request from Vercel's AI SDK. """ @@ -758,7 +758,7 @@ class LiteLLMProxyRequestSetup: return None @staticmethod - def _get_spend_logs_metadata_from_request_headers(headers: dict) -> Optional[dict]: + def _get_spend_logs_metadata_from_request_headers(headers: dict) -> dict | None: """ Get the `spend_logs_metadata` from the request headers. """ @@ -771,7 +771,7 @@ class LiteLLMProxyRequestSetup: @staticmethod def _get_forwardable_headers( - headers: Union[Headers, dict], + headers: Headers | dict, ): """ Get the headers that should be forwarded to the LLM Provider. @@ -782,17 +782,17 @@ class LiteLLMProxyRequestSetup: """ forwarded_headers = {} for header, value in headers.items(): - if header.lower().startswith("x-") and not header.lower().startswith( - "x-stainless" + if ( + header.lower().startswith("x-") + and not header.lower().startswith("x-stainless") + or header.lower().startswith("anthropic-beta") ): # causes openai sdk to fail forwarded_headers[header] = value - elif header.lower().startswith("anthropic-beta"): - forwarded_headers[header] = value return forwarded_headers @staticmethod - def _get_case_insensitive_header(headers: dict, key: str) -> Optional[str]: + def _get_case_insensitive_header(headers: dict, key: str) -> str | None: """ Get a case-insensitive header from the headers dictionary. """ @@ -803,7 +803,7 @@ class LiteLLMProxyRequestSetup: @staticmethod def add_internal_user_from_user_mapping( - general_settings: Optional[Dict], + general_settings: dict | None, user_api_key_dict: UserAPIKeyAuth, headers: dict, ) -> UserAPIKeyAuth: @@ -822,7 +822,7 @@ class LiteLLMProxyRequestSetup: return user_api_key_dict @staticmethod - def get_user_from_headers(headers: dict, general_settings: Optional[Dict] = None) -> Optional[str]: + def get_user_from_headers(headers: dict, general_settings: dict | None = None) -> str | None: """ Get the user from the specified header if `general_settings.user_header_name` is set. """ @@ -843,7 +843,7 @@ class LiteLLMProxyRequestSetup: return user @staticmethod - def get_openai_org_id_from_headers(headers: dict, general_settings: Optional[Dict] = None) -> Optional[str]: + def get_openai_org_id_from_headers(headers: dict, general_settings: dict | None = None) -> str | None: """ Get the OpenAI Org ID from the headers. """ @@ -877,11 +877,11 @@ class LiteLLMProxyRequestSetup: # to str and JSON-encode dict/list (e.g. user_api_key_spend is float, # user_api_key_auth_metadata is dict). See #27458. if isinstance(v, (dict, list)): - returned_headers["x-litellm-{}".format(k)] = json.dumps(v) + returned_headers[f"x-litellm-{k}"] = json.dumps(v) elif isinstance(v, (str, bytes)): - returned_headers["x-litellm-{}".format(k)] = v + returned_headers[f"x-litellm-{k}"] = v else: - returned_headers["x-litellm-{}".format(k)] = str(v) + returned_headers[f"x-litellm-{k}"] = str(v) return returned_headers @@ -913,7 +913,7 @@ class LiteLLMProxyRequestSetup: return data @staticmethod - def get_internal_user_header_from_mapping(user_header_mapping) -> Optional[str]: + def get_internal_user_header_from_mapping(user_header_mapping) -> str | None: if not user_header_mapping: return None items = user_header_mapping if isinstance(user_header_mapping, list) else [user_header_mapping] @@ -933,7 +933,7 @@ class LiteLLMProxyRequestSetup: *, headers: dict, user_api_key_dict: UserAPIKeyAuth, - general_settings: Optional[Dict[str, Any]] = None, + general_settings: dict[str, Any] | None = None, ) -> LitellmDataForBackendLLMCall: """ - Adds user from headers @@ -1103,7 +1103,7 @@ class LiteLLMProxyRequestSetup: return data @staticmethod - def add_key_level_controls(key_metadata: Optional[dict], data: dict, _metadata_variable_name: str): + def add_key_level_controls(key_metadata: dict | None, data: dict, _metadata_variable_name: str): if key_metadata is None: return data if "cache" in key_metadata: @@ -1146,7 +1146,7 @@ class LiteLLMProxyRequestSetup: return data @staticmethod - def _merge_tags(request_tags: Optional[list], tags_to_add: Optional[list]) -> list: + def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: """ Helper function to merge two lists of tags, ensuring no duplicates. @@ -1173,7 +1173,7 @@ class LiteLLMProxyRequestSetup: def add_team_based_callbacks_from_config( team_id: str, proxy_config: ProxyConfig, - ) -> Optional[TeamCallbackMetadata]: + ) -> TeamCallbackMetadata | None: """ Add team-based callbacks from the config """ @@ -1198,10 +1198,10 @@ class LiteLLMProxyRequestSetup: @staticmethod def add_request_tag_to_metadata( - llm_router: Optional[Router], + llm_router: Router | None, headers: dict, data: dict, - ) -> Optional[List[str]]: + ) -> list[str] | None: tags = None # Check request headers for tags @@ -1299,7 +1299,7 @@ class LiteLLMProxyRequestSetup: return if isinstance(raw_header_tags, str): - header_tags: List[str] = [t.strip() for t in raw_header_tags.split(",") if t.strip()] + header_tags: list[str] = [t.strip() for t in raw_header_tags.split(",") if t.strip()] elif isinstance(raw_header_tags, list): header_tags = [t for t in raw_header_tags if isinstance(t, str) and t] else: @@ -1337,8 +1337,8 @@ async def add_litellm_data_to_request( request: Request, user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig, - general_settings: Optional[Dict[str, Any]] = None, - version: Optional[str] = None, + general_settings: dict[str, Any] | None = None, + version: str | None = None, ): """ Adds LiteLLM-specific data to the request. @@ -1375,7 +1375,7 @@ async def add_litellm_data_to_request( if _mk.startswith("user_api_key_"): del _user_metadata[_mk] - _raw_headers: Dict[str, str] = RedactedDict(_safe_get_request_headers(request)) + _raw_headers: dict[str, str] = RedactedDict(_safe_get_request_headers(request)) forward_llm_auth = False if general_settings: @@ -1395,7 +1395,7 @@ async def add_litellm_data_to_request( # x-api-key or another header was used for auth authenticated_with_header = "x-api-key" - _headers: Dict[str, str] = clean_headers( + _headers: dict[str, str] = clean_headers( request.headers, litellm_key_header_name=( general_settings.get("litellm_key_header_name") if general_settings is not None else None @@ -1489,7 +1489,7 @@ async def add_litellm_data_to_request( query_dict = {} ## check for api version in query params - dynamic_api_version: Optional[str] = query_dict.get("api-version") + dynamic_api_version: str | None = query_dict.get("api-version") if dynamic_api_version is not None: # only pass, if set data["api_version"] = dynamic_api_version @@ -1947,14 +1947,13 @@ def _update_model_if_key_alias_exists( and _model in user_api_key_dict.aliases ): data["model"] = user_api_key_dict.aliases[_model] - return def _apply_credential_overrides_from_model_config( data: dict, user_api_key_dict: UserAPIKeyAuth, - pre_alias_model_name: Optional[str] = None, - llm_router: Optional[Router] = None, + pre_alias_model_name: str | None = None, + llm_router: Router | None = None, ) -> None: """ Walk the model_config precedence chain in team/project metadata. @@ -1994,7 +1993,7 @@ def _apply_credential_overrides_from_model_config( # When the user-facing name has no provider prefix, fall back to the # deployment's litellm_params so multi-provider defaultconfig entries # don't silently match the first dict key (#27516). - provider: Optional[str] = None + provider: str | None = None if "/" in model_name: provider = model_name.split("/", 1)[0] elif llm_router is not None: @@ -2041,8 +2040,8 @@ def _apply_credential_overrides_from_model_config( def _resolve_provider_from_deployment( llm_router: Router, model_name: str, - pre_alias_model_name: Optional[str] = None, -) -> Optional[str]: + pre_alias_model_name: str | None = None, +) -> str | None: """ Resolve a provider hint from the deployment's litellm_params when the user-facing model name has no provider prefix. @@ -2080,11 +2079,11 @@ def _resolve_provider_from_deployment( def _resolve_credential_from_model_config( model_name: str, - project_model_config: Optional[dict], - team_model_config: Optional[dict], - pre_alias_model_name: Optional[str] = None, - provider: Optional[str] = None, -) -> Optional[str]: + project_model_config: dict | None, + team_model_config: dict | None, + pre_alias_model_name: str | None = None, + provider: str | None = None, +) -> str | None: """ Walk the precedence chain and return the first matching credential name. @@ -2133,7 +2132,7 @@ def _resolve_credential_from_model_config( return None -def _extract_credential_from_entry(entry: dict, provider: Optional[str] = None) -> Optional[str]: +def _extract_credential_from_entry(entry: dict, provider: str | None = None) -> str | None: """ Extract litellm_credentials from a model_config entry. @@ -2162,8 +2161,8 @@ def _extract_credential_from_entry(entry: dict, provider: Optional[str] = None) return None -def _get_enforced_params(general_settings: Optional[dict], user_api_key_dict: UserAPIKeyAuth) -> Optional[list]: - enforced_params: Optional[list] = None +def _get_enforced_params(general_settings: dict | None, user_api_key_dict: UserAPIKeyAuth) -> list | None: + enforced_params: list | None = None if general_settings is not None: enforced_params = general_settings.get("enforced_params") if ( @@ -2198,14 +2197,14 @@ def check_if_token_is_service_account(valid_token: UserAPIKeyAuth) -> bool: def _enforced_params_check( request_body: dict, - general_settings: Optional[dict], + general_settings: dict | None, user_api_key_dict: UserAPIKeyAuth, premium_user: bool, ) -> bool: """ If enforced params are set, check if the request body contains the enforced params. """ - enforced_params: Optional[list] = _get_enforced_params( + enforced_params: list | None = _get_enforced_params( general_settings=general_settings, user_api_key_dict=user_api_key_dict ) if enforced_params is None: @@ -2236,11 +2235,11 @@ def _enforced_params_check( def _add_guardrails_from_key_or_team_metadata( - key_metadata: Optional[dict], - team_metadata: Optional[dict], + key_metadata: dict | None, + team_metadata: dict | None, data: dict, metadata_variable_name: str, - project_metadata: Optional[dict] = None, + project_metadata: dict | None = None, ) -> None: """ Helper add guardrails from key, team, or project metadata to request data @@ -2284,11 +2283,11 @@ def _add_guardrails_from_key_or_team_metadata( def _add_guardrails_from_policies_in_metadata( - key_metadata: Optional[dict], - team_metadata: Optional[dict], + key_metadata: dict | None, + team_metadata: dict | None, data: dict, metadata_variable_name: str, - project_metadata: Optional[dict] = None, + project_metadata: dict | None = None, ) -> None: """ Helper to resolve guardrails from policies attached to key/team/project metadata. @@ -2483,7 +2482,7 @@ def _is_policy_version_id(s: str) -> bool: return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) -def _extract_policy_id(s: str) -> Optional[str]: +def _extract_policy_id(s: str) -> str | None: """Extract raw UUID from policy_ string, or None if not a valid version ID.""" from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX @@ -2496,7 +2495,7 @@ def _match_and_track_policies( data: dict, context: "PolicyMatchContext", request_body_policies: Any, - policies_override: Optional[Dict[str, Any]] = None, + policies_override: dict[str, Any] | None = None, ) -> tuple[list[str], dict[str, str]]: """ Match policies via attachments and request body, track them in metadata. @@ -2553,8 +2552,8 @@ def _apply_resolved_guardrails_to_metadata( data: dict, metadata_variable_name: str, context: "PolicyMatchContext", - policy_names: Optional[List[str]] = None, - policies: Optional[Dict[str, Any]] = None, + policy_names: list[str] | None = None, + policies: dict[str, Any] | None = None, ) -> None: """Apply resolved guardrails and pipelines to request metadata.""" from litellm._logging import verbose_proxy_logger @@ -2664,8 +2663,8 @@ async def add_guardrails_from_policy_engine( ) # Separate policy names from policy version IDs (policy_) - request_body_names: List[str] = [] - request_body_version_ids: List[str] = [] + request_body_names: list[str] = [] + request_body_version_ids: list[str] = [] if request_body_policies_raw and isinstance(request_body_policies_raw, list): for item in request_body_policies_raw: if not isinstance(item, str): @@ -2678,8 +2677,8 @@ async def add_guardrails_from_policy_engine( request_body_names.append(item) # Resolve policy versions by ID from in-memory cache (populated by sync job; no DB in hot path) - merged_policies: Dict[str, Any] = dict(registry.get_all_policies()) - fetched_policy_names: List[str] = [] + merged_policies: dict[str, Any] = dict(registry.get_all_policies()) + fetched_policy_names: list[str] = [] for policy_id in request_body_version_ids: result = registry.get_policy_by_id_for_request(policy_id=policy_id) if result is not None: @@ -2743,8 +2742,6 @@ def add_provider_specific_headers_to_request( extra_headers=anthropic_headers, ) - return - def _add_otel_traceparent_to_data(data: dict, request: Request): from litellm.proxy.proxy_server import open_telemetry_logger diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 6ae97bfce85..13e45b17090 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -1,5 +1,3 @@ -from typing import List, Set - from fastapi import APIRouter, Depends, HTTPException, status from litellm._logging import verbose_proxy_logger @@ -111,7 +109,7 @@ async def _invalidate_cache_access_group(access_group_id: str) -> None: # --------------------------------------------------------------------------- -async def _sync_add_access_group_to_teams(tx, team_ids: List[str], access_group_id: str) -> None: +async def _sync_add_access_group_to_teams(tx, team_ids: list[str], access_group_id: str) -> None: """Add access_group_id to each team's access_group_ids (idempotent).""" for team_id in team_ids: team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id}) @@ -122,7 +120,7 @@ async def _sync_add_access_group_to_teams(tx, team_ids: List[str], access_group_ ) -async def _sync_remove_access_group_from_teams(tx, team_ids: List[str], access_group_id: str) -> None: +async def _sync_remove_access_group_from_teams(tx, team_ids: list[str], access_group_id: str) -> None: """Remove access_group_id from each team's access_group_ids (idempotent).""" for team_id in team_ids: team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id}) @@ -133,7 +131,7 @@ async def _sync_remove_access_group_from_teams(tx, team_ids: List[str], access_g ) -async def _sync_add_access_group_to_keys(tx, key_tokens: List[str], access_group_id: str) -> None: +async def _sync_add_access_group_to_keys(tx, key_tokens: list[str], access_group_id: str) -> None: """Add access_group_id to each key's access_group_ids (idempotent).""" for token in key_tokens: key = await tx.litellm_verificationtoken.find_unique(where={"token": token}) @@ -144,7 +142,7 @@ async def _sync_add_access_group_to_keys(tx, key_tokens: List[str], access_group ) -async def _sync_remove_access_group_from_keys(tx, key_tokens: List[str], access_group_id: str) -> None: +async def _sync_remove_access_group_from_keys(tx, key_tokens: list[str], access_group_id: str) -> None: """Remove access_group_id from each key's access_group_ids (idempotent).""" for token in key_tokens: key = await tx.litellm_verificationtoken.find_unique(where={"token": token}) @@ -161,7 +159,7 @@ async def _sync_remove_access_group_from_keys(tx, key_tokens: List[str], access_ async def _patch_team_caches_add_access_group( - team_ids: List[str], + team_ids: list[str], access_group_id: str, user_api_key_cache, proxy_logging_obj, @@ -169,7 +167,7 @@ async def _patch_team_caches_add_access_group( """Patch cached team objects to include access_group_id.""" for team_id in team_ids: cached_team = await _get_team_object_from_cache( - key="team_id:{}".format(team_id), + key=f"team_id:{team_id}", proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=None, @@ -191,7 +189,7 @@ async def _patch_team_caches_add_access_group( async def _patch_team_caches_remove_access_group( - team_ids: List[str], + team_ids: list[str], access_group_id: str, user_api_key_cache, proxy_logging_obj, @@ -199,7 +197,7 @@ async def _patch_team_caches_remove_access_group( """Patch cached team objects to remove access_group_id.""" for team_id in team_ids: cached_team = await _get_team_object_from_cache( - key="team_id:{}".format(team_id), + key=f"team_id:{team_id}", proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=None, @@ -215,7 +213,7 @@ async def _patch_team_caches_remove_access_group( async def _patch_key_caches_add_access_group( - key_tokens: List[str], + key_tokens: list[str], access_group_id: str, user_api_key_cache, proxy_logging_obj, @@ -243,7 +241,7 @@ async def _patch_key_caches_add_access_group( async def _patch_key_caches_remove_access_group( - key_tokens: List[str], + key_tokens: list[str], access_group_id: str, user_api_key_cache, proxy_logging_obj, @@ -341,11 +339,11 @@ async def create_access_group( @router.get( "/v1/access_group", - response_model=List[AccessGroupResponse], + response_model=list[AccessGroupResponse], ) async def list_access_groups( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> List[AccessGroupResponse]: +) -> list[AccessGroupResponse]: _require_admin_view(user_api_key_dict) prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value) @@ -404,10 +402,10 @@ async def update_access_group( # Initialize delta lists before the try block so they remain accessible # for cache updates after the transaction, even if an error path is added later. - teams_to_add: List[str] = [] - teams_to_remove: List[str] = [] - keys_to_add: List[str] = [] - keys_to_remove: List[str] = [] + teams_to_add: list[str] = [] + teams_to_remove: list[str] = [] + keys_to_add: list[str] = [] + keys_to_remove: list[str] = [] try: async with prisma_client.db.tx() as tx: @@ -420,12 +418,12 @@ async def update_access_group( detail=f"Access group '{access_group_id}' not found", ) - old_team_ids: Set[str] = set(existing.assigned_team_ids or []) - old_key_ids: Set[str] = set(existing.assigned_key_ids or []) - new_team_ids: Set[str] = ( + old_team_ids: set[str] = set(existing.assigned_team_ids or []) + old_key_ids: set[str] = set(existing.assigned_key_ids or []) + new_team_ids: set[str] = ( set(update_fields["assigned_team_ids"] or []) if "assigned_team_ids" in update_fields else old_team_ids ) - new_key_ids: Set[str] = ( + new_key_ids: set[str] = ( set(update_fields["assigned_key_ids"] or []) if "assigned_key_ids" in update_fields else old_key_ids ) @@ -479,8 +477,8 @@ async def delete_access_group( prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value) try: - affected_team_ids: List[str] = [] - affected_key_tokens: List[str] = [] + affected_team_ids: list[str] = [] + affected_key_tokens: list[str] = [] async with prisma_client.db.tx() as tx: existing = await tx.litellm_accessgrouptable.find_unique(where={"access_group_id": access_group_id}) @@ -495,7 +493,7 @@ async def delete_access_group( teams_with_group = await tx.litellm_teamtable.find_many( where={"access_group_ids": {"hasSome": [access_group_id]}} ) - all_affected_team_ids: Set[str] = {team.team_id for team in teams_with_group} | set( + all_affected_team_ids: set[str] = {team.team_id for team in teams_with_group} | set( existing.assigned_team_ids or [] ) affected_team_ids = list(all_affected_team_ids) @@ -505,7 +503,7 @@ async def delete_access_group( keys_with_group = await tx.litellm_verificationtoken.find_many( where={"access_group_ids": {"hasSome": [access_group_id]}} ) - all_affected_key_tokens: Set[str] = {key.token for key in keys_with_group} | set( + all_affected_key_tokens: set[str] = {key.token for key in keys_with_group} | set( existing.assigned_key_ids or [] ) affected_key_tokens = list(all_affected_key_tokens) @@ -578,7 +576,7 @@ router.add_api_route( "/v1/unified_access_group", list_access_groups, methods=["GET"], - response_model=List[AccessGroupResponse], + response_model=list[AccessGroupResponse], ) router.add_api_route( "/v1/unified_access_group/{access_group_id}", diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 6b70a9064df..945aea7e8ac 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -239,12 +239,7 @@ async def budget_settings( if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) ## get budget item from db @@ -304,12 +299,7 @@ async def list_budget( if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) response = await BudgetRepository(prisma_client).table.find_many() @@ -343,12 +333,7 @@ async def delete_budget( if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) response = await BudgetRepository(prisma_client).table.delete(where={"budget_id": data.id}) diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 66a79baa8d7..e08bc13a14d 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -12,7 +12,7 @@ import asyncio import json from collections.abc import Mapping from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field @@ -239,7 +239,7 @@ def _merge_over_saved(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> return merged -def _redact_settings(settings: Optional[Mapping[str, Any]]) -> Dict[str, Any]: +def _redact_settings(settings: Mapping[str, Any] | None) -> dict[str, Any]: """Replace every value in a settings map with a fixed marker. Cache config carries Redis credentials (passwords, connection strings). @@ -268,10 +268,10 @@ def _log_audit_task_exception(task: "asyncio.Task[None]") -> None: async def _emit_cache_settings_audit_log( *, action: AUDIT_ACTIONS, - before_settings: Optional[Mapping[str, Any]], - after_settings: Optional[Mapping[str, Any]], + before_settings: Mapping[str, Any] | None, + after_settings: Mapping[str, Any] | None, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, ) -> None: """Emit an audit-log row for a /cache/settings mutation. @@ -313,17 +313,17 @@ class CacheSettingsManager: Tracks last cache params to avoid unnecessary reinitialization. """ - _last_cache_params: Optional[Dict[str, Any]] = None + _last_cache_params: dict[str, Any] | None = None @staticmethod - def _cache_params_equal(params1: Dict[str, Any], params2: Dict[str, Any]) -> bool: + def _cache_params_equal(params1: dict[str, Any], params2: dict[str, Any]) -> bool: """ Compare two cache parameter dictionaries for equality. Normalizes values and filters out UI-only fields. """ # Normalize by removing None values and UI-only fields - def normalize(params: Dict[str, Any]) -> Dict[str, Any]: + def normalize(params: dict[str, Any]) -> dict[str, Any]: normalized = {} for k, v in params.items(): if k == "redis_type": # Skip UI-only field @@ -386,13 +386,11 @@ class CacheSettingsManager: verbose_proxy_logger.info("Cache settings initialized from database") except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.management_endpoints.cache_settings_endpoints.py::CacheSettingsManager::init_cache_settings_in_db - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.cache_settings_endpoints.py::CacheSettingsManager::init_cache_settings_in_db - {e!s}" ) @staticmethod - def update_cache_params(cache_params: Dict[str, Any]): + def update_cache_params(cache_params: dict[str, Any]): """ Update the last cache params after initialization. Called after cache settings are updated via the API. @@ -401,23 +399,23 @@ class CacheSettingsManager: class CacheSettingsResponse(BaseModel): - fields: List[CacheSettingsField] = Field(description="List of all configurable cache settings with metadata") - current_values: Dict[str, Any] = Field(description="Current values of cache settings") - redis_type_descriptions: Dict[str, str] = Field(description="Descriptions for each Redis type option") + fields: list[CacheSettingsField] = Field(description="List of all configurable cache settings with metadata") + current_values: dict[str, Any] = Field(description="Current values of cache settings") + redis_type_descriptions: dict[str, str] = Field(description="Descriptions for each Redis type option") class CacheTestRequest(BaseModel): - cache_settings: Dict[str, Any] = Field(description="Cache settings to test connection with") + cache_settings: dict[str, Any] = Field(description="Cache settings to test connection with") class CacheTestResponse(BaseModel): status: str = Field(description="Connection status: 'success' or 'failed'") message: str = Field(description="Connection result message") - error: Optional[str] = Field(default=None, description="Error message if connection failed") + error: str | None = Field(default=None, description="Error message if connection failed") class CacheSettingsUpdateRequest(BaseModel): - cache_settings: Dict[str, Any] = Field(description="Cache settings to save") + cache_settings: dict[str, Any] = Field(description="Cache settings to save") @router.get( @@ -482,8 +480,8 @@ async def get_cache_settings( redis_type_descriptions=REDIS_TYPE_DESCRIPTIONS, ) except Exception as e: - verbose_proxy_logger.error(f"Error fetching cache settings: {str(e)}") - raise HTTPException(status_code=500, detail=f"Error fetching cache settings: {str(e)}") + verbose_proxy_logger.error(f"Error fetching cache settings: {e!s}") + raise HTTPException(status_code=500, detail=f"Error fetching cache settings: {e!s}") @router.post( @@ -541,10 +539,10 @@ async def test_cache_connection( return CacheTestResponse(**result) except Exception as e: - verbose_proxy_logger.error(f"Error testing cache connection: {str(e)}") + verbose_proxy_logger.error(f"Error testing cache connection: {e!s}") return CacheTestResponse( status="failed", - message=f"Cache connection test failed: {str(e)}", + message=f"Cache connection test failed: {e!s}", error=str(e), ) @@ -557,7 +555,7 @@ async def test_cache_connection( async def update_cache_settings( request: CacheSettingsUpdateRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -592,7 +590,7 @@ async def update_cache_settings( # Read the stored row first: its decrypted values back any credential the # caller echoed back redacted, and its key set drives the audit diff. existing_row = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"}) - before_settings: Optional[Dict[str, Any]] = None + before_settings: dict[str, Any] | None = None saved_settings: dict[str, Any] = {} if existing_row is not None and existing_row.cache_settings: before_settings = _parse_stored_settings(existing_row.cache_settings) @@ -654,5 +652,5 @@ async def update_cache_settings( "settings": _redact_credentials(cache_settings), } except Exception as e: - verbose_proxy_logger.error(f"Error updating cache settings: {str(e)}") - raise HTTPException(status_code=500, detail=f"Error updating cache settings: {str(e)}") + verbose_proxy_logger.error(f"Error updating cache settings: {e!s}") + raise HTTPException(status_code=500, detail=f"Error updating cache settings: {e!s}") diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 8a5a31710cf..9bd8db16769 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -993,10 +993,10 @@ async def get_daily_activity( ) except Exception as e: - verbose_proxy_logger.exception(f"Error fetching daily activity: {str(e)}") + verbose_proxy_logger.exception(f"Error fetching daily activity: {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {str(e)}"}, + detail={"error": f"Failed to fetch analytics: {e!s}"}, ) @@ -1082,8 +1082,8 @@ async def get_daily_activity_aggregated( ) except Exception as e: - verbose_proxy_logger.exception(f"Error fetching aggregated daily activity: {str(e)}") + verbose_proxy_logger.exception(f"Error fetching aggregated daily activity: {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {str(e)}"}, + detail={"error": f"Failed to fetch analytics: {e!s}"}, ) diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 877130c2066..abb1b686d1d 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -1,5 +1,5 @@ import math -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Optional, Union from fastapi import HTTPException, status from pydantic import BaseModel @@ -417,7 +417,7 @@ def _is_set_budget_value(value: Any) -> bool: return True -def _has_meaningful_budget_limit(budget_values: Dict[str, Any]) -> bool: +def _has_meaningful_budget_limit(budget_values: dict[str, Any]) -> bool: """A budget is meaningful if at least one limit is actually set; an empty list (no model restriction) and None both count as unset.""" return any(_is_set_budget_value(budget_values.get(field)) for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS) @@ -428,10 +428,10 @@ async def _upsert_budget_and_membership( *, team_id: str, user_id: str, - existing_budget_id: Optional[str], + existing_budget_id: str | None, user_api_key_dict: UserAPIKeyAuth, - budget_patch: Dict[str, Any], - team_default_budget_id: Optional[str] = None, + budget_patch: dict[str, Any], + team_default_budget_id: str | None = None, ): """ Apply a merge-patch of per-member budget fields to a team membership. @@ -482,7 +482,7 @@ async def _upsert_budget_and_membership( ) return - create_data: Dict[str, Any] = { + create_data: dict[str, Any] = { "created_by": user_api_key_dict.user_id or "", "updated_by": user_api_key_dict.user_id or "", } diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 66c5f75b771..1c0f9fb85dd 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -3,7 +3,7 @@ import json import os from collections.abc import Mapping from datetime import datetime, timezone -from typing import Any, Dict, Optional, Set +from typing import Any from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import TypeAdapter @@ -44,7 +44,7 @@ router = APIRouter() _AUDIT_REDACTED = "***REDACTED***" -def _redact_config(config: Optional[Mapping[str, Any]]) -> Dict[str, Any]: +def _redact_config(config: Mapping[str, Any] | None) -> dict[str, Any]: """Strip values from a config snapshot before audit-log emission. Hashicorp Vault config carries ``vault_token``, ``approle_secret_id``, @@ -68,10 +68,10 @@ def _log_audit_task_exception(task: "asyncio.Task[None]") -> None: async def _emit_hashicorp_vault_audit_log( *, action: AUDIT_ACTIONS, - before_config: Optional[Mapping[str, Any]], - after_config: Optional[Mapping[str, Any]], + before_config: Mapping[str, Any] | None, + after_config: Mapping[str, Any] | None, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, ) -> None: """Emit an audit-log row for a /config_overrides/hashicorp_vault mutation. @@ -110,7 +110,7 @@ async def _emit_hashicorp_vault_audit_log( # --- Hashicorp Vault constants --- -HASHICORP_ENV_VAR_MAPPING: Dict[str, str] = { +HASHICORP_ENV_VAR_MAPPING: dict[str, str] = { "vault_addr": "HCP_VAULT_ADDR", "vault_token": "HCP_VAULT_TOKEN", "approle_role_id": "HCP_VAULT_APPROLE_ROLE_ID", @@ -124,7 +124,7 @@ HASHICORP_ENV_VAR_MAPPING: Dict[str, str] = { "vault_path_prefix": "HCP_VAULT_PATH_PREFIX", } -HASHICORP_SENSITIVE_FIELDS: Set[str] = { +HASHICORP_SENSITIVE_FIELDS: set[str] = { "vault_token", "approle_secret_id", "client_key", @@ -136,7 +136,7 @@ _sensitive_masker = SensitiveDataMasker() # --- Shared helpers --- -def _mask_sensitive_fields(data: Dict[str, Any], sensitive_fields: Set[str]) -> Dict[str, Any]: +def _mask_sensitive_fields(data: dict[str, Any], sensitive_fields: set[str]) -> dict[str, Any]: """Mask sensitive fields for API responses. Non-sensitive fields are left as-is.""" masked = {} for key, value in data.items(): @@ -147,7 +147,7 @@ def _mask_sensitive_fields(data: Dict[str, Any], sensitive_fields: Set[str]) -> return masked -def _get_current_env_values(env_var_mapping: Dict[str, str]) -> Dict[str, Any]: +def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, Any]: """Read current env var values as fallback when no DB record exists.""" values = {} for field_name, env_var_name in env_var_mapping.items(): @@ -156,7 +156,7 @@ def _get_current_env_values(env_var_mapping: Dict[str, str]) -> Dict[str, Any]: return values -def _extract_field_type(field_info: Dict[str, Any]) -> str: +def _extract_field_type(field_info: dict[str, Any]) -> str: """Extract the non-null type from a Pydantic v2 JSON schema field.""" if "type" in field_info: return field_info["type"] @@ -166,7 +166,7 @@ def _extract_field_type(field_info: Dict[str, Any]) -> str: return "string" -def _build_field_schema(model_class: type) -> Dict[str, Any]: +def _build_field_schema(model_class: type) -> dict[str, Any]: """Build field_schema dict from a Pydantic model for UI rendering.""" schema = TypeAdapter(model_class).json_schema(by_alias=True) properties = {} @@ -181,14 +181,14 @@ def _build_field_schema(model_class: type) -> Dict[str, Any]: } -def _parse_config_value(raw: Any) -> Dict[str, Any]: +def _parse_config_value(raw: Any) -> dict[str, Any]: """Parse a config_value from DB (may be JSON string or dict).""" if isinstance(raw, str): return safe_json_loads(raw, default={}) return dict(raw) -def _set_env_vars(config_data: Dict[str, Any]) -> None: +def _set_env_vars(config_data: dict[str, Any]) -> None: """Set HCP_VAULT_* env vars from config data. Unsets vars for missing/None/empty fields.""" for field_name, env_var_name in HASHICORP_ENV_VAR_MAPPING.items(): value = config_data.get(field_name) @@ -218,7 +218,7 @@ def _clear_hashicorp_vault_state(proxy_config: Any) -> None: async def update_hashicorp_vault_config( config: HashicorpVaultConfig, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -249,8 +249,8 @@ async def update_hashicorp_vault_config( existing_record = await ConfigOverridesRepository(prisma_client).table.find_unique( where={"config_type": "hashicorp_vault"} ) - existing_decrypted: Optional[Dict[str, Any]] = None - env_values: Dict[str, Any] = {} + existing_decrypted: dict[str, Any] | None = None + env_values: dict[str, Any] = {} if existing_record is not None and existing_record.config_value is not None: existing_data = _parse_config_value(existing_record.config_value) existing_decrypted = proxy_config._decrypt_db_variables(existing_data) @@ -412,7 +412,7 @@ async def get_hashicorp_vault_config( ) async def delete_hashicorp_vault_config( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -437,7 +437,7 @@ async def delete_hashicorp_vault_config( existing_record = await ConfigOverridesRepository(prisma_client).table.find_unique( where={"config_type": "hashicorp_vault"} ) - before_config: Optional[Dict[str, Any]] = None + before_config: dict[str, Any] | None = None if existing_record is not None and existing_record.config_value is not None: try: before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value)) diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index 7ab4e3019c3..42bb75ea752 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -16,7 +16,6 @@ import json from collections.abc import Mapping from contextlib import suppress from datetime import datetime, timezone -from typing import Optional from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field, TypeAdapter, ValidationError @@ -91,7 +90,7 @@ def _redact_credentials(settings: Mapping[str, object]) -> dict[str, object]: } -def _redact_all_values(settings: Optional[Mapping[str, object]]) -> dict[str, object]: +def _redact_all_values(settings: Mapping[str, object] | None) -> dict[str, object]: """Replace every value with a fixed marker, preserving the key set. The audit row shows *which* fields changed without the audit table becoming @@ -171,7 +170,7 @@ async def _read_general_settings() -> dict[str, object]: return _SETTINGS_ADAPTER.validate_python(config_param.param_value) -async def get_persisted_coordination_redis_settings() -> Optional[dict[str, object]]: +async def get_persisted_coordination_redis_settings() -> dict[str, object] | None: """The coordination_redis block saved to the database, if any. Read at startup so settings saved from the admin UI take effect on the next @@ -183,7 +182,7 @@ async def get_persisted_coordination_redis_settings() -> Optional[dict[str, obje return None -async def _current_coordination_redis_settings() -> Optional[dict[str, object]]: +async def _current_coordination_redis_settings() -> dict[str, object] | None: """The coordination_redis block the proxy would boot with. The persisted row wins over the yaml-loaded config state because startup @@ -205,7 +204,7 @@ async def _current_coordination_redis_settings() -> Optional[dict[str, object]]: return None -def _coordination_redis_source(settings: Optional[Mapping[str, object]]) -> Optional[CoordinationRedisSource]: +def _coordination_redis_source(settings: Mapping[str, object] | None) -> CoordinationRedisSource | None: """Which source the proxy's coordination Redis comes from, in startup precedence order. Mirrors `ProxyConfig._init_coordination_redis` -> `ProxyConfig._init_cache`: @@ -236,10 +235,10 @@ def _log_audit_task_exception(task: "asyncio.Task[None]") -> None: async def _emit_coordination_redis_audit_log( *, action: AUDIT_ACTIONS, - before_settings: Optional[Mapping[str, object]], - after_settings: Optional[Mapping[str, object]], + before_settings: Mapping[str, object] | None, + after_settings: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, ) -> None: """Emit an audit-log row for a /coordination_redis/settings mutation.""" if litellm.store_audit_logs is not True: @@ -271,7 +270,7 @@ class CoordinationRedisSettingsResponse(BaseModel): fields: list[CoordinationRedisSettingsField] = Field( description="List of all configurable coordination Redis settings with metadata" ) - source: Optional[CoordinationRedisSource] = Field( + source: CoordinationRedisSource | None = Field( description="Where the proxy's coordination Redis comes from; null when it has none" ) @@ -282,7 +281,7 @@ class CoordinationRedisSettingsRequest(BaseModel): class CoordinationRedisTestResponse(BaseModel): status: str = Field(description="Connection status: 'healthy' or 'unhealthy'") - error: Optional[str] = Field(default=None, description="Error message if the connection failed") + error: str | None = Field(default=None, description="Error message if the connection failed") @router.get( @@ -324,7 +323,7 @@ async def get_coordination_redis_settings( async def update_coordination_redis_settings( request: CoordinationRedisSettingsRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -413,7 +412,7 @@ async def check_coordination_redis_connection( settings = _merge_over_saved(request.settings, saved_settings or {}) params = _validated_params(settings) - redis_cache: Optional[RedisCache] = None + redis_cache: RedisCache | None = None try: redis_cache = _build_redis_usage_cache(params.model_dump(exclude_none=True)) await asyncio.wait_for(redis_cache.ping(), timeout=_PING_TIMEOUT_SECONDS) diff --git a/litellm/proxy/management_endpoints/cost_tracking_settings.py b/litellm/proxy/management_endpoints/cost_tracking_settings.py index cd2c5704778..9d985e48a60 100644 --- a/litellm/proxy/management_endpoints/cost_tracking_settings.py +++ b/litellm/proxy/management_endpoints/cost_tracking_settings.py @@ -10,8 +10,6 @@ PATCH /config/cost_margin_config - Update cost margin configuration POST /cost/estimate - Estimate cost for a given model and token counts """ -from typing import Dict, Optional, Tuple, Union - from fastapi import APIRouter, Depends, HTTPException import litellm @@ -29,7 +27,7 @@ from litellm.types.utils import LlmProvidersSet router = APIRouter() -def _resolve_model_for_cost_lookup(model: str) -> Tuple[str, Optional[str]]: +def _resolve_model_for_cost_lookup(model: str) -> tuple[str, str | None]: """ Resolve a model name (which may be a router alias/model_group) to the underlying litellm model name for cost lookup. @@ -45,7 +43,7 @@ def _resolve_model_for_cost_lookup(model: str) -> Tuple[str, Optional[str]]: """ from litellm.proxy.proxy_server import llm_router - custom_llm_provider: Optional[str] = None + custom_llm_provider: str | None = None # Try to resolve from router if available if llm_router is not None: @@ -131,7 +129,7 @@ async def get_cost_discount_config( return {"values": cost_discount_config} except Exception as e: - verbose_proxy_logger.error(f"Error fetching cost discount config: {str(e)}") + verbose_proxy_logger.error(f"Error fetching cost discount config: {e!s}") return {"values": {}} @@ -141,7 +139,7 @@ async def get_cost_discount_config( dependencies=[Depends(user_api_key_auth)], ) async def update_cost_discount_config( - cost_discount_config: Dict[str, float], + cost_discount_config: dict[str, float], user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -179,7 +177,7 @@ async def update_cost_discount_config( # Validate that all providers are valid LiteLLM providers invalid_providers = [] - for provider in cost_discount_config.keys(): + for provider in cost_discount_config: if provider not in LlmProvidersSet: invalid_providers.append(provider) @@ -226,10 +224,10 @@ async def update_cost_discount_config( "values": cost_discount_config, } except Exception as e: - verbose_proxy_logger.error(f"Error updating cost discount config: {str(e)}") + verbose_proxy_logger.error(f"Error updating cost discount config: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to update cost discount config: {str(e)}"}, + detail={"error": f"Failed to update cost discount config: {e!s}"}, ) @@ -264,7 +262,7 @@ async def get_cost_margin_config( return {"values": cost_margin_config} except Exception as e: - verbose_proxy_logger.error(f"Error fetching cost margin config: {str(e)}") + verbose_proxy_logger.error(f"Error fetching cost margin config: {e!s}") return {"values": {}} @@ -274,7 +272,7 @@ async def get_cost_margin_config( dependencies=[Depends(user_api_key_auth)], ) async def update_cost_margin_config( - cost_margin_config: Dict[str, Union[float, Dict[str, float]]], + cost_margin_config: dict[str, float | dict[str, float]], user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -317,7 +315,7 @@ async def update_cost_margin_config( # Validate that all providers are valid LiteLLM providers (except "global") invalid_providers = [] - for provider in cost_margin_config.keys(): + for provider in cost_margin_config: if provider != "global" and provider not in LlmProvidersSet: invalid_providers.append(provider) @@ -400,10 +398,10 @@ async def update_cost_margin_config( "values": cost_margin_config, } except Exception as e: - verbose_proxy_logger.error(f"Error updating cost margin config: {str(e)}") + verbose_proxy_logger.error(f"Error updating cost margin config: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to update cost margin config: {str(e)}"}, + detail={"error": f"Failed to update cost margin config: {e!s}"}, ) @@ -486,7 +484,7 @@ async def estimate_cost( raise HTTPException( status_code=404, detail={ - "error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {str(e)}" + "error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e!s}" }, ) diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 4d51295f8dc..c15e6bc7838 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -130,9 +130,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass: return "plaintext" if value.startswith(_V2_GCM_PREFIX): return "migrated" - decrypted = decrypt_value_helper( - value=value, key=key, exception_type="debug", return_original_value=False - ) + decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: # Did not decrypt under nacl and has no v2 marker: legacy plaintext. return "plaintext" @@ -151,9 +149,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: return value if value.startswith(_V2_GCM_PREFIX): return value # idempotent: already migrated - decrypted = decrypt_value_helper( - value=value, key=key, exception_type="debug", return_original_value=False - ) + decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: # Either legacy plaintext (no ciphertext to migrate) or corrupt. Either # way, do not overwrite — preserve the value as stored. @@ -161,9 +157,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: return encrypt_value_helper(decrypted) -def reencrypt_selective_dict( - data: dict[str, object], sensitive_keys: list[str] -) -> dict[str, object]: +def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]: """Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted. Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is. @@ -212,9 +206,7 @@ async def _migrate_config_settings_row( dict with selected sensitive fields (vantage_settings / cloudzero_settings). """ report = LocationReport(location=param_name) - record = await prisma_client.db.litellm_config.find_unique( - where={"param_name": param_name} - ) + record = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name}) if record is None or record.param_value is None: return report @@ -266,9 +258,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR every present string field. """ report = LocationReport(location="sso_config") - record = await prisma_client.db.litellm_ssoconfig.find_unique( - where={"id": "sso_config"} - ) + record = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"}) if record is None or record.sso_settings is None: return report @@ -344,9 +334,7 @@ async def _migrate_callback_vars_table( rows = await table.find_many() for row in rows or []: metadata = getattr(row, "metadata", None) - if not isinstance(metadata, dict) or ( - "logging" not in metadata and "callback_settings" not in metadata - ): + if not isinstance(metadata, dict) or ("logging" not in metadata and "callback_settings" not in metadata): continue # Classify every callback-var value directly (strip the litellm_enc:: @@ -435,8 +423,7 @@ def _classify_callback_value(value: object) -> ValueClass: if not isinstance(value, str): return "not-a-string" inner = value - if inner.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX): - inner = inner[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :] + inner = inner.removeprefix(_CALLBACK_VAR_ENCRYPTED_PREFIX) return classify_value(inner, key="callback") @@ -534,9 +521,7 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport: """Scan the ``environment_variables`` config row (``param_value`` dict).""" report = LocationReport(location="config_environment_variables") try: - record = await prisma_client.db.litellm_config.find_unique( - where={"param_name": "environment_variables"} - ) + record = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"}) except Exception as e: # pragma: no cover - defensive verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e)) return report @@ -557,11 +542,7 @@ async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]: """Read-only classification of every rotation-covered table. No writes.""" reports: list[LocationReport] = [] for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS: - reports.append( - await _scan_one_table( - prisma_client, location, db_attr, json_cols, scalar_cols - ) - ) + reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols)) reports.append(await _scan_config_env_vars(prisma_client)) return reports @@ -575,9 +556,7 @@ _VANTAGE_SENSITIVE = ["api_key", "integration_token"] _CLOUDZERO_SENSITIVE = ["api_key"] -async def _migrate_covered_tables( - prisma_client: object, user_api_key_dict: object -) -> list[LocationReport]: +async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]: """Re-encrypt the tables already covered by ``_rotate_master_key`` (model table, credentials, MCP credential/env tables, config environment_variables) by running that orchestrator in *same-key* mode. With the AES gate on, the @@ -597,8 +576,7 @@ async def _migrate_covered_tables( current_key = _get_salt_key() if current_key is None: raise RuntimeError( - "Cannot migrate covered tables: no salt key / master key is set. " - "Set LITELLM_SALT_KEY before migrating." + "Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating." ) await _rotate_master_key( prisma_client=cast("PrismaClient", prisma_client), @@ -648,19 +626,9 @@ async def migrate_encryption( # Net-new walkers (items 3, 4, 11, 12, 13). report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run)) - report.add( - await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run) - ) - report.add( - await _migrate_config_settings_row( - prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run - ) - ) - report.add( - await _migrate_config_settings_row( - prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run - ) - ) + report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)) + report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run)) + report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run)) report.add(await _migrate_sso_config(prisma_client, dry_run)) return report @@ -683,20 +651,10 @@ async def check_encryption(prisma_client: object) -> MigrationReport: # Net-new walker locations, in dry-run (read-only) mode. report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True)) + report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True)) + report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True)) report.add( - await _migrate_callback_vars_table( - prisma_client, "verification_token", dry_run=True - ) - ) - report.add( - await _migrate_config_settings_row( - prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True - ) - ) - report.add( - await _migrate_config_settings_row( - prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True - ) + await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True) ) report.add(await _migrate_sso_config(prisma_client, dry_run=True)) return report diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 84f67bdc3bc..09977fdce40 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -11,7 +11,6 @@ All /customer management endpoints #### END-USER/CUSTOMER MANAGEMENT #### from datetime import datetime, timedelta -from typing import List, Optional import fastapi from fastapi import APIRouter, Depends, HTTPException, Request @@ -104,7 +103,7 @@ async def block_user(data: BlockUsers): return {"blocked_users": records} except Exception as e: - verbose_proxy_logger.error(f"An error occurred - {str(e)}") + verbose_proxy_logger.error(f"An error occurred - {e!s}") raise HTTPException(status_code=500, detail={"error": str(e)}) @@ -167,7 +166,7 @@ async def unblock_user(data: BlockUsers): return {"blocked_users": litellm.blocked_user_list} -def new_budget_request(data: NewCustomerRequest) -> Optional[BudgetNewRequest]: +def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None: """ Return a new budget object if new budget params are passed. """ @@ -194,7 +193,7 @@ def new_budget_request(data: NewCustomerRequest) -> Optional[BudgetNewRequest]: async def _handle_customer_object_permission_update( non_default_values: dict, - end_user_table_data_typed: Optional[LiteLLM_EndUserTable], + end_user_table_data_typed: LiteLLM_EndUserTable | None, update_end_user_table_data: dict, prisma_client, ) -> None: @@ -338,13 +337,11 @@ async def new_end_user( raise HTTPException( status_code=422, detail={ - "error": "Default Model not on proxy. Configure via `/model/new` or config.yaml. Default_model={}, proxy_model_names={}".format( - data.default_model, set(llm_router.get_model_names()) - ) + "error": f"Default Model not on proxy. Configure via `/model/new` or config.yaml. Default_model={data.default_model}, proxy_model_names={set(llm_router.get_model_names())}" }, ) - new_end_user_obj: Dict = {} + new_end_user_obj: dict = {} ## CREATE BUDGET ## if set _new_budget = new_budget_request(data) @@ -393,9 +390,7 @@ async def new_end_user( return _to_customer_response(end_user_record) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.management_endpoints.customer_endpoints.new_end_user(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.customer_endpoints.new_end_user(): Exception occured - {e!s}" ) if "Unique constraint failed on the fields: (`user_id`)" in str(e): raise ProxyException( @@ -450,7 +445,7 @@ async def end_user_info( if user_info is None: raise ProxyException( - message="End User Id={} does not exist in db".format(end_user_id), + message=f"End User Id={end_user_id} does not exist in db", type="not_found", code=404, param="end_user_id", @@ -460,9 +455,7 @@ async def end_user_info( except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.management_endpoints.customer_endpoints.end_user_info(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.customer_endpoints.end_user_info(): Exception occured - {e!s}" ) raise handle_exception_on_proxy(e) @@ -560,7 +553,7 @@ async def update_end_user( if end_user_table_data is None: raise ProxyException( - message="End User Id={} does not exist in db".format(data.user_id), + message=f"End User Id={data.user_id} does not exist in db", type="not_found", code=404, param="user_id", @@ -643,9 +636,7 @@ async def update_end_user( # update based on remaining passed in values except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.update_end_user(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.update_end_user(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -720,9 +711,7 @@ async def delete_end_user( # update based on remaining passed in values except Exception as e: - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.delete_end_user(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.delete_end_user(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -730,7 +719,7 @@ async def delete_end_user( "/customer/list", tags=["Customer Management"], dependencies=[Depends(user_api_key_auth)], - response_model=List[CustomerResponse], + response_model=list[CustomerResponse], ) @router.get( "/end_user/list", @@ -741,7 +730,7 @@ async def delete_end_user( async def list_end_user( http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> List[CustomerResponse]: +) -> list[CustomerResponse]: """ [Admin-only] List all available customers @@ -761,7 +750,7 @@ async def list_end_user( ): raise HTTPException( status_code=401, - detail={"error": "Admin-only endpoint. Your user role={}".format(user_api_key_dict.user_role)}, + detail={"error": f"Admin-only endpoint. Your user role={user_api_key_dict.user_role}"}, ) if prisma_client is None: @@ -778,9 +767,7 @@ async def list_end_user( except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.management_endpoints.customer_endpoints.list_end_user(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.customer_endpoints.list_end_user(): Exception occured - {e!s}" ) raise handle_exception_on_proxy(e) @@ -798,14 +785,14 @@ async def list_end_user( dependencies=[Depends(user_api_key_auth)], ) async def get_customer_daily_activity( - end_user_ids: Optional[str] = None, - start_date: Optional[str] = None, - end_date: Optional[str] = None, - model: Optional[str] = None, - api_key: Optional[str] = None, + end_user_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, page: int = 1, page_size: int = 10, - exclude_end_user_ids: Optional[str] = None, + exclude_end_user_ids: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -817,7 +804,7 @@ async def get_customer_daily_activity( ): raise HTTPException( status_code=401, - detail={"error": "Admin-only endpoint. Your user role={}".format(user_api_key_dict.user_role)}, + detail={"error": f"Admin-only endpoint. Your user role={user_api_key_dict.user_role}"}, ) from litellm.proxy.proxy_server import prisma_client @@ -830,7 +817,7 @@ async def get_customer_daily_activity( # Parse comma-separated ids end_user_ids_list = end_user_ids.split(",") if end_user_ids else None - exclude_end_user_ids_list: Optional[List[str]] = None + exclude_end_user_ids_list: list[str] | None = None if exclude_end_user_ids: exclude_end_user_ids_list = exclude_end_user_ids.split(",") if exclude_end_user_ids else None diff --git a/litellm/proxy/management_endpoints/fallback_management_endpoints.py b/litellm/proxy/management_endpoints/fallback_management_endpoints.py index dc594923166..f765cf379e4 100644 --- a/litellm/proxy/management_endpoints/fallback_management_endpoints.py +++ b/litellm/proxy/management_endpoints/fallback_management_endpoints.py @@ -11,7 +11,7 @@ DELETE /fallback/{model} - Delete fallbacks for a specific model # pyright: reportMissingImports=false import json -from typing import TYPE_CHECKING, Dict, List, Literal +from typing import TYPE_CHECKING, Literal from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth @@ -134,7 +134,7 @@ async def create_fallback( fallback_key = "content_policy_fallbacks" # Get existing fallbacks - existing_fallbacks: List[Dict[str, List[str]]] = router_settings.get(fallback_key, []) + existing_fallbacks: list[dict[str, list[str]]] = router_settings.get(fallback_key, []) # Update or add the fallback configuration fallback_updated = False @@ -182,10 +182,10 @@ async def create_fallback( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error creating fallback: {str(e)}", exc_info=True) + verbose_proxy_logger.error(f"Error creating fallback: {e!s}", exc_info=True) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to create fallback: {str(e)}"}, + detail={"error": f"Failed to create fallback: {e!s}"}, ) @@ -239,10 +239,10 @@ async def get_fallback( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error getting fallback: {str(e)}", exc_info=True) + verbose_proxy_logger.error(f"Error getting fallback: {e!s}", exc_info=True) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to get fallback: {str(e)}"}, + detail={"error": f"Failed to get fallback: {e!s}"}, ) @@ -303,7 +303,7 @@ async def delete_fallback( fallback_key = "content_policy_fallbacks" # Get existing fallbacks - existing_fallbacks: List[Dict[str, List[str]]] = router_settings.get(fallback_key, []) + existing_fallbacks: list[dict[str, list[str]]] = router_settings.get(fallback_key, []) # Find and remove the fallback configuration fallback_found = False @@ -350,8 +350,8 @@ async def delete_fallback( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error deleting fallback: {str(e)}", exc_info=True) + verbose_proxy_logger.error(f"Error deleting fallback: {e!s}", exc_info=True) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to delete fallback: {str(e)}"}, + detail={"error": f"Failed to delete fallback: {e!s}"}, ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 1a0978c8eec..d87a0b3d096 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -181,9 +181,13 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d for key, value in litellm.default_internal_user_params.items(): if key == "available_teams": continue - elif key not in data_json or data_json[key] is None: - data_json[key] = value - elif key == "models" and isinstance(data_json[key], list) and len(data_json[key]) == 0: + elif ( + key not in data_json + or data_json[key] is None + or key == "models" + and isinstance(data_json[key], list) + and len(data_json[key]) == 0 + ): data_json[key] = value ## INTERNAL USER ROLE ONLY DEFAULT PARAMS ## @@ -326,9 +330,7 @@ async def _add_user_to_team( except HTTPException as e: if e.status_code == 400 and ("already exists" in str(e) or "doesn't exist" in str(e)): verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {e!s}" ) else: verbose_proxy_logger.error( @@ -339,17 +341,14 @@ async def _add_user_to_team( str(e), ) except Exception as e: - if "already exists" in str(e) or "doesn't exist" in str(e): + if ( + "already exists" in str(e) + or "doesn't exist" in str(e) + or isinstance(e, ProxyException) + and ProxyErrorTypes.team_member_already_in_team in e.type + ): verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( - str(e) - ) - ) - elif isinstance(e, ProxyException) and ProxyErrorTypes.team_member_already_in_team in e.type: - verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {e!s}" ) else: verbose_proxy_logger.error( @@ -606,7 +605,7 @@ async def new_user( return new_user_response except Exception as e: - verbose_proxy_logger.exception("/user/new: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"/user/new: Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -901,7 +900,7 @@ async def user_info( return response_data except Exception as e: - verbose_proxy_logger.exception("litellm.proxy.proxy_server.user_info(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.user_info(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -1051,9 +1050,7 @@ async def user_info_v2( object_permission=user_data.get("object_permission"), ) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.user_info_v2(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.user_info_v2(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -1323,7 +1320,7 @@ async def _invalidate_cached_user_entitlement(user_id: str | None, object_permis try: await user_api_key_cache.async_delete_cache(key=key) except Exception as e: # noqa: BLE001 # a cache we cannot clear still expires; never fail the write - verbose_proxy_logger.warning(f"Failed to invalidate cached entitlement key {key!r}: {str(e)}") + verbose_proxy_logger.warning(f"Failed to invalidate cached entitlement key {key!r}: {e!s}") async def _update_single_user_helper( @@ -1492,9 +1489,10 @@ def can_user_call_user_update( """ Helper to check if the user has access to the key's info """ - if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: - return True - elif user_api_key_dict.user_id == user_info.user_id: + if ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + or user_api_key_dict.user_id == user_info.user_id + ): return True return False @@ -1571,13 +1569,11 @@ async def user_update( ) return response except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.user_update(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.user_update(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -2399,7 +2395,7 @@ async def add_internal_user_to_organization( return new_membership except Exception as e: - raise Exception(f"Failed to add user to organization: {str(e)}") + raise Exception(f"Failed to add user to organization: {e!s}") async def _resolve_org_filter_for_user_search( @@ -2597,8 +2593,8 @@ async def ui_view_users( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error searching users: {str(e)}") - raise HTTPException(status_code=500, detail=f"Error searching users: {str(e)}") + verbose_proxy_logger.exception(f"Error searching users: {e!s}") + raise HTTPException(status_code=500, detail=f"Error searching users: {e!s}") # Using shared metric helper implementations from common_daily_activity @@ -2720,10 +2716,10 @@ async def get_user_daily_activity( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception("/spend/daily/analytics: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"/spend/daily/analytics: Exception occured - {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {str(e)}"}, + detail={"error": f"Failed to fetch analytics: {e!s}"}, ) @@ -2812,8 +2808,8 @@ async def get_user_daily_activity_aggregated( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception("/user/daily/activity/aggregated: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"/user/daily/activity/aggregated: Exception occured - {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {str(e)}"}, + detail={"error": f"Failed to fetch analytics: {e!s}"}, ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3ec9e1303c4..91b3d8a0130 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -20,7 +20,7 @@ import secrets import traceback from collections.abc import Callable, Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import Any, Dict, List, Literal, Optional, Protocol, Tuple, TypeVar, cast +from typing import Any, Literal, Optional, Protocol, TypeVar, cast import fastapi import yaml @@ -196,7 +196,7 @@ def _config_table(prisma_client: PrismaClient) -> _PrismaTableActions[ConfigPara return ConfigRepository(prisma_client).table -async def _check_custom_key_allowed(custom_key_value: Optional[str]) -> None: +async def _check_custom_key_allowed(custom_key_value: str | None) -> None: """Raise 403 if custom API keys are disabled and a custom key was provided.""" if custom_key_value is None: return @@ -214,7 +214,7 @@ def _is_team_key(data: Union[GenerateKeyRequest, LiteLLM_VerificationToken]): return data.team_id is not None -def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: Optional[str]) -> Optional[Member]: +def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: str | None) -> Member | None: if user_id is None: return None for member in team_table.members_with_roles: @@ -242,8 +242,8 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: def _set_key_rotation_fields( data: dict, auto_rotate: bool, - rotation_interval: Optional[str], - existing_key_alias: Optional[str] = None, + rotation_interval: str | None, + existing_key_alias: str | None = None, ) -> None: """ Helper function to set rotation fields in key data if auto_rotate is enabled. @@ -278,8 +278,8 @@ def _set_key_rotation_fields( def _is_allowed_to_make_key_request( user_api_key_dict: UserAPIKeyAuth, - user_id: Optional[str], - team_id: Optional[str], + user_id: str | None, + team_id: str | None, ) -> bool: """ Assert user only creates/updates keys for themselves @@ -292,9 +292,7 @@ def _is_allowed_to_make_key_request( if user_id is not None: assert user_id == user_api_key_dict.user_id, ( - "User can only create keys for themselves. Got user_id={}, Your ID={}".format( - user_id, user_api_key_dict.user_id - ) + f"User can only create keys for themselves. Got user_id={user_id}, Your ID={user_api_key_dict.user_id}" ) if team_id is not None: @@ -305,7 +303,7 @@ def _is_allowed_to_make_key_request( def _team_key_operation_team_member_check( - assigned_user_id: Optional[str], + assigned_user_id: str | None, team_table: LiteLLM_TeamTableCachedObj, user_api_key_dict: UserAPIKeyAuth, team_key_generation: TeamUIKeyGenerationConfig, @@ -350,7 +348,7 @@ def _team_key_operation_team_member_check( return True -def _key_generation_required_param_check(data: GenerateKeyRequest, required_params: Optional[List[str]]): +def _key_generation_required_param_check(data: GenerateKeyRequest, required_params: list[str] | None): if required_params is None: return True @@ -404,7 +402,7 @@ def _team_key_generation_check( def _personal_key_membership_check( user_api_key_dict: UserAPIKeyAuth, - personal_key_generation: Optional[PersonalUIKeyGenerationConfig], + personal_key_generation: PersonalUIKeyGenerationConfig | None, ): if personal_key_generation is None or "allowed_user_roles" not in personal_key_generation: return True @@ -419,8 +417,8 @@ def _personal_key_membership_check( def _object_permission_to_dict( - object_permission: Optional[LiteLLM_ObjectPermissionBase], -) -> Optional[ObjectPermissionDict]: + object_permission: LiteLLM_ObjectPermissionBase | None, +) -> ObjectPermissionDict | None: if object_permission is None: return None return cast(ObjectPermissionDict, object_permission.model_dump(exclude_unset=True)) @@ -455,7 +453,7 @@ def _personal_key_generation_check(user_api_key_dict: UserAPIKeyAuth, data: Gene def key_generation_check( - team_table: Optional[LiteLLM_TeamTableCachedObj], + team_table: LiteLLM_TeamTableCachedObj | None, user_api_key_dict: UserAPIKeyAuth, data: GenerateKeyRequest, route: KeyManagementRoutes, @@ -496,9 +494,9 @@ def key_generation_check( def common_key_access_checks( user_api_key_dict: UserAPIKeyAuth, data: Union[GenerateKeyRequest, UpdateKeyRequest], - llm_router: Optional[Router], + llm_router: Router | None, premium_user: bool, - user_id: Optional[str] = None, + user_id: str | None = None, ) -> Literal[True]: """ Check if user is allowed to make a key request, for this key @@ -553,7 +551,7 @@ _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS = frozenset({"llm_api_routes", "info_rout def _validate_caller_can_change_key_ownership( - data: Optional[BaseModel], + data: BaseModel | None, existing_key_row: LiteLLM_VerificationToken, user_api_key_dict: UserAPIKeyAuth, ) -> None: @@ -599,7 +597,7 @@ def _validate_caller_can_change_key_ownership( def _check_allowed_routes_caller_permission( - allowed_routes: Optional[list], + allowed_routes: list | None, user_api_key_dict: UserAPIKeyAuth, *, allowed_routes_was_provided: bool = False, @@ -664,11 +662,11 @@ def _check_permissions_caller_permission( def _check_budget_limits_delegation_ceiling( - budget_limits: Optional[List[BudgetLimitEntry]], - delegation_ceiling: Optional[float], + budget_limits: list[BudgetLimitEntry] | None, + delegation_ceiling: float | None, user_api_key_dict: UserAPIKeyAuth, is_ui_session_team_key: bool, - team_table: Optional[LiteLLM_TeamTableCachedObj], + team_table: LiteLLM_TeamTableCachedObj | None, ) -> None: """ Enforce three invariants on `budget_limits`: @@ -715,8 +713,8 @@ def _check_budget_limits_delegation_ceiling( async def validate_team_id_used_in_service_account_request( - team_id: Optional[str], - prisma_client: Optional[PrismaClient], + team_id: str | None, + prisma_client: PrismaClient | None, ): """ Validate team_id is used in the request body for generating a service account key @@ -813,8 +811,8 @@ def _enforce_upperbound_key_params( async def _common_key_generation_helper( data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], - team_table: Optional[LiteLLM_TeamTableCachedObj], + litellm_changed_by: str | None, + team_table: LiteLLM_TeamTableCachedObj | None, ) -> GenerateKeyResponse: from litellm.proxy.proxy_server import ( litellm_proxy_admin_name, @@ -938,9 +936,7 @@ async def _common_key_generation_helper( data = apply_enterprise_key_management_params(data, team_table) except Exception as e: verbose_proxy_logger.debug( - "litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - {}".format( - str(e) - ) + f"litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - {e!s}" ) # TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable @@ -1084,7 +1080,7 @@ async def _common_key_generation_helper( # Validate user-provided key format if data.key is not None and not data.key.startswith("sk-"): - _masked = "{}****{}".format(data.key[:4], data.key[-4:]) if len(data.key) > 8 else "****" + _masked = f"{data.key[:4]}****{data.key[-4:]}" if len(data.key) > 8 else "****" raise HTTPException( status_code=400, detail={"error": f"Invalid key format. LiteLLM Virtual Key must start with 'sk-'. Received: {_masked}"}, @@ -1159,12 +1155,12 @@ async def _common_key_generation_helper( def _check_key_model_specific_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], data: Union[GenerateKeyRequest, UpdateKeyRequest], - entity_rpm_limit: Optional[int], - entity_tpm_limit: Optional[int], - entity_model_rpm_limit_dict: Dict[str, int], - entity_model_tpm_limit_dict: Dict[str, int], + entity_rpm_limit: int | None, + entity_tpm_limit: int | None, + entity_model_rpm_limit_dict: dict[str, int], + entity_model_tpm_limit_dict: dict[str, int], entity_type: str, # "team" or "organization" ) -> None: """ @@ -1181,8 +1177,8 @@ def _check_key_model_specific_limits( return # get total model specific tpm/rpm limit - model_specific_rpm_limit: Dict[str, int] = {} - model_specific_tpm_limit: Dict[str, int] = {} + model_specific_rpm_limit: dict[str, int] = {} + model_specific_tpm_limit: dict[str, int] = {} for key in keys: if key.metadata.get("model_rpm_limit", None) is not None: @@ -1230,10 +1226,10 @@ def _check_key_model_specific_limits( def _check_key_rpm_tpm_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], data: Union[GenerateKeyRequest, UpdateKeyRequest], - entity_rpm_limit: Optional[int], - entity_tpm_limit: Optional[int], + entity_rpm_limit: int | None, + entity_tpm_limit: int | None, entity_type: str, # "team" or "organization" ) -> None: """ @@ -1268,7 +1264,7 @@ def _check_key_rpm_tpm_limits( def check_team_key_model_specific_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], team_table: LiteLLM_TeamTableCachedObj, data: Union[GenerateKeyRequest, UpdateKeyRequest], ) -> None: @@ -1293,7 +1289,7 @@ def check_team_key_model_specific_limits( def check_team_key_rpm_tpm_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], team_table: LiteLLM_TeamTableCachedObj, data: Union[GenerateKeyRequest, UpdateKeyRequest], ) -> None: @@ -1395,7 +1391,7 @@ async def _check_project_key_limits( def check_org_key_model_specific_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], org_table: LiteLLM_OrganizationTable, data: Union[GenerateKeyRequest, UpdateKeyRequest], ) -> None: @@ -1428,7 +1424,7 @@ def check_org_key_model_specific_limits( def check_org_key_rpm_tpm_limits( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], org_table: LiteLLM_OrganizationTable, data: Union[GenerateKeyRequest, UpdateKeyRequest], ) -> None: @@ -1537,7 +1533,7 @@ async def _check_org_key_limits( async def generate_key_fn( data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -1682,7 +1678,7 @@ async def generate_key_fn( user_api_key_dict.user_id, ) - team_table: Optional[LiteLLM_TeamTableCachedObj] = None + team_table: LiteLLM_TeamTableCachedObj | None = None if data.team_id is not None: try: team_table = await get_team_object( @@ -1732,9 +1728,7 @@ async def generate_key_fn( ) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.generate_key_fn(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.generate_key_fn(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -1747,7 +1741,7 @@ async def generate_key_fn( async def generate_service_account_key_fn( data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -1857,7 +1851,7 @@ async def generate_service_account_key_fn( message = result.get("message", "Authentication Failed - Custom Auth Rule") if not decision: raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message) - team_table: Optional[LiteLLM_TeamTableCachedObj] = None + team_table: LiteLLM_TeamTableCachedObj | None = None if data.team_id is not None: try: team_table = await get_team_object( @@ -1937,7 +1931,7 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_ except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.prepare_metadata_fields(): Exception occured - {}".format(str(e)) + f"litellm.proxy.proxy_server.prepare_metadata_fields(): Exception occured - {e!s}" ) non_default_values["metadata"] = encrypt_callback_vars(casted_metadata) @@ -2069,7 +2063,7 @@ def is_different_team(data: UpdateKeyRequest, existing_key_row: LiteLLM_Verifica return data.team_id != existing_key_row.team_id -def _validate_max_budget(max_budget: Optional[float]) -> None: +def _validate_max_budget(max_budget: float | None) -> None: """ Validate that max_budget is not negative. @@ -2087,7 +2081,7 @@ def _validate_max_budget(max_budget: Optional[float]) -> None: async def _get_and_validate_existing_key( - token: str | None, prisma_client: Optional[PrismaClient], key_alias: str | None = None + token: str | None, prisma_client: PrismaClient | None, key_alias: str | None = None ) -> LiteLLM_VerificationToken: """ Get existing key from database and validate it exists. @@ -2173,14 +2167,14 @@ def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_V async def _process_single_key_update( update_key_request: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], - prisma_client: Optional[PrismaClient], + litellm_changed_by: str | None, + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, - llm_router: Optional[Router], - user_custom_key_update: Optional[Callable] = None, - existing_key_row: Optional[LiteLLM_VerificationToken] = None, -) -> Dict[str, Any]: + llm_router: Router | None, + user_custom_key_update: Callable | None = None, + existing_key_row: LiteLLM_VerificationToken | None = None, +) -> dict[str, Any]: """ Process a single key update with all validations and checks. @@ -2243,7 +2237,7 @@ async def _process_single_key_update( _enforce_upperbound_key_params(update_key_request, fill_defaults=False) # Get team object and check team limits if team_id is provided - team_obj: Optional[LiteLLM_TeamTableCachedObj] = None + team_obj: LiteLLM_TeamTableCachedObj | None = None if update_key_request.team_id is not None: team_obj = await get_team_object( team_id=update_key_request.team_id, @@ -2333,7 +2327,7 @@ async def _validate_mcp_servers_for_key_update( prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, is_proxy_admin: bool, -) -> Optional[ObjectPermissionDict]: +) -> ObjectPermissionDict | None: """Validate MCP servers in object_permission against the effective team.""" effective_team_obj = team_obj # If team_id isn't being changed, resolve the existing key's team @@ -2489,7 +2483,7 @@ async def _validate_update_key_data( ) # Check team limits if key has a team_id (from request or existing key) - team_obj: Optional[LiteLLM_TeamTableCachedObj] = None + team_obj: LiteLLM_TeamTableCachedObj | None = None _team_id_to_check = data.team_id or getattr(existing_key_row, "team_id", None) if _team_id_to_check is not None: team_obj = await get_team_object( @@ -2610,7 +2604,7 @@ async def update_key_fn( request: Request, data: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -2801,12 +2795,10 @@ async def update_key_fn( return {"key": key, **response["data"]} # update based on remaining passed in values except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.update_key_fn(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.update_key_fn(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -2831,7 +2823,7 @@ async def update_key_fn( async def bulk_update_keys( data: BulkUpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -2910,8 +2902,8 @@ async def bulk_update_keys( detail={"error": f"Maximum {MAX_BATCH_SIZE} keys can be updated at once. Found {len(data.keys)} keys."}, ) - successful_updates: List[SuccessfulKeyUpdate] = [] - failed_updates: List[FailedKeyUpdate] = [] + successful_updates: list[SuccessfulKeyUpdate] = [] + failed_updates: list[FailedKeyUpdate] = [] for key_update_item in data.keys: try: @@ -2989,7 +2981,7 @@ async def bulk_update_keys( def _build_failed_team_key_update( token: str, exception: Exception, - existing_key_row: Optional[LiteLLM_VerificationToken], + existing_key_row: LiteLLM_VerificationToken | None, ) -> FailedKeyUpdate: """Normalize an exception from the per-key update loop into a FailedKeyUpdate.""" if isinstance(exception, HTTPException): @@ -3025,7 +3017,7 @@ def _build_failed_team_key_update( async def bulk_update_team_keys( data: BulkUpdateTeamKeysRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -3145,8 +3137,8 @@ async def bulk_update_team_keys( existing_by_token = {row.token: row for row in existing_keys} update_field_dict = data.update_fields.model_dump(exclude_unset=True) - successful_updates: List[SuccessfulKeyUpdate] = [] - failed_updates: List[FailedKeyUpdate] = [] + successful_updates: list[SuccessfulKeyUpdate] = [] + failed_updates: list[FailedKeyUpdate] = [] for token in requested_tokens: db_token = _hash_token_if_needed(token) @@ -3247,18 +3239,17 @@ async def validate_key_team_change( ) # Check if the person initiating the change is a Proxy Admin or Team Admin - if change_initiated_by.user_role == LitellmUserRoles.PROXY_ADMIN.value: - return - elif _is_user_team_admin( - user_api_key_dict=change_initiated_by, - team_obj=team, - ): - return - # this teams member permissions allow updating a - elif TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( - team_member_object=member_object, - team_table=cast(LiteLLM_TeamTableCachedObj, team), - route=KeyManagementRoutes.KEY_UPDATE.value, + if ( + change_initiated_by.user_role == LitellmUserRoles.PROXY_ADMIN.value + or _is_user_team_admin( + user_api_key_dict=change_initiated_by, + team_obj=team, + ) + or TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( + team_member_object=member_object, + team_table=cast(LiteLLM_TeamTableCachedObj, team), + route=KeyManagementRoutes.KEY_UPDATE.value, + ) ): return else: @@ -3273,7 +3264,7 @@ async def validate_key_team_change( async def delete_key_fn( data: KeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -3373,9 +3364,7 @@ async def delete_key_fn( return {"deleted_keys": deleted_keys} except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.delete_key_fn(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.delete_key_fn(): Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -3444,7 +3433,7 @@ async def _build_model_max_budget_usage( include_in_schema=False, ) async def info_key_fn_v2( - data: Optional[KeyRequest] = None, + data: KeyRequest | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -3529,7 +3518,7 @@ async def info_key_fn_v2( @router.get("/key/info", tags=["key management"], dependencies=[Depends(user_api_key_auth)]) @management_endpoint_wrapper async def info_key_fn( - key: Optional[str] = fastapi.Query(default=None, description="Key in the request parameters"), + key: str | None = fastapi.Query(default=None, description="Key in the request parameters"), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -3562,7 +3551,7 @@ async def info_key_fn( # default to using Auth token if no key is passed in key = key or user_api_key_dict.api_key - hashed_key: Optional[str] = key + hashed_key: str | None = key if key is not None: hashed_key = _hash_token_if_needed(token=key) key_info = await VerificationTokenRepository(prisma_client).table.find_unique( @@ -3587,9 +3576,7 @@ async def info_key_fn( ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail="You are not allowed to access this key's info. Your role={}".format( - user_api_key_dict.user_role - ), + detail=f"You are not allowed to access this key's info. Your role={user_api_key_dict.user_role}", ) ## REMOVE HASHED TOKEN INFO BEFORE RETURNING ## try: @@ -3618,9 +3605,7 @@ async def info_key_fn( raise handle_exception_on_proxy(e) -def _check_model_access_group( - models: Optional[List[str]], llm_router: Optional[Router], premium_user: bool -) -> Literal[True]: +def _check_model_access_group(models: list[str] | None, llm_router: Router | None, premium_user: bool) -> Literal[True]: """ if is_model_access_group is True + is_wildcard_route is True, check if user is a premium user @@ -3635,9 +3620,7 @@ def _check_model_access_group( raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={ - "error": "Setting a model access group on a wildcard model is only available for LiteLLM Enterprise users.{}".format( - CommonProxyErrors.not_premium_user.value - ) + "error": f"Setting a model access group on a wildcard model is only available for LiteLLM Enterprise users.{CommonProxyErrors.not_premium_user.value}" }, ) @@ -3646,63 +3629,62 @@ def _check_model_access_group( async def generate_key_helper_fn( request_type: Literal["user", "key"], # identifies if this request is from /user/new or /key/generate - duration: Optional[str] = None, + duration: str | None = None, models: list = [], aliases: dict = {}, config: dict = {}, spend: float = 0.0, - key_max_budget: Optional[float] = None, # key_max_budget is used to Budget Per key - key_budget_duration: Optional[str] = None, - budget_id: Optional[float] = None, # budget id <-> LiteLLM_BudgetTable - soft_budget: Optional[float] = None, # soft_budget is used to set soft Budgets Per user - max_budget: Optional[float] = None, # max_budget is used to Budget Per user - blocked: Optional[bool] = None, - budget_duration: Optional[str] = None, # max_budget is used to Budget Per user - token: Optional[str] = None, - key: Optional[ - str - ] = None, # dev-friendly alt param for 'token'. Exposed on `/key/generate` for setting key value yourself. - user_id: Optional[str] = None, - user_alias: Optional[str] = None, - team_id: Optional[str] = None, - agent_id: Optional[str] = None, - user_email: Optional[str] = None, - user_role: Optional[str] = None, - max_parallel_requests: Optional[int] = None, - metadata: Optional[dict] = {}, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, + key_max_budget: float | None = None, # key_max_budget is used to Budget Per key + key_budget_duration: str | None = None, + budget_id: float | None = None, # budget id <-> LiteLLM_BudgetTable + soft_budget: float | None = None, # soft_budget is used to set soft Budgets Per user + max_budget: float | None = None, # max_budget is used to Budget Per user + blocked: bool | None = None, + budget_duration: str | None = None, # max_budget is used to Budget Per user + token: str | None = None, + key: str + | None = None, # dev-friendly alt param for 'token'. Exposed on `/key/generate` for setting key value yourself. + user_id: str | None = None, + user_alias: str | None = None, + team_id: str | None = None, + agent_id: str | None = None, + user_email: str | None = None, + user_role: str | None = None, + max_parallel_requests: int | None = None, + metadata: dict | None = {}, + tpm_limit: int | None = None, + rpm_limit: int | None = None, query_type: Literal["insert_data", "update_data"] = "insert_data", - update_key_values: Optional[dict] = None, - key_alias: Optional[str] = None, - allowed_cache_controls: Optional[list] = [], - permissions: Optional[dict] = {}, - model_max_budget: Optional[dict] = {}, - budget_fallbacks: Optional[dict] = None, - model_rpm_limit: Optional[dict] = None, - model_tpm_limit: Optional[dict] = None, - mcp_rpm_limit: Optional[dict] = None, - tag_rpm_limit: Optional[dict] = None, - guardrails: Optional[list] = None, - policies: Optional[list] = None, - prompts: Optional[list] = None, - teams: Optional[list] = None, - organization_id: Optional[str] = None, - project_id: Optional[str] = None, - table_name: Optional[Literal["key", "user"]] = None, - send_invite_email: Optional[bool] = None, - created_by: Optional[str] = None, - updated_by: Optional[str] = None, - allowed_routes: Optional[list] = None, + update_key_values: dict | None = None, + key_alias: str | None = None, + allowed_cache_controls: list | None = [], + permissions: dict | None = {}, + model_max_budget: dict | None = {}, + budget_fallbacks: dict | None = None, + model_rpm_limit: dict | None = None, + model_tpm_limit: dict | None = None, + mcp_rpm_limit: dict | None = None, + tag_rpm_limit: dict | None = None, + guardrails: list | None = None, + policies: list | None = None, + prompts: list | None = None, + teams: list | None = None, + organization_id: str | None = None, + project_id: str | None = None, + table_name: Literal["key", "user"] | None = None, + send_invite_email: bool | None = None, + created_by: str | None = None, + updated_by: str | None = None, + allowed_routes: list | None = None, key_type: str | None = None, - sso_user_id: Optional[str] = None, - object_permission_id: Optional[str] = None, # object_permission_id <-> LiteLLM_ObjectPermissionTable - object_permission: Optional[LiteLLM_ObjectPermissionBase] = None, - auto_rotate: Optional[bool] = None, - rotation_interval: Optional[str] = None, - router_settings: Optional[dict] = None, - access_group_ids: Optional[list] = None, - budget_limits: Optional[list] = None, # multiple concurrent budget windows + sso_user_id: str | None = None, + object_permission_id: str | None = None, # object_permission_id <-> LiteLLM_ObjectPermissionTable + object_permission: LiteLLM_ObjectPermissionBase | None = None, + auto_rotate: bool | None = None, + rotation_interval: str | None = None, + router_settings: dict | None = None, + access_group_ids: list | None = None, + budget_limits: list | None = None, # multiple concurrent budget windows ): from litellm.proxy.proxy_server import premium_user, prisma_client @@ -3733,7 +3715,7 @@ async def generate_key_helper_fn( reset_at = get_budget_reset_time(budget_duration=budget_duration) # Initialize reset_at for each budget window - budget_limits_json: Optional[str] = None + budget_limits_json: str | None = None if budget_limits: initialized_windows = [] for window in budget_limits: @@ -3868,7 +3850,7 @@ async def generate_key_helper_fn( saved_token["permissions"] = json.loads(saved_token["permissions"]) if isinstance(saved_token["model_max_budget"], str): saved_token["model_max_budget"] = json.loads(saved_token["model_max_budget"]) - router_settings = cast(Optional[dict], saved_token.get("router_settings")) + router_settings = cast(dict | None, saved_token.get("router_settings")) if router_settings is not None and isinstance(router_settings, str): try: saved_token["router_settings"] = yaml.safe_load(router_settings) @@ -3922,9 +3904,7 @@ async def generate_key_helper_fn( # If it's not valid JSON/YAML, keep as is or set to empty dict key_data["router_settings"] = {} except Exception as e: - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.generate_key_helper_fn(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.generate_key_helper_fn(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise e @@ -4053,11 +4033,11 @@ async def can_modify_verification_token( async def delete_verification_tokens( - tokens: List, + tokens: list, user_api_key_cache: UserApiKeyCache, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, -) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]: + litellm_changed_by: str | None = None, +) -> tuple[dict | None, list[LiteLLM_VerificationToken]]: """ Helper that deletes the list of tokens from the database @@ -4078,11 +4058,11 @@ async def delete_verification_tokens( """ from litellm.proxy.proxy_server import prisma_client - failed_tokens: List = [] + failed_tokens: list = [] try: if prisma_client: tokens = [_hash_token_if_needed(token=key) for key in tokens] - _keys_being_deleted: List[LiteLLM_VerificationToken] = await VerificationTokenRepository( + _keys_being_deleted: list[LiteLLM_VerificationToken] = await VerificationTokenRepository( prisma_client ).table.find_many(where={"token": {"in": tokens}}) @@ -4131,7 +4111,7 @@ async def delete_verification_tokens( raise Exception("DB not connected. prisma_client is None") except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.delete_verification_tokens(): Exception occured - {}".format(str(e)) + f"litellm.proxy.proxy_server.delete_verification_tokens(): Exception occured - {e!s}" ) verbose_proxy_logger.debug(traceback.format_exc()) raise e @@ -4149,9 +4129,9 @@ async def delete_verification_tokens( def _transform_verification_tokens_to_deleted_records( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ) -> list[dict[str, object]]: """Transform verification tokens into deleted token records ready for persistence.""" if not keys: @@ -4215,10 +4195,10 @@ async def _save_deleted_verification_token_records( async def _persist_deleted_verification_tokens( - keys: List[LiteLLM_VerificationToken], + keys: list[LiteLLM_VerificationToken], prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ) -> None: """Persist deleted verification token records by transforming and saving them.""" records = _transform_verification_tokens_to_deleted_records( @@ -4233,12 +4213,12 @@ async def _persist_deleted_verification_tokens( async def delete_key_aliases( - key_aliases: List[str], + key_aliases: list[str], user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, -) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]: + litellm_changed_by: str | None = None, +) -> tuple[dict | None, list[LiteLLM_VerificationToken]]: _keys_being_deleted = await _prisma_table(VerificationTokenRepository(prisma_client)).find_many( where={"key_alias": {"in": key_aliases}} ) @@ -4276,7 +4256,7 @@ async def _rotate_master_key( from litellm.proxy.proxy_server import proxy_config try: - models: Optional[List] = await _prisma_table(ModelRepository(prisma_client)).find_many() + models: list | None = await _prisma_table(ModelRepository(prisma_client)).find_many() except Exception: models = None # 2. process model table @@ -4402,7 +4382,7 @@ async def _rotate_master_key( }, ) except Exception as e: - verbose_proxy_logger.error(f"Failed to re-encrypt credential {cred.credential_name}: {str(e)}") + verbose_proxy_logger.error(f"Failed to re-encrypt credential {cred.credential_name}: {e!s}") # Continue with next credential instead of failing entire rotation continue verbose_proxy_logger.debug(f"Successfully re-encrypted {len(credentials)} credentials with new master key") @@ -4488,7 +4468,7 @@ async def check_encryption_endpoint( return {"status": "success", "report": report.as_dict()} -async def get_new_token(data: Optional[RegenerateKeyRequest]) -> str: +async def get_new_token(data: RegenerateKeyRequest | None) -> str: if data and data.new_key is not None: # Reject custom key values if disabled by admin await _check_custom_key_allowed(data.new_key) @@ -4514,7 +4494,7 @@ async def _insert_deprecated_key( prisma_client: "PrismaClient", old_token_hash: str, new_token_hash: str, - grace_period: Optional[str], + grace_period: str | None, ) -> None: """ Insert old key into deprecated table so it remains valid during grace period. @@ -4577,9 +4557,9 @@ async def _execute_virtual_key_regeneration( key_in_db: LiteLLM_VerificationToken, hashed_api_key: str, key: str, - data: Optional[RegenerateKeyRequest], + data: RegenerateKeyRequest | None, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> GenerateKeyResponse: @@ -4675,14 +4655,14 @@ async def _execute_virtual_key_regeneration( ) @management_endpoint_wrapper async def regenerate_key_fn( - key: Optional[str] = None, - data: Optional[RegenerateKeyRequest] = None, + key: str | None = None, + data: RegenerateKeyRequest | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Optional[GenerateKeyResponse]: +) -> GenerateKeyResponse | None: """ Regenerate an existing API key while optionally updating its parameters. @@ -4867,7 +4847,7 @@ async def regenerate_key_fn( ) if data is not None and (data.access_group_ids or data.object_permission is not None): - regenerate_team_table: Optional[LiteLLM_TeamTableCachedObj] = None + regenerate_team_table: LiteLLM_TeamTableCachedObj | None = None if _key_in_db.team_id is not None: regenerate_team_table = await get_team_object( team_id=_key_in_db.team_id, @@ -5016,11 +4996,11 @@ async def reset_key_spend_fn( key: str, data: ResetSpendRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Dict[str, Any]: +) -> dict[str, Any]: try: from litellm.proxy.proxy_server import ( hash_token, @@ -5114,13 +5094,13 @@ async def reset_key_spend_fn( async def validate_key_list_check( user_api_key_dict: UserAPIKeyAuth, - user_id: Optional[str], - team_id: Optional[str], - organization_id: Optional[str], - key_alias: Optional[str], - key_hash: Optional[str], + user_id: str | None, + team_id: str | None, + organization_id: str | None, + key_alias: str | None, + key_hash: str | None, prisma_client: PrismaClient, -) -> Optional[LiteLLM_UserTable]: +) -> LiteLLM_UserTable | None: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return None @@ -5131,7 +5111,7 @@ async def validate_key_list_check( param="user_id", code=status.HTTP_403_FORBIDDEN, ) - complete_user_info_db_obj: Optional[BaseModel] = await _prisma_table(UserRepository(prisma_client)).find_unique( + complete_user_info_db_obj: BaseModel | None = await _prisma_table(UserRepository(prisma_client)).find_unique( where={"user_id": user_api_key_dict.user_id}, include={"organization_memberships": True}, ) @@ -5196,22 +5176,20 @@ async def validate_key_list_check( if not can_user_query_key_info: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail="You are not allowed to access this key's info. Your role={}".format( - user_api_key_dict.user_role - ), + detail=f"You are not allowed to access this key's info. Your role={user_api_key_dict.user_role}", ) return complete_user_info async def _fetch_user_team_objects( - complete_user_info: Optional[LiteLLM_UserTable], + complete_user_info: LiteLLM_UserTable | None, prisma_client: PrismaClient, -) -> List[LiteLLM_TeamTable]: +) -> list[LiteLLM_TeamTable]: """Fetch team objects for all teams a user belongs to (single DB query).""" if complete_user_info is None or not complete_user_info.teams: return [] - teams: Optional[List[BaseModel]] = await TeamRepository(prisma_client).table.find_many( + teams: list[BaseModel] | None = await TeamRepository(prisma_client).table.find_many( where={"team_id": {"in": complete_user_info.teams}} ) if teams is None: @@ -5222,8 +5200,8 @@ async def _fetch_user_team_objects( def _get_admin_team_ids_from_objects( user_api_key_dict: UserAPIKeyAuth, - team_objects: List[LiteLLM_TeamTable], -) -> List[str]: + team_objects: list[LiteLLM_TeamTable], +) -> list[str]: """Filter team objects to those where the user is an admin.""" return [ team.team_id for team in team_objects if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team) @@ -5232,8 +5210,8 @@ def _get_admin_team_ids_from_objects( def _get_team_ids_with_key_list_permission_from_objects( user_api_key_dict: UserAPIKeyAuth, - team_objects: List[LiteLLM_TeamTable], -) -> List[str]: + team_objects: list[LiteLLM_TeamTable], +) -> list[str]: """Filter team objects to non-admin teams where the caller has /key/list permission via team_member_permissions. These teams should grant the caller full key visibility (same as a team admin), so other members' @@ -5252,8 +5230,8 @@ def _get_team_ids_with_key_list_permission_from_objects( def _get_member_team_ids_from_objects( user_api_key_dict: UserAPIKeyAuth, - team_objects: List[LiteLLM_TeamTable], -) -> List[str]: + team_objects: list[LiteLLM_TeamTable], +) -> list[str]: """Filter team objects to those where the user is a member (any role).""" return [ team.team_id @@ -5266,20 +5244,20 @@ def _get_member_team_ids_from_objects( async def get_admin_team_ids( - complete_user_info: Optional[LiteLLM_UserTable], + complete_user_info: LiteLLM_UserTable | None, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, -) -> List[str]: +) -> list[str]: """Get all team IDs where the user is an admin.""" team_objects = await _fetch_user_team_objects(complete_user_info, prisma_client) return _get_admin_team_ids_from_objects(user_api_key_dict, team_objects) async def get_member_team_ids( - complete_user_info: Optional[LiteLLM_UserTable], + complete_user_info: LiteLLM_UserTable | None, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, -) -> List[str]: +) -> list[str]: """ Get all team IDs where the user is a member (any role, including admin). @@ -5304,30 +5282,30 @@ async def list_keys( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), page: int = Query(1, description="Page number", ge=1), size: int = Query(10, description="Page size", ge=1, le=100), - user_id: Optional[str] = Query( + user_id: str | None = Query( None, description="Filter keys by user ID. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.", ), - team_id: Optional[str] = Query(None, description="Filter keys by team ID"), - organization_id: Optional[str] = Query(None, description="Filter keys by organization ID"), - key_hash: Optional[str] = Query(None, description="Filter keys by key hash"), - key_alias: Optional[str] = Query( + team_id: str | None = Query(None, description="Filter keys by team ID"), + organization_id: str | None = Query(None, description="Filter keys by organization ID"), + key_hash: str | None = Query(None, description="Filter keys by key hash"), + key_alias: str | None = Query( None, description="Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.", ), return_full_object: bool = Query(False, description="Return full key object"), include_team_keys: bool = Query(False, description="Include all keys for teams that user is an admin of."), include_created_by_keys: bool = Query(False, description="Include keys created by the user"), - sort_by: Optional[str] = Query( + sort_by: str | None = Query( default=None, description="Column to sort by (e.g. 'user_id', 'created_at', 'spend')", ), sort_order: str = Query(default="desc", description="Sort order ('asc' or 'desc')"), - expand: Optional[List[str]] = Query(None, description="Expand related objects (e.g. 'user')"), - status: Optional[str] = Query(None, description="Filter by status (e.g. 'deleted')"), - project_id: Optional[str] = Query(None, description="Filter keys by project ID"), - access_group_id: Optional[str] = Query(None, description="Filter keys by access group ID"), - agent_id: Optional[str] = Query(None, description="Filter keys by agent ID"), + expand: list[str] | None = Query(None, description="Expand related objects (e.g. 'user')"), + status: str | None = Query(None, description="Filter by status (e.g. 'deleted')"), + project_id: str | None = Query(None, description="Filter keys by project ID"), + access_group_id: str | None = Query(None, description="Filter keys by access group ID"), + agent_id: str | None = Query(None, description="Filter keys by agent ID"), substring_matching: bool = Query( False, description="If true (proxy admins only), match user_id/key_alias as case-insensitive substrings instead of exact values. Defaults to false: /key/list matched these exactly before substring search was added, and an exact user_id/key_alias filter must never return another user's keys.", @@ -5468,7 +5446,7 @@ async def list_keys( verbose_proxy_logger.exception(f"Error in list_keys: {e}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"error({str(e)})"), + message=getattr(e, "detail", f"error({e!s})"), type=ProxyErrorTypes.internal_server_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -5487,17 +5465,17 @@ async def _apply_non_admin_alias_scope( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, query_params: list[object], - where_parts: List[str], + where_parts: list[str], ) -> None: """Append SQL scope conditions so non-admin users only see aliases for keys they own or keys belonging to teams they are members of.""" - scope_conditions: List[str] = [] + scope_conditions: list[str] = [] if user_api_key_dict.user_id: query_params.append(user_api_key_dict.user_id) scope_conditions.append(f"user_id = ${len(query_params)}") # Look up the user's teams from the user table - user_teams: List[str] = [] + user_teams: list[str] = [] if user_api_key_dict.user_id: user_row = await _prisma_table(UserRepository(prisma_client)).find_unique( where={"user_id": user_api_key_dict.user_id} @@ -5527,9 +5505,9 @@ async def key_aliases( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), page: int = Query(1, ge=1, description="Page number"), size: int = Query(50, ge=1, le=100, description="Page size"), - search: Optional[str] = Query(None, description="Search key aliases (case-insensitive partial match)"), - team_id: Optional[str] = Query(None, description="Filter aliases to keys belonging to this team"), -) -> Dict[str, Any]: + search: str | None = Query(None, description="Search key aliases (case-insensitive partial match)"), + team_id: str | None = Query(None, description="Filter aliases to keys belonging to this team"), +) -> dict[str, Any]: """ Lists key aliases with pagination and optional search. @@ -5600,7 +5578,7 @@ async def key_aliases( f" LIMIT ${limit_idx} OFFSET ${offset_idx}" ) alias_rows = await prisma_client.db.query_raw(aliases_sql, *aliases_params) - aliases: List[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")] + aliases: list[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")] total_pages = -(-total_count // size) if total_count > 0 else 0 verbose_proxy_logger.debug( @@ -5620,7 +5598,7 @@ async def key_aliases( verbose_proxy_logger.exception(f"Error in key_aliases: {e}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"error({str(e)})"), + message=getattr(e, "detail", f"error({e!s})"), type=ProxyErrorTypes.internal_server_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -5635,8 +5613,8 @@ async def key_aliases( ) -def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[Dict[str, str]]: - order_by: Dict[str, str] = {} +def _validate_sort_params(sort_by: str | None, sort_order: str) -> dict[str, str] | None: + order_by: dict[str, str] = {} if sort_by is None: return None @@ -5674,21 +5652,21 @@ def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str, def _build_key_filter_conditions( - user_id: Optional[str], - team_id: Optional[str], - organization_id: Optional[str], - key_alias: Optional[str], - key_hash: Optional[str], - exclude_team_id: Optional[str], - admin_team_ids: Optional[List[str]], - member_team_ids: Optional[List[str]] = None, + user_id: str | None, + team_id: str | None, + organization_id: str | None, + key_alias: str | None, + key_hash: str | None, + exclude_team_id: str | None, + admin_team_ids: list[str] | None, + member_team_ids: list[str] | None = None, include_created_by_keys: bool = False, - project_id: Optional[str] = None, - access_group_id: Optional[str] = None, - agent_id: Optional[str] = None, + project_id: str | None = None, + access_group_id: str | None = None, + agent_id: str | None = None, use_substring_matching: bool = False, expires_filter: str | None = None, -) -> Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]]: +) -> dict[str, Union[str, dict[str, Any], list[dict[str, Any]]]]: """Build filter conditions for key listing. Visibility rules: @@ -5700,14 +5678,14 @@ def _build_key_filter_conditions( so former members cannot see service accounts they created after leaving. """ # Prepare filter conditions - where: Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]] = {} + where: dict[str, Union[str, dict[str, Any], list[dict[str, Any]]]] = {} where.update(_get_condition_to_filter_out_ui_session_tokens()) # Build the OR conditions for user's keys and admin team keys - or_conditions: List[Dict[str, Any]] = [] + or_conditions: list[dict[str, Any]] = [] # Base conditions for user's own keys - user_condition: Dict[str, Any] = {} + user_condition: dict[str, Any] = {} if user_id and isinstance(user_id, str): if use_substring_matching: user_condition["user_id"] = { @@ -5806,25 +5784,24 @@ async def _list_key_helper( prisma_client: PrismaClient, page: int, size: int, - user_id: Optional[str], - team_id: Optional[str], - organization_id: Optional[str], - key_alias: Optional[str], - key_hash: Optional[str], - exclude_team_id: Optional[str] = None, + user_id: str | None, + team_id: str | None, + organization_id: str | None, + key_alias: str | None, + key_hash: str | None, + exclude_team_id: str | None = None, return_full_object: bool = False, - admin_team_ids: Optional[List[str]] = None, # New parameter for teams where user is admin - member_team_ids: Optional[ - List[str] - ] = None, # Team IDs where user is a member (any role) - for service account visibility + admin_team_ids: list[str] | None = None, # New parameter for teams where user is admin + member_team_ids: list[str] + | None = None, # Team IDs where user is a member (any role) - for service account visibility include_created_by_keys: bool = False, - sort_by: Optional[str] = None, + sort_by: str | None = None, sort_order: str = "desc", - expand: Optional[List[str]] = None, - status: Optional[str] = None, - project_id: Optional[str] = None, - access_group_id: Optional[str] = None, - agent_id: Optional[str] = None, + expand: list[str] | None = None, + status: str | None = None, + project_id: str | None = None, + access_group_id: str | None = None, + agent_id: str | None = None, use_substring_matching: bool = False, expires_filter: str | None = None, ) -> KeyListResponseObject: @@ -5872,7 +5849,7 @@ async def _list_key_helper( verbose_proxy_logger.debug(f"Pagination: skip={skip}, take={size}") - order_by: Optional[Dict[str, str]] = ( + order_by: dict[str, str] | None = ( _validate_sort_params(sort_by, sort_order) if sort_by is not None and isinstance(sort_by, str) else None ) @@ -5938,7 +5915,7 @@ async def _list_key_helper( user_map = {user.user_id: user for user in users} # Prepare response - key_list: List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]] = [] + key_list: list[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]] = [] for key in keys: # Convert Prisma model to dict (supports both Pydantic v1 and v2) try: @@ -5983,7 +5960,7 @@ async def _list_key_helper( ) -def _get_condition_to_filter_out_ui_session_tokens() -> Dict[str, Any]: +def _get_condition_to_filter_out_ui_session_tokens() -> dict[str, Any]: """ Condition to filter out UI session tokens """ @@ -6055,11 +6032,11 @@ async def block_key( data: BlockKeyRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Optional[LiteLLM_VerificationToken]: +) -> LiteLLM_VerificationToken | None: """ Block an Virtual key from making any requests. @@ -6091,7 +6068,7 @@ async def block_key( ) if prisma_client is None: - raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value)) + raise Exception(f"{CommonProxyErrors.db_not_connected_error.value}") if not is_valid_api_key(data.key): raise ProxyException( @@ -6168,7 +6145,7 @@ async def unblock_key( data: BlockKeyRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -6204,7 +6181,7 @@ async def unblock_key( ) if prisma_client is None: - raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value)) + raise Exception(f"{CommonProxyErrors.db_not_connected_error.value}") if not is_valid_api_key(data.key): raise ProxyException( @@ -6358,7 +6335,7 @@ async def key_health( except Exception as e: raise ProxyException( - message=f"Key health check failed: {str(e)}", + message=f"Key health check failed: {e!s}", type=ProxyErrorTypes.internal_server_error, param=getattr(e, "param", "None"), code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -6367,25 +6344,23 @@ async def key_health( async def _can_user_query_key_info( user_api_key_dict: UserAPIKeyAuth, - key: Optional[str], + key: str | None, key_info: LiteLLM_VerificationToken, ) -> bool: """ Helper to check if the user has access to the key's info """ if ( - user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value - or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value - ): - return True - elif user_api_key_dict.api_key == key: - return True - # user can query their own key info - elif key_info.user_id == user_api_key_dict.user_id: - return True - elif await TeamMemberPermissionChecks.user_belongs_to_keys_team( - user_api_key_dict=user_api_key_dict, - existing_key_row=key_info, + ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value + ) + or user_api_key_dict.api_key == key + or key_info.user_id == user_api_key_dict.user_id + or await TeamMemberPermissionChecks.user_belongs_to_keys_team( + user_api_key_dict=user_api_key_dict, + existing_key_row=key_info, + ) ): return True return False @@ -6394,7 +6369,7 @@ async def _can_user_query_key_info( async def test_key_logging( user_api_key_dict: UserAPIKeyAuth, request: Request, - key_logging: List[Dict[str, Any]], + key_logging: list[dict[str, Any]], ) -> LoggingCallbackStatus: """ Test the key-based logging @@ -6409,7 +6384,7 @@ async def test_key_logging( from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.proxy_server import general_settings, proxy_config - logging_callbacks: List[str] = [] + logging_callbacks: list[str] = [] for callback in key_logging: if callback.get("callback_name") is not None: logging_callbacks.append(callback["callback_name"]) @@ -6445,7 +6420,7 @@ async def test_key_logging( return LoggingCallbackStatus( callbacks=logging_callbacks, status="unhealthy", - details=f"Logging test failed: {str(e)}", + details=f"Logging test failed: {e!s}", ) await asyncio.sleep(2) # wait for callbacks to run, callbacks use batching so wait for the flush event @@ -6470,7 +6445,7 @@ async def test_key_logging( _KEY_ALIAS_PATTERN = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9_\-/\.@]{0,253}[a-zA-Z0-9]$") -def _validate_key_alias_format(key_alias: Optional[str]) -> None: +def _validate_key_alias_format(key_alias: str | None) -> None: """ Validate the format of the key_alias. @@ -6514,9 +6489,9 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None: async def _enforce_unique_key_alias( - key_alias: Optional[str], + key_alias: str | None, prisma_client: PrismaClient | None, - existing_key_token: Optional[str] = None, + existing_key_token: str | None = None, ) -> None: """ Helper to enforce unique key aliases across all keys. @@ -6546,7 +6521,7 @@ async def _enforce_unique_key_alias( ) -def validate_model_max_budget(model_max_budget: Optional[Dict]) -> None: +def validate_model_max_budget(model_max_budget: dict | None) -> None: """ Validate the model_max_budget is GenericBudgetConfigType + enforce user has an enterprise license @@ -6576,5 +6551,5 @@ def validate_model_max_budget(model_max_budget: Optional[Dict]) -> None: BudgetConfig(**_info) except Exception as e: raise ValueError( - f"Invalid model_max_budget: {str(e)}. Example of valid model_max_budget: https://docs.litellm.ai/docs/proxy/users" + f"Invalid model_max_budget: {e!s}. Example of valid model_max_budget: https://docs.litellm.ai/docs/proxy/users" ) diff --git a/litellm/proxy/management_endpoints/management_v1/budgets.py b/litellm/proxy/management_endpoints/management_v1/budgets.py index bc1521caf0b..8ecd7b1fa30 100644 --- a/litellm/proxy/management_endpoints/management_v1/budgets.py +++ b/litellm/proxy/management_endpoints/management_v1/budgets.py @@ -108,7 +108,7 @@ def _serialize(row: BudgetListItem) -> BudgetListItem: def _scope(caller: UserAPIKeyAuth) -> Scope: if user_api_key_has_admin_view(caller): return ScopeAll() - return ScopeDenied(reason="Only proxy admins can list budgets, your role={}".format(caller.user_role)) + return ScopeDenied(reason=f"Only proxy admins can list budgets, your role={caller.user_role}") # budget_duration is deliberately absent from `sortable`: the column holds strings @@ -191,9 +191,7 @@ async def list_budgets( raise except Exception as e: # noqa: BLE001 # a driver error answers as a problem document, not the OpenAI error shape verbose_proxy_logger.exception( - "litellm.proxy.management_endpoints.management_v1.budgets.list_budgets(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.management_endpoints.management_v1.budgets.list_budgets(): Exception occured - {e!s}" ) raise ManagementProblem( ProblemDetail( diff --git a/litellm/proxy/management_endpoints/management_v1/list_framework.py b/litellm/proxy/management_endpoints/management_v1/list_framework.py index 8f25c45016b..353064d4144 100644 --- a/litellm/proxy/management_endpoints/management_v1/list_framework.py +++ b/litellm/proxy/management_endpoints/management_v1/list_framework.py @@ -359,10 +359,7 @@ def _parse_sort(spec: ListSpec[TRow, TOut], params: Mapping[str, str]) -> tuple[ if raw is None: return spec.default_sort segments = tuple(segment.strip() for segment in raw.split(",")) - keys = tuple( - SortKey(field=segment[1:] if segment.startswith("-") else segment, descending=segment.startswith("-")) - for segment in segments - ) + keys = tuple(SortKey(field=segment.removeprefix("-"), descending=segment.startswith("-")) for segment in segments) rejected = tuple(sorted(frozenset(key.field for key in keys) - spec.sortable)) if rejected: return _problem( diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index ccde3c4112c..1927e94d01b 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -188,7 +188,7 @@ async def list_spend_log_end_users( except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.management_endpoints.management_v1.spend_logs.list_spend_log_end_users(): " - "Exception occured - {}".format(str(e)) + f"Exception occured - {e!s}" ) raise ManagementProblem( ProblemDetail( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index dcc4dc36d83..2ae9da576b2 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -399,7 +399,7 @@ if MCP_AVAILABLE: try: encrypted_payload = encrypt_value_helper(payload_json) except Exception as e: - verbose_proxy_logger.debug(f"Failed to encrypt temporary MCP server payload for Redis cache: {str(e)}") + verbose_proxy_logger.debug(f"Failed to encrypt temporary MCP server payload for Redis cache: {e!s}") return if not isinstance(encrypted_payload, str): @@ -413,7 +413,7 @@ if MCP_AVAILABLE: ttl=max(1, ttl_seconds), ) except Exception as e: - verbose_proxy_logger.debug(f"Failed to write temporary MCP server to Redis cache: {str(e)}") + verbose_proxy_logger.debug(f"Failed to write temporary MCP server to Redis cache: {e!s}") async def _get_temporary_mcp_server_from_redis( server_id: str, @@ -435,7 +435,7 @@ if MCP_AVAILABLE: key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}" ) except Exception as e: - verbose_proxy_logger.debug(f"Failed reading temporary MCP server from Redis cache: {str(e)}") + verbose_proxy_logger.debug(f"Failed reading temporary MCP server from Redis cache: {e!s}") return None if not isinstance(cached_server, str): @@ -454,7 +454,7 @@ if MCP_AVAILABLE: try: loaded = json.loads(decrypted_json) except Exception as e: - verbose_proxy_logger.debug(f"Invalid decrypted temporary MCP payload in Redis cache: {str(e)}") + verbose_proxy_logger.debug(f"Invalid decrypted temporary MCP payload in Redis cache: {e!s}") return None if not isinstance(loaded, dict): return None @@ -463,7 +463,7 @@ if MCP_AVAILABLE: try: return MCPServer.model_validate(payload_dict) except Exception as e: - verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}") + verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {e!s}") return None async def get_cached_temporary_mcp_server( @@ -1183,10 +1183,10 @@ if MCP_AVAILABLE: touched_by=user_api_key_dict.user_id or user_api_key_dict.team_id, ) except Exception as e: - verbose_proxy_logger.exception(f"Error registering mcp server: {str(e)}") + verbose_proxy_logger.exception(f"Error registering mcp server: {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Error registering mcp server: {str(e)}"}, + detail={"error": f"Error registering mcp server: {e!s}"}, ) # Do NOT add to runtime registry — pending servers are not active return _redact_mcp_credentials(new_mcp_server) @@ -1483,10 +1483,10 @@ if MCP_AVAILABLE: touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, ) except Exception as e: - verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}") + verbose_proxy_logger.exception(f"Error creating mcp server: {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Error creating mcp server: {str(e)}"}, + detail={"error": f"Error creating mcp server: {e!s}"}, ) # Registry refresh is best-effort: the row is already committed, so a @@ -1498,7 +1498,7 @@ if MCP_AVAILABLE: await global_mcp_server_manager.reload_servers_from_database() except Exception as e: verbose_proxy_logger.exception( - f"MCP server {new_mcp_server.server_id} created but in-memory registry refresh failed: {str(e)}" + f"MCP server {new_mcp_server.server_id} created but in-memory registry refresh failed: {e!s}" ) return _redact_mcp_credentials(new_mcp_server) @@ -1559,10 +1559,10 @@ if MCP_AVAILABLE: ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS, ) except Exception as e: - verbose_proxy_logger.exception(f"Error caching temporary mcp server: {str(e)}") + verbose_proxy_logger.exception(f"Error caching temporary mcp server: {e!s}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Error caching temporary mcp server: {str(e)}"}, + detail={"error": f"Error caching temporary mcp server: {e!s}"}, ) return _redact_mcp_credentials(temp_record) @@ -2539,9 +2539,7 @@ if MCP_AVAILABLE: raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can update public mcp servers. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can update public mcp servers. Your role={user_api_key_dict.user_role}" }, ) @@ -2625,9 +2623,7 @@ if MCP_AVAILABLE: raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can access MCP discovery. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can access MCP discovery. Your role={user_api_key_dict.user_role}" }, ) @@ -2683,9 +2679,7 @@ if MCP_AVAILABLE: raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can access the OpenAPI registry. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can access the OpenAPI registry. Your role={user_api_key_dict.user_role}" }, ) try: diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index 5e3ff8eb7f8..2ac0b32ec13 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -7,7 +7,7 @@ Endpoints here: import json from collections.abc import Mapping, Sequence -from typing import Any, Dict, List, Tuple +from typing import Any from fastapi import APIRouter, Depends, HTTPException @@ -17,10 +17,10 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth # Clear cache and reload models to pick up the access group changes from litellm.proxy.management_endpoints.model_management_endpoints import ( + clear_cache, live_model_ids_snapshot, model_info_as_mapping, reload_serving_verdict, - clear_cache, ) from litellm.proxy.utils import PrismaClient from litellm.repositories.model_repository import ModelRepository @@ -36,7 +36,7 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import router = APIRouter() -def validate_models_exist(model_names: List[str], llm_router) -> Tuple[bool, List[str]]: +def validate_models_exist(model_names: list[str], llm_router) -> tuple[bool, list[str]]: """ Validate that all requested model names exist in the router. Checks only exact model name matches. @@ -52,7 +52,7 @@ def validate_models_exist(model_names: List[str], llm_router) -> Tuple[bool, Lis return (len(missing) == 0, missing) -def add_access_group_to_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]: +def add_access_group_to_deployment(model_info: dict[str, Any], access_group: str) -> tuple[dict[str, Any], bool]: """ Add an access group to a deployment's model_info. @@ -158,7 +158,7 @@ async def _strip_access_group_from_deployment( async def update_deployments_with_access_group( - model_names: List[str], + model_names: list[str], access_group: str, prisma_client: PrismaClient, ) -> tuple[tuple[str, Mapping[str, object]], ...]: @@ -200,7 +200,7 @@ async def update_deployments_with_access_group( async def update_specific_deployments_with_access_group( - model_ids: List[str], + model_ids: list[str], access_group: str, prisma_client: PrismaClient, ) -> tuple[tuple[str, Mapping[str, object]], ...]: @@ -235,7 +235,7 @@ async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> return deployment.model_info -def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]: +def remove_access_group_from_deployment(model_info: dict[str, Any], access_group: str) -> tuple[dict[str, Any], bool]: """ Remove an access group from a deployment's model_info. @@ -261,7 +261,7 @@ def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group async def get_all_access_groups_from_db( prisma_client: PrismaClient, -) -> Dict[str, AccessGroupInfo]: +) -> dict[str, AccessGroupInfo]: """ Get all access groups from the database. @@ -272,7 +272,7 @@ async def get_all_access_groups_from_db( deployments = await ModelRepository(prisma_client).table.find_many() # Build access group map - access_group_map: Dict[str, Dict[str, Any]] = {} + access_group_map: dict[str, dict[str, Any]] = {} for deployment in deployments: model_info = deployment.model_info or {} @@ -439,10 +439,10 @@ async def create_model_group( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error creating access group '{data.access_group}': {str(e)}") + verbose_proxy_logger.exception(f"Error creating access group '{data.access_group}': {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to create access group: {str(e)}"}, + detail={"error": f"Failed to create access group: {e!s}"}, ) @@ -489,10 +489,10 @@ async def list_access_groups( return ListAccessGroupsResponse(access_groups=access_groups_list) except Exception as e: - verbose_proxy_logger.exception(f"Error listing access groups: {str(e)}") + verbose_proxy_logger.exception(f"Error listing access groups: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to list access groups: {str(e)}"}, + detail={"error": f"Failed to list access groups: {e!s}"}, ) @@ -546,10 +546,10 @@ async def get_access_group_info( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error getting access group info for '{access_group}': {str(e)}") + verbose_proxy_logger.exception(f"Error getting access group info for '{access_group}': {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to get access group info: {str(e)}"}, + detail={"error": f"Failed to get access group info: {e!s}"}, ) @@ -627,7 +627,7 @@ async def update_access_group( except Exception as e: raise HTTPException( status_code=500, - detail={"error": f"Failed to check access group existence: {str(e)}"}, + detail={"error": f"Failed to check access group existence: {e!s}"}, ) # Validation: Check if all new models exist (only if using model_names path) @@ -699,10 +699,10 @@ async def update_access_group( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error updating access group '{access_group}': {str(e)}") + verbose_proxy_logger.exception(f"Error updating access group '{access_group}': {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to update access group: {str(e)}"}, + detail={"error": f"Failed to update access group: {e!s}"}, ) @@ -759,7 +759,7 @@ async def delete_access_group( except Exception as e: raise HTTPException( status_code=500, - detail={"error": f"Failed to check access group existence: {str(e)}"}, + detail={"error": f"Failed to check access group existence: {e!s}"}, ) try: @@ -800,8 +800,8 @@ async def delete_access_group( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error deleting access group '{access_group}': {str(e)}") + verbose_proxy_logger.exception(f"Error deleting access group '{access_group}': {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to delete access group: {str(e)}"}, + detail={"error": f"Failed to delete access group: {e!s}"}, ) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 4885ec42578..1cb8bcafc33 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -14,7 +14,7 @@ import asyncio import datetime import json from collections.abc import Mapping, Sequence -from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast +from typing import Any, Literal, cast from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from pydantic import BaseModel, ConfigDict, Field @@ -86,14 +86,14 @@ async def update_team(*args, **kwargs): class UpdatePublicModelGroupsRequest(BaseModel): """Request model for updating public model groups""" - model_groups: List[str] = Field(description="List of model group names to make public") + model_groups: list[str] = Field(description="List of model group names to make public") model_config = ConfigDict(extra="forbid") -async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[Deployment]: +async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Deployment | None: db_model = cast( - Optional[BaseModel], + BaseModel | None, await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id}), ) @@ -351,13 +351,13 @@ async def patch_model( return updated_model except Exception as e: - verbose_proxy_logger.exception(f"Error in patch_model: {str(e)}") + verbose_proxy_logger.exception(f"Error in patch_model: {e!s}") if isinstance(e, (HTTPException, ProxyException)): raise e raise ProxyException( - message=f"Error updating model: {str(e)}", + message=f"Error updating model: {e!s}", type=ProxyErrorTypes.internal_server_error, code=status.HTTP_500_INTERNAL_SERVER_ERROR, param=None, @@ -369,8 +369,8 @@ async def _set_model_blocked_status( user_api_key_dict: UserAPIKeyAuth, blocked: bool, action: Literal["blocked", "unblocked"], - litellm_changed_by: Optional[str], -) -> Optional[LiteLLM_ProxyModelTable]: + litellm_changed_by: str | None, +) -> LiteLLM_ProxyModelTable | None: from litellm.proxy.proxy_server import ( litellm_proxy_admin_name, llm_router, @@ -458,13 +458,13 @@ async def _set_model_blocked_status( return updated_model except Exception as e: - verbose_proxy_logger.exception(f"Error in model {action}: {str(e)}") + verbose_proxy_logger.exception(f"Error in model {action}: {e!s}") if isinstance(e, (HTTPException, ProxyException)): raise e raise ProxyException( - message=f"Error updating model blocked status: {str(e)}", + message=f"Error updating model blocked status: {e!s}", type=ProxyErrorTypes.internal_server_error, code=status.HTTP_500_INTERNAL_SERVER_ERROR, param=None, @@ -480,11 +480,11 @@ async def block_model( data: BlockModelRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Optional[LiteLLM_ProxyModelTable]: +) -> LiteLLM_ProxyModelTable | None: """ Block a DB-stored model deployment from serving requests. @@ -509,11 +509,11 @@ async def unblock_model( data: BlockModelRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Optional[LiteLLM_ProxyModelTable]: +) -> LiteLLM_ProxyModelTable | None: """ Unblock a DB-stored model deployment so it can serve requests again. @@ -539,9 +539,9 @@ async def _add_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, - new_encryption_key: Optional[str] = None, + new_encryption_key: str | None = None, should_create_model_in_db: bool = True, -) -> Optional[LiteLLM_ProxyModelTable]: +) -> LiteLLM_ProxyModelTable | None: # encrypt litellm params # _litellm_params_dict = model_params.litellm_params.dict(exclude_none=True) _original_litellm_model_name = model_params.litellm_params.model @@ -573,7 +573,7 @@ async def _add_team_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, -) -> Optional[LiteLLM_ProxyModelTable]: +) -> LiteLLM_ProxyModelTable | None: """ If 'team_id' is provided, @@ -717,7 +717,7 @@ def _get_public_model_name( db_model.model_info.team_id if db_model.model_info else None ) - def _is_internal_shape(name: Optional[str]) -> bool: + def _is_internal_shape(name: str | None) -> bool: if team_id is None or not name: return False return name.startswith(f"model_name_{team_id}_") @@ -753,8 +753,8 @@ async def _setup_new_team_model_assignment( async def _get_team_deployments( - team_id: str, prisma_client: PrismaClient, table: Optional[Any] = None -) -> List[LiteLLM_ProxyModelTable]: + team_id: str, prisma_client: PrismaClient, table: Any | None = None +) -> list[LiteLLM_ProxyModelTable]: """ Fetch all deployments for a given team_id from the database. @@ -788,10 +788,10 @@ async def _get_team_deployments( async def delete_team_models( - team_ids: List[str], + team_ids: list[str], prisma_client: PrismaClient, - llm_router: Optional[Any], -) -> List[str]: + llm_router: Any | None, +) -> list[str]: """ Delete every BYOK model owned by the given teams, from the DB and the router. @@ -803,7 +803,7 @@ async def delete_team_models( Returns the model_ids that were deleted. """ - deleted_model_ids: List[str] = [] + deleted_model_ids: list[str] = [] async with prisma_client.db.tx() as tx: for team_id in team_ids: rows = await _get_team_deployments(team_id, prisma_client, table=tx.litellm_proxymodeltable) @@ -822,7 +822,7 @@ async def delete_team_models( async def _get_team_public_model_names( team_id: str, prisma_client: PrismaClient, -) -> Set[str]: +) -> set[str]: """ Public model names currently backed by a deployment in the team. @@ -831,7 +831,7 @@ async def _get_team_public_model_names( still serves it. """ deployments = await _get_team_deployments(team_id, prisma_client) - public_names: Set[str] = set() + public_names: set[str] = set() for row in deployments: model_info = model_info_as_mapping(row.model_info) if model_info is not None: @@ -874,7 +874,7 @@ async def _remove_unbacked_team_models( deleted_name_still_served = ( llm_router is not None and model_params.model_name in llm_router.model_name_to_deployment_indices ) - removed_model_aliases: List[Tuple[str, str]] = ( + removed_model_aliases: list[tuple[str, str]] = ( [] if deleted_name_still_served else await delete_team_model_alias( @@ -923,7 +923,7 @@ async def _update_existing_team_model_assignment( db_model: Deployment, patch_data: updateDeployment, user_api_key_dict: UserAPIKeyAuth, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, ) -> None: """Update an existing team model if the public name changed. @@ -934,8 +934,8 @@ async def _update_existing_team_model_assignment( """ def _get_team_public_model_name( - model_info: Optional[Union[dict, str]], - ) -> Optional[str]: + model_info: dict | str | None, + ) -> str | None: parsed = model_info_as_mapping(model_info) if parsed is None: return None @@ -1009,7 +1009,7 @@ class ModelManagementAuthChecks: def can_user_make_team_model_call( team_id: str, user_api_key_dict: UserAPIKeyAuth, - team_obj: Optional[LiteLLM_TeamTable] = None, + team_obj: LiteLLM_TeamTable | None = None, premium_user: bool = False, ) -> Literal[True]: if premium_user is False: @@ -1023,16 +1023,14 @@ class ModelManagementAuthChecks: raise HTTPException( status_code=403, detail={ - "error": "Team ID={} does not match the API key's team ID={}, OR you are not the admin for this team. Check `/user/info` to verify your team admin status.".format( - team_id, user_api_key_dict.team_id - ) + "error": f"Team ID={team_id} does not match the API key's team ID={user_api_key_dict.team_id}, OR you are not the admin for this team. Check `/user/info` to verify your team admin status." }, ) return True @staticmethod async def allow_team_model_action( - model_params: Union[Deployment, updateDeployment], + model_params: Deployment | updateDeployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, premium_user: bool, @@ -1052,7 +1050,7 @@ class ModelManagementAuthChecks: if _existing_team_row is None: raise HTTPException( status_code=400, - detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, + detail={"error": f"Team id={model_params.model_info.team_id} does not exist in db"}, ) existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) @@ -1090,7 +1088,7 @@ class ModelManagementAuthChecks: ) raise HTTPException( status_code=400, - detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, + detail={"error": f"Team id={model_params.model_info.team_id} does not exist in db"}, ) team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump()) @@ -1105,9 +1103,7 @@ class ModelManagementAuthChecks: raise HTTPException( status_code=403, detail={ - "error": "User does not have permission to make this model call. Your role={}. You can only make model calls if you are a PROXY_ADMIN or if you are a team admin, by specifying a team_id in the model_info.".format( - user_api_key_dict.user_role - ) + "error": f"User does not have permission to make this model call. Your role={user_api_key_dict.user_role}. You can only make model calls if you are a PROXY_ADMIN or if you are a team admin, by specifying a team_id in the model_info." }, ) else: @@ -1220,10 +1216,10 @@ async def delete_model( ) except Exception as e: - verbose_proxy_logger.exception(f"Failed to delete model. Due to error - {str(e)}") + verbose_proxy_logger.exception(f"Failed to delete model. Due to error - {e!s}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -1241,7 +1237,7 @@ async def delete_model( async def delete_team_model_alias( public_model_name: str, prisma_client: PrismaClient, -) -> List[Tuple[str, str]]: +) -> list[tuple[str, str]]: """ Delete a team model alias @@ -1349,7 +1345,7 @@ async def add_new_model( existing_params=None, ) - model_response: Optional[LiteLLM_ProxyModelTable] = None + model_response: LiteLLM_ProxyModelTable | None = None # update DB if store_model_in_db is True: """ @@ -1426,12 +1422,10 @@ async def add_new_model( return model_response except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.add_new_model(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.add_new_model(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -1581,12 +1575,10 @@ async def update_model( return model_response except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.update_model(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.update_model(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -1637,9 +1629,7 @@ async def update_public_model_groups( raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can update public model groups. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can update public model groups. Your role={user_api_key_dict.user_role}" }, ) @@ -1678,13 +1668,13 @@ async def update_public_model_groups( } except Exception as e: - verbose_proxy_logger.exception(f"Error updating public model groups: {str(e)}") + verbose_proxy_logger.exception(f"Error updating public model groups: {e!s}") if isinstance(e, HTTPException): raise e raise ProxyException( - message=f"Error updating public model groups: {str(e)}", + message=f"Error updating public model groups: {e!s}", type=ProxyErrorTypes.internal_server_error, code=status.HTTP_500_INTERNAL_SERVER_ERROR, param=None, @@ -1714,9 +1704,7 @@ async def update_useful_links( raise HTTPException( status_code=403, detail={ - "error": "Only proxy admins can update public model groups. Your role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only proxy admins can update public model groups. Your role={user_api_key_dict.user_role}" }, ) @@ -1748,20 +1736,20 @@ async def update_useful_links( } except Exception as e: - verbose_proxy_logger.exception(f"Error updating public model groups: {str(e)}") + verbose_proxy_logger.exception(f"Error updating public model groups: {e!s}") if isinstance(e, HTTPException): raise e raise ProxyException( - message=f"Error updating public model groups: {str(e)}", + message=f"Error updating public model groups: {e!s}", type=ProxyErrorTypes.internal_server_error, code=status.HTTP_500_INTERNAL_SERVER_ERROR, param=None, ) -def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]: +def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: """ Deduplicate models based on their model_info.id field. Returns a list of unique models keeping only the first occurrence of each model ID. @@ -1975,5 +1963,5 @@ async def clear_cache() -> frozenset[str] | None: ) return still_desired_ids except Exception as e: - verbose_proxy_logger.exception(f"Failed to clear cache and reload models. Due to error - {str(e)}") + verbose_proxy_logger.exception(f"Failed to clear cache and reload models. Due to error - {e!s}") return None diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index e6b2124a594..3ec728c7f79 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -559,7 +559,7 @@ async def get_organization_daily_activity( if org_id not in admin_org_ids: raise HTTPException( status_code=403, - detail={"error": "User is not org_admin for Organization= {}.".format(org_id)}, + detail={"error": f"User is not org_admin for Organization= {org_id}."}, ) # Fetch organization aliases for metadata @@ -1261,7 +1261,7 @@ async def organization_member_add( verbose_proxy_logger.exception(f"Error adding member to organization: {e}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), diff --git a/litellm/proxy/management_endpoints/policy_endpoints/ai_policy_suggester.py b/litellm/proxy/management_endpoints/policy_endpoints/ai_policy_suggester.py index 00ca52bb081..99fef328b1c 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/ai_policy_suggester.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/ai_policy_suggester.py @@ -4,7 +4,6 @@ based on user-provided attack examples and descriptions. """ import json -from typing import List, Optional import litellm from litellm._logging import verbose_proxy_logger @@ -53,9 +52,9 @@ class AiPolicySuggester: async def suggest( self, templates: list, - attack_examples: List[str], + attack_examples: list[str], description: str, - model: Optional[str] = None, + model: str | None = None, ) -> dict: system_prompt = self._build_system_prompt(templates) user_prompt = self._build_user_prompt(attack_examples, description) @@ -117,7 +116,7 @@ class AiPolicySuggester: "Available templates:\n\n" + "\n\n".join(template_descriptions) ) - def _build_user_prompt(self, attack_examples: List[str], description: str) -> str: + def _build_user_prompt(self, attack_examples: list[str], description: str) -> str: parts = [] filtered_examples = [e for e in attack_examples if e.strip()] if filtered_examples: diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index 7d9a57fa02d..fb139902978 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -13,7 +13,7 @@ import copy import json import os from collections.abc import AsyncIterator -from typing import TYPE_CHECKING, Any, List, Literal, Optional, cast +from typing import TYPE_CHECKING, Any, Literal, cast from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse @@ -84,7 +84,7 @@ class _ApplyPoliciesResultBase(TypedDict): """Base result of apply_policies: inputs plus any guardrail failures.""" inputs: GenericGuardrailAPIInputs - guardrail_errors: List[GuardrailErrorEntry] + guardrail_errors: list[GuardrailErrorEntry] class ApplyPoliciesResult(_ApplyPoliciesResultBase, total=False): @@ -97,7 +97,7 @@ class _ApplyPoliciesPerItemResultBase(TypedDict): """Base result for one input when using inputs_list.""" inputs: GenericGuardrailAPIInputs - guardrail_errors: List[GuardrailErrorEntry] + guardrail_errors: list[GuardrailErrorEntry] class ApplyPoliciesPerItemResult(_ApplyPoliciesPerItemResultBase, total=False): @@ -109,16 +109,16 @@ class ApplyPoliciesPerItemResult(_ApplyPoliciesPerItemResultBase, total=False): class ApplyPoliciesListResult(TypedDict): """Result when using inputs_list: one result per input.""" - results: List[ApplyPoliciesPerItemResult] + results: list[ApplyPoliciesPerItemResult] async def apply_policies( - policy_names: Optional[list[str]], + policy_names: list[str] | None, inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: Literal["request", "response"], proxy_logging_obj: "LiteLLMLoggingObj", - guardrail_names: Optional[list[str]] = None, + guardrail_names: list[str] | None = None, ) -> ApplyPoliciesResult: """ Apply guardrails to inputs from policy names and/or a direct list of guardrail names. @@ -135,7 +135,7 @@ async def apply_policies( ApplyPoliciesResult with "inputs" (final GenericGuardrailAPIInputs) and "guardrail_errors" (list of {"guardrail_name", "message"} for each failure). """ - guardrail_errors: List[GuardrailErrorEntry] = [] + guardrail_errors: list[GuardrailErrorEntry] = [] guardrail_name_set: set[str] = set() @@ -208,7 +208,7 @@ async def apply_policies( def _chat_body_from_inputs(inputs: GenericGuardrailAPIInputs, agent_id: str, request_data: dict) -> dict: """Build a chat completion request body from guardrail inputs and agent_id.""" - messages: List[dict] + messages: list[dict] structured = inputs.get("structured_messages") texts = inputs.get("texts") if structured: @@ -229,7 +229,7 @@ def _chat_body_from_inputs(inputs: GenericGuardrailAPIInputs, agent_id: str, req def _request_with_json_body(body: dict) -> Request: """Create a Starlette Request that will return the given dict as parsed JSON body.""" body_bytes = json.dumps(body).encode() - received: List[bool] = [False] + received: list[bool] = [False] async def receive() -> dict: if received[0]: @@ -256,9 +256,9 @@ def _request_with_json_body(body: dict) -> Request: class TestPoliciesAndGuardrailsRequest(BaseModel): """Request body for POST /utils/test_policies_and_guardrails.""" - policy_names: Optional[List[str]] = Field(default=None, description="Policy names to resolve guardrails from") - guardrail_names: Optional[List[str]] = Field(default=None, description="Guardrail names to apply directly") - inputs_list: List[GenericGuardrailAPIInputs] = Field( + policy_names: list[str] | None = Field(default=None, description="Policy names to resolve guardrails from") + guardrail_names: list[str] | None = Field(default=None, description="Guardrail names to apply directly") + inputs_list: list[GenericGuardrailAPIInputs] = Field( default=[], description="List of GenericGuardrailAPIInputs; each item processed separately (for batch compliance testing).", ) @@ -266,7 +266,7 @@ class TestPoliciesAndGuardrailsRequest(BaseModel): input_type: Literal["request", "response"] = Field( default="request", description="Whether inputs are request or response" ) - agent_id: Optional[str] = Field( + agent_id: str | None = Field( default=None, description="When set, call chat completion with this model/agent for each input and include the response in the result.", ) @@ -321,7 +321,7 @@ async def test_policies_and_guardrails( try: logging_obj = cast(LiteLLMLoggingObj, proxy_logging_obj) - results: List[ApplyPoliciesPerItemResult] = [] + results: list[ApplyPoliciesPerItemResult] = [] for inp in data.inputs_list: item_result = await apply_policies( policy_names=data.policy_names, @@ -662,13 +662,13 @@ async def get_policy_templates( class EnrichTemplateRequest(BaseModel): template_id: str parameters: dict - model: Optional[str] = None - competitors: Optional[List[str]] = Field( + model: str | None = None + competitors: list[str] | None = Field( default=None, max_length=MAX_COMPETITOR_NAMES, description="Optional list of competitor names", ) - instruction: Optional[str] = Field( + instruction: str | None = Field( default=None, description="Refinement instruction for modifying the competitor list (e.g. 'add 10 more from Asia')", ) @@ -769,7 +769,7 @@ async def _stream_llm_competitor_names( prompt: str, model: str, existing: list[str], -) -> AsyncIterator[tuple[Optional[str], bool]]: +) -> AsyncIterator[tuple[str | None, bool]]: """ Stream competitor names from LLM. Yields (name, is_error) tuples. @@ -889,7 +889,7 @@ async def enrich_policy_template_stream( ) -def _clean_competitor_line(line: str) -> Optional[str]: +def _clean_competitor_line(line: str) -> str | None: """Strip numbering, bullets, and whitespace from a competitor name line.""" name = line.strip().strip(".-) ").strip() return name if name and len(name) > 1 else None @@ -982,7 +982,7 @@ def _build_competitor_guardrail_definitions( definitions: list, competitors: list, brand_name: str, - variations_map: Optional[dict] = None, + variations_map: dict | None = None, ) -> list: """Build enriched guardrailDefinitions with competitor names and variations populated.""" variations_map = variations_map or {} @@ -1076,9 +1076,9 @@ def _build_comparison_blocked_words( class SuggestTemplatesRequest(BaseModel): - attack_examples: List[str] = Field(default_factory=list) + attack_examples: list[str] = Field(default_factory=list) description: str = Field(default="") - model: Optional[str] = None + model: str | None = None @router.post( @@ -1119,13 +1119,13 @@ class GuardrailTestResultEntry(TypedDict): class TestPolicyTemplateRequest(BaseModel): - guardrail_definitions: List[dict] = Field(description="All guardrailDefinitions from the policy template") + guardrail_definitions: list[dict] = Field(description="All guardrailDefinitions from the policy template") text: str = Field(description="Test input text to run guardrails against") class TestPolicyTemplateResponse(TypedDict): overall_action: str # worst-case across all guardrails - results: List[GuardrailTestResultEntry] + results: list[GuardrailTestResultEntry] @router.post( @@ -1163,15 +1163,15 @@ async def test_policy_template( async def _test_guardrail_definitions( - guardrail_definitions: List[dict], + guardrail_definitions: list[dict], text: str, -) -> List[GuardrailTestResultEntry]: +) -> list[GuardrailTestResultEntry]: """Instantiate and run each guardrail definition against the text.""" from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( ContentFilterGuardrail, ) - results: List[GuardrailTestResultEntry] = [] + results: list[GuardrailTestResultEntry] = [] for guardrail_def in guardrail_definitions: guardrail_name = guardrail_def.get("guardrail_name", "unknown") @@ -1246,7 +1246,7 @@ async def _test_guardrail_definitions( return results -def _compute_overall_action(results: List[GuardrailTestResultEntry]) -> str: +def _compute_overall_action(results: list[GuardrailTestResultEntry]) -> str: """Return the worst-case action: blocked > masked > error > unsupported > passed.""" priority = {"blocked": 4, "masked": 3, "error": 2, "unsupported": 1, "passed": 0} worst = "passed" diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index d46ad41eee3..061b820093c 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -8,7 +8,7 @@ GET /router/fields - Get router settings field definitions without values (for U """ import inspect -from typing import Any, Dict, List, get_args +from typing import Any, get_args from fastapi import APIRouter, Depends from pydantic import BaseModel, Field @@ -27,19 +27,19 @@ router = APIRouter() class RouterSettingsResponse(BaseModel): - fields: List[RouterSettingsField] = Field(description="List of all configurable router settings with metadata") - current_values: Dict[str, Any] = Field(description="Current values of router settings") - routing_strategy_descriptions: Dict[str, str] = Field(description="Descriptions for each routing strategy option") + fields: list[RouterSettingsField] = Field(description="List of all configurable router settings with metadata") + current_values: dict[str, Any] = Field(description="Current values of router settings") + routing_strategy_descriptions: dict[str, str] = Field(description="Descriptions for each routing strategy option") class RouterFieldsResponse(BaseModel): - fields: List[RouterSettingsField] = Field( + fields: list[RouterSettingsField] = Field( description="List of all configurable router settings with metadata (without field values)" ) - routing_strategy_descriptions: Dict[str, str] = Field(description="Descriptions for each routing strategy option") + routing_strategy_descriptions: dict[str, str] = Field(description="Descriptions for each routing strategy option") -def _get_routing_strategies_from_router_class() -> List[str]: +def _get_routing_strategies_from_router_class() -> list[str]: """ Dynamically extract routing strategies from the Router class __init__ method. """ @@ -94,7 +94,7 @@ async def get_router_settings( config = await proxy_config.get_config() router_settings_from_config = config.get("router_settings", {}) - current_values: Dict[str, Any] = {} + current_values: dict[str, Any] = {} if llm_router is not None: # Router exposes routing groups as private `_routing_groups`; the # generic `hasattr` loop below would miss them. @@ -120,7 +120,7 @@ async def get_router_settings( routing_strategy_descriptions=ROUTING_STRATEGY_DESCRIPTIONS, ) except Exception as e: - verbose_proxy_logger.error(f"Error fetching router settings: {str(e)}") + verbose_proxy_logger.error(f"Error fetching router settings: {e!s}") raise @@ -168,5 +168,5 @@ async def get_router_fields( routing_strategy_descriptions=ROUTING_STRATEGY_DESCRIPTIONS, ) except Exception as e: - verbose_proxy_logger.error(f"Error fetching router fields: {str(e)}") + verbose_proxy_logger.error(f"Error fetching router fields: {e!s}") raise diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 531716bb1c4..a559e5a845d 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import List, TypeVar, Union +from typing import TypeVar from pydantic import ValidationError @@ -24,7 +24,7 @@ class ScimTransformations: @staticmethod async def transform_litellm_user_to_scim_user( - user: Union[LiteLLM_UserTable, NewUserResponse], + user: LiteLLM_UserTable | NewUserResponse, ) -> SCIMUser: from litellm.proxy.proxy_server import prisma_client @@ -90,7 +90,7 @@ class ScimTransformations: @staticmethod def _parse_directory_metadata( - user: Union[LiteLLM_UserTable, NewUserResponse], + user: LiteLLM_UserTable | NewUserResponse, key: str, validate: Callable[[object], T], ) -> T | None: @@ -114,7 +114,7 @@ class ScimTransformations: return None @staticmethod - def _get_scim_user_name(user: Union[LiteLLM_UserTable, NewUserResponse]) -> str: + def _get_scim_user_name(user: LiteLLM_UserTable | NewUserResponse) -> str: """ SCIM requires a display name with length > 0 @@ -125,7 +125,7 @@ class ScimTransformations: return ScimTransformations.DEFAULT_SCIM_DISPLAY_NAME @staticmethod - def _get_scim_family_name(user: Union[LiteLLM_UserTable, NewUserResponse]) -> str: + def _get_scim_family_name(user: LiteLLM_UserTable | NewUserResponse) -> str: """ SCIM requires a family name with length > 0 """ @@ -140,7 +140,7 @@ class ScimTransformations: return ScimTransformations.DEFAULT_SCIM_FAMILY_NAME @staticmethod - def _get_scim_given_name(user: Union[LiteLLM_UserTable, NewUserResponse]) -> str: + def _get_scim_given_name(user: LiteLLM_UserTable | NewUserResponse) -> str: """ SCIM requires a given name with length > 0 """ @@ -156,7 +156,7 @@ class ScimTransformations: @staticmethod async def transform_litellm_team_to_scim_group( - team: Union[LiteLLM_TeamTable, dict], + team: LiteLLM_TeamTable | dict, ) -> SCIMGroup: from litellm.proxy.proxy_server import prisma_client @@ -167,7 +167,7 @@ class ScimTransformations: team = LiteLLM_TeamTable(**team) # Get team members with proper display names - scim_members: List[SCIMMember] = [] + scim_members: list[SCIMMember] = [] for member in team.members_with_roles or []: if isinstance(member, dict): member = Member(**member) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 9f0b12b6fea..b3140ec911e 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -9,13 +9,8 @@ from collections.abc import Iterable, Mapping, Sequence from itertools import chain from typing import ( TYPE_CHECKING, - Dict, - List, NamedTuple, - Optional, Protocol, - Set, - Tuple, overload, ) @@ -169,8 +164,8 @@ class UserProvisionerHelpers: async def handle_existing_user_by_email( prisma_client: PrismaClient, new_user_request: NewUserRequest, - admin_group: Optional[str] = None, - ) -> Optional[SCIMUser]: + admin_group: str | None = None, + ) -> SCIMUser | None: """ Check if a user with the given email already exists and update them if found. @@ -228,14 +223,14 @@ class UserProvisionerHelpers: class ScimUserData(TypedDict): """Typed structure for extracted SCIM user data.""" - user_email: Optional[str] - user_alias: Optional[str] - sso_user_id: Optional[str] - teams: List[str] - given_name: Optional[str] - family_name: Optional[str] - active: Optional[bool] - enterprise: Optional[SCIMEnterpriseUser] + user_email: str | None + user_alias: str | None + sso_user_id: str | None + teams: list[str] + given_name: str | None + family_name: str | None + active: bool | None + enterprise: SCIMEnterpriseUser | None entitlements: list[SCIMMultiValuedAttribute] | None roles: list[SCIMMultiValuedAttribute] | None @@ -247,9 +242,9 @@ class GroupMemberExtractionResult(BaseModel): so a repeated resolved id appears once in the former and twice in the latter. """ - existing_member_ids: List[str] - created_users: List[NewUserResponse] - all_member_ids: List[str] # existing + newly created + existing_member_ids: list[str] + created_users: list[NewUserResponse] + all_member_ids: list[str] # existing + newly created scim_router = APIRouter( @@ -322,10 +317,10 @@ def _extract_scim_user_data(user: SCIMUser) -> ScimUserData: def _build_scim_metadata( - given_name: Optional[str], - family_name: Optional[str], - active: Optional[bool] = None, - enterprise: Optional[SCIMEnterpriseUser] = None, + given_name: str | None, + family_name: str | None, + active: bool | None = None, + enterprise: SCIMEnterpriseUser | None = None, entitlements: list[SCIMMultiValuedAttribute] | None = None, roles: list[SCIMMultiValuedAttribute] | None = None, ) -> dict[str, object]: @@ -392,7 +387,7 @@ def _default_scim_user_role() -> ScimUserRole: return LitellmUserRoles.INTERNAL_USER_VIEW_ONLY -async def _get_scim_admin_group() -> Optional[str]: +async def _get_scim_admin_group() -> str | None: """ Get the scim_admin_group setting from litellm_settings. @@ -412,9 +407,9 @@ async def _get_scim_admin_group() -> Optional[str]: def _resolve_scim_user_role( groups: list[SCIMUserGroup], - admin_group: Optional[str], + admin_group: str | None, default_role: ScimUserRole, -) -> Optional[LitellmUserRoles]: +) -> LitellmUserRoles | None: """ Resolve a user's global proxy role from their SCIM groups. @@ -513,7 +508,7 @@ def _normalized_member_type(member: SCIMMember) -> str | None: return normalized or None -_JSON_OBJECT_ADAPTER = TypeAdapter(Dict[str, object]) +_JSON_OBJECT_ADAPTER = TypeAdapter(dict[str, object]) def _json_object_fields(raw: object) -> Mapping[str, object] | None: @@ -707,10 +702,10 @@ async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionRe ) -async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: +async def _get_team_members_display(member_ids: list[str]) -> list[SCIMMember]: """Get SCIMMember objects with display names for a list of member IDs.""" prisma_client = await _get_prisma_client_or_raise_exception() - members: List[SCIMMember] = [] + members: list[SCIMMember] = [] for member_id in member_ids: user = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": member_id}) @@ -723,8 +718,8 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: async def _handle_team_membership_changes( user_id: str, - existing_teams: List[str], - new_teams: List[str], + existing_teams: list[str], + new_teams: list[str], raise_on_error: bool = False, ) -> None: """Handle adding/removing user from teams based on changes.""" @@ -833,7 +828,7 @@ async def _delete_rows_referencing_user(prisma_client: PrismaClient, *, user_id: await _table(TeamMembershipRepository(prisma_client)).delete_many(where={"user_id": user_id}) -def _scim_active_value(metadata: Optional[Mapping[str, object]]) -> Optional[bool]: +def _scim_active_value(metadata: Mapping[str, object] | None) -> bool | None: """Read the SCIM active flag from a user's metadata dict, if present.""" if not metadata: return None @@ -843,13 +838,13 @@ def _scim_active_value(metadata: Optional[Mapping[str, object]]) -> Optional[boo return bool(value) -def _user_scim_active(user: LiteLLM_UserTable) -> Optional[bool]: +def _user_scim_active(user: LiteLLM_UserTable) -> bool | None: """Read the SCIM active flag off a user row's metadata, if present.""" metadata: dict[str, object] | None = user.metadata return _scim_active_value(metadata) -async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_group") -> Optional[NewUserResponse]: +async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_group") -> NewUserResponse | None: """ Helper function to create a user if they don't exist. @@ -864,14 +859,15 @@ async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_grou try: # Get default role for new internal users - default_role: Optional[ + default_role: ( Literal[ LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - ] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + | None + ) = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY if litellm.default_internal_user_params: default_role = litellm.default_internal_user_params.get("user_role") @@ -894,14 +890,14 @@ async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_grou return None -async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[str]: +async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> list[str]: """ Get the IDs of the members from a team. Use one source of truth for the member IDs: team.members_with_roles """ - member_user_ids: List[str] = [] + member_user_ids: list[str] = [] for member in team.members_with_roles or []: if hasattr(member, "user_id") and member.user_id is not None: member_user_ids.append(member.user_id) @@ -1305,7 +1301,7 @@ async def get_service_provider_config(request: Request): return SCIMServiceProviderConfig(meta=meta) -def _parse_scim_eq_filter(scim_filter: str) -> Optional[Tuple[str, str]]: +def _parse_scim_eq_filter(scim_filter: str) -> tuple[str, str] | None: """Parse the SCIM equality filters Okta uses before user lifecycle changes.""" match = re.match( r"""\s*([\w.]+)\s+eq\s+(['"]?)(.*?)\2\s*$""", @@ -1327,7 +1323,7 @@ def _parse_scim_eq_filter(scim_filter: str) -> Optional[Tuple[str, str]]: async def get_users( startIndex: int = Query(1, ge=1), count: int = Query(10, ge=1, le=100), - filter: Optional[str] = Query(None), + filter: str | None = Query(None), ): """ Get a list of users according to SCIM v2 protocol @@ -1369,7 +1365,7 @@ async def get_users( total_count = await _table(UserRepository(prisma_client)).count(where=where_conditions) # Convert to SCIM format - scim_users: List[SCIMUser] = [] + scim_users: list[SCIMUser] = [] for user in users: scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(user=user) scim_users.append(scim_user) @@ -1652,12 +1648,12 @@ def _parse_member_entries(value: object) -> tuple[SCIMMember, ...]: return tuple(member for member in (_parse_member_entry(entry) for entry in entries) if member is not None) -def _extract_group_values(value: object) -> List[str]: +def _extract_group_values(value: object) -> list[str]: """Return group ids from a SCIM patch value.""" return [member.value for member in _parse_member_entries(value)] -def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]: +def _extract_ids_from_path_filter(path: str | None, attribute: str) -> list[str]: """Return ids from a SCIM filtered path like ``members[value eq "id"]``. Okta commonly sends membership removals as a filtered path and omits the @@ -1725,7 +1721,7 @@ def _handle_name_update(path: str, op_type: str, value: object, scim_metadata: d scim_metadata["familyName"] = str(value) -def _handle_group_operations(op_type: str, value: object, teams_set: Set[str], path: str | None) -> Set[str] | None: +def _handle_group_operations(op_type: str, value: object, teams_set: set[str], path: str | None) -> set[str] | None: """Handle group/team membership operations.""" group_values = _extract_group_values(value) if not group_values and value is None: @@ -1793,14 +1789,14 @@ def _handle_generic_metadata(path: str, op_type: str, value: object, metadata: d def _apply_patch_ops( existing_user: LiteLLM_UserTable, patch_ops: SCIMPatchOp, -) -> Tuple[dict[str, object], Set[str]]: +) -> tuple[dict[str, object], set[str]]: """Apply patch operations and return update data and final team set.""" update_data: dict[str, object] = {} metadata = existing_user.metadata or {} scim_metadata = metadata.get("scim_metadata", {}) - teams_set: Set[str] = set(existing_user.teams or []) - replace_team_set: Optional[Set[str]] = None + teams_set: set[str] = set(existing_user.teams or []) + replace_team_set: set[str] | None = None for op in patch_ops.Operations: path = (op.path or "").lower() @@ -1863,8 +1859,8 @@ def _is_user_not_in_team_error(exc: HTTPException) -> bool: async def patch_team_membership( user_id: str, - teams_ids_to_add_user_to: List[str], - teams_ids_to_remove_user_from: List[str], + teams_ids_to_add_user_to: list[str], + teams_ids_to_remove_user_from: list[str], raise_on_error: bool = False, ) -> bool: """ @@ -2009,7 +2005,7 @@ class _TeamWhereConditions(TypedDict, total=False): async def get_groups( startIndex: int = Query(1, ge=1), count: int = Query(10, ge=1, le=100), - filter: Optional[str] = Query(None), + filter: str | None = Query(None), ): """ Get a list of groups according to SCIM v2 protocol @@ -2042,7 +2038,7 @@ async def get_groups( total_count = await _table(TeamRepository(prisma_client)).count(where=where_conditions) # Convert to SCIM format - scim_groups: List[SCIMGroup] = [] + scim_groups: list[SCIMGroup] = [] for team in teams: # Get team members with display names. members_with_roles is the # source of truth; the legacy `members` column is not populated by @@ -2275,7 +2271,7 @@ async def delete_group( async def _process_group_patch_operations( patch_ops: SCIMPatchOp, existing_team: LiteLLM_TeamTable, prisma_client: PrismaClient -) -> Tuple[dict[str, object], Set[str], Set[str] | None]: +) -> tuple[dict[str, object], set[str], set[str] | None]: """Process patch operations for a group and return update data, final members and, when the request contained a member ``replace`` op, the absolute target roster it declared (``None`` otherwise). @@ -2379,7 +2375,7 @@ async def _apply_group_patch_updates(group_id: str, update_data: dict[str, objec return await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) -async def _handle_group_membership_changes(group_id: str, current_members: Set[str], final_members: Set[str]): +async def _handle_group_membership_changes(group_id: str, current_members: set[str], final_members: set[str]): """Handle adding/removing members from the group.""" members_to_add = final_members - current_members members_to_remove = current_members - final_members diff --git a/litellm/proxy/management_endpoints/sso/custom_microsoft_sso.py b/litellm/proxy/management_endpoints/sso/custom_microsoft_sso.py index 00d021efd80..2f900f9b6b6 100644 --- a/litellm/proxy/management_endpoints/sso/custom_microsoft_sso.py +++ b/litellm/proxy/management_endpoints/sso/custom_microsoft_sso.py @@ -14,7 +14,6 @@ If these are not set, the default Microsoft endpoints are used. """ import os -from typing import List, Optional, Union import pydantic from fastapi_sso.sso.base import DiscoveryDocument @@ -37,10 +36,10 @@ class CustomMicrosoftSSO(MicrosoftSSO): self, client_id: str, client_secret: str, - redirect_uri: Optional[Union[pydantic.AnyHttpUrl, str]] = None, + redirect_uri: pydantic.AnyHttpUrl | str | None = None, allow_insecure_http: bool = False, - scope: Optional[List[str]] = None, - tenant: Optional[str] = None, + scope: list[str] | None = None, + tenant: str | None = None, ): super().__init__( client_id=client_id, diff --git a/litellm/proxy/management_endpoints/sso_helper_utils.py b/litellm/proxy/management_endpoints/sso_helper_utils.py index fc10ab11658..11f4184437b 100644 --- a/litellm/proxy/management_endpoints/sso_helper_utils.py +++ b/litellm/proxy/management_endpoints/sso_helper_utils.py @@ -1,9 +1,7 @@ -from typing import Dict, Union - from litellm.proxy._types import LitellmUserRoles -def check_is_admin_only_access(ui_access_mode: Union[str, Dict]) -> bool: +def check_is_admin_only_access(ui_access_mode: str | dict) -> bool: """Checks ui access mode is admin_only""" if isinstance(ui_access_mode, str): return ui_access_mode == "admin_only" diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index ac53f254981..fecb14b08d3 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -201,7 +201,7 @@ async def _get_model_names(prisma_client: "PrismaClient", model_ids: Sequence[st models = await _table(ModelRepository(prisma_client)).find_many(where={"model_id": {"in": model_ids}}) return {model.model_id: model.model_name for model in models} except Exception as e: - verbose_proxy_logger.error(f"Error getting model names: {str(e)}") + verbose_proxy_logger.error(f"Error getting model names: {e!s}") return {} @@ -331,7 +331,7 @@ async def new_tag( "tag": tag_config, } except Exception as e: - verbose_proxy_logger.exception(f"Error creating tag: {str(e)}") + verbose_proxy_logger.exception(f"Error creating tag: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -372,7 +372,7 @@ async def _add_tag_to_deployment(deployment: "Deployment", tag: str): data={"litellm_params": json.dumps(existing_params)}, ) except Exception as e: - verbose_proxy_logger.exception(f"Error adding tag to deployment: {str(e)}") + verbose_proxy_logger.exception(f"Error adding tag to deployment: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -461,7 +461,7 @@ async def update_tag( "tag": tag_config, } except Exception as e: - verbose_proxy_logger.exception(f"Error updating tag: {str(e)}") + verbose_proxy_logger.exception(f"Error updating tag: {e!s}") raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 025b7c4210e..152c27202b4 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -9,7 +9,7 @@ import copy import json import traceback from datetime import datetime, timezone -from typing import Any, List, Optional +from typing import Any from fastapi import APIRouter, Depends, Header, HTTPException, Request, status @@ -85,7 +85,7 @@ async def _emit_team_callback_audit_log( before_metadata: Any, after_metadata: Any, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, ) -> None: """Emit an audit-log row for a team-callback mutation. @@ -138,7 +138,7 @@ async def add_team_callbacks( http_request: Request, team_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -210,7 +210,7 @@ async def add_team_callbacks( # store team callback settings in metadata team_metadata = _existing_team.metadata - team_callback_settings: List[dict] = team_metadata.get("logging") # will be dict of type AddTeamCallback + team_callback_settings: list[dict] = team_metadata.get("logging") # will be dict of type AddTeamCallback if team_callback_settings is None or not isinstance(team_callback_settings, list): team_callback_settings = [] @@ -257,9 +257,7 @@ async def add_team_callbacks( except ProxyException as e: raise e except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.add_team_callbacks(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.add_team_callbacks(): Exception occured - {e!s}") raise ProxyException( message="Internal Server Error, " + str(e), type=ProxyErrorTypes.internal_server_error.value, @@ -278,7 +276,7 @@ async def disable_team_logging( http_request: Request, team_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -375,7 +373,7 @@ async def disable_team_logging( except ProxyException: raise except Exception as e: - verbose_proxy_logger.error(f"litellm.proxy.proxy_server.disable_team_logging(): Exception occurred - {str(e)}") + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.disable_team_logging(): Exception occurred - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) raise ProxyException( message="Internal Server Error, " + str(e), @@ -467,13 +465,11 @@ async def get_team_callbacks( except ProxyException: raise except Exception as e: - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.get_team_callbacks(): Exception occurred - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.get_team_callbacks(): Exception occurred - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Internal Server Error({str(e)})"), + message=getattr(e, "detail", f"Internal Server Error({e!s})"), type=ProxyErrorTypes.internal_server_error.value, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index a4c81a9d13e..def58794040 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -17,13 +17,8 @@ from collections.abc import Mapping, Sequence from datetime import datetime, timezone from typing import ( Annotated, - Dict, - List, - Optional, Protocol, - Tuple, TypeVar, - Union, cast, ) @@ -351,7 +346,7 @@ class TeamMemberBudgetHandler: SYSTEM_MANAGED_METADATA_KEYS = ("team_member_budget_id",) @staticmethod - def strip_system_managed_metadata_keys(metadata: Optional[dict]) -> None: + def strip_system_managed_metadata_keys(metadata: dict | None) -> None: """Remove server-owned metadata keys from a caller-supplied dict.""" if not isinstance(metadata, dict): return @@ -360,10 +355,10 @@ class TeamMemberBudgetHandler: @staticmethod def should_create_budget( - team_member_budget: Optional[float] = None, - team_member_rpm_limit: Optional[int] = None, - team_member_tpm_limit: Optional[int] = None, - team_member_budget_duration: Optional[str] = None, + team_member_budget: float | None = None, + team_member_rpm_limit: int | None = None, + team_member_tpm_limit: int | None = None, + team_member_budget_duration: str | None = None, ) -> bool: """Check if any team member limits are provided""" return any( @@ -377,13 +372,13 @@ class TeamMemberBudgetHandler: @staticmethod async def create_team_member_budget_table( - data: Union[NewTeamRequest, LiteLLM_TeamTable], + data: NewTeamRequest | LiteLLM_TeamTable, new_team_data_json: dict, user_api_key_dict: UserAPIKeyAuth, - team_member_budget: Optional[float] = None, - team_member_rpm_limit: Optional[int] = None, - team_member_tpm_limit: Optional[int] = None, - team_member_budget_duration: Optional[str] = None, + team_member_budget: float | None = None, + team_member_rpm_limit: int | None = None, + team_member_tpm_limit: int | None = None, + team_member_budget_duration: str | None = None, ) -> dict: """Create team member budget table with provided limits""" from litellm.proxy._types import BudgetNewRequest @@ -431,10 +426,10 @@ class TeamMemberBudgetHandler: team_table: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth, updated_kv: dict, - team_member_budget: Optional[float] = None, - team_member_rpm_limit: Optional[int] = None, - team_member_tpm_limit: Optional[int] = None, - team_member_budget_duration: Optional[str] = None, + team_member_budget: float | None = None, + team_member_rpm_limit: int | None = None, + team_member_tpm_limit: int | None = None, + team_member_budget_duration: str | None = None, ) -> dict: """Upsert team member budget table with provided limits""" from litellm.proxy._types import BudgetNewRequest @@ -532,7 +527,7 @@ class TeamMemberBudgetHandler: @staticmethod async def backfill_team_member_budget_entries( team_id: str, - members_with_roles: Sequence[Union[Member, dict[str, object]]], + members_with_roles: Sequence[Member | dict[str, object]], team_member_budget_id: str, prisma_client: PrismaClient, ) -> None: @@ -628,11 +623,11 @@ def _is_available_team(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> bool: async def get_all_team_memberships( - prisma_client: PrismaClient, team_ids: List[str], user_id: Optional[str] = None -) -> List[LiteLLM_TeamMembership]: + prisma_client: PrismaClient, team_ids: list[str], user_id: str | None = None +) -> list[LiteLLM_TeamMembership]: """Get all team memberships for a given user""" ## GET ALL MEMBERSHIPS ## - where_obj: Dict[str, Dict[str, List[str]]] = {"team_id": {"in": team_ids}} + where_obj: dict[str, dict[str, list[str]]] = {"team_id": {"in": team_ids}} if user_id is not None: where_obj["user_id"] = {"in": [user_id]} # if user_id is None: @@ -645,7 +640,7 @@ async def get_all_team_memberships( include={"litellm_budget_table": True}, ) - returned_tm: List[LiteLLM_TeamMembership] = [] + returned_tm: list[LiteLLM_TeamMembership] = [] for tm in team_memberships: returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump())) @@ -653,12 +648,12 @@ async def get_all_team_memberships( def _check_team_model_specific_limits( - teams: List[LiteLLM_TeamTable], - data: Union[NewTeamRequest, UpdateTeamRequest], - entity_rpm_limit: Optional[int], - entity_tpm_limit: Optional[int], - entity_model_rpm_limit_dict: Dict[str, int], - entity_model_tpm_limit_dict: Dict[str, int], + teams: list[LiteLLM_TeamTable], + data: NewTeamRequest | UpdateTeamRequest, + entity_rpm_limit: int | None, + entity_tpm_limit: int | None, + entity_model_rpm_limit_dict: dict[str, int], + entity_model_tpm_limit_dict: dict[str, int], entity_type: str, # "organization" ) -> None: """ @@ -675,8 +670,8 @@ def _check_team_model_specific_limits( return # get total model specific tpm/rpm limit - model_specific_rpm_limit: Dict[str, int] = {} - model_specific_tpm_limit: Dict[str, int] = {} + model_specific_rpm_limit: dict[str, int] = {} + model_specific_tpm_limit: dict[str, int] = {} for team in teams: if team.metadata and team.metadata.get("model_rpm_limit", None) is not None: @@ -724,10 +719,10 @@ def _check_team_model_specific_limits( def _check_team_rpm_tpm_limits( - teams: List[LiteLLM_TeamTable], - data: Union[NewTeamRequest, UpdateTeamRequest], - entity_rpm_limit: Optional[int], - entity_tpm_limit: Optional[int], + teams: list[LiteLLM_TeamTable], + data: NewTeamRequest | UpdateTeamRequest, + entity_rpm_limit: int | None, + entity_tpm_limit: int | None, entity_type: str, # "organization" ) -> None: """ @@ -762,9 +757,9 @@ def _check_team_rpm_tpm_limits( def check_org_team_model_specific_limits( - teams: List[LiteLLM_TeamTable], + teams: list[LiteLLM_TeamTable], org_table: LiteLLM_OrganizationTable, - data: Union[NewTeamRequest, UpdateTeamRequest], + data: NewTeamRequest | UpdateTeamRequest, ) -> None: """ Check if the organization team is allocating model specific limits. If so, raise an error if we're overallocating. @@ -796,9 +791,9 @@ def check_org_team_model_specific_limits( def check_org_team_rpm_tpm_limits( - teams: List[LiteLLM_TeamTable], + teams: list[LiteLLM_TeamTable], org_table: LiteLLM_OrganizationTable, - data: Union[NewTeamRequest, UpdateTeamRequest], + data: NewTeamRequest | UpdateTeamRequest, ) -> None: """ Check if the organization team is allocating rpm/tpm limits. If so, raise an error if we're overallocating. @@ -822,7 +817,7 @@ def check_org_team_rpm_tpm_limits( async def _check_org_team_limits( org_table: LiteLLM_OrganizationTable, - data: Union[NewTeamRequest, UpdateTeamRequest], + data: NewTeamRequest | UpdateTeamRequest, prisma_client: PrismaClient, ) -> None: """ @@ -907,7 +902,7 @@ async def _check_org_team_limits( ) # Convert teams to LiteLLM_TeamTable objects - team_objs: List[LiteLLM_TeamTable] = [] + team_objs: list[LiteLLM_TeamTable] = [] for team in teams: team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump())) @@ -924,7 +919,7 @@ async def _check_org_team_limits( async def _check_user_team_limits( - data: Union[NewTeamRequest, UpdateTeamRequest], + data: NewTeamRequest | UpdateTeamRequest, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, @@ -998,7 +993,7 @@ async def _check_user_team_limits( def _check_team_budget_update_authority( data: UpdateTeamRequest, user_api_key_dict: UserAPIKeyAuth, - existing_team_max_budget: Optional[float], + existing_team_max_budget: float | None, ) -> None: """ Restrict who can grow a standalone team's spend ceiling on /team/update. @@ -1055,7 +1050,7 @@ async def new_team( data: NewTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -1372,7 +1367,7 @@ async def new_team( complete_team_data.budget_limits = initialized_windows # type: ignore[assignment] ## Add Team Member Budget Table - members_with_roles: List[Member] = [] + members_with_roles: list[Member] = [] if complete_team_data.members_with_roles is not None: members_with_roles = complete_team_data.members_with_roles complete_team_data.members_with_roles = [] @@ -1447,7 +1442,7 @@ async def _create_team_update_audit_log( existing_team_row: LiteLLM_TeamTable, updated_kv: dict, team_id: str, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, ) -> None: @@ -1494,11 +1489,11 @@ async def _create_team_update_audit_log( async def _update_model_table( data: UpdateTeamRequest, - model_id: Optional[int], + model_id: int | None, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, -) -> Optional[int]: +) -> int | None: """ Upsert model table and return the model id """ @@ -1570,9 +1565,9 @@ async def _auto_add_team_members_to_organization( async def fetch_and_validate_organization( organization_id: str, existing_team_row: LiteLLM_TeamTable, - llm_router: Optional[Router], + llm_router: Router | None, prisma_client: PrismaClient, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ) -> LiteLLM_OrganizationTable: """ Fetch and validate an organization for team update operations. @@ -1731,7 +1726,7 @@ async def update_team( data: UpdateTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -2045,7 +2040,7 @@ async def update_team( updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"]) updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) - team_row: Optional[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.update( + team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data=updated_kv, # `object_permission` is included so `_refresh_cached_team` @@ -2060,7 +2055,7 @@ async def update_team( if team_row is None or team_row.team_id is None: raise HTTPException( status_code=400, - detail={"error": "Team doesn't exist. Got={}".format(team_row)}, + detail={"error": f"Team doesn't exist. Got={team_row}"}, ) verbose_proxy_logger.info("Successfully updated team - %s, info", team_row.team_id) @@ -2098,7 +2093,7 @@ async def patch_team( http_request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], litellm_changed_by: Annotated[ - Optional[str], + str | None, Header( description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -2209,13 +2204,13 @@ async def handle_update_object_permission(data_json: dict, existing_team_row: Li def _check_team_member_admin_add( - member: Union[Member, List[Member]], + member: Member | list[Member], premium_user: bool, ): if isinstance(member, Member) and member.role == "admin": if premium_user is not True: raise ValueError(f"Assigning team admins is a premium feature. {CommonProxyErrors.not_premium_user.value}") - elif isinstance(member, List): + elif isinstance(member, list): for m in member: if m.role == "admin": if premium_user is not True: @@ -2225,7 +2220,7 @@ def _check_team_member_admin_add( def team_call_validation_checks( - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, data: TeamMemberAddRequest, premium_user: bool, ): @@ -2274,7 +2269,7 @@ def team_member_add_duplication_check( # First, populate the invalid_team_members list by checking for duplicates if isinstance(data.member, Member): _check_member_duplication(data.member) - elif isinstance(data.member, List): + elif isinstance(data.member, list): for m in data.member: _check_member_duplication(m) @@ -2385,10 +2380,10 @@ async def _process_team_members( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, -) -> Tuple[List[LiteLLM_UserTable], List[LiteLLM_TeamMembership]]: +) -> tuple[list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: """Process and add new team members.""" - updated_users: List[LiteLLM_UserTable] = [] - updated_team_memberships: List[LiteLLM_TeamMembership] = [] + updated_users: list[LiteLLM_UserTable] = [] + updated_team_memberships: list[LiteLLM_TeamMembership] = [] default_team_budget_id = ( complete_team_data.metadata.get("team_member_budget_id") if complete_team_data.metadata is not None else None @@ -2416,16 +2411,12 @@ async def _process_team_members( except Exception as e: raise HTTPException( status_code=500, - detail={ - "error": "Unable to add user - {}, to team - {}, for reason - {}".format( - data.member, data.team_id, str(e) - ) - }, + detail={"error": f"Unable to add user - {data.member}, to team - {data.team_id}, for reason - {e!s}"}, ) updated_users.append(updated_user) if updated_tm is not None: updated_team_memberships.append(updated_tm) - elif isinstance(data.member, List): + elif isinstance(data.member, list): for m in data.member: try: updated_user, updated_tm = await add_new_member( @@ -2442,11 +2433,7 @@ async def _process_team_members( except Exception as e: raise HTTPException( status_code=500, - detail={ - "error": "Unable to add user - {}, to team - {}, for reason - {}".format( - m, data.team_id, str(e) - ) - }, + detail={"error": f"Unable to add user - {m}, to team - {data.team_id}, for reason - {e!s}"}, ) updated_users.append(updated_user) if updated_tm is not None: @@ -2458,7 +2445,7 @@ async def _process_team_members( async def _update_team_members_list( data: TeamMemberAddRequest, complete_team_data: LiteLLM_TeamTable, - updated_users: List[LiteLLM_UserTable], + updated_users: list[LiteLLM_UserTable], ) -> None: """Update the team's members_with_roles list.""" if isinstance(data.member, Member): @@ -2482,7 +2469,7 @@ async def _update_team_members_list( if not member_already_exists: complete_team_data.members_with_roles.append(new_member) - elif isinstance(data.member, List): + elif isinstance(data.member, list): for nm in data.member: if nm.user_id is None and nm.user_email is not None: for user in updated_users: @@ -2508,7 +2495,7 @@ async def _add_team_members_to_team( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, -) -> Tuple[LiteLLM_TeamTable, List[LiteLLM_UserTable], List[LiteLLM_TeamMembership]]: +) -> tuple[LiteLLM_TeamTable, list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: """Add team members to the team. The members_with_roles reconciliation runs inside a transaction that locks @@ -2638,7 +2625,7 @@ def _validate_member_user_id_provisioning( "error": ( "Only proxy admins can add a user_id that does not exist yet: {}{}. " "Add the member by user_email to invite a new user, or ask a proxy admin " - "to create the user first.".format(listed, " and {} more".format(remaining) if remaining > 0 else "") + "to create the user first.".format(listed, f" and {remaining} more" if remaining > 0 else "") ) }, ) @@ -2901,7 +2888,7 @@ async def team_member_add( member=data.member, prisma_client=prisma_client, ) - elif isinstance(data.member, List): + elif isinstance(data.member, list): for m in data.member: await _validate_and_populate_member_user_info( member=m, @@ -2954,15 +2941,19 @@ async def team_member_add( def _cleanup_members_with_roles( existing_team_row: LiteLLM_TeamTable, data: TeamMemberDeleteRequest, -) -> Tuple[bool, List[Member]]: +) -> tuple[bool, list[Member]]: """Cleanup members_with_roles list for a team.""" is_member_in_team = False - new_team_members: List[Member] = [] + new_team_members: list[Member] = [] for m in existing_team_row.members_with_roles: - if data.user_id is not None and m.user_id is not None and data.user_id == m.user_id: - is_member_in_team = True - continue - elif data.user_email is not None and m.user_email is not None and data.user_email == m.user_email: + if ( + data.user_id is not None + and m.user_id is not None + and data.user_id == m.user_id + or data.user_email is not None + and m.user_email is not None + and data.user_email == m.user_email + ): is_member_in_team = True continue new_team_members.append(m) @@ -3017,7 +3008,7 @@ async def team_member_delete( if _existing_team_row is None: raise HTTPException( status_code=400, - detail={"error": "Team id={} does not exist in db".format(data.team_id)}, + detail={"error": f"Team id={data.team_id} does not exist in db"}, ) existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) @@ -3048,7 +3039,7 @@ async def team_member_delete( existing_team_row.members_with_roles = new_team_members - _db_new_team_members: List[dict] = [m.model_dump() for m in new_team_members] + _db_new_team_members: list[dict] = [m.model_dump() for m in new_team_members] _ = await _team_db(prisma_client).update( where={ @@ -3104,7 +3095,7 @@ async def team_member_delete( ) # Fetch keys before deletion to persist them - keys_to_delete: List[LiteLLM_VerificationToken] = await VerificationTokenRepository( + keys_to_delete: list[LiteLLM_VerificationToken] = await VerificationTokenRepository( prisma_client ).table.find_many( where={ @@ -3140,7 +3131,7 @@ _MEMBER_BUDGET_PATCH_FIELDS = { } -def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, object]: +def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, object]: """Map the budget fields the request actually set (merge-patch: a sent value updates, an explicit null clears, an absent field is left untouched) to their budget-table columns.""" @@ -3152,7 +3143,7 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, objec } -def _validate_budget_duration(budget_duration: Optional[str]) -> None: +def _validate_budget_duration(budget_duration: str | None) -> None: """Reject budget durations that can't be parsed, are non-positive, or overflow date math, so a bad value can't be persisted and later crash the budget reset job.""" @@ -3170,9 +3161,7 @@ def _validate_budget_duration(budget_duration: Optional[str]) -> None: raise HTTPException( status_code=400, detail={ - "error": "Invalid budget_duration '{}'. Use a format like '1h', '24h', '7d', or '30d'.".format( - budget_duration - ) + "error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'." }, ) @@ -3221,7 +3210,7 @@ async def team_member_update( if _existing_team_row is None: raise HTTPException( status_code=400, - detail={"error": "Team id={} does not exist in db".format(data.team_id)}, + detail={"error": f"Team id={data.team_id} does not exist in db"}, ) existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) @@ -3251,7 +3240,7 @@ async def team_member_update( team_table = returned_team_info["team_info"] ## get user id - received_user_id: Optional[str] = None + received_user_id: str | None = None if data.user_id is not None: received_user_id = data.user_id elif data.user_email is not None: @@ -3263,10 +3252,10 @@ async def team_member_update( if received_user_id is None: raise HTTPException( status_code=400, - detail={"error": "User id doesn't exist in team table. Data={}".format(data)}, + detail={"error": f"User id doesn't exist in team table. Data={data}"}, ) ## find the relevant team membership - identified_budget_id: Optional[str] = None + identified_budget_id: str | None = None for tm in returned_team_info["team_memberships"]: if tm.user_id == received_user_id: identified_budget_id = tm.budget_id @@ -3275,7 +3264,7 @@ async def team_member_update( # If this membership still points at the team's shared default member # budget, _upsert_budget_and_membership will clone-on-write so that the # update only touches this user (not every member sharing the default). - team_default_budget_id: Optional[str] = None + team_default_budget_id: str | None = None if team_table.metadata is not None: raw_default_budget_id = team_table.metadata.get("team_member_budget_id") if isinstance(raw_default_budget_id, str): @@ -3296,7 +3285,7 @@ async def team_member_update( ### update team member role if data.role is not None: - team_members: List[Member] = [] + team_members: list[Member] = [] for member in team_table.members_with_roles: if member.user_id == received_user_id: team_members.append( @@ -3311,7 +3300,7 @@ async def team_member_update( team_table.members_with_roles = team_members - _db_team_members: List[dict] = [m.model_dump() for m in team_members] + _db_team_members: list[dict] = [m.model_dump() for m in team_members] await _team_db(prisma_client).update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore @@ -3330,13 +3319,13 @@ async def team_member_update( def _create_results_from_response( - members: List[Member], + members: list[Member], response: TeamAddMemberResponse, -) -> List[TeamMemberAddResult]: +) -> list[TeamMemberAddResult]: """ Convert TeamAddMemberResponse into individual TeamMemberAddResult objects """ - results: List[TeamMemberAddResult] = [] + results: list[TeamMemberAddResult] = [] for member in members: # Find corresponding updated user @@ -3520,7 +3509,7 @@ async def delete_team( data: DeleteTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_changed_by: Optional[str] = Header( + litellm_changed_by: str | None = Header( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), @@ -3556,10 +3545,10 @@ async def delete_team( raise HTTPException(status_code=400, detail={"error": "No team id passed in"}) # check that all teams passed exist - team_rows: List[LiteLLM_TeamTable] = [] + team_rows: list[LiteLLM_TeamTable] = [] for team_id in data.team_ids: try: - team_row_base: Optional[BaseModel] = await _team_db(prisma_client).find_unique(where={"team_id": team_id}) + team_row_base: BaseModel | None = await _team_db(prisma_client).find_unique(where={"team_id": team_id}) if team_row_base is None: raise Exception except Exception: @@ -3589,7 +3578,7 @@ async def delete_team( if litellm.store_audit_logs is True: # make an audit log for each team deleted for team_id in data.team_ids: - team_row: Optional[LiteLLM_TeamTable] = await prisma_client.get_data( # type: ignore + team_row: LiteLLM_TeamTable | None = await prisma_client.get_data( # type: ignore team_id=team_id, table_name="team", query_type="find_unique" ) @@ -3626,7 +3615,7 @@ async def delete_team( _persist_deleted_verification_tokens, ) - keys_to_delete: List[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + keys_to_delete: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( where={"team_id": {"in": data.team_ids}} ) @@ -3679,10 +3668,10 @@ async def delete_team( def _transform_teams_to_deleted_records( - teams: List[LiteLLM_TeamTable], + teams: list[LiteLLM_TeamTable], user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, -) -> List[Dict[str, object]]: + litellm_changed_by: str | None = None, +) -> list[dict[str, object]]: """Transform teams into deleted team records ready for persistence.""" if not teams: return [] @@ -3727,7 +3716,7 @@ def _transform_teams_to_deleted_records( async def _save_deleted_team_records( - records: List[Dict[str, object]], + records: list[dict[str, object]], prisma_client: PrismaClient, ) -> None: """Save deleted team records to the database.""" @@ -3737,10 +3726,10 @@ async def _save_deleted_team_records( async def _persist_deleted_team_records( - teams: List[LiteLLM_TeamTable], + teams: list[LiteLLM_TeamTable], prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, - litellm_changed_by: Optional[str] = None, + litellm_changed_by: str | None = None, ) -> None: """Persist deleted team records by transforming and saving them.""" records = _transform_teams_to_deleted_records( @@ -3770,18 +3759,14 @@ async def validate_membership(user_api_key_dict: UserAPIKeyAuth, team_table: Lit raise HTTPException( status_code=403, detail={ - "error": "Team key for team={} not authorized to access this team={}".format( - user_api_key_dict.team_id, team_table.team_id - ) + "error": f"Team key for team={user_api_key_dict.team_id} not authorized to access this team={team_table.team_id}" }, ) else: raise HTTPException( status_code=403, detail={ - "error": "API key not authorized to access this team={}. No user_id or team_id associated with this key.".format( - team_table.team_id - ) + "error": f"API key not authorized to access this team={team_table.team_id}. No user_id or team_id associated with this key." }, ) @@ -3795,11 +3780,7 @@ async def validate_membership(user_api_key_dict: UserAPIKeyAuth, team_table: Lit raise HTTPException( status_code=403, - detail={ - "error": "User={} not authorized to access this team={}".format( - user_api_key_dict.user_id, team_table.team_id - ) - }, + detail={"error": f"User={user_api_key_dict.user_id} not authorized to access this team={team_table.team_id}"}, ) @@ -3873,7 +3854,7 @@ async def team_info( ) try: - team_info: Optional[BaseModel] = await _team_db(prisma_client).find_unique( + team_info: BaseModel | None = await _team_db(prisma_client).find_unique( where={"team_id": team_id}, include={"litellm_model_table": True, "object_permission": True}, ) @@ -3951,13 +3932,11 @@ async def team_info( except Exception as e: verbose_proxy_logger.error( - "litellm.proxy.management_endpoints.team_endpoints.py::team_info - Exception occurred - {}\n{}".format( - e, traceback.format_exc() - ) + f"litellm.proxy.management_endpoints.team_endpoints.py::team_info - Exception occurred - {e}\n{traceback.format_exc()}" ) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -4024,7 +4003,7 @@ async def team_member_me( ) caller_user_email = user_api_key_dict.user_email - member_role: Optional[str] = None + member_role: str | None = None for m in team_table.members_with_roles: # Match by user_id when present, else fall back to email — members # added by email may have user_id=None on the stored entry. @@ -4194,7 +4173,7 @@ async def unblock_team( async def list_available_teams( http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - response_model=List[LiteLLM_TeamTable], + response_model=list[LiteLLM_TeamTable], ): from litellm.proxy.proxy_server import prisma_client @@ -4205,7 +4184,7 @@ async def list_available_teams( ) available_teams = cast( - Optional[List[str]], + list[str] | None, ( litellm.default_internal_user_params.get("available_teams") if litellm.default_internal_user_params is not None @@ -4238,7 +4217,7 @@ async def _get_org_admin_org_ids( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, -) -> Optional[List[str]]: +) -> list[str] | None: """ Return the list of organization IDs where the user is an org admin. Returns None if the user is not an org admin of any organization or if @@ -4269,24 +4248,24 @@ async def _get_org_admin_org_ids( async def _build_team_list_where_conditions( prisma_client: PrismaClient, - team_id: Optional[str], - team_alias: Optional[str], - organization_id: Optional[str], - user_id: Optional[str], + team_id: str | None, + team_alias: str | None, + organization_id: str | None, + user_id: str | None, use_deleted_table: bool, - search: Optional[str] = None, + search: str | None = None, search_team_id_match: TeamIdSearchMatch = "exact", - org_admin_org_ids: Optional[List[str]] = None, - user_api_key_cache: Optional[UserApiKeyCache] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, -) -> Optional[Dict[str, object]]: + org_admin_org_ids: list[str] | None = None, + user_api_key_cache: UserApiKeyCache | None = None, + proxy_logging_obj: ProxyLogging | None = None, +) -> dict[str, object] | None: """ Build where conditions for team list query. Returns None when the query is guaranteed to yield no results (e.g. user has no team memberships), allowing the caller to skip the DB round-trip. """ - where_conditions: Dict[str, object] = {} + where_conditions: dict[str, object] = {} if team_id: where_conditions["team_id"] = team_id @@ -4353,8 +4332,8 @@ async def _build_team_list_where_conditions( async def _batch_resolve_access_group_resources( - all_access_group_ids: List[str], -) -> Dict[str, Dict[str, List[str]]]: + all_access_group_ids: list[str], +) -> dict[str, dict[str, list[str]]]: """ Batch-fetch access groups in a single DB query and return a per-group resource map. @@ -4372,7 +4351,7 @@ async def _batch_resolve_access_group_resources( where={"access_group_id": {"in": unique_ids}}, ) - result: Dict[str, Dict[str, List[str]]] = {} + result: dict[str, dict[str, list[str]]] = {} for row in rows: result[row.access_group_id] = { "models": list(row.access_model_names or []), @@ -4385,10 +4364,10 @@ async def _batch_resolve_access_group_resources( def _convert_teams_to_response_models( teams: list, use_deleted_table: bool, - keys_count_by_team: Optional[Dict[str, int]] = None, -) -> List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]: + keys_count_by_team: dict[str, int] | None = None, +) -> list[TeamListItem | LiteLLM_TeamTable | LiteLLM_DeletedTeamTable]: """Convert raw Prisma team rows to response models.""" - team_list: List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]] = [] + team_list: list[TeamListItem | LiteLLM_TeamTable | LiteLLM_DeletedTeamTable] = [] counts = keys_count_by_team or {} for team in teams: try: @@ -4418,7 +4397,7 @@ def _convert_teams_to_response_models( async def _get_keys_count_by_team( prisma_client: PrismaClient, teams: Sequence[LiteLLM_TeamTable], -) -> Dict[str, int]: +) -> dict[str, int]: """Aggregate virtual-key counts per team for the given page of teams. Runs a single GROUP BY against LiteLLM_VerificationToken. The IN clause is @@ -4439,12 +4418,12 @@ async def _get_keys_count_by_team( async def _enforce_list_team_v2_access( user_api_key_dict: UserAPIKeyAuth, - user_id: Optional[str], - organization_id: Optional[str], + user_id: str | None, + organization_id: str | None, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, -) -> Tuple[Optional[str], Optional[List[str]]]: +) -> tuple[str | None, list[str] | None]: """Enforce access control for list_team_v2. - Proxy admins and admin viewers can query any teams. @@ -4454,7 +4433,7 @@ async def _enforce_list_team_v2_access( Returns the (possibly overridden) user_id and org_admin_org_ids. """ is_proxy_admin = _user_has_admin_view(user_api_key_dict) - org_admin_org_ids: Optional[List[str]] = None + org_admin_org_ids: list[str] | None = None if is_proxy_admin: return user_id, org_admin_org_ids @@ -4496,9 +4475,7 @@ async def _enforce_list_team_v2_access( raise HTTPException( status_code=401, detail={ - "error": "Only admin users can query all teams/other teams. Your user role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}" }, ) # Regular user — auto-inject caller's user_id @@ -4517,21 +4494,17 @@ async def _enforce_list_team_v2_access( @management_endpoint_wrapper async def list_team_v2( http_request: Request, - user_id: Optional[str] = fastapi.Query( - default=None, description="Only return teams which this 'user_id' belongs to" - ), - organization_id: Optional[str] = fastapi.Query( + user_id: str | None = fastapi.Query(default=None, description="Only return teams which this 'user_id' belongs to"), + organization_id: str | None = fastapi.Query( default=None, description="Only return teams which this 'organization_id' belongs to", ), - team_id: Optional[str] = fastapi.Query( - default=None, description="Only return teams which this 'team_id' belongs to" - ), - team_alias: Optional[str] = fastapi.Query( + team_id: str | None = fastapi.Query(default=None, description="Only return teams which this 'team_id' belongs to"), + team_alias: str | None = fastapi.Query( default=None, description="Only return teams which this 'team_alias' belongs to. Supports partial matching.", ), - search: Optional[str] = fastapi.Query( + search: str | None = fastapi.Query( default=None, description="Combined search: matches teams whose 'team_id' matches the value OR whose 'team_alias' contains it (case-insensitive).", ), @@ -4543,12 +4516,12 @@ async def list_team_v2( ] = "exact", page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1), page_size: int = fastapi.Query(default=10, description="Number of teams per page", ge=1, le=100), - sort_by: Optional[str] = fastapi.Query( + sort_by: str | None = fastapi.Query( default=None, description="Column to sort by (e.g. 'team_id', 'team_alias', 'created_at')", ), sort_order: str = fastapi.Query(default="asc", description="Sort order ('asc' or 'desc')"), - status: Optional[str] = fastapi.Query(default=None, description="Filter by status (e.g. 'deleted')"), + status: str | None = fastapi.Query(default=None, description="Filter by status (e.g. 'deleted')"), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -4665,7 +4638,7 @@ async def list_team_v2( # Aggregate virtual-key counts per team for the current page. The deleted # table does not carry keys_count, so it is skipped. - keys_count_by_team: Dict[str, int] = {} + keys_count_by_team: dict[str, int] = {} if not use_deleted_table: keys_count_by_team = await _get_keys_count_by_team(prisma_client, teams) @@ -4700,7 +4673,7 @@ async def list_team_v2( async def _authorize_and_filter_teams( user_api_key_dict: UserAPIKeyAuth, - user_id: Optional[str], + user_id: str | None, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, @@ -4714,7 +4687,7 @@ async def _authorize_and_filter_teams( - Others: 401. """ is_proxy_admin = _user_has_admin_view(user_api_key_dict) - allowed_org_ids: Optional[List[str]] = None + allowed_org_ids: list[str] | None = None if not is_proxy_admin: is_own_query = ( @@ -4743,9 +4716,7 @@ async def _authorize_and_filter_teams( raise HTTPException( status_code=401, detail={ - "error": "Only admin users can query all teams/other teams. Your user role={}".format( - user_api_key_dict.user_role - ) + "error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}" }, ) @@ -4780,10 +4751,8 @@ async def _authorize_and_filter_teams( @management_endpoint_wrapper async def list_team( http_request: Request, - user_id: Optional[str] = fastapi.Query( - default=None, description="Only return teams which this 'user_id' belongs to" - ), - organization_id: Optional[str] = None, + user_id: str | None = fastapi.Query(default=None, description="Only return teams which this 'user_id' belongs to"), + organization_id: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -4819,9 +4788,9 @@ async def list_team( _team_ids = [team.team_id for team in filtered_response] returned_tm = await get_all_team_memberships(prisma_client, _team_ids, user_id=user_id) - returned_responses: List[TeamListResponseObject] = [] + returned_responses: list[TeamListResponseObject] = [] for team in filtered_response: - _team_memberships: List[LiteLLM_TeamMembership] = [] + _team_memberships: list[LiteLLM_TeamMembership] = [] for tm in returned_tm: if tm.team_id == team.team_id: _team_memberships.append(tm) @@ -4838,9 +4807,9 @@ async def list_team( ) ) except Exception as e: - team_exception = """Invalid team object for team_id: {}. team_object={}. - Error: {} - """.format(team.team_id, team.model_dump(), str(e)) + team_exception = f"""Invalid team object for team_id: {team.team_id}. team_object={team.model_dump()}. + Error: {e!s} + """ verbose_proxy_logger.exception(team_exception) continue # Sort the responses by team_alias @@ -4859,7 +4828,7 @@ async def get_paginated_teams( prisma_client: PrismaClient, page_size: int = 10, page: int = 1, -) -> Tuple[List[LiteLLM_TeamTable], int]: +) -> tuple[list[LiteLLM_TeamTable], int]: """ Get paginated list of teams from team table @@ -4895,12 +4864,12 @@ async def get_paginated_teams( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, responses={ - 200: {"model": List[LiteLLM_TeamTable]}, + 200: {"model": list[LiteLLM_TeamTable]}, }, ) async def ui_view_teams( - team_id: Optional[str] = fastapi.Query(default=None, description="Team ID in the request parameters"), - team_alias: Optional[str] = fastapi.Query(default=None, description="Team alias in the request parameters"), + team_id: str | None = fastapi.Query(default=None, description="Team ID in the request parameters"), + team_alias: str | None = fastapi.Query(default=None, description="Team alias in the request parameters"), page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1), page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=100), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -4956,10 +4925,10 @@ async def ui_view_teams( return teams except Exception as e: - raise HTTPException(status_code=500, detail=f"Error searching teams: {str(e)}") + raise HTTPException(status_code=500, detail=f"Error searching teams: {e!s}") -def add_new_models_to_team(team_obj: LiteLLM_TeamTable, new_models: List[str]) -> List[str]: +def add_new_models_to_team(team_obj: LiteLLM_TeamTable, new_models: list[str]) -> list[str]: """ Add new models to a team's allowed model list. """ @@ -5367,7 +5336,7 @@ async def _compute_and_batch_updates(prisma_client, teams: Sequence[LiteLLM_Team async def _append_permissions_to_specific_teams( - prisma_client: PrismaClient, team_ids: List[str], permissions_to_add: set + prisma_client: PrismaClient, team_ids: list[str], permissions_to_add: set ) -> int: """Fetch specific teams by ID and append permissions.""" teams = await _team_db(prisma_client).find_many( @@ -5421,14 +5390,14 @@ async def _append_permissions_to_all_teams(prisma_client: PrismaClient, permissi tags=["team management"], ) async def get_team_daily_activity( - team_ids: Optional[str] = None, - start_date: Optional[str] = None, - end_date: Optional[str] = None, - model: Optional[str] = None, - api_key: Optional[str] = None, + team_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, page: int = 1, page_size: int = 10, - exclude_team_ids: Optional[str] = None, + exclude_team_ids: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -5460,7 +5429,7 @@ async def get_team_daily_activity( # Convert comma-separated tags string to list if provided team_ids_list = team_ids.split(",") if team_ids else None - exclude_team_ids_list: Optional[List[str]] = None + exclude_team_ids_list: list[str] | None = None if exclude_team_ids: exclude_team_ids_list = exclude_team_ids.split(",") if exclude_team_ids else None @@ -5478,7 +5447,7 @@ async def get_team_daily_activity( if user_info is None: raise HTTPException( status_code=404, - detail={"error": "User= {} not found".format(user_api_key_dict.user_id)}, + detail={"error": f"User= {user_api_key_dict.user_id} not found"}, ) if team_ids_list is None: @@ -5490,9 +5459,7 @@ async def get_team_daily_activity( raise HTTPException( status_code=404, detail={ - "error": "User does not belong to Team= {}. Call `/user/info` to see user's teams".format( - team_id - ) + "error": f"User does not belong to Team= {team_id}. Call `/user/info` to see user's teams" }, ) @@ -5513,7 +5480,7 @@ async def get_team_daily_activity( # only has admin/permission for a strict subset, fall back to # filtering the entire response by their own API keys (they can re- # request the admin-only teams separately to get the wider view). - user_api_keys: Optional[List[str]] = None + user_api_keys: list[str] | None = None if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: @@ -5538,7 +5505,7 @@ async def get_team_daily_activity( user_api_keys = [""] # Use empty string to ensure no matches # If api_key parameter is provided, use it; otherwise use user_api_keys if set - final_api_key_filter: Optional[Union[str, List[str]]] = api_key + final_api_key_filter: str | list[str] | None = api_key if final_api_key_filter is None and user_api_keys is not None: final_api_key_filter = user_api_keys diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 11ec19a8581..b08393708c6 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -11,7 +11,7 @@ POST /v1/tool/policy - Update the input_policy / output_policy for a import uuid from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Annotated, Any, List, Optional +from typing import TYPE_CHECKING, Annotated, Any from fastapi import APIRouter, Depends, HTTPException, Query from pydantic import BaseModel, Field, TypeAdapter @@ -107,7 +107,7 @@ async def get_tool_policy_options( response_model=ToolListResponse, ) async def list_tools( - input_policy: Optional[ToolInputPolicy] = None, + input_policy: ToolInputPolicy | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -270,7 +270,7 @@ async def get_tool_detail( raise HTTPException(status_code=500, detail=str(e)) -def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]: +def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> str | None: """Short snippet from messages or proxy_server_request for tool usage log row.""" if sl is None: return None @@ -299,7 +299,7 @@ def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]: return _snippet_str(psr, max_len) -def _snippet_str(text: Any, max_len: int = 200) -> Optional[str]: +def _snippet_str(text: Any, max_len: int = 200) -> str | None: if text is None: return None if isinstance(text, str): @@ -330,8 +330,8 @@ async def get_tool_usage_logs( tool_name: str, page: int = Query(1, ge=1), page_size: int = Query(50, ge=1, le=100), - start_date: Optional[str] = Query(None, description="YYYY-MM-DD"), - end_date: Optional[str] = Query(None, description="YYYY-MM-DD"), + start_date: str | None = Query(None, description="YYYY-MM-DD"), + end_date: str | None = Query(None, description="YYYY-MM-DD"), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -346,8 +346,8 @@ async def get_tool_usage_logs( try: where: dict = {"tool_name": tool_name} if start_date or end_date: - start_time_filter: Optional[datetime] = None - end_time_filter: Optional[datetime] = None + start_time_filter: datetime | None = None + end_time_filter: datetime | None = None if start_date: try: start_time_filter = datetime.strptime(start_date + "T00:00:00", "%Y-%m-%dT%H:%M:%S").replace( @@ -383,7 +383,7 @@ async def get_tool_usage_logs( spend_logs = await SpendLogsRepository(prisma_client).table.find_many(where={"request_id": {"in": request_ids}}) log_by_id = {s.request_id: s for s in spend_logs} - logs_out: List[ToolUsageLogEntry] = [] + logs_out: list[ToolUsageLogEntry] = [] for r in index_rows: sl = log_by_id.get(r.request_id) if not sl: @@ -442,7 +442,7 @@ async def get_tool( async def _resolve_key_hash_to_object_permission_id( prisma_client: "PrismaClient", key_hash: str, -) -> Optional[str]: +) -> str | None: """Resolve key (hash or raw) to object_permission_id; create permission if key has none.""" from litellm.proxy.proxy_server import hash_token @@ -473,7 +473,7 @@ async def _resolve_key_hash_to_object_permission_id( async def _resolve_team_id_to_object_permission_id( prisma_client: "PrismaClient", team_id: str, -) -> Optional[str]: +) -> str | None: """Resolve team_id to object_permission_id; create permission if team has none.""" if not team_id or not team_id.strip(): return None @@ -619,8 +619,8 @@ async def update_tool_policy( ) async def delete_tool_policy_override( tool_name: str, - team_id: Optional[str] = Query(None, description="Team ID of the override to remove"), - key_hash: Optional[str] = Query(None, description="Key hash of the override to remove"), + team_id: str | None = Query(None, description="Team ID of the override to remove"), + key_hash: str | None = Query(None, description="Key hash of the override to remove"), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/management_endpoints/types.py b/litellm/proxy/management_endpoints/types.py index 295c2ad50b3..b59f189ea27 100644 --- a/litellm/proxy/management_endpoints/types.py +++ b/litellm/proxy/management_endpoints/types.py @@ -4,7 +4,7 @@ Types for the management endpoints Might include fastapi/proxy requirements.txt related imports """ -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast from fastapi_sso.sso.base import OpenID @@ -28,7 +28,7 @@ def is_valid_litellm_user_role(role_str: str) -> bool: return False -def get_litellm_user_role(role_str) -> Optional[LitellmUserRoles]: +def get_litellm_user_role(role_str) -> LitellmUserRoles | None: """ Convert a string (or list of strings) to a LitellmUserRoles enum if valid (case-insensitive). @@ -48,12 +48,12 @@ def get_litellm_user_role(role_str) -> Optional[LitellmUserRoles]: role_str = role_str[0] # Use _value2member_map_ for O(1) lookup, case-insensitive result = LitellmUserRoles._value2member_map_.get(role_str.lower()) - return cast(Optional[LitellmUserRoles], result) + return cast(LitellmUserRoles | None, result) except Exception: return None class CustomOpenID(OpenID): - team_ids: List[str] - user_role: Optional[LitellmUserRoles] = None - extra_fields: Optional[Dict[str, Any]] = None + team_ids: list[str] + user_role: LitellmUserRoles | None = None + extra_fields: dict[str, Any] | None = None diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 110e5883485..d274879f82a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -21,12 +21,9 @@ from html import escape from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, NoReturn, Optional, - Tuple, Union, cast, ) @@ -184,7 +181,7 @@ def _get_cli_sso_flow_cache_key(login_id: str) -> str: return f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:{login_id}" -def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool: +def _is_valid_cli_sso_login_id(login_id: str | None) -> bool: return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id)) @@ -221,7 +218,7 @@ def _cli_sso_start_response_body( } -def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: Optional[bool] = False) -> str: +def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: bool | None = False) -> str: client_ip = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" client_ip_hash = _hash_cli_sso_secret(client_ip) return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}" @@ -230,7 +227,7 @@ def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_fo def _check_cli_sso_start_rate_limit( request: Request, cache: DualCache, - use_x_forwarded_for: Optional[bool] = False, + use_x_forwarded_for: bool | None = False, ) -> None: rate_limit_cache_key = _get_cli_sso_start_rate_limit_cache_key( request=request, use_x_forwarded_for=use_x_forwarded_for @@ -247,7 +244,7 @@ def _check_cli_sso_start_rate_limit( ) -def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dict: +def _get_cli_sso_flow_or_raise(login_id: str | None, cache: DualCache) -> dict: if isinstance(login_id, str) and login_id.startswith("sk-"): raise HTTPException( status_code=400, @@ -297,7 +294,7 @@ def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None: cache.set_cache(key=cache_key, value=flow, ttl=CLI_SSO_SESSION_TTL_SECONDS) -def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool: +def _verify_cli_sso_poll_secret(flow: dict, poll_secret: str | None) -> bool: expected_poll_secret_hash = flow.get("poll_secret_hash") if not isinstance(expected_poll_secret_hash, str) or not isinstance(poll_secret, str): return False @@ -305,7 +302,7 @@ def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool: return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash) -def _parse_cli_sso_claim_map() -> List[Tuple[str, str]]: +def _parse_cli_sso_claim_map() -> list[tuple[str, str]]: """ Parse CLI_SSO_CLAIM_MAP / LITELLM_CLI_SSO_CLAIM_MAP. @@ -318,7 +315,7 @@ def _parse_cli_sso_claim_map() -> List[Tuple[str, str]]: if not claim_map_raw: return [] - parsed: List[Tuple[str, str]] = [] + parsed: list[tuple[str, str]] = [] for entry in claim_map_raw.split(","): entry = entry.strip() if not entry or "->" not in entry: @@ -326,8 +323,7 @@ def _parse_cli_sso_claim_map() -> List[Tuple[str, str]]: source_claim, dest_key = entry.split("->", 1) source_claim = source_claim.strip() dest_key = dest_key.strip() - if dest_key.startswith("metadata."): - dest_key = dest_key[len("metadata.") :] + dest_key = dest_key.removeprefix("metadata.") if source_claim and dest_key: parsed.append((source_claim, dest_key)) return parsed @@ -351,17 +347,17 @@ def _is_safe_cli_sso_scalar_claim_value(value: Any) -> bool: return True -def _sso_result_to_dict(result: Union[CustomOpenID, OpenID, dict]) -> Dict[str, Any]: +def _sso_result_to_dict(result: CustomOpenID | OpenID | dict) -> dict[str, Any]: if isinstance(result, dict): return result if hasattr(result, "model_dump"): dumped = result.model_dump() if isinstance(dumped, dict): - return cast(Dict[str, Any], dumped) + return cast(dict[str, Any], dumped) return {} -def _get_nested_claim_value(data: Dict[str, Any], claim_path: str) -> Any: +def _get_nested_claim_value(data: dict[str, Any], claim_path: str) -> Any: """Resolve a dot-notation claim path against an SSO result dict. Unlike ``get_nested_value``, this does not strip a leading ``metadata.`` @@ -384,7 +380,7 @@ def _get_nested_claim_value(data: Dict[str, Any], claim_path: str) -> Any: return current -def _extract_sso_claim_value(result: Union[CustomOpenID, OpenID, dict], claim_path: str) -> Any: +def _extract_sso_claim_value(result: CustomOpenID | OpenID | dict, claim_path: str) -> Any: extra_fields = getattr(result, "extra_fields", None) if isinstance(extra_fields, dict): if claim_path in extra_fields: @@ -400,7 +396,7 @@ def _extract_sso_claim_value(result: Union[CustomOpenID, OpenID, dict], claim_pa return _get_nested_claim_value(result_dict, claim_path) -def _set_nested_metadata_value(metadata: Dict[str, Any], key_path: str, value: Any) -> None: +def _set_nested_metadata_value(metadata: dict[str, Any], key_path: str, value: Any) -> None: placeholder = "\x00" parts = key_path.replace("\\.", placeholder).split(".") parts = [p.replace(placeholder, ".") for p in parts] @@ -415,11 +411,11 @@ def _set_nested_metadata_value(metadata: Dict[str, Any], key_path: str, value: A def _flatten_cli_sso_metadata_for_poll( - metadata: Dict[str, Any], -) -> Dict[str, Union[str, int, float, bool]]: + metadata: dict[str, Any], +) -> dict[str, str | int | float | bool]: """Expose scalar attribution metadata as a flat dict for CLI poll responses.""" - flattened: Dict[str, Union[str, int, float, bool]] = {} - stack: List[Tuple[str, Any]] = [("", metadata)] + flattened: dict[str, str | int | float | bool] = {} + stack: list[tuple[str, Any]] = [("", metadata)] while stack: prefix, value = stack.pop() if isinstance(value, dict): @@ -432,8 +428,8 @@ def _flatten_cli_sso_metadata_for_poll( def build_cli_sso_attribution_metadata( - result: Union[CustomOpenID, OpenID, dict], -) -> Dict[str, Any]: + result: CustomOpenID | OpenID | dict, +) -> dict[str, Any]: """ Build allowlisted, non-secret scalar attribution metadata from an SSO result. @@ -444,7 +440,7 @@ def build_cli_sso_attribution_metadata( if not claim_map: return {} - metadata: Dict[str, Any] = {} + metadata: dict[str, Any] = {} for source_claim, dest_key in claim_map: if not _is_safe_cli_sso_metadata_dest_key(dest_key): verbose_proxy_logger.debug(f"Skipping unsafe CLI SSO metadata destination key: {dest_key}") @@ -460,8 +456,8 @@ def build_cli_sso_attribution_metadata( def _merge_cli_sso_attribution_metadata( - existing_metadata: Dict[str, Any], attribution_metadata: Dict[str, Any] -) -> Dict[str, Any]: + existing_metadata: dict[str, Any], attribution_metadata: dict[str, Any] +) -> dict[str, Any]: """Merge attribution metadata into existing user metadata in-place. Preserves original value types (in particular, string claim values that @@ -469,7 +465,7 @@ def _merge_cli_sso_attribution_metadata( are merged iteratively so attribution claims do not clobber unrelated keys under the same parent. """ - pending: List[Tuple[Dict[str, Any], Dict[str, Any]]] = [(existing_metadata, attribution_metadata)] + pending: list[tuple[dict[str, Any], dict[str, Any]]] = [(existing_metadata, attribution_metadata)] while pending: target, source = pending.pop() for key, value in source.items(): @@ -486,14 +482,14 @@ def _merge_cli_sso_attribution_metadata( async def _persist_cli_sso_user_metadata( prisma_client: PrismaClient, user_id: str, - attribution_metadata: Dict[str, Any], + attribution_metadata: dict[str, Any], ) -> None: if not attribution_metadata: return try: user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - existing_metadata: Dict[str, Any] = {} + existing_metadata: dict[str, Any] = {} if user_row is not None: row_metadata = user_row.metadata if isinstance(row_metadata, dict): @@ -516,8 +512,8 @@ async def _persist_cli_sso_user_metadata( def _cli_poll_attribution_metadata_from_session( - session_data: Dict[str, Any], -) -> Dict[str, Union[str, int, float, bool]]: + session_data: dict[str, Any], +) -> dict[str, str | int | float | bool]: stored = session_data.get("attribution_metadata") if isinstance(stored, dict): return _flatten_cli_sso_metadata_for_poll(stored) @@ -688,7 +684,7 @@ async def cli_sso_complete(request: Request, login_id: str): return HTMLResponse(content=html_content, status_code=200) -def normalize_email(email: Optional[str]) -> Optional[str]: +def normalize_email(email: str | None) -> str | None: """ Normalize email address to lowercase for consistent storage and comparison. @@ -709,9 +705,9 @@ def normalize_email(email: Optional[str]) -> Optional[str]: def determine_role_from_groups( - user_groups: List[str], + user_groups: list[str], role_mappings: "RoleMappings", -) -> Optional[LitellmUserRoles]: +) -> LitellmUserRoles | None: """ Determine the highest privilege role for a user based on their groups. @@ -761,11 +757,11 @@ def determine_role_from_groups( def process_sso_jwt_access_token( - access_token_str: Optional[str], - sso_jwt_handler: Optional[JWTHandler], - result: Union[OpenID, dict, None], + access_token_str: str | None, + sso_jwt_handler: JWTHandler | None, + result: OpenID | dict | None, role_mappings: Optional["RoleMappings"] = None, -) -> Optional[dict]: +) -> dict | None: """ Process SSO JWT access token and extract team IDs and user role if available. @@ -800,7 +796,7 @@ def process_sso_jwt_access_token( # Extract team IDs from access token if sso_jwt_handler is available if sso_jwt_handler: if isinstance(result, dict): - result_team_ids: Optional[List[str]] = result.get("team_ids", []) + result_team_ids: list[str] | None = result.get("team_ids", []) if not result_team_ids: team_ids = sso_jwt_handler.get_team_ids_from_jwt(access_token_payload) result["team_ids"] = team_ids @@ -813,14 +809,14 @@ def process_sso_jwt_access_token( # Extract user role from access token if not already set from UserInfo existing_role = result.get("user_role") if isinstance(result, dict) else getattr(result, "user_role", None) if existing_role is None: - user_role: Optional[LitellmUserRoles] = None + user_role: LitellmUserRoles | None = None # Try role_mappings first (group-based role determination) if role_mappings is not None and role_mappings.roles: group_claim = role_mappings.group_claim user_groups_raw: Any = get_nested_value(access_token_payload, group_claim) - user_groups: List[str] = [] + user_groups: list[str] = [] if isinstance(user_groups_raw, list): user_groups = [str(g) for g in user_groups_raw] elif isinstance(user_groups_raw, str): @@ -882,10 +878,10 @@ async def _raise_if_sso_exceeds_free_user_limit(premium_user: bool, prisma_clien @router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False) async def google_login( request: Request, - source: Optional[str] = None, - key: Optional[str] = None, - existing_key: Optional[str] = None, - return_to: Optional[str] = None, + source: str | None = None, + key: str | None = None, + existing_key: str | None = None, + return_to: str | None = None, user_code: str | None = None, ): """ @@ -937,7 +933,7 @@ async def google_login( _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache) # Store CLI login handle in state for OAuth flow - cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( + cli_state: str | None = SSOAuthenticationHandler._get_cli_state( source=source, key=key, user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None), @@ -1017,7 +1013,7 @@ async def google_login( def generic_response_convertor( response, jwt_handler: JWTHandler, - sso_jwt_handler: Optional[JWTHandler] = None, + sso_jwt_handler: JWTHandler | None = None, role_mappings: Optional["RoleMappings"] = None, team_mappings: Optional["TeamMappings"] = None, ) -> CustomOpenID: @@ -1044,7 +1040,7 @@ def generic_response_convertor( all_teams.extend(team_ids) if team_mappings is not None and team_mappings.team_ids_jwt_field is not None: - team_ids_from_db_mapping: Optional[List[str]] = get_nested_value( + team_ids_from_db_mapping: list[str] | None = get_nested_value( data=cast(dict, response), key_path=team_mappings.team_ids_jwt_field, default=[], @@ -1060,7 +1056,7 @@ def generic_response_convertor( # Determine user role based on role_mappings if available # Only apply role_mappings for GENERIC SSO provider - user_role: Optional[LitellmUserRoles] = None + user_role: LitellmUserRoles | None = None if role_mappings is not None and role_mappings.provider.lower() in [ "generic", @@ -1071,7 +1067,7 @@ def generic_response_convertor( user_groups_raw: Any = get_nested_value(response, group_claim) # Handle different formats: could be a list, string (comma-separated), or single value - user_groups: List[str] = [] + user_groups: list[str] = [] if isinstance(user_groups_raw, list): user_groups = [str(g) for g in user_groups_raw] elif isinstance(user_groups_raw, str): @@ -1105,7 +1101,7 @@ def generic_response_convertor( ) # Build extra_fields dict from GENERIC_USER_EXTRA_ATTRIBUTES if specified - extra_fields: Optional[Dict[str, Any]] = None + extra_fields: dict[str, Any] | None = None if generic_user_extra_attributes: extra_fields = {} for attr_name in generic_user_extra_attributes.split(","): @@ -1127,7 +1123,7 @@ def generic_response_convertor( def _setup_generic_sso_env_vars( generic_client_id: str, redirect_url: str -) -> Tuple[str, List[str], str, str, str, bool]: +) -> tuple[str, list[str], str, str, str, bool]: """Setup and validate Generic SSO environment variables.""" generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None) generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ") @@ -1183,7 +1179,9 @@ def _setup_generic_sso_env_vars( async def _setup_team_mappings() -> Optional["TeamMappings"]: """Setup team mappings from SSO database settings.""" - team_mappings: Optional["TeamMappings"] = None + from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings + + team_mappings: TeamMappings | None = None try: from litellm.proxy.utils import get_prisma_client_or_throw @@ -1196,8 +1194,6 @@ async def _setup_team_mappings() -> Optional["TeamMappings"]: team_mappings_data = sso_settings_dict.get("team_mappings") if team_mappings_data: - from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings - if isinstance(team_mappings_data, dict): team_mappings = TeamMappings(**team_mappings_data) elif isinstance(team_mappings_data, TeamMappings): @@ -1217,7 +1213,7 @@ async def _setup_team_mappings() -> Optional["TeamMappings"]: async def _setup_role_mappings() -> Optional["RoleMappings"]: """Setup role mappings from SSO database settings.""" - role_mappings: Optional["RoleMappings"] = None + role_mappings: RoleMappings | None = None try: from litellm.proxy.utils import get_prisma_client_or_throw @@ -1230,8 +1226,6 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: role_mappings_data = sso_settings_dict.get("role_mappings") if role_mappings_data: - from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings - if isinstance(role_mappings_data, dict): role_mappings = RoleMappings(**role_mappings_data) elif isinstance(role_mappings_data, RoleMappings): @@ -1252,10 +1246,8 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: import ast try: - generic_user_role_mappings_data: Dict[LitellmUserRoles, List[str]] = ast.literal_eval(generic_role_mappings) + generic_user_role_mappings_data: dict[LitellmUserRoles, list[str]] = ast.literal_eval(generic_role_mappings) if isinstance(generic_user_role_mappings_data, dict): - from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings - role_mappings_data = { "provider": "generic", "group_claim": generic_role_mappings_group_claim, @@ -1280,7 +1272,7 @@ def _parse_generic_sso_headers() -> dict: raw = os.getenv("GENERIC_SSO_HEADERS", None) if raw is None: return {} - result: Dict[str, str] = {} + result: dict[str, str] = {} for header in raw.split(","): header = header.strip() if header: @@ -1291,8 +1283,8 @@ def _parse_generic_sso_headers() -> dict: def _handle_generic_sso_error( e: Exception, - generic_authorization_endpoint: Optional[str], - generic_token_endpoint: Optional[str], + generic_authorization_endpoint: str | None, + generic_token_endpoint: str | None, additional_headers: dict, ) -> NoReturn: """Handle errors from generic SSO verify_and_process. Always re-raises.""" @@ -1347,17 +1339,17 @@ def _handle_generic_sso_error( async def get_generic_sso_response( request: Request, jwt_handler: JWTHandler, - sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control + sso_jwt_handler: JWTHandler | None, # sso specific jwt handler - used for restricted sso group access control generic_client_id: str, redirect_url: str, ) -> tuple[ - Union[OpenID, dict], dict | None, dict | None, SSOIdentityAssertion | None + OpenID | dict, dict | None, dict | None, SSOIdentityAssertion | None ]: # (result, received_response, access_token_payload, sso_assertion) # make generic sso provider from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider - received_response: Optional[dict] = None + received_response: dict | None = None sso_assertion: SSOIdentityAssertion | None = None # Setup environment variables @@ -1405,8 +1397,8 @@ async def get_generic_sso_response( verbose_proxy_logger.debug("calling generic_sso.verify_and_process") additional_generic_sso_headers_dict = _parse_generic_sso_headers() - code_verifier: Optional[str] = None # assigned inside try; initialized for type tracking - access_token_payload: Optional[dict] = None # decoded JWT access token claims + code_verifier: str | None = None # assigned inside try; initialized for type tracking + access_token_payload: dict | None = None # decoded JWT access token claims try: token_exchange_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters( @@ -1498,7 +1490,7 @@ async def get_generic_sso_response( # In the PKCE path verify_and_process is skipped, so generic_sso.access_token # is never set. Read the token directly from the exchange response instead so # process_sso_jwt_access_token can extract JWT-embedded roles/teams. - access_token_str: Optional[str] = combined_response.get("access_token") + access_token_str: str | None = combined_response.get("access_token") else: result = await generic_sso.verify_and_process( request, @@ -1545,7 +1537,7 @@ async def create_team_member_add_task(team_id, user_info): verbose_proxy_logger.debug(f"[Non-Blocking] Error trying to add sso user to db: {e}") -async def add_missing_team_member(user_info: Union[NewUserResponse, LiteLLM_UserTable], sso_teams: List[str]): +async def add_missing_team_member(user_info: NewUserResponse | LiteLLM_UserTable, sso_teams: list[str]): """ - Get missing teams (diff b/w user_info.team_ids and sso_teams) - Add missing user to missing teams @@ -1573,12 +1565,12 @@ def get_disabled_non_admin_personal_key_creation(): async def get_existing_user_info_from_db( - user_id: Optional[str], - user_email: Optional[str], + user_id: str | None, + user_email: str | None, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, -) -> Optional[LiteLLM_UserTable]: +) -> LiteLLM_UserTable | None: try: user_info = await get_user_object( user_id=user_id, @@ -1598,14 +1590,14 @@ async def get_existing_user_info_from_db( async def get_user_info_from_db( - result: Union[CustomOpenID, OpenID, dict], + result: CustomOpenID | OpenID | dict, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, - user_email: Optional[str], - user_defined_values: Optional[SSOUserDefinedValues], - alternate_user_id: Optional[str] = None, -) -> Optional[Union[LiteLLM_UserTable, NewUserResponse]]: + user_email: str | None, + user_defined_values: SSOUserDefinedValues | None, + alternate_user_id: str | None = None, +) -> LiteLLM_UserTable | NewUserResponse | None: try: potential_user_ids = [] if alternate_user_id is not None: @@ -1623,7 +1615,7 @@ async def get_user_info_from_db( getattr(result, "email", None) if not isinstance(result, dict) else result.get("email", None) ) - user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]] = None + user_info: LiteLLM_UserTable | NewUserResponse | None = None for user_id in potential_user_ids: user_info = await get_existing_user_info_from_db( @@ -1661,7 +1653,7 @@ async def get_user_info_from_db( return None -def _should_use_role_from_sso_response(sso_role: Optional[str]) -> bool: +def _should_use_role_from_sso_response(sso_role: str | None) -> bool: """returns true if SSO upsert should use the 'role' defined on the SSO response""" if sso_role is None: return False @@ -1676,9 +1668,9 @@ def _should_use_role_from_sso_response(sso_role: Optional[str]) -> bool: def _build_sso_user_update_data( - result: Optional[Union["CustomOpenID", OpenID, dict]], - user_email: Optional[str], - user_id: Optional[str], + result: Union["CustomOpenID", OpenID, dict] | None, + user_email: str | None, + user_id: str | None, ) -> dict: """ Build the update data dictionary for SSO user upsert. @@ -1708,12 +1700,12 @@ def _build_sso_user_update_data( async def _sync_user_role_from_jwt_role_map( - jwt_handler: Optional[JWTHandler], - received_response: Optional[dict], - user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], + jwt_handler: JWTHandler | None, + received_response: dict | None, + user_info: LiteLLM_UserTable | NewUserResponse | None, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, - user_defined_values: Optional[SSOUserDefinedValues], + user_defined_values: SSOUserDefinedValues | None, ) -> None: """ Apply jwt_litellm_role_map during SSO login. @@ -1755,9 +1747,9 @@ async def _sync_user_role_from_jwt_role_map( def apply_user_info_values_to_sso_user_defined_values( - user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], - user_defined_values: Optional[SSOUserDefinedValues], -) -> Optional[SSOUserDefinedValues]: + user_info: LiteLLM_UserTable | NewUserResponse | None, + user_defined_values: SSOUserDefinedValues | None, +) -> SSOUserDefinedValues | None: if user_defined_values is None: return None if user_info is not None and user_info.user_id is not None: @@ -1787,7 +1779,7 @@ def apply_user_info_values_to_sso_user_defined_values( return user_defined_values -async def check_and_update_if_proxy_admin_id(user_role: str, user_id: str, prisma_client: Optional[PrismaClient]): +async def check_and_update_if_proxy_admin_id(user_role: str, user_id: str, prisma_client: PrismaClient | None): """ - Check if user role in DB is admin - If not, update user role in DB to admin role @@ -1809,7 +1801,7 @@ async def check_and_update_if_proxy_admin_id(user_role: str, user_id: str, prism @router.get("/sso/callback", tags=["experimental"], include_in_schema=False) -async def auth_callback(request: Request, state: Optional[str] = None): +async def auth_callback(request: Request, state: str | None = None): """Verify login""" verbose_proxy_logger.info(f"Starting SSO callback with state: {state}") @@ -1840,7 +1832,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) - sso_jwt_handler: Optional[JWTHandler] = None + sso_jwt_handler: JWTHandler | None = None ui_access_mode = general_settings.get("ui_access_mode", None) if ui_access_mode is not None and isinstance(ui_access_mode, dict): sso_jwt_handler = JWTHandler() @@ -1856,8 +1848,8 @@ async def auth_callback(request: Request, state: Optional[str] = None): microsoft_client_id = os.getenv("MICROSOFT_CLIENT_ID", None) google_client_id = os.getenv("GOOGLE_CLIENT_ID", None) generic_client_id = os.getenv("GENERIC_CLIENT_ID", None) - received_response: Optional[dict] = None - access_token_payload: Optional[dict] = None + received_response: dict | None = None + access_token_payload: dict | None = None sso_assertion: SSOIdentityAssertion | None = None # get url from request if master_key is None: @@ -1922,7 +1914,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # Control-plane cross-origin: read return_to from cookie. # Starlette's cookie_parser already handles RFC 2109 unquoting. - cp_return_to: Optional[str] = request.cookies.get("litellm_cp_return_to") + cp_return_to: str | None = request.cookies.get("litellm_cp_return_to") return await SSOAuthenticationHandler.get_redirect_response_from_openid( result=result, @@ -2013,9 +2005,9 @@ async def saml_callback(request: Request): async def _build_cli_sso_user_defined_values( - result: Union[OpenID, dict], + result: OpenID | dict, parsed_openid_result: ParsedOpenIDResult, -) -> Optional[SSOUserDefinedValues]: +) -> SSOUserDefinedValues | None: from litellm.proxy.proxy_server import user_custom_sso user_id = parsed_openid_result.get("user_id") @@ -2037,9 +2029,9 @@ async def _build_cli_sso_user_defined_values( async def _fetch_cli_sso_team_details( prisma_client: PrismaClient, - teams: List[str], -) -> List[Dict[str, Any]]: - team_details: List[Dict[str, Any]] = [] + teams: list[str], +) -> list[dict[str, Any]]: + team_details: list[dict[str, Any]] = [] try: if teams: prisma_teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": teams}}) @@ -2061,9 +2053,9 @@ async def _complete_cli_sso_callback_session( request: Request, key: str, flow: dict, - result: Union[OpenID, dict], + result: OpenID | dict, parsed_openid_result: ParsedOpenIDResult, - user_defined_values: Optional[SSOUserDefinedValues], + user_defined_values: SSOUserDefinedValues | None, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, cli_sso_session_cache: DualCache, @@ -2091,7 +2083,7 @@ async def _complete_cli_sso_callback_session( await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion) - teams: List[str] = [] + teams: list[str] = [] if hasattr(user_info, "teams") and user_info.teams: teams = user_info.teams if isinstance(user_info.teams, list) else [] @@ -2137,9 +2129,9 @@ async def _complete_cli_sso_callback_session( async def cli_sso_callback( request: Request, - key: Optional[str] = None, - result: Optional[Union[OpenID, dict]] = None, - received_response: Optional[dict] = None, + key: str | None = None, + result: OpenID | dict | None = None, + received_response: dict | None = None, prefill_user_code: str | None = None, sso_assertion: SSOIdentityAssertion | None = None, ): @@ -2166,7 +2158,7 @@ async def cli_sso_callback( ) # After None check, cast to non-None type for type checker - result_non_none: Union[OpenID, dict] = cast(Union[OpenID, dict], result) + result_non_none: OpenID | dict = cast(OpenID | dict, result) try: parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result( @@ -2205,14 +2197,14 @@ async def cli_sso_callback( raise except Exception as e: verbose_proxy_logger.error(f"Error with CLI SSO callback: {e}") - raise HTTPException(status_code=500, detail=f"Failed to process CLI SSO: {str(e)}") + raise HTTPException(status_code=500, detail=f"Failed to process CLI SSO: {e!s}") @router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False) async def cli_poll_key( key_id: str, - team_id: Optional[str] = None, - x_litellm_cli_poll_secret: Optional[str] = Header(default=None), + team_id: str | None = None, + x_litellm_cli_poll_secret: str | None = Header(default=None), ): """ CLI polling endpoint - retrieves session from cache and generates JWT. @@ -2255,12 +2247,12 @@ async def cli_poll_key( verbose_proxy_logger.info(f"Returning teams list for user {user_id} to select from: {user_teams}") # Best-effort construction of team_details if it wasn't # already cached for some reason. - team_details_response: Optional[List[Dict[str, Any]]] = None + team_details_response: list[dict[str, Any]] | None = None if isinstance(user_team_details, list) and user_team_details: team_details_response = user_team_details elif user_teams: team_details_response = [{"team_id": t, "team_alias": None} for t in user_teams] - poll_response: Dict[str, Any] = { + poll_response: dict[str, Any] = { "status": "ready", "user_id": user_id, "teams": user_teams, @@ -2328,12 +2320,12 @@ async def cli_poll_key( raise except Exception as e: verbose_proxy_logger.error(f"Error polling for CLI JWT: {e}") - raise HTTPException(status_code=500, detail=f"Error checking session status: {str(e)}") + raise HTTPException(status_code=500, detail=f"Error checking session status: {e!s}") async def insert_sso_user( - result_openid: Optional[Union[OpenID, dict]], - user_defined_values: Optional[SSOUserDefinedValues] = None, + result_openid: OpenID | dict | None, + user_defined_values: SSOUserDefinedValues | None = None, ) -> NewUserResponse: """ Helper function to create a New User in LiteLLM DB after a successful SSO login @@ -2641,12 +2633,12 @@ class SSOAuthenticationHandler: @staticmethod async def get_sso_login_redirect( redirect_url: str, - google_client_id: Optional[str] = None, - microsoft_client_id: Optional[str] = None, - generic_client_id: Optional[str] = None, - state: Optional[str] = None, - request: Optional[Request] = None, - ) -> Optional[RedirectResponse]: + google_client_id: str | None = None, + microsoft_client_id: str | None = None, + generic_client_id: str | None = None, + state: str | None = None, + request: Request | None = None, + ) -> RedirectResponse | None: """ Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url` @@ -2772,10 +2764,10 @@ class SSOAuthenticationHandler: @staticmethod async def get_generic_sso_redirect_response( generic_sso: Any, - state: Optional[str] = None, - generic_authorization_endpoint: Optional[str] = None, - request: Optional[Request] = None, - ) -> Optional[RedirectResponse]: + state: str | None = None, + generic_authorization_endpoint: str | None = None, + request: Request | None = None, + ) -> RedirectResponse | None: """ Get the redirect response for Generic SSO """ @@ -2884,9 +2876,9 @@ class SSOAuthenticationHandler: @staticmethod def _get_generic_sso_redirect_params( - state: Optional[str] = None, - generic_authorization_endpoint: Optional[str] = None, - ) -> Tuple[dict, Optional[str]]: + state: str | None = None, + generic_authorization_endpoint: str | None = None, + ) -> tuple[dict, str | None]: """ Get redirect parameters for Generic SSO with proper state priority handling. Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled. @@ -2907,7 +2899,7 @@ class SSOAuthenticationHandler: - code_verifier (if PKCE is enabled, None otherwise) """ redirect_params = {} - code_verifier: Optional[str] = None + code_verifier: str | None = None if state: # CLI state takes priority @@ -2937,9 +2929,9 @@ class SSOAuthenticationHandler: @staticmethod def should_use_sso_handler( - google_client_id: Optional[str] = None, - microsoft_client_id: Optional[str] = None, - generic_client_id: Optional[str] = None, + google_client_id: str | None = None, + microsoft_client_id: str | None = None, + generic_client_id: str | None = None, ) -> bool: if google_client_id is not None or microsoft_client_id is not None or generic_client_id is not None: return True @@ -2949,7 +2941,7 @@ class SSOAuthenticationHandler: def get_redirect_url_for_sso( request: Request, sso_callback_route: str, - existing_key: Optional[str] = None, + existing_key: str | None = None, ) -> str: """ Get the redirect URL for SSO @@ -2969,10 +2961,10 @@ class SSOAuthenticationHandler: @staticmethod async def upsert_sso_user( - result: Optional[Union[CustomOpenID, OpenID, dict]], - user_info: Optional[Union[NewUserResponse, LiteLLM_UserTable]], - user_email: Optional[str], - user_defined_values: Optional[SSOUserDefinedValues], + result: CustomOpenID | OpenID | dict | None, + user_info: NewUserResponse | LiteLLM_UserTable | None, + user_email: str | None, + user_defined_values: SSOUserDefinedValues | None, prisma_client: PrismaClient, ): """ @@ -3005,8 +2997,8 @@ class SSOAuthenticationHandler: @staticmethod async def add_user_to_teams_from_sso_response( - result: Optional[Union[CustomOpenID, OpenID, dict]], - user_info: Optional[Union[NewUserResponse, LiteLLM_UserTable]], + result: CustomOpenID | OpenID | dict | None, + user_info: NewUserResponse | LiteLLM_UserTable | None, ): """ Adds the user as a team member to the teams specified in the SSO responses `team_ids` field @@ -3022,9 +3014,9 @@ class SSOAuthenticationHandler: @staticmethod def verify_user_in_restricted_sso_group( - general_settings: Dict, - result: Optional[Union[CustomOpenID, OpenID, dict]], - received_response: Optional[dict], + general_settings: dict, + result: CustomOpenID | OpenID | dict | None, + received_response: dict | None, ) -> Literal[True]: """ when ui_access_mode.type == "restricted_sso_group": @@ -3037,7 +3029,7 @@ class SSOAuthenticationHandler: - if result.team_ids is a list, return True if the restricted_sso_group is in the list, otherwise return False """ - ui_access_mode = cast(Optional[Union[Dict, str]], general_settings.get("ui_access_mode")) + ui_access_mode = cast(dict | str | None, general_settings.get("ui_access_mode")) if ui_access_mode is None: return True @@ -3059,7 +3051,7 @@ class SSOAuthenticationHandler: @staticmethod async def create_litellm_team_from_sso_group( litellm_team_id: str, - litellm_team_name: Optional[str] = None, + litellm_team_name: str | None = None, ): """ Creates a Litellm Team from a SSO Group ID @@ -3116,10 +3108,10 @@ class SSOAuthenticationHandler: @staticmethod def _cast_and_deepcopy_litellm_default_team_params( - default_team_params: Union[DefaultTeamSSOParams, Dict], + default_team_params: DefaultTeamSSOParams | dict, team_request: NewTeamRequest, litellm_team_id: str, - litellm_team_name: Optional[str] = None, + litellm_team_name: str | None = None, ) -> NewTeamRequest: """ Casts and deepcopies the litellm.default_team_params to a NewTeamRequest object @@ -3146,7 +3138,7 @@ class SSOAuthenticationHandler: key: str | None, existing_key: str | None = None, user_code: str | None = None, - ) -> Optional[str]: + ) -> str | None: """ Checks the request 'source' if a cli state token was passed in @@ -3169,15 +3161,15 @@ class SSOAuthenticationHandler: @staticmethod def _get_user_email_and_id_from_result( - result: Optional[Union[OpenID, dict]], - generic_client_id: Optional[str] = None, + result: OpenID | dict | None, + generic_client_id: str | None = None, ) -> ParsedOpenIDResult: """ Gets the user email and id from the OpenID result after validating the email domain """ - user_email: Optional[str] = normalize_email(getattr(result, "email", None)) - user_id: Optional[str] = getattr(result, "id", None) if result is not None else None - user_role: Optional[str] = None + user_email: str | None = normalize_email(getattr(result, "email", None)) + user_id: str | None = getattr(result, "id", None) if result is not None else None + user_role: str | None = None if user_email is not None and os.getenv("ALLOWED_EMAIL_DOMAINS") is not None: email_domain = user_email.split("@")[1] @@ -3186,9 +3178,7 @@ class SSOAuthenticationHandler: raise HTTPException( status_code=401, detail={ - "message": "The email domain={}, is not an allowed email domain={}. Contact your admin to change this.".format( - email_domain, allowed_domains - ) + "message": f"The email domain={email_domain}, is not an allowed email domain={allowed_domains}. Contact your admin to change this." }, ) @@ -3229,14 +3219,14 @@ class SSOAuthenticationHandler: @staticmethod async def get_redirect_response_from_openid( - result: Union[OpenID, dict, CustomOpenID], + result: OpenID | dict | CustomOpenID, request: Request, - received_response: Optional[dict] = None, - generic_client_id: Optional[str] = None, - ui_access_mode: Optional[Dict] = None, - access_token_payload: Optional[dict] = None, - jwt_handler: Optional[JWTHandler] = None, - return_to: Optional[str] = None, + received_response: dict | None = None, + generic_client_id: str | None = None, + ui_access_mode: dict | None = None, + access_token_payload: dict | None = None, + jwt_handler: JWTHandler | None = None, + return_to: str | None = None, sso_assertion: SSOIdentityAssertion | None = None, ) -> RedirectResponse: @@ -3265,13 +3255,13 @@ class SSOAuthenticationHandler: verbose_proxy_logger.info(f"SSO callback result: {result}") user_info = None - user_id_models: List = [] + user_id_models: list = [] max_internal_user_budget = litellm.max_internal_user_budget internal_user_budget_duration = litellm.internal_user_budget_duration # User might not be already created on first generation of key # But if it is, we want their models preferences - user_defined_values: Optional[SSOUserDefinedValues] = None + user_defined_values: SSOUserDefinedValues | None = None if user_custom_sso is not None: if inspect.iscoroutinefunction(user_custom_sso): @@ -3376,7 +3366,7 @@ class SSOAuthenticationHandler: litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/") if get_secret_bool("EXPERIMENTAL_UI_LOGIN"): - _user_info: Optional[LiteLLM_UserTable] = None + _user_info: LiteLLM_UserTable | None = None if user_defined_values is not None and user_defined_values["user_id"] is not None: _user_info = LiteLLM_UserTable( user_id=user_defined_values["user_id"], @@ -3443,7 +3433,7 @@ class SSOAuthenticationHandler: dict: Token exchange parameters """ # Prepare token exchange parameters (may add code_verifier: str later) - token_params: Dict[str, Any] = {"include_client_id": generic_include_client_id} + token_params: dict[str, Any] = {"include_client_id": generic_include_client_id} # Retrieve PKCE code_verifier if PKCE was used in authorization. # Gate on GENERIC_CLIENT_USE_PKCE to avoid an unnecessary Redis round-trip @@ -3526,7 +3516,7 @@ class SSOAuthenticationHandler: @staticmethod async def _handle_missing_pkce_verifier( - state: Optional[str], + state: str | None, cache_key: str, cached_data: object, empty_value_in_dict: bool, @@ -3628,7 +3618,7 @@ class SSOAuthenticationHandler: ) @staticmethod - def generate_pkce_params() -> Tuple[str, str]: + def generate_pkce_params() -> tuple[str, str]: """ Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0. @@ -3715,12 +3705,12 @@ class SSOAuthenticationHandler: authorization_code: str, code_verifier: str, client_id: str, - client_secret: Optional[str], + client_secret: str | None, token_endpoint: str, - userinfo_endpoint: Optional[str], + userinfo_endpoint: str | None, include_client_id: bool, - redirect_url: Optional[str], - additional_headers: Dict[str, str], + redirect_url: str | None, + additional_headers: dict[str, str], ) -> dict: """ Performs a direct OAuth token exchange including the PKCE code_verifier. @@ -3736,7 +3726,7 @@ class SSOAuthenticationHandler: len(code_verifier), ) - token_data: Dict[str, str] = { + token_data: dict[str, str] = { "grant_type": "authorization_code", "code": authorization_code, "code_verifier": code_verifier, @@ -3840,9 +3830,9 @@ class SSOAuthenticationHandler: @staticmethod async def _get_pkce_userinfo( access_token: str, - id_token: Optional[str], - userinfo_endpoint: Optional[str], - additional_headers: Dict[str, str], + id_token: str | None, + userinfo_endpoint: str | None, + additional_headers: dict[str, str], ) -> dict: """ Fetches user info from the userinfo endpoint. @@ -3850,7 +3840,7 @@ class SSOAuthenticationHandler: """ # None = request not yet attempted, failed, or returned empty/null (treated as failure # so the id_token fallback can be attempted instead of returning a session with no claims). - userinfo: Optional[dict] = None + userinfo: dict | None = None if userinfo_endpoint: try: @@ -3978,7 +3968,7 @@ class MicrosoftSSOHandler: microsoft_client_id: str, redirect_url: str, return_raw_sso_response: bool = False, - ) -> Union[CustomOpenID, OpenID, dict]: + ) -> CustomOpenID | OpenID | dict: """ Get the Microsoft SSO callback response @@ -4025,7 +4015,7 @@ class MicrosoftSSOHandler: verbose_proxy_logger.debug(f"Extracted app roles from id_token: {app_roles}") # Combine groups and app roles - user_role: Optional[LitellmUserRoles] = None + user_role: LitellmUserRoles | None = None if app_roles: # Check if any app role is a valid LitellmUserRoles for role_str in app_roles: @@ -4052,9 +4042,9 @@ class MicrosoftSSOHandler: @staticmethod def openid_from_response( - response: Optional[dict], - team_ids: List[str], - user_role: Optional[LitellmUserRoles], + response: dict | None, + team_ids: list[str], + user_role: LitellmUserRoles | None, ) -> CustomOpenID: response = response or {} verbose_proxy_logger.debug(f"Microsoft SSO Callback Response: {response}") @@ -4072,7 +4062,7 @@ class MicrosoftSSOHandler: return openid_response @staticmethod - def get_app_roles_from_id_token(id_token: Optional[str]) -> List[str]: + def get_app_roles_from_id_token(id_token: str | None) -> list[str]: """ Extract app roles from the Microsoft Entra ID (Azure AD) id_token JWT. @@ -4113,8 +4103,8 @@ class MicrosoftSSOHandler: @staticmethod async def get_user_groups_from_graph_api( - access_token: Optional[str] = None, - ) -> List[str]: + access_token: str | None = None, + ) -> list[str]: """ Returns a list of `team_ids` the user belongs to from the Microsoft Graph API @@ -4129,8 +4119,8 @@ class MicrosoftSSOHandler: # Handle MSFT Enterprise Application Groups service_principal_id = os.getenv("MICROSOFT_SERVICE_PRINCIPAL_ID", None) - service_principal_group_ids: Optional[List[str]] = [] - service_principal_teams: Optional[List[MicrosoftServicePrincipalTeam]] = [] + service_principal_group_ids: list[str] | None = [] + service_principal_teams: list[MicrosoftServicePrincipalTeam] | None = [] if service_principal_id: ( service_principal_group_ids, @@ -4148,7 +4138,7 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[str] = MicrosoftSSOHandler.get_graph_api_user_groups_endpoint() + next_link: str | None = MicrosoftSSOHandler.get_graph_api_user_groups_endpoint() auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 @@ -4177,7 +4167,7 @@ class MicrosoftSSOHandler: @staticmethod async def fetch_and_parse_groups( url: str, headers: dict, async_client: AsyncHTTPHandler - ) -> Tuple[List[str], Optional[str]]: + ) -> tuple[list[str], str | None]: """Helper function to fetch and parse group data from a URL""" response = await async_client.get(url, headers=headers) response_json = response.json() @@ -4188,7 +4178,7 @@ class MicrosoftSSOHandler: @staticmethod def _get_group_ids_from_graph_api_response( response: MicrosoftGraphAPIUserGroupResponse, - ) -> List[str]: + ) -> list[str]: group_ids = [] for _object in response.get("value", []) or []: _group_id = _object.get("id") @@ -4200,7 +4190,7 @@ class MicrosoftSSOHandler: async def _cast_graph_api_response_dict( response: dict, ) -> MicrosoftGraphAPIUserGroupResponse: - directory_objects: List[MicrosoftGraphAPIUserGroupDirectoryObject] = [] + directory_objects: list[MicrosoftGraphAPIUserGroupDirectoryObject] = [] for _object in response.get("value", []): directory_objects.append( MicrosoftGraphAPIUserGroupDirectoryObject( @@ -4222,8 +4212,8 @@ class MicrosoftSSOHandler: async def get_group_ids_from_service_principal( service_principal_id: str, async_client: AsyncHTTPHandler, - access_token: Optional[str] = None, - ) -> Tuple[List[str], List[MicrosoftServicePrincipalTeam]]: + access_token: str | None = None, + ) -> tuple[list[str], list[MicrosoftServicePrincipalTeam]]: """ Gets the groups belonging to the Service Principal Application @@ -4241,8 +4231,8 @@ class MicrosoftSSOHandler: "Content-Type": "application/json", } - group_ids: List[str] = [] - service_principal_teams: List[MicrosoftServicePrincipalTeam] = [] + group_ids: list[str] = [] + service_principal_teams: list[MicrosoftServicePrincipalTeam] = [] page_count = 0 while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES: @@ -4274,7 +4264,7 @@ class MicrosoftSSOHandler: @staticmethod async def create_litellm_teams_from_service_principal_team_ids( - service_principal_teams: List[MicrosoftServicePrincipalTeam], + service_principal_teams: list[MicrosoftServicePrincipalTeam], ): """ Creates Litellm Teams from the Service Principal Group IDs @@ -4283,8 +4273,8 @@ class MicrosoftSSOHandler: """ verbose_proxy_logger.debug(f"Creating Litellm Teams from Service Principal Teams: {service_principal_teams}") for service_principal_team in service_principal_teams: - litellm_team_id: Optional[str] = service_principal_team.get("principalId") - litellm_team_name: Optional[str] = service_principal_team.get("principalDisplayName") + litellm_team_id: str | None = service_principal_team.get("principalId") + litellm_team_name: str | None = service_principal_team.get("principalDisplayName") if not litellm_team_id: verbose_proxy_logger.debug( f"Skipping team creation for {litellm_team_name} because it has no principalId" @@ -4308,7 +4298,7 @@ class GoogleSSOHandler: google_client_id: str, redirect_url: str, return_raw_sso_response: bool = False, - ) -> Union[OpenID, dict]: + ) -> OpenID | dict: """ Get the Google SSO callback response @@ -4410,7 +4400,7 @@ async def debug_sso_callback(request: Request): user_api_key_cache, ) - sso_jwt_handler: Optional[JWTHandler] = None + sso_jwt_handler: JWTHandler | None = None ui_access_mode = general_settings.get("ui_access_mode", None) if ui_access_mode is not None and isinstance(ui_access_mode, dict): sso_jwt_handler = JWTHandler() @@ -4434,8 +4424,8 @@ async def debug_sso_callback(request: Request): redirect_url += "/sso/debug/callback" result = None - received_response: Optional[dict] = None - access_token_payload: Optional[dict] = None + received_response: dict | None = None + access_token_payload: dict | None = None if google_client_id is not None: result = await GoogleSSOHandler.get_google_callback_response( request=request, @@ -4489,7 +4479,7 @@ async def debug_sso_callback(request: Request): # Try to convert to string or another JSON serializable format filtered_result[key] = str(value) except Exception as e: - filtered_result[key] = f"Complex value (not displayable): {str(e)}" + filtered_result[key] = f"Complex value (not displayable): {e!s}" # Defense-in-depth: ensure no bearer tokens leak into the rendered HTML even if # a non-conforming IdP places them in its userinfo response. diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index 9b8eb435803..67df88da941 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -6,7 +6,7 @@ usage/spend data by querying the aggregated daily activity endpoints. import json from collections.abc import AsyncIterator, Callable from datetime import date -from typing import Any, Dict, List, Literal, Optional, cast +from typing import Any, Literal, cast from typing_extensions import TypedDict @@ -51,7 +51,7 @@ class SSEToolCallEvent(TypedDict, total=False): type: Literal["tool_call"] tool_name: str tool_label: str - arguments: Dict[str, str] + arguments: dict[str, str] status: Literal["running", "complete", "error"] error: str @@ -75,7 +75,7 @@ SSEEvent = SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SS class ToolHandler(TypedDict): fetch: Callable[..., Any] - summarise: Callable[[Dict[str, Any]], str] + summarise: Callable[[dict[str, Any]], str] label: str @@ -159,7 +159,7 @@ TOOLS_BASE = [_TOOL_USAGE] TOOLS_ADMIN = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG] -def get_tools_for_role(is_admin: bool) -> List[Dict[str, Any]]: +def get_tools_for_role(is_admin: bool) -> list[dict[str, Any]]: """Return the tool list appropriate for the user's role.""" return TOOLS_ADMIN if is_admin else TOOLS_BASE @@ -205,7 +205,7 @@ SYSTEM_PROMPT = _SYSTEM_PROMPT_BASE # --------------------------------------------------------------------------- -def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]: +def _parse_csv_ids(raw: str | None) -> list[str] | None: if not raw: return None return [t.strip() for t in raw.split(",") if t.strip()] @@ -214,7 +214,7 @@ def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]: async def _query_activity( table_name: str, entity_id_field: str, - entity_id: Optional[Any], + entity_id: Any | None, start_date: str, end_date: str, *, @@ -254,7 +254,7 @@ async def _query_activity( ) -async def _fetch_usage_data(start_date: str, end_date: str, user_id: Optional[str] = None) -> Dict[str, Any]: +async def _fetch_usage_data(start_date: str, end_date: str, user_id: str | None = None) -> dict[str, Any]: resp = await _query_activity( TABLE_DAILY_USER_SPEND, ENTITY_FIELD_USER, @@ -266,7 +266,7 @@ async def _fetch_usage_data(start_date: str, end_date: str, user_id: Optional[st return resp.model_dump(mode="json") -async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: Optional[str] = None) -> Dict[str, Any]: +async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: str | None = None) -> dict[str, Any]: resp = await _query_activity( TABLE_DAILY_TEAM_SPEND, ENTITY_FIELD_TEAM, @@ -277,7 +277,7 @@ async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: Optio return resp.model_dump(mode="json") -async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: Optional[str] = None) -> Dict[str, Any]: +async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: str | None = None) -> dict[str, Any]: resp = await _query_activity( TABLE_DAILY_TAG_SPEND, ENTITY_FIELD_TAG, @@ -294,10 +294,10 @@ async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: Optional[s def _accumulate_breakdown( - results: List[Dict[str, Any]], dimension: str, fields: List[str] -) -> Dict[str, Dict[str, float]]: + results: list[dict[str, Any]], dimension: str, fields: list[str] +) -> dict[str, dict[str, float]]: """Aggregate a single breakdown dimension across days.""" - totals: Dict[str, Dict[str, float]] = {} + totals: dict[str, dict[str, float]] = {} for day in results: for key, entry in day.get("breakdown", {}).get(dimension, {}).items(): if key not in totals: @@ -309,15 +309,15 @@ def _accumulate_breakdown( def _ranked_lines( - totals: Dict[str, Dict[str, float]], - fmt: Callable[[str, Dict[str, float]], str], + totals: dict[str, dict[str, float]], + fmt: Callable[[str, dict[str, float]], str], limit: int, -) -> List[str]: +) -> list[str]: """Sort by spend descending, format each entry, and truncate.""" return [fmt(name, vals) for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[:limit]] -def _summarise_usage_data(data: Dict[str, Any]) -> str: +def _summarise_usage_data(data: dict[str, Any]) -> str: meta = data.get("metadata", {}) results = data.get("results", []) @@ -349,13 +349,13 @@ def _summarise_usage_data(data: Dict[str, Any]) -> str: return "\n".join(sections) -def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str: +def _summarise_entity_data(data: dict[str, Any], entity_label: str) -> str: """Summarise team/tag entity usage data.""" results = data.get("results", []) if not results: return f"No {entity_label} usage data found for the given date range." - totals: Dict[str, Dict[str, Any]] = {} + totals: dict[str, dict[str, Any]] = {} for day in results: for eid, entry in day.get("breakdown", {}).get("entities", {}).items(): if eid not in totals: @@ -379,7 +379,7 @@ def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str: # Tool dispatch registry # --------------------------------------------------------------------------- -TOOL_HANDLERS: Dict[str, ToolHandler] = { +TOOL_HANDLERS: dict[str, ToolHandler] = { "get_usage_data": ToolHandler( fetch=_fetch_usage_data, summarise=_summarise_usage_data, @@ -409,16 +409,16 @@ def _sse(event: SSEEvent) -> str: def _resolve_fetch_kwargs( fn_name: str, - fn_args: Dict[str, str], - user_id: Optional[str], + fn_args: dict[str, str], + user_id: str | None, is_admin: bool, -) -> Dict[str, Any]: +) -> dict[str, Any]: """Build keyword arguments for a tool's fetch function.""" start_date = fn_args.get("start_date", "") end_date = fn_args.get("end_date", "") if not start_date or not end_date: raise ValueError("Missing required start_date or end_date from tool arguments") - kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date} + kwargs: dict[str, Any] = {"start_date": start_date, "end_date": end_date} if fn_name == "get_usage_data": if not is_admin: if user_id is None: @@ -443,8 +443,8 @@ def _resolve_fetch_kwargs( async def _execute_tool_call( handler: ToolHandler, fn_name: str, - fn_args: Dict[str, str], - user_id: Optional[str], + fn_args: dict[str, str], + user_id: str | None, is_admin: bool, ) -> str: """Run a single tool and return the summarised result text.""" @@ -455,8 +455,8 @@ async def _execute_tool_call( async def _process_tool_call( tc: Any, - chat_messages: List[Dict[str, Any]], - user_id: Optional[str], + chat_messages: list[dict[str, Any]], + user_id: str | None, is_admin: bool, ) -> AsyncIterator[str]: """Execute a single tool call, yielding SSE events for status.""" @@ -495,7 +495,7 @@ async def _process_tool_call( chat_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_result}) -async def _stream_final_response(model: str, chat_messages: List[Dict[str, Any]]) -> AsyncIterator[str]: +async def _stream_final_response(model: str, chat_messages: list[dict[str, Any]]) -> AsyncIterator[str]: """Stream the final LLM response after tool results are appended.""" yield _sse({"type": "status", "message": "Analyzing results..."}) @@ -512,15 +512,15 @@ async def _stream_final_response(model: str, chat_messages: List[Dict[str, Any]] async def stream_usage_ai_chat( - messages: List[Dict[str, str]], - model: Optional[str] = None, - user_id: Optional[str] = None, + messages: list[dict[str, str]], + model: str | None = None, + user_id: str | None = None, is_admin: bool = False, ) -> AsyncIterator[str]: """Stream SSE events: status → tool_call → chunk → done.""" resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL truncated = messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages - chat_messages: List[Dict[str, Any]] = [ + chat_messages: list[dict[str, Any]] = [ {"role": "system", "content": _build_system_prompt(is_admin)}, *truncated, ] diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index 26515c749fb..272cc99e57f 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -4,7 +4,7 @@ USAGE AI CHAT ENDPOINTS /usage/ai/chat - Stream AI chat responses about usage data """ -from typing import List, Literal, Optional +from typing import Literal from fastapi import APIRouter, Depends, Request from fastapi.responses import StreamingResponse @@ -22,8 +22,8 @@ class ChatMessage(BaseModel): class UsageAIChatRequest(BaseModel): - messages: List[ChatMessage] = Field(..., description="Chat messages (user/assistant history)") - model: Optional[str] = Field(default=None, description="Model to use for AI chat") + messages: list[ChatMessage] = Field(..., description="Chat messages (user/assistant history)") + model: str | None = Field(default=None, description="Model to use for AI chat") @router.post( diff --git a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py index af65eb6c3d8..20f4e91b030 100644 --- a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py +++ b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py @@ -12,7 +12,7 @@ user metrics from tag activity data and return time series for dashboard visuali """ from datetime import datetime, timedelta -from typing import Any, Dict, List, Optional +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Query from pydantic import BaseModel @@ -40,14 +40,14 @@ class TagActiveUsersResponse(BaseModel): tag: str active_users: int date: str # The specific date or period identifier - period_start: Optional[str] = None # For WAU/MAU, this will be the start of the period - period_end: Optional[str] = None # For WAU/MAU, this will be the end of the period + period_start: str | None = None # For WAU/MAU, this will be the start of the period + period_end: str | None = None # For WAU/MAU, this will be the end of the period class ActiveUsersAnalyticsResponse(BaseModel): """Response for active users analytics""" - results: List[TagActiveUsersResponse] + results: list[TagActiveUsersResponse] class TagSummaryMetrics(BaseModel): @@ -65,7 +65,7 @@ class TagSummaryMetrics(BaseModel): class TagSummaryResponse(BaseModel): """Response for tag summary analytics""" - results: List[TagSummaryMetrics] + results: list[TagSummaryMetrics] class DistinctTagResponse(BaseModel): @@ -77,15 +77,15 @@ class DistinctTagResponse(BaseModel): class DistinctTagsResponse(BaseModel): """Response for all distinct user agent tags""" - results: List[DistinctTagResponse] + results: list[DistinctTagResponse] class PerUserMetrics(BaseModel): """Metrics for individual user""" user_id: str - user_email: Optional[str] = None - user_agent: Optional[str] = None + user_email: str | None = None + user_agent: str | None = None successful_requests: int = 0 failed_requests: int = 0 total_requests: int = 0 @@ -96,7 +96,7 @@ class PerUserMetrics(BaseModel): class PerUserAnalyticsResponse(BaseModel): """Response for per-user analytics""" - results: List[PerUserMetrics] + results: list[PerUserMetrics] total_count: int page: int page_size: int @@ -150,7 +150,7 @@ async def get_distinct_user_agent_tags( except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch distinct user agent tags: {str(e)}", + detail=f"Failed to fetch distinct user agent tags: {e!s}", ) @@ -161,11 +161,11 @@ async def get_distinct_user_agent_tags( dependencies=[Depends(user_api_key_auth)], ) async def get_daily_active_users( - tag_filter: Optional[str] = Query( + tag_filter: str | None = Query( default=None, description="Filter by specific tag (optional)", ), - tag_filters: Optional[List[str]] = Query( + tag_filters: list[str] | None = Query( default=None, description="Filter by multiple specific tags (optional, takes precedence over tag_filter)", ), @@ -243,7 +243,7 @@ async def get_daily_active_users( except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch DAU analytics: {str(e)}", + detail=f"Failed to fetch DAU analytics: {e!s}", ) @@ -254,11 +254,11 @@ async def get_daily_active_users( dependencies=[Depends(user_api_key_auth)], ) async def get_weekly_active_users( - tag_filter: Optional[str] = Query( + tag_filter: str | None = Query( default=None, description="Filter by specific tag (optional)", ), - tag_filters: Optional[List[str]] = Query( + tag_filters: list[str] | None = Query( default=None, description="Filter by multiple specific tags (optional, takes precedence over tag_filter)", ), @@ -364,7 +364,7 @@ async def get_weekly_active_users( except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch WAU analytics: {str(e)}", + detail=f"Failed to fetch WAU analytics: {e!s}", ) @@ -375,11 +375,11 @@ async def get_weekly_active_users( dependencies=[Depends(user_api_key_auth)], ) async def get_monthly_active_users( - tag_filter: Optional[str] = Query( + tag_filter: str | None = Query( default=None, description="Filter by specific tag (optional)", ), - tag_filters: Optional[List[str]] = Query( + tag_filters: list[str] | None = Query( default=None, description="Filter by multiple specific tags (optional, takes precedence over tag_filter)", ), @@ -485,7 +485,7 @@ async def get_monthly_active_users( except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch MAU analytics: {str(e)}", + detail=f"Failed to fetch MAU analytics: {e!s}", ) @@ -498,11 +498,11 @@ async def get_monthly_active_users( async def get_tag_summary( start_date: str = Query(description="Start date in YYYY-MM-DD format"), end_date: str = Query(description="End date in YYYY-MM-DD format"), - tag_filter: Optional[str] = Query( + tag_filter: str | None = Query( default=None, description="Filter by specific tag (optional)", ), - tag_filters: Optional[List[str]] = Query( + tag_filters: list[str] | None = Query( default=None, description="Filter by multiple specific tags (optional, takes precedence over tag_filter)", ), @@ -585,12 +585,12 @@ async def get_tag_summary( except ValueError as e: raise HTTPException( status_code=400, - detail=f"Invalid date format. Use YYYY-MM-DD: {str(e)}", + detail=f"Invalid date format. Use YYYY-MM-DD: {e!s}", ) except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch tag summary analytics: {str(e)}", + detail=f"Failed to fetch tag summary analytics: {e!s}", ) @@ -601,11 +601,11 @@ async def get_tag_summary( dependencies=[Depends(user_api_key_auth)], ) async def get_per_user_analytics( - tag_filter: Optional[str] = Query( + tag_filter: str | None = Query( default=None, description="Filter by specific tag (optional)", ), - tag_filters: Optional[List[str]] = Query( + tag_filters: list[str] | None = Query( default=None, description="Filter by multiple specific tags (optional, takes precedence over tag_filter)", ), @@ -648,7 +648,7 @@ async def get_per_user_analytics( start_date = start_dt.strftime("%Y-%m-%d") # Build where clause with date range - where_clause: Dict[str, Any] = {"date": {"gte": start_date, "lte": end_date}} + where_clause: dict[str, Any] = {"date": {"gte": start_date, "lte": end_date}} # Add tag filtering if provided if tag_filters and len(tag_filters) > 0: @@ -687,7 +687,7 @@ async def get_per_user_analytics( user_id_to_email = {record.user_id: record.user_email for record in user_records} # Aggregate metrics by user - user_metrics: Dict[str, PerUserMetrics] = {} + user_metrics: dict[str, PerUserMetrics] = {} for record in tag_records: if record.api_key in api_key_to_user_id: @@ -740,5 +740,5 @@ async def get_per_user_analytics( except Exception as e: raise HTTPException( status_code=500, - detail=f"Failed to fetch per-user analytics: {str(e)}", + detail=f"Failed to fetch per-user analytics: {e!s}", ) diff --git a/litellm/proxy/management_endpoints/workflow_management_endpoints.py b/litellm/proxy/management_endpoints/workflow_management_endpoints.py index b2488d20127..568cbbb7b4f 100644 --- a/litellm/proxy/management_endpoints/workflow_management_endpoints.py +++ b/litellm/proxy/management_endpoints/workflow_management_endpoints.py @@ -14,7 +14,7 @@ GET /v1/workflows/runs/{run_id}/messages - Fetch conversation history """ import json -from typing import Any, Dict, Literal, Optional +from typing import Any, Literal from fastapi import APIRouter, Depends, HTTPException, Query @@ -47,13 +47,13 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value -def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> Optional[str]: +def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> str | None: """Return the hashed key token that identifies this caller, or None for master key.""" return user_api_key_dict.token # Status transitions driven by event_type -_EVENT_STATUS_MAP: Dict[str, str] = { +_EVENT_STATUS_MAP: dict[str, str] = { "step.started": "running", "step.failed": "failed", "hook.waiting": "paused", @@ -68,29 +68,29 @@ _EVENT_STATUS_MAP: Dict[str, str] = { class WorkflowRunCreateRequest(BaseModel): workflow_type: str - input: Optional[Dict[str, Any]] = None - metadata: Optional[Dict[str, Any]] = None + input: dict[str, Any] | None = None + metadata: dict[str, Any] | None = None WorkflowRunStatus = Literal["pending", "running", "paused", "completed", "failed"] class WorkflowRunUpdateRequest(BaseModel): - status: Optional[WorkflowRunStatus] = None - output: Optional[Dict[str, Any]] = None - metadata: Optional[Dict[str, Any]] = None + status: WorkflowRunStatus | None = None + output: dict[str, Any] | None = None + metadata: dict[str, Any] | None = None class WorkflowEventCreateRequest(BaseModel): event_type: str step_name: str - data: Optional[Dict[str, Any]] = None + data: dict[str, Any] | None = None class WorkflowMessageCreateRequest(BaseModel): role: str content: str - session_id: Optional[str] = None + session_id: str | None = None # --------------------------------------------------------------------------- @@ -118,7 +118,7 @@ async def _get_next_sequence_number(prisma_client: Any, run_id: str, table: str) async def _require_run( prisma_client: Any, run_id: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ) -> Any: """Return the run or raise 404. For non-admin callers, also enforce key ownership.""" run = await WorkflowRunRepository(prisma_client).table.find_unique(where={"run_id": run_id}) @@ -156,7 +156,7 @@ async def create_workflow_run( raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) try: - create_data: Dict[str, Any] = { + create_data: dict[str, Any] = { "workflow_type": data.workflow_type, "created_by": _caller_key(user_api_key_dict), } @@ -177,8 +177,8 @@ async def create_workflow_run( dependencies=[Depends(user_api_key_auth)], ) async def list_workflow_runs( - workflow_type: Optional[str] = Query(None), - status: Optional[str] = Query(None), + workflow_type: str | None = Query(None), + status: str | None = Query(None), limit: int = Query(50, ge=1, le=250), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): @@ -191,7 +191,7 @@ async def list_workflow_runs( if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) - where: Dict[str, Any] = {} + where: dict[str, Any] = {} if workflow_type: where["workflow_type"] = workflow_type if status: @@ -266,7 +266,7 @@ async def update_workflow_run( if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) - update: Dict[str, Any] = {} + update: dict[str, Any] = {} if data.status is not None: update["status"] = data.status if data.output is not None: @@ -323,7 +323,7 @@ async def append_workflow_event( for attempt in range(_MAX_SEQUENCE_RETRIES): try: seq = await _get_next_sequence_number(prisma_client, run_id, "events") - event_data: Dict[str, Any] = { + event_data: dict[str, Any] = { "run_id": run_id, "event_type": data.event_type, "step_name": data.step_name, @@ -415,7 +415,7 @@ async def append_workflow_message( for attempt in range(_MAX_SEQUENCE_RETRIES): try: seq = await _get_next_sequence_number(prisma_client, run_id, "messages") - msg_data: Dict[str, Any] = { + msg_data: dict[str, Any] = { "run_id": run_id, "role": data.role, "content": data.content, diff --git a/litellm/proxy/management_helpers/audit_logs.py b/litellm/proxy/management_helpers/audit_logs.py index c184ce6bca5..8677627c607 100644 --- a/litellm/proxy/management_helpers/audit_logs.py +++ b/litellm/proxy/management_helpers/audit_logs.py @@ -5,7 +5,6 @@ Functions to create audit logs for LiteLLM Proxy import asyncio import json from datetime import datetime, timezone -from typing import Dict import litellm from litellm._logging import verbose_proxy_logger @@ -15,13 +14,12 @@ from litellm.proxy._types import ( AUDIT_ACTIONS, LiteLLM_AuditLogs, LitellmTableNames, - Optional, UserAPIKeyAuth, ) from litellm.repositories.table_repositories import AuditLogRepository from litellm.types.utils import StandardAuditLogPayload -_audit_log_callback_cache: Dict[str, CustomLogger] = {} +_audit_log_callback_cache: dict[str, CustomLogger] = {} ALLOW_LITELLM_CHANGED_BY_HEADER_METADATA_KEY = "allow_litellm_changed_by_header" @@ -37,16 +35,16 @@ def _allows_litellm_changed_by_header(user_api_key_dict: UserAPIKeyAuth) -> bool def get_audit_log_changed_by( *, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, user_api_key_dict: UserAPIKeyAuth, - litellm_proxy_admin_name: Optional[str], -) -> Optional[str]: + litellm_proxy_admin_name: str | None, +) -> str | None: if litellm_changed_by and _allows_litellm_changed_by_header(user_api_key_dict): return litellm_changed_by return user_api_key_dict.user_id or litellm_proxy_admin_name -def _resolve_audit_log_callback(name: str) -> Optional[CustomLogger]: +def _resolve_audit_log_callback(name: str) -> CustomLogger | None: """Resolve a string callback name to a CustomLogger instance, with caching. For "s3_v2" with `litellm.s3_audit_callback_params` set, constructs a @@ -56,7 +54,7 @@ def _resolve_audit_log_callback(name: str) -> Optional[CustomLogger]: if name in _audit_log_callback_cache: return _audit_log_callback_cache[name] - instance: Optional[CustomLogger] + instance: CustomLogger | None if name == "s3_v2" and getattr(litellm, "s3_audit_callback_params", None) is not None: from litellm.integrations.s3_v2 import S3Logger as S3V2Logger @@ -130,7 +128,7 @@ async def _dispatch_audit_log_to_callbacks( for callback in litellm.audit_log_callbacks: try: - resolved: Optional[CustomLogger] = callback if isinstance(callback, CustomLogger) else None + resolved: CustomLogger | None = callback if isinstance(callback, CustomLogger) else None if isinstance(callback, str): resolved = _resolve_audit_log_callback(callback) if resolved is None: @@ -147,12 +145,12 @@ async def _dispatch_audit_log_to_callbacks( async def create_object_audit_log( object_id: str, action: AUDIT_ACTIONS, - litellm_changed_by: Optional[str], + litellm_changed_by: str | None, user_api_key_dict: UserAPIKeyAuth, - litellm_proxy_admin_name: Optional[str], + litellm_proxy_admin_name: str | None, table_name: LitellmTableNames, - before_value: Optional[str] = None, - after_value: Optional[str] = None, + before_value: str | None = None, + after_value: str | None = None, ): """ Create an audit log for an internal user. @@ -167,7 +165,7 @@ async def create_object_audit_log( """ from litellm.secret_managers.main import get_secret_bool - _store_audit_logs: Optional[bool] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS") + _store_audit_logs: bool | None = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS") if _store_audit_logs is not True: return @@ -199,7 +197,7 @@ async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs): """ from litellm.secret_managers.main import get_secret_bool - _store_audit_logs: Optional[bool] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS") + _store_audit_logs: bool | None = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS") if _store_audit_logs is not True: return diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 72b1e38f406..24bfbe2f9a9 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -6,7 +6,7 @@ organizations, teams, and keys. import json from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union +from typing import TYPE_CHECKING, Any, Optional from fastapi import HTTPException, status @@ -26,9 +26,9 @@ if TYPE_CHECKING: async def attach_object_permission_to_dict( - data_dict: Dict, + data_dict: dict, prisma_client: PrismaClient, -) -> Dict: +) -> dict: """ Helper method to attach object_permission to a dictionary if object_permission_id is set. @@ -118,10 +118,10 @@ async def prepare_object_permission_upsert( async def handle_update_object_permission_common( - data_json: Dict, - existing_object_permission_id: Optional[str], - prisma_client: Optional[PrismaClient], -) -> Optional[str]: + data_json: dict, + existing_object_permission_id: str | None, + prisma_client: PrismaClient | None, +) -> str | None: """ Common logic for handling object permission updates across organizations, teams, and keys. @@ -146,7 +146,7 @@ async def handle_update_object_permission_common( if prisma_client is None: raise ValueError("Prisma client not found") - new_object_permission: Union[dict, str, None] = data_json.pop("object_permission", None) + new_object_permission: dict | str | None = data_json.pop("object_permission", None) if new_object_permission is None: return None @@ -173,7 +173,7 @@ async def handle_update_object_permission_common( async def _set_object_permission( data_json: dict, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, ): """ Creates the LiteLLM_ObjectPermissionTable record for the key/team. @@ -201,9 +201,9 @@ async def _set_object_permission( return data_json -def _dedupe_preserving_order(values: List[str]) -> List[str]: - seen: Set[str] = set() - result: List[str] = [] +def _dedupe_preserving_order(values: list[str]) -> list[str]: + seen: set[str] = set() + result: list[str] = [] for value in values: if value in seen: continue @@ -222,9 +222,9 @@ def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool: async def _get_db_mcp_servers_by_identifiers( - identifiers: Set[str], - prisma_client: Optional[PrismaClient], -) -> List[Any]: + identifiers: set[str], + prisma_client: PrismaClient | None, +) -> list[Any]: if prisma_client is None or not identifiers: return [] @@ -241,9 +241,9 @@ async def _get_db_mcp_servers_by_identifiers( async def _resolve_mcp_server_identifiers_to_ids( - identifiers: Set[str], - prisma_client: Optional[PrismaClient], -) -> Dict[str, Set[str]]: + identifiers: set[str], + prisma_client: PrismaClient | None, +) -> dict[str, set[str]]: """ Resolve MCP permission entries written as server_id, alias, or server_name to canonical server IDs. @@ -258,7 +258,7 @@ async def _resolve_mcp_server_identifiers_to_ids( global_mcp_server_manager, ) - resolved: Dict[str, Set[str]] = {identifier: set() for identifier in identifiers} + resolved: dict[str, set[str]] = {identifier: set() for identifier in identifiers} for server in await _get_db_mcp_servers_by_identifiers( identifiers=identifiers, @@ -284,13 +284,13 @@ async def _resolve_mcp_server_identifiers_to_ids( def _rewrite_object_permission_mcp_servers( object_permission: ObjectPermissionDict, - identifier_to_server_ids: Dict[str, Set[str]], + identifier_to_server_ids: dict[str, set[str]], ) -> None: mcp_servers = object_permission.get("mcp_servers") if not isinstance(mcp_servers, list): return - normalized_servers: List[str] = [] + normalized_servers: list[str] = [] for identifier in mcp_servers: if identifier == SpecialMCPServerNames.no_mcp_servers.value: normalized_servers.append(SpecialMCPServerNames.no_mcp_servers.value) @@ -301,13 +301,13 @@ def _rewrite_object_permission_mcp_servers( def _rewrite_object_permission_mcp_tool_permissions( object_permission: ObjectPermissionDict, - identifier_to_server_ids: Dict[str, Set[str]], + identifier_to_server_ids: dict[str, set[str]], ) -> None: mcp_tool_permissions = object_permission.get("mcp_tool_permissions") if not isinstance(mcp_tool_permissions, dict): return - normalized_tool_permissions: Dict[str, List[str]] = {} + normalized_tool_permissions: dict[str, list[str]] = {} for identifier, tools in mcp_tool_permissions.items(): if not isinstance(tools, list): tools = [] @@ -321,8 +321,8 @@ def _rewrite_object_permission_mcp_tool_permissions( def _rewrite_object_permission_mcp_identifiers( - object_permission: Optional[ObjectPermissionDict], - identifier_to_server_ids: Dict[str, Set[str]], + object_permission: ObjectPermissionDict | None, + identifier_to_server_ids: dict[str, set[str]], ) -> None: if not object_permission or not isinstance(object_permission, dict): return @@ -338,15 +338,15 @@ def _rewrite_object_permission_mcp_identifiers( def _flatten_resolved_mcp_server_ids( - identifier_to_server_ids: Dict[str, Set[str]], -) -> Set[str]: + identifier_to_server_ids: dict[str, set[str]], +) -> set[str]: return {server_id for server_ids in identifier_to_server_ids.values() for server_id in server_ids} async def _resolve_team_allowed_mcp_servers( team_object_permission: "LiteLLM_ObjectPermissionTable", - prisma_client: Optional[PrismaClient] = None, -) -> Set[str]: + prisma_client: PrismaClient | None = None, +) -> set[str]: """ Resolve the full set of MCP server IDs a team has access to. @@ -359,16 +359,16 @@ async def _resolve_team_allowed_mcp_servers( MCPRequestHandler, ) - direct_servers: List[str] = team_object_permission.mcp_servers or [] + direct_servers: list[str] = team_object_permission.mcp_servers or [] if SpecialMCPServerName.all_proxy_servers.value in direct_servers: return _get_all_mcp_server_ids() - access_group_servers: List[str] = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: list[str] = await MCPRequestHandler._get_mcp_servers_from_access_groups( team_object_permission.mcp_access_groups or [] ) raw_tool_perms = team_object_permission.mcp_tool_permissions or {} if isinstance(raw_tool_perms, str): raw_tool_perms = json.loads(raw_tool_perms) - tool_perm_servers: List[str] = list(raw_tool_perms.keys()) + tool_perm_servers: list[str] = list(raw_tool_perms.keys()) raw_servers = set(direct_servers + access_group_servers + tool_perm_servers) resolved_servers = await _resolve_mcp_server_identifiers_to_ids( identifiers=raw_servers, @@ -378,7 +378,7 @@ async def _resolve_team_allowed_mcp_servers( return _flatten_resolved_mcp_server_ids(resolved_servers) | unresolved_servers -def _get_allow_all_keys_server_ids() -> Set[str]: +def _get_allow_all_keys_server_ids() -> set[str]: """Return the set of MCP server IDs marked with allow_all_keys=True.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -397,8 +397,8 @@ def _get_all_mcp_server_ids() -> set[str]: async def _existing_object_permission_mcp_servers( - object_permission_id: Optional[str], - prisma_client: Optional[PrismaClient], + object_permission_id: str | None, + prisma_client: PrismaClient | None, ) -> list[str]: if not object_permission_id or prisma_client is None: return [] @@ -411,10 +411,10 @@ async def _existing_object_permission_mcp_servers( async def enforce_all_proxy_mcp_servers_grant_is_admin_only( - requested_mcp_servers: Optional[list[str]], - existing_object_permission_id: Optional[str], + requested_mcp_servers: list[str] | None, + existing_object_permission_id: str | None, is_proxy_admin: bool, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, ) -> None: """ Only a proxy admin may newly grant the all-proxy MCP sentinel. @@ -445,8 +445,8 @@ async def enforce_all_proxy_mcp_servers_grant_is_admin_only( async def _get_team_allowed_mcp_servers( team_obj: Optional["LiteLLM_TeamTableCachedObj"], - prisma_client: Optional[PrismaClient] = None, -) -> Set[str]: + prisma_client: PrismaClient | None = None, +) -> set[str]: """ Get the full set of MCP server IDs a team allows. @@ -467,8 +467,8 @@ async def _get_team_allowed_mcp_servers( def _extract_requested_mcp_server_ids( - object_permission: Optional[ObjectPermissionDict], -) -> Set[str]: + object_permission: ObjectPermissionDict | None, +) -> set[str]: """ Extract all MCP server IDs referenced in a key's object_permission dict. @@ -479,7 +479,7 @@ def _extract_requested_mcp_server_ids( if not object_permission or not isinstance(object_permission, dict): return set() - server_ids: Set[str] = set() + server_ids: set[str] = set() mcp_servers = object_permission.get("mcp_servers") if isinstance(mcp_servers, list): server_ids.update(mcp_servers) @@ -493,8 +493,8 @@ def _extract_requested_mcp_server_ids( def _extract_requested_mcp_access_groups( - object_permission: Optional[ObjectPermissionDict], -) -> Set[str]: + object_permission: ObjectPermissionDict | None, +) -> set[str]: """Extract MCP access groups from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): return set() @@ -506,8 +506,8 @@ def _extract_requested_mcp_access_groups( def _extract_requested_mcp_toolsets( - object_permission: Optional[ObjectPermissionDict], -) -> Set[str]: + object_permission: ObjectPermissionDict | None, +) -> set[str]: """Extract MCP toolset IDs from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): return set() @@ -519,11 +519,11 @@ def _extract_requested_mcp_toolsets( async def validate_key_mcp_servers_against_team( - object_permission: Optional[ObjectPermissionDict], + object_permission: ObjectPermissionDict | None, team_obj: Optional["LiteLLM_TeamTableCachedObj"], - prisma_client: Optional[PrismaClient] = None, + prisma_client: PrismaClient | None = None, is_proxy_admin: bool = False, -) -> Optional[ObjectPermissionDict]: +) -> ObjectPermissionDict | None: """ Validate that MCP servers requested on a key are within the allowed scope. @@ -608,7 +608,7 @@ async def validate_key_mcp_servers_against_team( # Validate requested access groups (must be subset of team's access groups) if requested_access_groups: - team_access_groups: Set[str] = set() + team_access_groups: set[str] = set() if ( team_obj is not None and team_obj.object_permission is not None @@ -694,7 +694,7 @@ def _validate_requested_toolsets( def _extract_requested_vector_stores( - object_permission: Optional[ObjectPermissionDict], + object_permission: ObjectPermissionDict | None, ) -> set[str]: """Return vector_store IDs from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -706,7 +706,7 @@ def _extract_requested_vector_stores( async def validate_key_vector_stores_against_team( - object_permission: Optional[ObjectPermissionDict], + object_permission: ObjectPermissionDict | None, team_obj: Optional["LiteLLM_TeamTableCachedObj"], is_proxy_admin: bool = False, ) -> None: @@ -734,7 +734,7 @@ async def validate_key_vector_stores_against_team( def _extract_requested_search_tools( - object_permission: Optional[ObjectPermissionDict], + object_permission: ObjectPermissionDict | None, ) -> list[str]: """Return search_tool_name values from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -746,7 +746,7 @@ def _extract_requested_search_tools( async def validate_key_search_tools_against_team( - object_permission: Optional[ObjectPermissionDict], + object_permission: ObjectPermissionDict | None, team_obj: Optional["LiteLLM_TeamTableCachedObj"], is_proxy_admin: bool = False, ) -> None: @@ -772,7 +772,7 @@ async def validate_key_search_tools_against_team( }, ) - team_tools: List[str] = [] + team_tools: list[str] = [] if team_obj is not None and team_obj.object_permission is not None: st = team_obj.object_permission.search_tools if st: diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index 1532668ed19..5d93c965ad4 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -1,5 +1,3 @@ -from typing import List, Optional - from litellm.proxy._types import ( KeyManagementRoutes, LiteLLM_TeamTableCachedObj, @@ -11,9 +9,9 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import PrismaClient BASELINE_TEAM_MEMBER_PERMISSIONS = [ @@ -29,7 +27,7 @@ class TeamMemberPermissionChecks: def get_permissions_for_team_member( team_member_object: Member, team_table: LiteLLM_TeamTableCachedObj, - ) -> List[KeyManagementRoutes]: + ) -> list[KeyManagementRoutes]: """ Returns the permissions for a team member. @@ -48,8 +46,8 @@ class TeamMemberPermissionChecks: @staticmethod def _get_list_of_route_enum_as_str( - route_enum: List[KeyManagementRoutes], - ) -> List[str]: + route_enum: list[KeyManagementRoutes], + ) -> list[str]: """ Returns a list of the route enum as a list of strings """ @@ -106,10 +104,10 @@ class TeamMemberPermissionChecks: @staticmethod def does_team_member_have_permissions_for_endpoint( - team_member_object: Optional[Member], + team_member_object: Member | None, team_table: LiteLLM_TeamTableCachedObj, route: str, - ) -> Optional[bool]: + ) -> bool | None: """ Raises an exception if the team member does not have permissions for calling the endpoint for a team """ @@ -140,8 +138,8 @@ class TeamMemberPermissionChecks: @staticmethod def enforce_member_can_assign_access_groups( user_api_key_dict: UserAPIKeyAuth, - team_table: Optional[LiteLLM_TeamTableCachedObj], - access_group_ids: Optional[List[str]], + team_table: LiteLLM_TeamTableCachedObj | None, + access_group_ids: list[str] | None, ) -> None: """ Field-level opt-in gate: a non-admin team member may only set @@ -233,7 +231,7 @@ class TeamMemberPermissionChecks: return team_member_object is not None @staticmethod - def get_all_available_team_member_permissions() -> List[str]: + def get_all_available_team_member_permissions() -> list[str]: """ Returns all available team member permissions """ @@ -243,5 +241,5 @@ class TeamMemberPermissionChecks: return all_available_permissions @staticmethod - def default_team_member_permissions() -> List[str]: + def default_team_member_permissions() -> list[str]: return [route.value for route in DEFAULT_TEAM_MEMBER_PERMISSIONS] diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 15ebc4e6990..5ca5a8aeb89 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -3,7 +3,7 @@ from collections.abc import Callable from datetime import datetime from functools import wraps -from typing import Any, List, Optional, Tuple +from typing import Any from fastapi import HTTPException, Request from pydantic import BaseModel @@ -39,7 +39,7 @@ from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.repositories.user_repository import UserRepository -def get_new_internal_user_defaults(user_id: str, user_email: Optional[str] = None) -> dict: +def get_new_internal_user_defaults(user_id: str, user_email: str | None = None) -> dict: user_info = litellm.default_internal_user_params or {} returned_dict: SSOUserDefinedValues = { @@ -60,11 +60,11 @@ def get_new_internal_user_defaults(user_id: str, user_email: Optional[str] = Non async def handle_budget_for_entity( data, - existing_budget_id: Optional[str], + existing_budget_id: str | None, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, litellm_proxy_admin_name: str, -) -> Optional[str]: +) -> str | None: """ Common helper to handle budget creation/updates for entities (organizations, tags, etc). @@ -141,7 +141,7 @@ async def handle_budget_for_entity( # (i.e. the values an admin sets). We copy these when cloning a team's # default member-budget into an individual member-budget so that the new # row starts with the same limits as the default. -_CLONABLE_BUDGET_FIELDS: Tuple[str, ...] = ( +_CLONABLE_BUDGET_FIELDS: tuple[str, ...] = ( "max_budget", "soft_budget", "max_parallel_requests", @@ -158,8 +158,8 @@ async def _clone_team_default_budget_for_member( default_team_budget_id: str, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, - budget_duration_override: Optional[str] = None, -) -> Optional[str]: + budget_duration_override: str | None = None, +) -> str | None: """ Create a new budget row that copies the values from the team's default member budget. Returns the new budget_id, or None if the default budget @@ -210,11 +210,11 @@ async def _resolve_member_budget_id( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, - max_budget_in_team: Optional[float], - allowed_models: Optional[list[str]], - budget_duration: Optional[str], - default_team_budget_id: Optional[str], -) -> Optional[str]: + max_budget_in_team: float | None, + allowed_models: list[str] | None, + budget_duration: str | None, + default_team_budget_id: str | None, +) -> str | None: """ Resolve the budget a new team member should be linked to. @@ -270,15 +270,15 @@ async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, t async def add_new_member( new_member: Member, - max_budget_in_team: Optional[float], + max_budget_in_team: float | None, prisma_client: PrismaClient, team_id: str, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, - default_team_budget_id: Optional[str] = None, - allowed_models: Optional[List[str]] = None, - budget_duration: Optional[str] = None, -) -> Tuple[LiteLLM_UserTable, Optional[LiteLLM_TeamMembership]]: + default_team_budget_id: str | None = None, + allowed_models: list[str] | None = None, + budget_duration: str | None = None, +) -> tuple[LiteLLM_UserTable, LiteLLM_TeamMembership | None]: """ Add a new member to a team @@ -287,8 +287,8 @@ async def add_new_member( Returns created/existing user + team membership w/ budget id """ - returned_user: Optional[LiteLLM_UserTable] = None - returned_team_membership: Optional[LiteLLM_TeamMembership] = None + returned_user: LiteLLM_UserTable | None = None + returned_team_membership: LiteLLM_TeamMembership | None = None ## ADD TEAM ID, to USER TABLE IF NEW ## if new_member.user_id is not None: new_user_defaults = get_new_internal_user_defaults(user_id=new_member.user_id) @@ -314,7 +314,7 @@ async def add_new_member( new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement ### for now: check if it exists in db, if not - insert it - existing_user_row: Optional[list] = await prisma_client.get_data( + existing_user_row: list | None = await prisma_client.get_data( key_val={"user_email": new_member.user_email}, table_name="user", query_type="find_all", @@ -375,7 +375,6 @@ def _delete_user_id_from_cache(kwargs): if isinstance(update_user_request, DeleteUserRequest): for user_id in update_user_request.user_ids: user_api_key_cache.delete_cache(key=user_id) - pass def _delete_api_key_from_cache(kwargs): @@ -390,7 +389,6 @@ def _delete_api_key_from_cache(kwargs): if isinstance(update_request, KeyRequest) and update_request.keys: for key in update_request.keys: user_api_key_cache.delete_cache(key=key) - pass def _delete_team_id_from_cache(kwargs): @@ -405,7 +403,6 @@ def _delete_team_id_from_cache(kwargs): if isinstance(update_request, DeleteTeamRequest): for team_id in update_request.team_ids: user_api_key_cache.delete_cache(key=team_id) - pass def _delete_customer_id_from_cache(kwargs): @@ -420,7 +417,6 @@ def _delete_customer_id_from_cache(kwargs): if isinstance(update_request, DeleteCustomerRequest): for user_id in update_request.user_ids: user_api_key_cache.delete_cache(key=user_id) - pass async def send_management_endpoint_alert( @@ -526,7 +522,7 @@ async def _emit_management_endpoint_otel_span( start_time: datetime, end_time: datetime, result: Any = None, - exception: Optional[Exception] = None, + exception: Exception | None = None, ) -> None: """Stamp + end the parent OTEL SERVER span for a management endpoint. @@ -548,7 +544,7 @@ async def _emit_management_endpoint_otel_span( if is_otel_v2_enabled(): return - http_request: Optional[Request] = kwargs.get("http_request") + http_request: Request | None = kwargs.get("http_request") if http_request is not None: # Inline import — auth_utils participates in a proxy import cycle. from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415 @@ -575,7 +571,7 @@ async def _emit_management_endpoint_otel_span( } ) - _response: Optional[dict] = None + _response: dict | None = None if exception is None and result is not None: try: raw = dict(result) @@ -646,7 +642,6 @@ def management_endpoint_wrapper(func): except Exception as e: # Non-Blocking Exception verbose_logger.debug("Error in management endpoint wrapper: %s", str(e)) - pass return result except Exception as e: diff --git a/litellm/proxy/mcp_tools.py b/litellm/proxy/mcp_tools.py index eb63e6b43bc..eaf9e4b2665 100644 --- a/litellm/proxy/mcp_tools.py +++ b/litellm/proxy/mcp_tools.py @@ -1,7 +1,7 @@ -from typing import Any, Dict, Optional +from typing import Any -def get_current_time(params: Optional[Dict[str, Any]] = None) -> str: +def get_current_time(params: dict[str, Any] | None = None) -> str: """ Get the current time (hardcoded sample implementation) @@ -18,7 +18,7 @@ def get_current_time(params: Optional[Dict[str, Any]] = None) -> str: return "10:30:45 AM" -def get_current_date(params: Optional[Dict[str, Any]] = None) -> str: +def get_current_date(params: dict[str, Any] | None = None) -> str: """ Get the current date (hardcoded sample implementation) diff --git a/litellm/proxy/memory/memory_endpoints.py b/litellm/proxy/memory/memory_endpoints.py index 1d9704ba619..fbf00f10ead 100644 --- a/litellm/proxy/memory/memory_endpoints.py +++ b/litellm/proxy/memory/memory_endpoints.py @@ -18,7 +18,7 @@ Scoping: """ import json -from typing import Any, List, Optional +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Query @@ -61,14 +61,14 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN -def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Optional[dict]: +def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None: """ Prisma `where` fragment restricting rows to those the caller can see. Returns None for admins (no restriction). """ if _is_admin(user_api_key_dict): return None - ors: List[dict] = [] + ors: list[dict] = [] if user_api_key_dict.user_id: ors.append({"user_id": user_api_key_dict.user_id}) if user_api_key_dict.team_id: @@ -206,9 +206,9 @@ def _is_unique_violation(exc: Exception) -> bool: def _resolve_scope( user_api_key_dict: UserAPIKeyAuth, - requested_user_id: Optional[str], - requested_team_id: Optional[str], -) -> tuple[Optional[str], Optional[str]]: + requested_user_id: str | None, + requested_team_id: str | None, +) -> tuple[str | None, str | None]: """ Resolve the (user_id, team_id) to stamp on a new row. @@ -304,8 +304,8 @@ async def create_memory( response_model=MemoryListResponse, ) async def list_memory( - key: Optional[str] = Query(None, description="Filter by exact key match."), - key_prefix: Optional[str] = Query( + key: str | None = Query(None, description="Filter by exact key match."), + key_prefix: str | None = Query( None, description=( "Filter by key prefix (Redis-style namespace scan). " diff --git a/litellm/proxy/middleware/billable_request_metrics_middleware.py b/litellm/proxy/middleware/billable_request_metrics_middleware.py index ac3d27cd849..81efd01a8a0 100644 --- a/litellm/proxy/middleware/billable_request_metrics_middleware.py +++ b/litellm/proxy/middleware/billable_request_metrics_middleware.py @@ -12,7 +12,7 @@ import re import threading from collections.abc import Callable, Sequence from enum import Enum -from typing import Optional, Protocol, runtime_checkable +from typing import Protocol, runtime_checkable from starlette.types import ASGIApp, Message, Receive, Scope, Send @@ -28,7 +28,7 @@ class BillableCategory(str, Enum): @runtime_checkable class BillingRecorder(Protocol): - def record(self, *, category: BillableCategory, route: str, status_code: int, model_id: Optional[str]) -> None: ... + def record(self, *, category: BillableCategory, route: str, status_code: int, model_id: str | None) -> None: ... _MODEL_ID_HEADER = b"x-litellm-model-id" @@ -85,7 +85,7 @@ _PASSTHROUGH_PREFIXES: tuple[str, ...] = tuple( ) -def _classify_llm_route(path: str) -> Optional[str]: +def _classify_llm_route(path: str) -> str | None: exact_match = next((route for route in _LLM_ROUTE_EXACT if path == route), None) if exact_match is not None: return exact_match @@ -114,7 +114,7 @@ _A2A_TRANSPORT_PREFIXES: tuple[str, ...] = ("/v1/a2a/", "/a2a/") # method-agnostic by contrast because its list path logs a SpendLogs row too. -def _classify_mcp_route(path: str) -> Optional[str]: +def _classify_mcp_route(path: str) -> str | None: if path == _MCP_MANAGEMENT_PREFIX or path.startswith(f"{_MCP_MANAGEMENT_PREFIX}/"): return None if path == "/mcp" or path.startswith("/mcp/"): @@ -126,13 +126,13 @@ def _classify_mcp_route(path: str) -> Optional[str]: return None -def _classify_a2a_route(path: str) -> Optional[str]: +def _classify_a2a_route(path: str) -> str | None: if path.endswith(_A2A_INVOKE_SUFFIX) and any(path.startswith(prefix) for prefix in _A2A_TRANSPORT_PREFIXES): return "/a2a" return None -def classify_billable_request(path: str, method: str = "POST") -> Optional[tuple[BillableCategory, str]]: +def classify_billable_request(path: str, method: str = "POST") -> tuple[BillableCategory, str] | None: """Map a request path to its (category, normalized route), or None if not billable.""" normalized = path.rstrip("/") or "/" @@ -156,7 +156,7 @@ def classify_billable_request(path: str, method: str = "POST") -> Optional[tuple return None -def _extract_model_id(headers: Sequence[tuple[bytes, bytes]]) -> Optional[str]: +def _extract_model_id(headers: Sequence[tuple[bytes, bytes]]) -> str | None: return next( (value.decode("latin-1") for name, value in headers if name.lower() == _MODEL_ID_HEADER and value), None, @@ -174,8 +174,8 @@ class BillableRequestMetricsMiddleware: def __init__( self, app: ASGIApp, - recorder: Optional[BillingRecorder] = None, - recorder_factory: Optional[Callable[[], Optional[BillingRecorder]]] = None, + recorder: BillingRecorder | None = None, + recorder_factory: Callable[[], BillingRecorder | None] | None = None, ) -> None: self.app = app self.recorder = recorder @@ -188,7 +188,7 @@ class BillableRequestMetricsMiddleware: self._resolved = recorder_factory is None self._resolve_lock = threading.Lock() - def _resolve_recorder(self) -> Optional[BillingRecorder]: + def _resolve_recorder(self) -> BillingRecorder | None: if self._resolved: return self.recorder # The lock keeps concurrent first requests from each building their own @@ -217,7 +217,7 @@ class BillableRequestMetricsMiddleware: category, route = classification status_code = 0 - model_id: Optional[str] = None + model_id: str | None = None async def send_wrapper(message: Message) -> None: nonlocal status_code, model_id diff --git a/litellm/proxy/middleware/in_flight_requests_middleware.py b/litellm/proxy/middleware/in_flight_requests_middleware.py index 3b93e3a3992..e8add2c2fa2 100644 --- a/litellm/proxy/middleware/in_flight_requests_middleware.py +++ b/litellm/proxy/middleware/in_flight_requests_middleware.py @@ -6,7 +6,7 @@ Prometheus gauge `litellm_in_flight_requests`. """ import os -from typing import Any, Optional +from typing import Any from starlette.types import ASGIApp, Receive, Scope, Send @@ -27,7 +27,7 @@ class InFlightRequestsMiddleware: """ _in_flight: int = 0 - _gauge: Optional[Any] = None + _gauge: Any | None = None _gauge_init_attempted: bool = False def __init__(self, app: ASGIApp) -> None: @@ -55,7 +55,7 @@ class InFlightRequestsMiddleware: return InFlightRequestsMiddleware._in_flight @staticmethod - def _get_gauge() -> Optional[Any]: + def _get_gauge() -> Any | None: if InFlightRequestsMiddleware._gauge_init_attempted: return InFlightRequestsMiddleware._gauge InFlightRequestsMiddleware._gauge_init_attempted = True diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index 52be388c95e..5ff562abddc 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -4,7 +4,7 @@ Prometheus Auth Middleware - Pure ASGI implementation import json from collections.abc import MutableMapping -from typing import Any, List +from typing import Any from fastapi import Request from starlette.types import ASGIApp, Receive, Scope, Send @@ -45,7 +45,7 @@ class PrometheusAuthMiddleware: # user_api_key_auth reads the request body, which consumes ASGI `receive`. # Buffer those messages and replay them for the inner app; otherwise a # successful auth would forward an exhausted receive and /metrics hangs. - buffered_messages: List[MutableMapping[str, Any]] = [] + buffered_messages: list[MutableMapping[str, Any]] = [] async def receive_for_auth() -> MutableMapping[str, Any]: message = await receive() diff --git a/litellm/proxy/middleware/request_size_limit_middleware.py b/litellm/proxy/middleware/request_size_limit_middleware.py index 75e2f9a523f..d2fdf5b93e8 100644 --- a/litellm/proxy/middleware/request_size_limit_middleware.py +++ b/litellm/proxy/middleware/request_size_limit_middleware.py @@ -1,10 +1,9 @@ import json from collections.abc import Callable -from typing import Optional, Union from starlette.types import ASGIApp, Message, Receive, Scope, Send -MaxRequestSizeGetter = Callable[[], Optional[Union[int, float]]] +MaxRequestSizeGetter = Callable[[], int | float | None] RequestSizeLimitEnabledGetter = Callable[[], bool] @@ -77,7 +76,7 @@ class RequestSizeLimitMiddleware: await _send_request_too_large(send=send, max_request_size_mb=max_request_size_mb) -def _mb_to_bytes(max_request_size_mb: Optional[Union[int, float]]) -> Optional[int]: +def _mb_to_bytes(max_request_size_mb: float | None) -> int | None: if max_request_size_mb is None: return None if max_request_size_mb <= 0: @@ -85,7 +84,7 @@ def _mb_to_bytes(max_request_size_mb: Optional[Union[int, float]]) -> Optional[i return int(max_request_size_mb * 1024 * 1024) -def _get_content_length(scope: Scope) -> Optional[int]: +def _get_content_length(scope: Scope) -> int | None: headers = dict(scope.get("headers") or []) raw_content_length = headers.get(b"content-length") if raw_content_length is None: @@ -99,7 +98,7 @@ def _get_content_length(scope: Scope) -> Optional[int]: async def _send_request_too_large( send: Send, - max_request_size_mb: Optional[Union[int, float]], + max_request_size_mb: float | None, ) -> None: body = json.dumps( {"error": f"Request size is too large. Max size is {max_request_size_mb} MB"}, diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 91699ad829a..df8b0725257 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -1,7 +1,7 @@ #### OCR Endpoints ##### import json -from typing import Any, Dict, Optional, cast +from typing import Any, cast import orjson from fastapi import APIRouter, Depends, Request, Response, UploadFile @@ -18,9 +18,9 @@ router = APIRouter() def _build_document_from_upload( file_content: bytes, - filename: Optional[str], - content_type: Optional[str], -) -> Dict[str, str]: + filename: str | None, + content_type: str | None, +) -> dict[str, str]: """ Convert uploaded file bytes into a Mistral-format document dict with base64 data URI. @@ -41,7 +41,7 @@ def _build_document_from_upload( ) -async def _parse_multipart_form(request: Request) -> Dict[str, Any]: +async def _parse_multipart_form(request: Request) -> dict[str, Any]: """ Extract OCR data from a multipart form request. @@ -55,7 +55,7 @@ async def _parse_multipart_form(request: Request) -> Dict[str, Any]: form = await request.form() except Exception as e: raise ValueError( - f"Failed to parse multipart form data: {str(e)}. " + f"Failed to parse multipart form data: {e!s}. " "When using curl with --form/-F, do NOT set the Content-Type header " "manually — curl will set it automatically with the required boundary." ) @@ -81,7 +81,7 @@ async def _parse_multipart_form(request: Request) -> Dict[str, Any]: content_type=uploaded_file.content_type, ) - data: Dict[str, Any] = {"document": document} + data: dict[str, Any] = {"document": document} for field_name, field_value in form.items(): if field_name in ("file", "document"): @@ -104,7 +104,7 @@ async def _parse_multipart_form(request: Request) -> Dict[str, Any]: return data -async def _parse_ocr_request(request: Request) -> Dict[str, Any]: +async def _parse_ocr_request(request: Request) -> dict[str, Any]: """ Parse an OCR request, supporting both JSON and multipart form data. diff --git a/litellm/proxy/openai_evals_endpoints/endpoints.py b/litellm/proxy/openai_evals_endpoints/endpoints.py index 565ba607a43..4b1e7921e5d 100644 --- a/litellm/proxy/openai_evals_endpoints/endpoints.py +++ b/litellm/proxy/openai_evals_endpoints/endpoints.py @@ -2,8 +2,6 @@ OpenAI Evals API endpoints - /v1/evals """ -from typing import Optional - import orjson from fastapi import APIRouter, Depends, Request, Response @@ -33,7 +31,7 @@ router = APIRouter() async def create_eval( fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -124,12 +122,12 @@ async def create_eval( async def list_evals( fastapi_response: Response, request: Request, - limit: Optional[int] = 20, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - order_by: Optional[str] = None, - custom_llm_provider: Optional[str] = "openai", + limit: int | None = 20, + after: str | None = None, + before: str | None = None, + order: str | None = None, + order_by: str | None = None, + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -226,7 +224,7 @@ async def get_eval( eval_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -314,7 +312,7 @@ async def update_eval( eval_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -404,7 +402,7 @@ async def delete_eval( eval_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -492,7 +490,7 @@ async def cancel_eval( eval_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -585,7 +583,7 @@ async def create_run( eval_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -684,11 +682,11 @@ async def list_runs( eval_id: str, fastapi_response: Response, request: Request, - limit: Optional[int] = 20, - after: Optional[str] = None, - before: Optional[str] = None, - order: Optional[str] = None, - custom_llm_provider: Optional[str] = "openai", + limit: int | None = 20, + after: str | None = None, + before: str | None = None, + order: str | None = None, + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -778,7 +776,7 @@ async def get_run( run_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -865,7 +863,7 @@ async def cancel_run( run_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -954,7 +952,7 @@ async def delete_run( run_id: str, fastapi_response: Response, request: Request, - custom_llm_provider: Optional[str] = "openai", + custom_llm_provider: str | None = "openai", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 2960b031cd3..87514b46dbd 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -3,7 +3,7 @@ import mimetypes import re from dataclasses import dataclass, field from types import MappingProxyType -from typing import TYPE_CHECKING, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Literal, Optional from litellm.repositories.table_repositories import ( ManagedFileRepository, @@ -18,7 +18,7 @@ if TYPE_CHECKING: from litellm.router import Router -def _is_base64_encoded_unified_file_id(b64_uid: str) -> Union[str, Literal[False]]: +def _is_base64_encoded_unified_file_id(b64_uid: str) -> str | Literal[False]: # Ensure b64_uid is a string and not a mock object if not isinstance(b64_uid, str): return False @@ -43,7 +43,7 @@ def convert_b64_uid_to_unified_uid(b64_uid: str) -> str: return b64_uid -def get_models_from_unified_file_id(unified_file_id: str) -> List[str]: +def get_models_from_unified_file_id(unified_file_id: str) -> list[str]: """ Extract model names from unified file ID. @@ -64,7 +64,7 @@ def get_models_from_unified_file_id(unified_file_id: str) -> List[str]: return [] -def get_model_id_from_unified_batch_id(file_id: str) -> Optional[str]: +def get_model_id_from_unified_batch_id(file_id: str) -> str | None: """ Get the model_id from the file_id @@ -150,7 +150,7 @@ def encode_batch_response_ids(response, model: str) -> None: ) -def decode_model_from_file_id(encoded_id: str) -> Optional[str]: +def decode_model_from_file_id(encoded_id: str) -> str | None: """ Extract model name from an encoded file/batch ID. Handles IDs that start with "file-" or "batch_" prefix. @@ -224,8 +224,8 @@ def is_model_embedded_id(file_id: str) -> bool: def extract_model_from_sources( file_id: str, request, # FastAPI Request object - data: Optional[dict] = None, -) -> tuple[Optional[str], Optional[str]]: + data: dict | None = None, +) -> tuple[str | None, str | None]: """ Extract model information from multiple sources in priority order: 1. Embedded in file_id (highest priority) @@ -297,7 +297,7 @@ def get_team_provider_credentials( llm_router: Optional["Router"], user_api_key_dict: "UserAPIKeyAuth", custom_llm_provider: str, -) -> Optional[dict]: +) -> dict | None: """ Resolve upstream credentials for a provider-scoped file operation (e.g. GET /v1/files), which doesn't pin a model. @@ -352,12 +352,12 @@ def get_team_provider_credentials( ) key_model_allowlist_set = frozenset(key_model_allowlist) - def _key_may_use(public_model_name: Optional[str]) -> bool: + def _key_may_use(public_model_name: str | None) -> bool: if not key_model_allowlist_set: return True return public_model_name is not None and public_model_name in key_model_allowlist_set - def _provider_credentials(model_id: str) -> Optional[dict]: + def _provider_credentials(model_id: str) -> dict | None: credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id) if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider: return credentials @@ -438,7 +438,7 @@ def apply_team_provider_credentials( def prepare_data_with_credentials( data: dict, credentials: dict, - file_id: Optional[str] = None, + file_id: str | None = None, include_internal_credentials: bool = False, ) -> None: """ @@ -466,7 +466,7 @@ def handle_model_based_routing( llm_router, # Router instance data: dict, check_file_id_encoding: bool = True, -) -> tuple[bool, Optional[str], Optional[str], Optional[dict]]: +) -> tuple[bool, str | None, str | None, dict | None]: """ Orchestrate model-based credential routing for file operations. @@ -600,7 +600,7 @@ def detect_content_type_from_filename(filename: str) -> str: return "application/octet-stream" -def normalize_mime_type_for_provider(mime_type: str, provider: Optional[str] = None) -> str: +def normalize_mime_type_for_provider(mime_type: str, provider: str | None = None) -> str: """ Normalize MIME type for specific provider requirements. @@ -653,7 +653,7 @@ def is_gemini_supported_mime_type(mime_type: str) -> bool: ) -def get_content_type_from_file_object(file_object: Optional[dict]) -> str: +def get_content_type_from_file_object(file_object: dict | None) -> str: """ Determine content type from file object (from database or API response). @@ -706,8 +706,8 @@ class FileCreationParams: """ target_storage: str = "default" - target_model_names: List[str] = field(default_factory=list) - model: Optional[str] = None + target_model_names: list[str] = field(default_factory=list) + model: str | None = None def __post_init__(self): """Normalize and validate parameters after initialization.""" @@ -724,9 +724,9 @@ class FileCreationParams: async def extract_file_creation_params( request: "Request", - request_body: Optional[dict] = None, - target_model_names_form: Optional[str] = None, - target_storage_form: Optional[str] = None, + request_body: dict | None = None, + target_model_names_form: str | None = None, + target_storage_form: str | None = None, ) -> FileCreationParams: """ Extract file creation parameters from request. @@ -763,7 +763,7 @@ async def extract_file_creation_params( ) -def _extract_target_storage_simple(target_storage_form: Optional[str] = None) -> str: +def _extract_target_storage_simple(target_storage_form: str | None = None) -> str: """ Extract target_storage parameter from form field. @@ -779,8 +779,8 @@ def _extract_target_storage_simple(target_storage_form: Optional[str] = None) -> def _extract_target_model_names_simple( - target_model_names_form: Optional[str] = None, -) -> List[str]: + target_model_names_form: str | None = None, +) -> list[str]: """ Extract target_model_names parameter from form field. """ @@ -800,7 +800,7 @@ def _is_target_model_names_key(key: str) -> bool: return key == "target_model_names" or (key.startswith("target_model_names[") and key.endswith("]")) -async def _extract_target_model_names_from_form(request: "Request") -> List[str]: +async def _extract_target_model_names_from_form(request: "Request") -> list[str]: """ Collect target_model_names from the raw multipart form. @@ -812,13 +812,13 @@ async def _extract_target_model_names_from_form(request: "Request") -> List[str] """ form_data = await request.form() - names: List[str] = [] + names: list[str] = [] for key, value in form_data.multi_items(): if _is_target_model_names_key(key) and isinstance(value, str): names.extend(_extract_target_model_names_simple(value)) seen = set() - result: List[str] = [] + result: list[str] = [] for name in names: if name and name not in seen: seen.add(name) @@ -827,8 +827,8 @@ async def _extract_target_model_names_from_form(request: "Request") -> List[str] def validate_managed_files_requirement( - target_model_names: List[str], - model: Optional[str] = None, + target_model_names: list[str], + model: str | None = None, ) -> None: """ Enforce proxy-level managed files when litellm.require_managed_files is enabled. @@ -838,9 +838,10 @@ def validate_managed_files_requirement( target_model_names is missing or a model parameter routes the request through the direct provider path instead of the managed-files hook. """ - import litellm from fastapi import HTTPException + import litellm + if litellm.require_managed_files is not True: return @@ -865,7 +866,7 @@ def validate_managed_files_requirement( ) -def _extract_model_param(request: "Request", request_body: dict) -> Optional[str]: +def _extract_model_param(request: "Request", request_body: dict) -> str | None: """ Extract model parameter from request. @@ -991,7 +992,7 @@ async def ensure_batch_response_managed_file_ids( async def get_batch_from_database( batch_id: str, - unified_batch_id: Union[str, Literal[False]], + unified_batch_id: str | Literal[False], managed_files_obj, prisma_client, verbose_proxy_logger, @@ -1052,7 +1053,7 @@ async def get_batch_from_database( async def update_batch_in_database( batch_id: str, - unified_batch_id: Union[str, Literal[False]], + unified_batch_id: str | Literal[False], response, managed_files_obj, prisma_client, 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 5cbdc530f60..58c22b72b08 100644 --- a/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py +++ b/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast from fastapi.responses import StreamingResponse @@ -18,11 +18,11 @@ class FileContentStreamingHandler: *, custom_llm_provider: str, file_id: str, - data: Dict[str, Any], + data: dict[str, Any], should_route: bool, - original_file_id: Optional[str], - credentials: Optional[Dict[str, Any]], - ) -> Tuple[str, str, Dict[str, Any]]: + original_file_id: str | None, + credentials: dict[str, Any] | None, + ) -> tuple[str, str, dict[str, Any]]: """ Resolve the provider, file ID, and request payload to use for streaming. @@ -71,7 +71,7 @@ class FileContentStreamingHandler: stream_iterator: AsyncIterator[bytes], proxy_logging_obj: "ProxyLogging", user_api_key_dict: "UserAPIKeyAuth", - data: Dict[str, Any], + data: dict[str, Any], ): try: async for chunk in stream_iterator: @@ -95,7 +95,7 @@ class FileContentStreamingHandler: *, custom_llm_provider: str, file_id: str, - data: Dict[str, Any], + data: dict[str, Any], proxy_logging_obj: "ProxyLogging", user_api_key_dict: "UserAPIKeyAuth", version: str, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index f1bcfbafe58..4e4718272bd 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -7,7 +7,7 @@ import asyncio import traceback -from typing import Any, BinaryIO, Optional, Union, cast, get_args +from typing import Any, BinaryIO, cast, get_args import httpx from fastapi import ( @@ -25,6 +25,9 @@ from fastapi import ( import litellm from litellm import CreateFileRequest, get_secret_str from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.cloud_storage_security import ( + is_managed_cloud_storage_uri, +) from litellm.llms.base_llm.files.transformation import BaseFileEndpoints from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -38,9 +41,6 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.litellm_core_utils.cloud_storage_security import ( - is_managed_cloud_storage_uri, -) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, apply_team_provider_credentials, @@ -97,7 +97,7 @@ def get_files_provider_config( return None -def get_first_json_object(file_source: Union[bytes, BinaryIO]) -> Optional[dict]: +def get_first_json_object(file_source: bytes | BinaryIO) -> dict | None: try: if isinstance(file_source, (bytes, bytearray)): newline = file_source.find(b"\n") @@ -112,7 +112,7 @@ def get_first_json_object(file_source: Union[bytes, BinaryIO]) -> Optional[dict] return None -def get_model_from_json_obj(json_object: dict) -> Optional[str]: +def get_model_from_json_obj(json_object: dict) -> str | None: body = json_object.get("body", {}) or {} model = body.get("model") @@ -120,7 +120,7 @@ def get_model_from_json_obj(json_object: dict) -> Optional[str]: async def _deprecated_loadbalanced_create_file( - llm_router: Optional[Router], + llm_router: Router | None, router_model: str, _create_file_request: CreateFileRequest, ) -> OpenAIFileObject: @@ -135,17 +135,17 @@ async def _deprecated_loadbalanced_create_file( async def route_create_file( - llm_router: Optional[Router], + llm_router: Router | None, _create_file_request: CreateFileRequest, purpose: OpenAIFilesPurpose, proxy_logging_obj: ProxyLogging, user_api_key_dict: UserAPIKeyAuth, - target_model_names_list: List[str], + target_model_names_list: list[str], is_router_model: bool, - router_model: Optional[str], + router_model: str | None, custom_llm_provider: str, - model: Optional[str] = None, - target_storage: Optional[str] = "default", + model: str | None = None, + target_storage: str | None = "default", ) -> OpenAIFileObject: """ Route file creation request to the appropriate provider. @@ -292,10 +292,10 @@ async def create_file( purpose: str = Form(...), target_model_names: str = Form(default=""), target_storage: str = Form(default="default"), - provider: Optional[str] = None, + provider: str | None = None, custom_llm_provider: str = Form(default="openai"), file: UploadFile = File(...), - litellm_metadata: Optional[str] = Form(default=None), + litellm_metadata: str | None = Form(default=None), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -323,12 +323,12 @@ async def create_file( version, ) - data: Dict = {} + data: dict = {} try: # Batch uploads can be gigabytes. Starlette has already spooled the upload # to disk, so stream from that handle instead of reading it into memory. # Other uploads are small and stay in-memory bytes. - file_source: Union[bytes, BinaryIO] + file_source: bytes | BinaryIO if purpose == "batch": await file.seek(0) file_source = file.file @@ -374,10 +374,10 @@ async def create_file( data = {} # Parse expires_after if provided - expires_after: Optional[FileExpiresAfter] = None + expires_after: FileExpiresAfter | None = None form_data_raw = await request.form() - form_data_dict: Dict[str, Any] = dict(form_data_raw) - extracted_litellm_metadata: Optional[Dict[str, Any]] = extract_nested_form_metadata( + form_data_dict: dict[str, Any] = dict(form_data_raw) + extracted_litellm_metadata: dict[str, Any] | None = extract_nested_form_metadata( form_data=form_data_dict, prefix="litellm_metadata[" ) expires_after_anchor = form_data_raw.get("expires_after[anchor]") @@ -457,7 +457,7 @@ async def create_file( file_data = (file.filename, file_source, file.content_type) ## check if model is a loadbalanced model - router_model: Optional[str] = None + router_model: str | None = None is_router_model = False if litellm.enable_loadbalancing_on_batch_endpoints is True: json_obj = get_first_json_object(file_source) @@ -549,9 +549,7 @@ async def create_file( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.create_file(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.create_file(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e.detail)), @@ -560,7 +558,7 @@ async def create_file( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -588,7 +586,7 @@ async def get_file_content( request: Request, fastapi_response: Response, file_id: str, - provider: Optional[str] = None, + provider: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -612,7 +610,7 @@ async def get_file_content( version, ) - data: Dict = {"file_id": file_id} + data: dict = {"file_id": file_id} try: # Include original request and headers in the data base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) @@ -692,13 +690,13 @@ async def get_file_content( ) except ValueError as e: raise ProxyException( - message=f"Storage backend error: {str(e)}", + message=f"Storage backend error: {e!s}", type="invalid_request_error", param="file_id", code=400, ) - model = cast(Optional[str], data.get("model")) + model = cast(str | None, data.get("model")) if model: response = await llm_router.afile_content( **{ @@ -833,7 +831,7 @@ async def get_file_content( model_region=getattr(user_api_key_dict, "allowed_model_region", ""), ) ) - httpx_response: Optional[httpx.Response] = getattr(response, "response", None) + httpx_response: httpx.Response | None = getattr(response, "response", None) if httpx_response is None: raise ValueError(f"Invalid response - response.response is None - got {response}") @@ -847,9 +845,7 @@ async def get_file_content( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.retrieve_file_content(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.retrieve_file_content(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -859,7 +855,7 @@ async def get_file_content( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -887,7 +883,7 @@ async def get_file( request: Request, fastapi_response: Response, file_id: str, - provider: Optional[str] = None, + provider: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -910,7 +906,7 @@ async def get_file( version, ) - data: Dict = {"file_id": file_id} + data: dict = {"file_id": file_id} try: custom_llm_provider = ( provider @@ -1036,7 +1032,7 @@ async def get_file( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.retrieve_file(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.retrieve_file(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -1046,7 +1042,7 @@ async def get_file( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -1074,7 +1070,7 @@ async def delete_file( request: Request, fastapi_response: Response, file_id: str, - provider: Optional[str] = None, + provider: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -1100,7 +1096,7 @@ async def delete_file( version, ) - data: Dict = {"file_id": file_id} + data: dict = {"file_id": file_id} try: custom_llm_provider = ( provider @@ -1242,9 +1238,7 @@ async def delete_file( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.delete_file(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.delete_file(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e.detail)), @@ -1253,7 +1247,7 @@ async def delete_file( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -1281,9 +1275,9 @@ async def list_files( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - provider: Optional[str] = None, - target_model_names: Optional[str] = None, - purpose: Optional[str] = None, + provider: str | None = None, + target_model_names: str | None = None, + purpose: str | None = None, ): """ Returns information about a specific file. that can be used across - Assistants API, Batch API @@ -1306,7 +1300,7 @@ async def list_files( version, ) - data: Dict = {} + data: dict = {} try: # Include original request and headers in the data base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) @@ -1323,7 +1317,7 @@ async def list_files( route_type=CallTypes.alist_fine_tuning_jobs.value, ) - response: Optional[Any] = None + response: Any | None = None # Check for model-based credential routing (no file_id encoding check for list) should_route, model_used, _, credentials = handle_model_based_routing( @@ -1433,7 +1427,7 @@ async def list_files( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.list_files(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.list_files(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -1443,7 +1437,7 @@ async def list_files( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index 1f7df846e66..b9658d8efee 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -8,7 +8,7 @@ storage backends (e.g., Azure Blob Storage) and managing associated metadata. import base64 import time from collections.abc import Mapping -from typing import Any, List, cast +from typing import Any, cast from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid as uuid_module @@ -35,7 +35,7 @@ class StorageBackendFileService: async def upload_file_to_storage_backend( file_data: Mapping[str, Any], target_storage: str, - target_model_names: List[str], + target_model_names: list[str], purpose: OpenAIFilesPurpose, proxy_logging_obj: ProxyLogging, user_api_key_dict: UserAPIKeyAuth, @@ -154,7 +154,7 @@ class StorageBackendFileService: @staticmethod def _create_unified_file_id( file_type: str, - target_model_names: List[str], + target_model_names: list[str], file_id: str, ) -> str: """ @@ -184,7 +184,7 @@ class StorageBackendFileService: async def _store_in_managed_files( file_object: OpenAIFileObject, file_data: Mapping[str, Any], - target_model_names: List[str], + target_model_names: list[str], target_storage: str, storage_url: str, proxy_logging_obj: ProxyLogging, diff --git a/litellm/proxy/pass_through_endpoints/jsonpath_extractor.py b/litellm/proxy/pass_through_endpoints/jsonpath_extractor.py index 6456bc71594..f3c4835a878 100644 --- a/litellm/proxy/pass_through_endpoints/jsonpath_extractor.py +++ b/litellm/proxy/pass_through_endpoints/jsonpath_extractor.py @@ -4,7 +4,7 @@ JSONPath Extractor Module Extracts field values from data using simple JSONPath-like expressions. """ -from typing import Any, List, Union +from typing import Any from litellm._logging import verbose_proxy_logger @@ -15,7 +15,7 @@ class JsonPathExtractor: @staticmethod def extract_fields( data: dict, - jsonpath_expressions: List[str], + jsonpath_expressions: list[str], ) -> str: """ Extract field values from data using JSONPath-like expressions. @@ -27,7 +27,7 @@ class JsonPathExtractor: Returns concatenated string of all extracted values. """ - extracted_values: List[str] = [] + extracted_values: list[str] = [] for expr in jsonpath_expressions: try: @@ -43,7 +43,7 @@ class JsonPathExtractor: return "\n".join(extracted_values) @staticmethod - def evaluate(data: dict, expr: str) -> Union[str, List[str], None]: + def evaluate(data: dict, expr: str) -> str | list[str] | None: """ Evaluate a simple JSONPath-like expression. diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 28d2c62f1f1..1395fc9d32f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -9,7 +9,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc. import json import os import re -from typing import Any, Optional, Tuple, Union, cast +from typing import Any, cast import httpx from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket @@ -42,10 +42,6 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( create_websocket_passthrough_route, websocket_passthrough_request, ) -from litellm.types.passthrough_endpoints.pass_through_endpoints import ( - LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, - LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, -) from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( assert_user_can_access_vector_store, @@ -53,6 +49,10 @@ from litellm.proxy.vector_store_endpoints.utils import ( is_allowed_to_call_vector_store_endpoint, ) from litellm.secret_managers.main import get_secret_str +from litellm.types.passthrough_endpoints.pass_through_endpoints import ( + LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, +) from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -75,7 +75,7 @@ def create_request_copy(request: Request): } -def is_passthrough_request_using_router_model(request_body: dict, llm_router: Optional[litellm.Router]) -> bool: +def is_passthrough_request_using_router_model(request_body: dict, llm_router: litellm.Router | None) -> bool: """ Returns True if the model is in the llm_router model names """ @@ -213,7 +213,7 @@ async def gemini_proxy_route( ) # Add or update query parameters - gemini_api_key: Optional[str] = passthrough_endpoint_router.get_credentials( + gemini_api_key: str | None = passthrough_endpoint_router.get_credentials( custom_llm_provider="gemini", region_name=None, ) @@ -290,7 +290,7 @@ async def cohere_proxy_route( endpoint_func = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers={"Authorization": "Bearer {}".format(cohere_api_key)}, + custom_headers={"Authorization": f"Bearer {cohere_api_key}"}, is_streaming_request=is_streaming_request, ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( @@ -411,7 +411,7 @@ async def mistral_proxy_route( endpoint_func = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers={"Authorization": "Bearer {}".format(mistral_api_key)}, + custom_headers={"Authorization": f"Bearer {mistral_api_key}"}, is_streaming_request=is_streaming_request, ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( @@ -449,9 +449,9 @@ async def milvus_proxy_route( request_body = await get_request_body(request) # check collectionName - collection_name = cast(Optional[str], request_body.get("collectionName")) + collection_name = cast(str | None, request_body.get("collectionName")) extra_headers = {} - base_target_url: Optional[str] = None + base_target_url: str | None = None if not collection_name: raise HTTPException( status_code=400, @@ -719,13 +719,13 @@ async def handle_bedrock_passthrough_router_model( general_settings: dict, proxy_config, select_data_generator, - user_model: Optional[str], - user_temperature: Optional[float], - user_request_timeout: Optional[float], - user_max_tokens: Optional[int], - user_api_base: Optional[str], - version: Optional[str], -) -> Union[Response, StreamingResponse]: + user_model: str | None, + user_temperature: float | None, + user_request_timeout: float | None, + user_max_tokens: int | None, + user_api_base: str | None, + version: str | None, +) -> Response | StreamingResponse: """ Handle Bedrock passthrough for router models (models defined in config.yaml). @@ -757,7 +757,7 @@ async def handle_bedrock_passthrough_router_model( # Use the common processing path (same as non-router models) # This ensures all metadata, hooks, and logging are properly initialized - data: Dict[str, Any] = {} + data: dict[str, Any] = {} base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) data["model"] = model @@ -801,8 +801,8 @@ async def handle_bedrock_count_tokens( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth, - request_body: Dict[str, Any], -) -> Dict[str, Any]: + request_body: dict[str, Any], +) -> dict[str, Any]: """ Handle AWS Bedrock CountTokens API requests. @@ -857,14 +857,14 @@ async def handle_bedrock_count_tokens( except BedrockError as e: # Convert BedrockError to HTTPException for FastAPI - verbose_proxy_logger.error(f"BedrockError in handle_bedrock_count_tokens: {str(e)}") + verbose_proxy_logger.error(f"BedrockError in handle_bedrock_count_tokens: {e!s}") raise HTTPException(status_code=e.status_code, detail={"error": e.message}) except HTTPException: # Re-raise HTTP exceptions as-is raise except Exception as e: - verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {str(e)}") - raise HTTPException(status_code=500, detail={"error": f"CountTokens processing error: {str(e)}"}) + verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {e!s}") + raise HTTPException(status_code=500, detail={"error": f"CountTokens processing error: {e!s}"}) async def bedrock_llm_proxy_route( @@ -949,7 +949,7 @@ async def bedrock_llm_proxy_route( # Fall back to existing implementation for direct Bedrock models verbose_proxy_logger.debug(f"Bedrock passthrough: Using direct Bedrock model '{model}' for endpoint '{endpoint}'") - data: Dict[str, Any] = {} + data: dict[str, Any] = {} base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) data["method"] = request.method @@ -1077,12 +1077,12 @@ async def bedrock_proxy_route( def _resolve_vertex_model_from_router( model_id: str, - llm_router: Optional[litellm.Router], + llm_router: litellm.Router | None, encoded_endpoint: str, endpoint: str, - vertex_project: Optional[str], - vertex_location: Optional[str], -) -> Tuple[str, str, Optional[str], Optional[str]]: + vertex_project: str | None, + vertex_location: str | None, +) -> tuple[str, str, str | None, str | None]: """ Resolve Vertex AI model configuration from router. @@ -1219,7 +1219,7 @@ async def assemblyai_proxy_route( endpoint_func = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers={"Authorization": "{}".format(assemblyai_api_key)}, + custom_headers={"Authorization": f"{assemblyai_api_key}"}, is_streaming_request=is_streaming_request, ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( @@ -1410,36 +1410,36 @@ from abc import ABC, abstractmethod class BaseVertexAIPassThroughHandler(ABC): @staticmethod @abstractmethod - def get_default_base_target_url(vertex_location: Optional[str]) -> str: + def get_default_base_target_url(vertex_location: str | None) -> str: pass @staticmethod @abstractmethod - def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: + def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str: pass class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler): @staticmethod - def get_default_base_target_url(vertex_location: Optional[str]) -> str: + def get_default_base_target_url(vertex_location: str | None) -> str: return "https://discoveryengine.googleapis.com/" @staticmethod - def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: + def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str: return base_target_url class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler): @staticmethod - def get_default_base_target_url(vertex_location: Optional[str]) -> str: + def get_default_base_target_url(vertex_location: str | None) -> str: return get_vertex_base_url(vertex_location) @staticmethod - def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: + def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str: return get_vertex_base_url(vertex_location) -def get_vertex_base_url(vertex_location: Optional[str]) -> str: +def get_vertex_base_url(vertex_location: str | None) -> str: """ Base URL for Vertex AI pass-through (trailing slash for URL joining). @@ -1489,10 +1489,10 @@ def get_vertex_pass_through_handler( def _override_vertex_params_from_router_credentials( - router_credentials: Optional[Any], - vertex_project: Optional[str], - vertex_location: Optional[str], -) -> Tuple[Optional[str], Optional[str]]: + router_credentials: Any | None, + vertex_project: str | None, + vertex_location: str | None, +) -> tuple[str | None, str | None]: """ Override vertex_project and vertex_location with values from router_credentials if available. @@ -1543,13 +1543,13 @@ def _override_vertex_params_from_router_credentials( async def _prepare_vertex_auth_headers( request: Request, - vertex_credentials: Optional[Any], - router_credentials: Optional[Any], - vertex_project: Optional[str], - vertex_location: Optional[str], - base_target_url: Optional[str], + vertex_credentials: Any | None, + router_credentials: Any | None, + vertex_project: str | None, + vertex_location: str | None, + base_target_url: str | None, get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler, -) -> Tuple[dict, Optional[str], bool, Optional[str], Optional[str]]: +) -> tuple[dict, str | None, bool, str | None, str | None]: """ Prepare authentication headers for Vertex AI pass-through requests. @@ -1637,8 +1637,8 @@ async def _base_vertex_proxy_route( request: Request, fastapi_response: Response, get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - router_credentials: Optional[Any] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, + router_credentials: Any | None = None, ): """ Base function for Vertex AI passthrough routes. @@ -1684,8 +1684,8 @@ async def _base_vertex_proxy_route( user_api_key_dict=user_api_key_dict, ) - vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint) - vertex_location: Optional[str] = get_vertex_location_from_url(endpoint) + vertex_project: str | None = get_vertex_project_id_from_url(endpoint) + vertex_location: str | None = get_vertex_location_from_url(endpoint) # Override with vector store credentials if available vertex_project, vertex_location = _override_vertex_params_from_router_credentials( @@ -1810,7 +1810,7 @@ async def vertex_discovery_proxy_route( from litellm.types.vector_stores import LiteLLM_ManagedVectorStore # Extract vector store ID from endpoint if present (e.g., dataStores/test-litellm-app_1761094730750) - vector_store_credentials: Optional[LiteLLM_ManagedVectorStore] = None + vector_store_credentials: LiteLLM_ManagedVectorStore | None = None vector_store_id_match = re.search(r"dataStores/([^/]+)", endpoint) if vector_store_id_match: @@ -1940,9 +1940,9 @@ class BaseOpenAIPassThroughHandler: fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth, base_target_url: str, - api_key: Optional[str], + api_key: str | None, custom_llm_provider: litellm.LlmProviders, - extra_headers: Optional[dict] = None, + extra_headers: dict | None = None, ): encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -1996,12 +1996,12 @@ class BaseOpenAIPassThroughHandler: return headers @staticmethod - def _assemble_headers(api_key: Optional[str], request: Request, extra_headers: Optional[dict] = None) -> dict: + def _assemble_headers(api_key: str | None, request: Request, extra_headers: dict | None = None) -> dict: base_headers = {} if api_key is not None: base_headers = { - "authorization": "Bearer {}".format(api_key), - "api-key": "{}".format(api_key), + "authorization": f"Bearer {api_key}", + "api-key": f"{api_key}", } if extra_headers is not None: base_headers.update(extra_headers) @@ -2096,7 +2096,7 @@ async def cursor_proxy_route( path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, encoded_endpoint) ) - auth_value = base64.b64encode(f"{cursor_api_key}:".encode("utf-8")).decode("ascii") + auth_value = base64.b64encode(f"{cursor_api_key}:".encode()).decode("ascii") endpoint_func = create_pass_through_route( endpoint=endpoint, @@ -2115,10 +2115,10 @@ async def cursor_proxy_route( async def vertex_ai_live_websocket_passthrough( websocket: WebSocket, - model: Optional[str] = None, - vertex_project: Optional[str] = None, - vertex_location: Optional[str] = None, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + model: str | None = None, + vertex_project: str | None = None, + vertex_location: str | None = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ): """ Vertex AI Live API WebSocket Pass-through Function @@ -2150,16 +2150,14 @@ async def vertex_ai_live_websocket_passthrough( ) resolved_project = vertex_project - resolved_location: Optional[str] = vertex_location - credentials_value: Optional[str] = None + resolved_location: str | None = vertex_location + credentials_value: str | None = None if vertex_credentials_config is not None: resolved_project = resolved_project or vertex_credentials_config.vertex_project temp_location = resolved_location or vertex_credentials_config.vertex_location # Ensure resolved_location is a string - if isinstance(temp_location, dict): - resolved_location = str(temp_location) - elif temp_location is not None: + if isinstance(temp_location, dict) or temp_location is not None: resolved_location = str(temp_location) else: resolved_location = None @@ -2249,9 +2247,9 @@ def create_vertex_ai_live_websocket_endpoint(): def create_generic_websocket_passthrough_endpoint( provider: str, target_url: str, - custom_headers: Optional[dict] = None, + custom_headers: dict | None = None, forward_headers: bool = False, - cost_per_request: Optional[float] = None, + cost_per_request: float | None = None, ): """ Create a generic WebSocket passthrough endpoint for any provider. 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 529a52daf9c..ddf86e9cd80 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 @@ -1,7 +1,7 @@ import json from collections.abc import Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast import httpx @@ -51,7 +51,7 @@ class AnthropicPassthroughLoggingHandler: start_time: datetime, end_time: datetime, cache_hit: bool, - request_body: Optional[dict] = None, + request_body: dict | None = None, **kwargs, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -106,7 +106,7 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _get_user_from_metadata( passthrough_logging_payload: PassthroughStandardLoggingPayload, - ) -> Optional[str]: + ) -> str | None: request_body = passthrough_logging_payload.get("request_body") if request_body: return get_end_user_id_from_request_body(request_body) @@ -127,8 +127,8 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _extract_model_from_anthropic_chunks( - all_chunks: Sequence[Union[str, bytes]], - ) -> Optional[str]: + all_chunks: Sequence[str | bytes], + ) -> str | None: for raw in all_chunks: text = raw.decode("utf-8") if isinstance(raw, bytes) else raw for line in text.splitlines(): @@ -148,7 +148,7 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _stream_was_interrupted( - all_chunks: Sequence[Union[str, bytes]], + all_chunks: Sequence[str | bytes], ) -> bool: """ Anthropic ends a stream with ``content_block_stop`` -> ``message_delta`` @@ -181,8 +181,8 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _recover_interrupted_stream_output_tokens( - response: Union[ModelResponse, TextCompletionResponse], - all_chunks: Sequence[Union[str, bytes]], + response: ModelResponse | TextCompletionResponse, + all_chunks: Sequence[str | bytes], model: str, ) -> None: """ @@ -223,7 +223,7 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _create_anthropic_response_logging_payload( - litellm_model_response: Union[ModelResponse, TextCompletionResponse], + litellm_model_response: ModelResponse | TextCompletionResponse, model: str, kwargs: dict, start_time: datetime, @@ -269,7 +269,7 @@ class AnthropicPassthroughLoggingHandler: # the pass-through success path reads spend from # model_call_details["response_cost"], not from kwargs logging_obj.model_call_details["response_cost"] = response_cost - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = ( # type: ignore kwargs.get("passthrough_logging_payload") ) if passthrough_logging_payload: @@ -305,7 +305,7 @@ class AnthropicPassthroughLoggingHandler: request_body: dict, endpoint_type: EndpointType, start_time: datetime, - all_chunks: List[str], + all_chunks: list[str], end_time: datetime, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -392,7 +392,7 @@ class AnthropicPassthroughLoggingHandler: } @staticmethod - def _split_sse_chunk_into_events(chunk: Union[str, bytes]) -> List[str]: + def _split_sse_chunk_into_events(chunk: str | bytes) -> list[str]: """ Split a chunk that may contain multiple SSE events into individual events. @@ -417,10 +417,10 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _build_complete_streaming_response( - all_chunks: Sequence[Union[str, bytes]], + all_chunks: Sequence[str | bytes], litellm_logging_obj: LiteLLMLoggingObj, model: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: """ Builds complete response from raw Anthropic chunks. @@ -465,8 +465,8 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _collapse_pure_text_chunks( - all_chunks: Sequence[Union[str, bytes]], - ) -> Optional[List[str]]: + all_chunks: Sequence[str | bytes], + ) -> list[str] | None: """ Return a new chunk list with the contiguous run of text-only ``content_block_delta`` events replaced by a single equivalent event, @@ -478,7 +478,7 @@ class AnthropicPassthroughLoggingHandler: ``message_delta`` / ``message_stop`` / ``ping`` events are accepted. Any other content-block type or delta type returns ``None``. """ - normalized: List[str] = [] + normalized: list[str] = [] for raw in all_chunks: line = raw.decode("utf-8") if isinstance(raw, bytes) else raw for ev in line.split("\n\n"): @@ -487,9 +487,9 @@ class AnthropicPassthroughLoggingHandler: normalized.append(ev) text_block_indexes: set = set() - out: List[str] = [] - pending_text: List[str] = [] - pending_index: Optional[int] = None + out: list[str] = [] + pending_text: list[str] = [] + pending_index: int | None = None saw_any_text_delta = False def flush() -> None: @@ -573,10 +573,10 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _build_complete_streaming_response_legacy( - all_chunks: Sequence[Union[str, bytes]], + all_chunks: Sequence[str | bytes], litellm_logging_obj: LiteLLMLoggingObj, model: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: """ Original reconstruction: convert every SSE event to a generic chunk and assemble via stream_chunk_builder. Kept verbatim as the fallback @@ -632,7 +632,7 @@ class AnthropicPassthroughLoggingHandler: return complete_streaming_response @staticmethod - def _extract_sse_data(event_str: str) -> Optional[dict]: + def _extract_sse_data(event_str: str) -> dict | None: """Parse the JSON object from the ``data:`` line of an Anthropic SSE event.""" for line in event_str.splitlines(): stripped = line.strip() @@ -648,9 +648,9 @@ class AnthropicPassthroughLoggingHandler: @staticmethod def _build_usage_only_response_from_chunks( - all_chunks: Sequence[Union[str, bytes]], + all_chunks: Sequence[str | bytes], model: str, - ) -> Optional[ModelResponse]: + ) -> ModelResponse | None: """ Build a usage-bearing ModelResponse from Anthropic SSE token-usage events, for cost tracking when stream_chunk_builder cannot reassemble the stream. @@ -663,13 +663,13 @@ class AnthropicPassthroughLoggingHandler: input_tokens = 0 cache_read = 0 cache_creation = 0 - cache_creation_5m: Optional[int] = None - cache_creation_1h: Optional[int] = None + cache_creation_5m: int | None = None + cache_creation_1h: int | None = None output_tokens = 0 - web_search_requests: Optional[int] = None - tool_search_requests: Optional[int] = None - inference_geo: Optional[str] = None - stop_reason: Optional[str] = None + web_search_requests: int | None = None + tool_search_requests: int | None = None + inference_geo: str | None = None + stop_reason: str | None = None found_usage = False resolved_model = model for _chunk_str in all_chunks: @@ -765,7 +765,7 @@ class AnthropicPassthroughLoggingHandler: start_time: datetime, end_time: datetime, cache_hit: bool, - request_body: Optional[dict] = None, + request_body: dict | None = None, **kwargs, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -935,7 +935,7 @@ class AnthropicPassthroughLoggingHandler: index=0, message={ "role": "assistant", - "content": f"Error creating batch job: {str(e)}", + "content": f"Error creating batch job: {e!s}", "tool_calls": None, "function_call": None, "provider_specific_fields": { diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py index 7d7bc889120..93bcac704e5 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py @@ -3,7 +3,7 @@ import json import time import urllib.parse from datetime import datetime -from typing import Literal, Optional +from typing import Literal from urllib.parse import urlparse import httpx @@ -102,7 +102,7 @@ class AssemblyAIPassthroughLoggingHandler: verbose_proxy_logger.debug("response body %s", json.dumps(response_body, indent=4)) kwargs["model"] = model kwargs["custom_llm_provider"] = "assemblyai" - response_cost: Optional[float] = None + response_cost: float | None = None transcript_id = response_body.get("id") if transcript_id is None: @@ -131,7 +131,7 @@ class AssemblyAIPassthroughLoggingHandler: status="success", ) - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = ( # type: ignore kwargs.get("passthrough_logging_payload") ) @@ -158,9 +158,7 @@ class AssemblyAIPassthroughLoggingHandler: ) ) - pass - - def _get_response_to_log(self, transcript_response: Optional[AssemblyAITranscriptResponse]) -> dict: + def _get_response_to_log(self, transcript_response: AssemblyAITranscriptResponse | None) -> dict: if transcript_response is None: return {} return dict(transcript_response) @@ -168,8 +166,8 @@ class AssemblyAIPassthroughLoggingHandler: def _get_assembly_transcript( self, transcript_id: str, - request_region: Optional[Literal["eu"]] = None, - ) -> Optional[dict]: + request_region: Literal["eu"] | None = None, + ) -> dict | None: """ Get the transcript details from AssemblyAI API @@ -205,16 +203,14 @@ class AssemblyAIPassthroughLoggingHandler: return response.json() except Exception as e: - verbose_proxy_logger.exception( - f"[Non blocking logging error] Error getting AssemblyAI transcript: {str(e)}" - ) + verbose_proxy_logger.exception(f"[Non blocking logging error] Error getting AssemblyAI transcript: {e!s}") return None def _poll_assembly_for_transcript_response( self, transcript_id: str, - url_route: Optional[str] = None, - ) -> Optional[AssemblyAITranscriptResponse]: + url_route: str | None = None, + ) -> AssemblyAITranscriptResponse | None: """ Poll the status of the transcript until it is completed or timeout (30 minutes) """ @@ -234,7 +230,7 @@ class AssemblyAIPassthroughLoggingHandler: def get_cost_for_assembly_transcript( transcript_response: AssemblyAITranscriptResponse, speech_model: str, - ) -> Optional[float]: + ) -> float | None: """ Get the cost for the assembly transcript """ @@ -249,7 +245,7 @@ class AssemblyAIPassthroughLoggingHandler: return _audio_duration * _cost_per_second @staticmethod - def get_cost_per_second_for_assembly_model(speech_model: str) -> Optional[float]: + def get_cost_per_second_for_assembly_model(speech_model: str) -> float | None: """ Get the cost per second for the assembly model. Falls back to assemblyai/nano if the specific speech model info cannot be found. @@ -279,9 +275,7 @@ class AssemblyAIPassthroughLoggingHandler: return None except Exception as e: - verbose_proxy_logger.exception( - f"[Non blocking logging error] Error getting AssemblyAI model info: {str(e)}" - ) + verbose_proxy_logger.exception(f"[Non blocking logging error] Error getting AssemblyAI model info: {e!s}") return None @staticmethod @@ -292,7 +286,7 @@ class AssemblyAIPassthroughLoggingHandler: return request_method == "POST" @staticmethod - def _get_assembly_region_from_url(url: Optional[str]) -> Optional[Literal["eu"]]: + def _get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: """ Get the region from the URL """ @@ -303,7 +297,7 @@ class AssemblyAIPassthroughLoggingHandler: return None @staticmethod - def _get_assembly_base_url_from_region(region: Optional[Literal["eu"]]) -> str: + def _get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: """ Get the base URL for the AssemblyAI API if region == "eu", return "https://api.eu.assemblyai.com" diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/base_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/base_passthrough_logging_handler.py index 980c05fa412..6a882ad78ac 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/base_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/base_passthrough_logging_handler.py @@ -1,6 +1,6 @@ import json from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -87,7 +87,7 @@ class BasePassthroughLoggingHandler(ABC): def _get_user_from_metadata( self, passthrough_logging_payload: PassthroughStandardLoggingPayload, - ) -> Optional[str]: + ) -> str | None: request_body = passthrough_logging_payload.get("request_body") if request_body: return get_end_user_id_from_request_body(request_body) @@ -95,7 +95,7 @@ class BasePassthroughLoggingHandler(ABC): def _create_response_logging_payload( self, - litellm_model_response: Union[ModelResponse, TextCompletionResponse], + litellm_model_response: ModelResponse | TextCompletionResponse, model: str, kwargs: dict, start_time: datetime, @@ -119,7 +119,7 @@ class BasePassthroughLoggingHandler(ABC): # the pass-through success path reads spend from # model_call_details["response_cost"], not from kwargs logging_obj.model_call_details["response_cost"] = response_cost - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = ( # type: ignore kwargs.get("passthrough_logging_payload") ) if passthrough_logging_payload: @@ -159,10 +159,10 @@ class BasePassthroughLoggingHandler(ABC): @abstractmethod def _build_complete_streaming_response( self, - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: LiteLLMLoggingObj, model: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: """ Builds complete response from raw chunks @@ -170,7 +170,6 @@ class BasePassthroughLoggingHandler(ABC): - Converts generic chunks to litellm chunks (OpenAI format) - Builds complete response from litellm chunks """ - pass def _handle_logging_llm_collected_chunks( self, @@ -180,7 +179,7 @@ class BasePassthroughLoggingHandler(ABC): request_body: dict, endpoint_type: EndpointType, start_time: datetime, - all_chunks: List[str], + all_chunks: list[str], end_time: datetime, ) -> PassThroughEndpointLoggingTypedDict: """ diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py index 70b09e101fc..a9561992952 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py @@ -1,5 +1,4 @@ from datetime import datetime -from typing import List, Optional, Union import httpx @@ -39,10 +38,10 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler): def _build_complete_streaming_response( self, - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: LiteLLMLoggingObj, model: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: cohere_model_response_iterator = CohereModelResponseIterator( streaming_response=None, sync_stream=False, @@ -123,7 +122,7 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler): kwargs["custom_llm_provider"] = "cohere" # Extract user information for tracking - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = kwargs.get( + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = kwargs.get( "passthrough_logging_payload" ) if passthrough_logging_payload: diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cursor_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cursor_passthrough_logging_handler.py index 63907e9638b..b5dfcc8f6a6 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cursor_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cursor_passthrough_logging_handler.py @@ -6,7 +6,6 @@ so they appear cleanly in the LiteLLM Logs page. """ from datetime import datetime -from typing import Dict import httpx @@ -18,7 +17,7 @@ from litellm.litellm_core_utils.litellm_logging import ( from litellm.proxy._types import PassThroughEndpointLoggingTypedDict from litellm.types.utils import StandardPassThroughResponseObject -CURSOR_AGENT_ENDPOINTS: Dict[str, str] = { +CURSOR_AGENT_ENDPOINTS: dict[str, str] = { "POST /v0/agents": "cursor:agent:create", "GET /v0/agents": "cursor:agent:list", "POST /v0/agents/{id}/followup": "cursor:agent:followup", diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py index 716dff21efc..96ec7291389 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py @@ -1,6 +1,6 @@ import re from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any import httpx @@ -122,8 +122,8 @@ class GeminiPassthroughLoggingHandler: request_body: dict, endpoint_type: EndpointType, start_time: datetime, - all_chunks: List[str], - model: Optional[str], + all_chunks: list[str], + model: str | None, end_time: datetime, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -133,7 +133,7 @@ class GeminiPassthroughLoggingHandler: - Creates standard logging object - Logs in litellm callbacks """ - kwargs: Dict[str, Any] = {} + kwargs: dict[str, Any] = {} model = model or GeminiPassthroughLoggingHandler.extract_model_from_url(url_route) complete_streaming_response = GeminiPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, @@ -168,11 +168,11 @@ class GeminiPassthroughLoggingHandler: @staticmethod def _build_complete_streaming_response( - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: LiteLLMLoggingObj, model: str, url_route: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: parsed_chunks = [] if "generateContent" in url_route or "streamGenerateContent" in url_route: gemini_iterator: Any = GeminiModelResponseIterator( @@ -208,7 +208,7 @@ class GeminiPassthroughLoggingHandler: @staticmethod def _create_gemini_response_logging_payload_for_generate_content( - litellm_model_response: Union[ModelResponse, TextCompletionResponse], + litellm_model_response: ModelResponse | TextCompletionResponse, model: str, kwargs: dict, start_time: datetime, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index f71bc167fb9..e878f2a544d 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -5,7 +5,6 @@ Handles cost tracking and logging for OpenAI passthrough endpoints, specifically """ from datetime import datetime -from typing import List, Optional, Tuple, Union from urllib.parse import urlparse import httpx @@ -26,11 +25,11 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.passthrough_endpoints.pass_through_endpoints import ( EndpointType, PassthroughStandardLoggingPayload, ) -from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import ImageResponse, LlmProviders, PassthroughCallTypes from litellm.utils import ModelResponse, TextCompletionResponse @@ -62,7 +61,7 @@ def _hostname_matches(hostname: str, suffixes: tuple) -> bool: return any(hostname == suffix or hostname.endswith("." + suffix) for suffix in suffixes) -def _is_openai_compatible_host(hostname: Optional[str]) -> bool: +def _is_openai_compatible_host(hostname: str | None) -> bool: """True if the hostname is OpenAI proper or one of the Azure OpenAI domains. Hostname-only check, kept for the route-level helpers that additionally @@ -75,7 +74,7 @@ def _is_openai_compatible_host(hostname: Optional[str]) -> bool: return _hostname_matches(hostname, _OPENAI_HOSTNAMES) or _hostname_matches(hostname, _AZURE_OPENAI_HOSTNAMES) -def _is_openai_compatible_url(url_route: Optional[str]) -> bool: +def _is_openai_compatible_url(url_route: str | None) -> bool: """True if the URL targets an OpenAI-compatible API surface. For the shared Azure Cognitive Services domains we additionally require an @@ -146,7 +145,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): def _get_user_from_metadata( self, passthrough_logging_payload: PassthroughStandardLoggingPayload, - ) -> Optional[str]: + ) -> str | None: """Extract user information from passthrough logging payload.""" request_body = passthrough_logging_payload.get("request_body") if request_body: @@ -184,7 +183,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return cost except Exception as e: - verbose_proxy_logger.warning(f"Error calculating image generation cost: {str(e)}") + verbose_proxy_logger.warning(f"Error calculating image generation cost: {e!s}") return 0.0 @staticmethod @@ -218,7 +217,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return cost except Exception as e: - verbose_proxy_logger.warning(f"Error calculating image editing cost: {str(e)}") + verbose_proxy_logger.warning(f"Error calculating image editing cost: {e!s}") return 0.0 @staticmethod @@ -227,7 +226,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): httpx_response: httpx.Response, logging_obj: LiteLLMLoggingObj, custom_llm_provider: str, - ) -> Tuple[ResponsesAPIResponse, float]: + ) -> tuple[ResponsesAPIResponse, float]: """Transform a Responses API raw response into a ResponsesAPIResponse and compute its cost. @@ -306,14 +305,9 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): try: response_cost = 0.0 - litellm_model_response: Optional[ - Union[ - ModelResponse, - TextCompletionResponse, - ImageResponse, - ResponsesAPIResponse, - ] - ] = None + litellm_model_response: ( + ModelResponse | TextCompletionResponse | ImageResponse | ResponsesAPIResponse | None + ) = None handler_instance = OpenAIPassthroughLoggingHandler() custom_llm_provider = kwargs.get("custom_llm_provider", "openai") @@ -406,7 +400,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): kwargs["custom_llm_provider"] = custom_llm_provider # Extract user information for tracking - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = kwargs.get( + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = kwargs.get( "passthrough_logging_payload" ) if passthrough_logging_payload: @@ -451,7 +445,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): } except Exception as e: - verbose_proxy_logger.error(f"Error in OpenAI passthrough cost tracking: {str(e)}") + verbose_proxy_logger.error(f"Error in OpenAI passthrough cost tracking: {e!s}") # Fall back to base handler without cost tracking base_handler = OpenAIPassthroughLoggingHandler() return base_handler.passthrough_chat_handler( @@ -472,7 +466,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): all_chunks: list, litellm_logging_obj: LiteLLMLoggingObj, model: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: """ Builds complete response from raw chunks for OpenAI streaming responses. @@ -520,7 +514,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return complete_streaming_response except Exception as e: - verbose_proxy_logger.error(f"Error building complete streaming response: {str(e)}") + verbose_proxy_logger.error(f"Error building complete streaming response: {e!s}") return None @staticmethod @@ -531,7 +525,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): request_body: dict, endpoint_type: EndpointType, start_time: datetime, - all_chunks: List[str], + all_chunks: list[str], end_time: datetime, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -577,7 +571,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): } # Extract user information for tracking - passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( + passthrough_logging_payload: PassthroughStandardLoggingPayload | None = ( litellm_logging_obj.model_call_details.get("passthrough_logging_payload") ) if passthrough_logging_payload: @@ -614,7 +608,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): } except Exception as e: - verbose_proxy_logger.error(f"Error in OpenAI streaming passthrough cost tracking: {str(e)}") + verbose_proxy_logger.error(f"Error in OpenAI streaming passthrough cost tracking: {e!s}") return { "result": None, "kwargs": {}, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py index 158b629ad27..11672571f3f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py @@ -6,7 +6,7 @@ Supports different modalities: text, audio, video, and web search. """ from datetime import datetime -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import ( @@ -33,7 +33,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): def _build_complete_streaming_response(self, *args, **kwargs): """Not applicable for WebSocket passthrough.""" - return None + return def get_provider_config(self, model: str): """Return Vertex AI provider configuration.""" @@ -50,8 +50,8 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): @staticmethod def _extract_usage_metadata_from_websocket_messages( - websocket_messages: List[Dict], - ) -> Optional[Dict]: + websocket_messages: list[dict], + ) -> dict | None: """ Extract and aggregate usage metadata from a list of WebSocket messages. @@ -76,7 +76,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): return all_usage_metadata[0] # Aggregate multiple usage metadata messages - aggregated: Dict[str, Any] = { + aggregated: dict[str, Any] = { "promptTokenCount": 0, "candidatesTokenCount": 0, "totalTokenCount": 0, @@ -130,7 +130,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): @staticmethod def _calculate_live_api_cost( model: str, - usage_metadata: Dict, + usage_metadata: dict, custom_llm_provider: str = "vertex_ai", ) -> float: """ @@ -226,7 +226,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): @staticmethod def _create_usage_object_from_metadata( - usage_metadata: Dict, + usage_metadata: dict, model: str, ) -> Usage: """ @@ -275,7 +275,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): def vertex_ai_live_passthrough_handler( self, - websocket_messages: List[Dict], + websocket_messages: list[dict], logging_obj, url_route: str, start_time: datetime, 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 c61d48eda8c..a2f17eb8911 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -1,6 +1,6 @@ import re from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast from urllib.parse import urlparse import httpx @@ -49,7 +49,7 @@ class VertexPassthroughLoggingHandler: start_time: datetime, end_time: datetime, cache_hit: bool, - request_body: Optional[dict] = None, + request_body: dict | None = None, **kwargs, ) -> PassThroughEndpointLoggingTypedDict: if "predictLongRunning" in url_route: @@ -253,7 +253,7 @@ class VertexPassthroughLoggingHandler: _json_response = httpx_response.json() - litellm_prediction_response: Union[ModelResponse, EmbeddingResponse, ImageResponse] = ModelResponse() + litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse() if vertex_image_generation_class.is_image_generation_response(_json_response): litellm_prediction_response = vertex_image_generation_class.process_image_generation_response( _json_response, @@ -307,7 +307,7 @@ class VertexPassthroughLoggingHandler: } @staticmethod - def _extract_embed_content_input(request_body: Optional[dict], batch: bool) -> str: + def _extract_embed_content_input(request_body: dict | None, batch: bool) -> str: """Extract raw input text from an :embedContent or :batchEmbedContents request body for token counting.""" if not request_body: return "" @@ -327,11 +327,13 @@ class VertexPassthroughLoggingHandler: logging_obj: LiteLLMLoggingObj, url_route: str, kwargs: dict, - request_body: Optional[dict] = None, + request_body: dict | None = None, ) -> PassThroughEndpointLoggingTypedDict: """Handle Vertex :embedContent and :batchEmbedContents endpoint responses.""" from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( process_embed_content_response, + ) + from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( process_response as process_batch_embed_response, ) @@ -391,8 +393,8 @@ class VertexPassthroughLoggingHandler: request_body: dict, endpoint_type: EndpointType, start_time: datetime, - all_chunks: List[str], - model: Optional[str], + all_chunks: list[str], + model: str | None, end_time: datetime, ) -> PassThroughEndpointLoggingTypedDict: """ @@ -402,7 +404,7 @@ class VertexPassthroughLoggingHandler: - Creates standard logging object - Logs in litellm callbacks """ - kwargs: Dict[str, Any] = {} + kwargs: dict[str, Any] = {} model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route) complete_streaming_response = VertexPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, @@ -437,11 +439,11 @@ class VertexPassthroughLoggingHandler: @staticmethod def _build_complete_streaming_response( - all_chunks: List[str], + all_chunks: list[str], litellm_logging_obj: LiteLLMLoggingObj, model: str, url_route: str, - ) -> Optional[Union[ModelResponse, TextCompletionResponse]]: + ) -> ModelResponse | TextCompletionResponse | None: parsed_chunks = [] if "generateContent" in url_route or "streamGenerateContent" in url_route: vertex_iterator: Any = VertexModelResponseIterator( @@ -505,14 +507,12 @@ class VertexPassthroughLoggingHandler: The extracted model name for use with LiteLLM """ # Handle publishers/google/models/ format - if "publishers/" in vertex_model_path and "models/" in vertex_model_path: - # Extract everything after the last models/ - parts = vertex_model_path.split("models/") - if len(parts) > 1: - return parts[-1] - - # Handle projects/PROJECT_ID/locations/LOCATION/models/MODEL_ID format - elif "projects/" in vertex_model_path and "models/" in vertex_model_path: + if ( + "publishers/" in vertex_model_path + and "models/" in vertex_model_path + or "projects/" in vertex_model_path + and "models/" in vertex_model_path + ): # Extract everything after the last models/ parts = vertex_model_path.split("models/") if len(parts) > 1: @@ -522,7 +522,7 @@ class VertexPassthroughLoggingHandler: return vertex_model_path @staticmethod - def _get_vertex_publisher_or_api_spec_from_url(url: str) -> Optional[str]: + def _get_vertex_publisher_or_api_spec_from_url(url: str) -> str | None: # Check for specific Vertex AI partner publishers if "/publishers/mistralai/" in url: return "mistralai" @@ -576,7 +576,7 @@ class VertexPassthroughLoggingHandler: @staticmethod def _create_vertex_response_logging_payload_for_generate_content( - litellm_model_response: Union[ModelResponse, TextCompletionResponse], + litellm_model_response: ModelResponse | TextCompletionResponse, model: str, kwargs: dict, start_time: datetime, @@ -759,7 +759,7 @@ class VertexPassthroughLoggingHandler: index=0, message={ "role": "assistant", - "content": f"Error creating batch prediction job: {str(e)}", + "content": f"Error creating batch prediction job: {e!s}", "tool_calls": None, "function_call": None, "provider_specific_fields": { diff --git a/litellm/proxy/pass_through_endpoints/managed_id_codec.py b/litellm/proxy/pass_through_endpoints/managed_id_codec.py index f0c24bbaf39..ed91252d679 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_codec.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_codec.py @@ -20,7 +20,6 @@ from __future__ import annotations import base64 import uuid as _uuid_mod from dataclasses import dataclass -from typing import Optional from litellm.types.utils import SpecialEnums @@ -45,7 +44,7 @@ def encode(provider: str, unified_uuid: str, raw_provider_id: str) -> str: return base64.urlsafe_b64encode(plaintext.encode()).decode().rstrip("=") -def decode(managed_id: str) -> Optional[ManagedIdPayload]: +def decode(managed_id: str) -> ManagedIdPayload | None: """ Decode *managed_id*. diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index f6970c2a287..78fa732a67b 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -32,7 +32,7 @@ from __future__ import annotations import json import re -from typing import Any, Dict, FrozenSet, List, Optional, Tuple +from typing import Any from urllib.parse import quote, unquote from fastapi import HTTPException @@ -55,13 +55,13 @@ from .managed_id_codec import ManagedIdPayload, decode, is_managed, new_managed_ # Field map # --------------------------------------------------------------------------- -_FieldSpec = Tuple[str, str] # (field_name, expected_raw_id_prefix) -_MapKey = Tuple[str, str, str] # (provider, HTTP_METHOD, canonical_path) +_FieldSpec = tuple[str, str] # (field_name, expected_raw_id_prefix) +_MapKey = tuple[str, str, str] # (provider, HTTP_METHOD, canonical_path) # ``canonical_path`` uses ``/v1/...`` form without any ``/openai/`` prefix. # Both ``/openai/...`` and ``/openai_passthrough/...`` are normalised by # ``_canonical_path()`` before the lookup so only one set of entries is needed. -BUILTIN_OUTPUT_ID_FIELD_MAP: Dict[_MapKey, List[_FieldSpec]] = { +BUILTIN_OUTPUT_ID_FIELD_MAP: dict[_MapKey, list[_FieldSpec]] = { # ------------------------------------------------------------------ files ("openai", "POST", "/v1/files"): [ ("id", "file-"), @@ -146,10 +146,10 @@ BUILTIN_OUTPUT_ID_FIELD_MAP: Dict[_MapKey, List[_FieldSpec]] = { } # Prefixes that live in the *file* table rather than the object table. -_FILE_PREFIXES: FrozenSet[str] = frozenset({"file-"}) +_FILE_PREFIXES: frozenset[str] = frozenset({"file-"}) # Raw provider-ID prefixes that live in the object table (batches, responses). -_OBJECT_PREFIXES: FrozenSet[str] = frozenset({"batch_", "resp_"}) +_OBJECT_PREFIXES: frozenset[str] = frozenset({"batch_", "resp_"}) # Guards request-body rewriting against stack exhaustion from adversarially # deep payloads. Real OpenAI files/batches bodies nest only a few levels. @@ -197,7 +197,7 @@ class _RawIdGuardBudget: # --------------------------------------------------------------------------- # Maps (provider, canonical_path) -> "files" | "batches" -_LIST_ROUTE_TABLE: Dict[Tuple[str, str], str] = { +_LIST_ROUTE_TABLE: dict[tuple[str, str], str] = { ("openai", "/v1/files"): "files", ("openai", "/v1/batches"): "batches", ("azure", "/v1/files"): "files", @@ -276,7 +276,7 @@ async def _resolve_one( Raises ``HTTPException(404)`` on unknown / forged managed IDs — never forwarded upstream as a literal string. """ - payload: Optional[ManagedIdPayload] = decode(managed_id) + payload: ManagedIdPayload | None = decode(managed_id) if payload is None: return managed_id # not a passthrough managed ID; pass through verbose_proxy_logger.debug( @@ -296,8 +296,8 @@ async def _resolve_one( detail=(f"Managed ID was minted for provider '{payload.provider}', not '{provider}'."), ) - row_created_by: Optional[str] = None - row_team_id: Optional[str] = None + row_created_by: str | None = None + row_team_id: str | None = None found = False raw_id = payload.raw_provider_id @@ -373,7 +373,7 @@ async def _guard_raw_provider_id( provider: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: Any, - budget: Optional[_RawIdGuardBudget] = None, + budget: _RawIdGuardBudget | None = None, ) -> None: """Deny a raw provider ID that maps to a managed resource the caller does not own, before it is forwarded upstream. @@ -434,7 +434,7 @@ async def _guard_raw_provider_id( # --------------------------------------------------------------------------- -def _build_managed_file_object(snapshot: Optional[Dict[str, Any]], managed_id: str) -> Optional[OpenAIFileObject]: +def _build_managed_file_object(snapshot: dict[str, Any] | None, managed_id: str) -> OpenAIFileObject | None: """Build an ``OpenAIFileObject`` (with the managed ID swapped in) from an upstream file response so the DB-served list returns the same metadata as a direct file GET. Returns ``None`` when no usable snapshot is available, in @@ -457,7 +457,7 @@ async def _mint_or_reuse_file( user_api_key_dict: UserAPIKeyAuth, prisma_client: Any, managed_files_hook: Any, - file_object_snapshot: Optional[Dict[str, Any]] = None, + file_object_snapshot: dict[str, Any] | None = None, is_create_route: bool = True, ) -> str: """Return an existing managed file ID or mint + store a new one.""" @@ -809,7 +809,7 @@ def _parse_file_object(file_object: Any) -> Any: return file_object -def _empty_list_response() -> Dict[str, Any]: +def _empty_list_response() -> dict[str, Any]: return { "object": "list", "data": [], @@ -819,7 +819,7 @@ def _empty_list_response() -> Dict[str, Any]: } -def _parse_list_limit(query_params: Optional[Dict[str, Any]]) -> Tuple[int, int]: +def _parse_list_limit(query_params: dict[str, Any] | None) -> tuple[int, int]: params = query_params or {} try: raw_limit = int(params.get("limit", 20)) @@ -833,14 +833,14 @@ async def _build_list_where_with_cursor( prisma_client: Any, resource_kind: str, provider: str, - owner_filter: Dict[str, Any], - query_params: Optional[Dict[str, Any]], -) -> Tuple[Dict[str, Any], str]: + owner_filter: dict[str, Any], + query_params: dict[str, Any] | None, +) -> tuple[dict[str, Any], str]: """Return a Prisma ``where`` clause and fetch order for a list query.""" params = query_params or {} - after_id: Optional[str] = params.get("after") - before_id: Optional[str] = params.get("before") - where: Dict[str, Any] = dict(owner_filter) + after_id: str | None = params.get("after") + before_id: str | None = params.get("before") + where: dict[str, Any] = dict(owner_filter) fetch_order = "desc" cursor_id = after_id or before_id @@ -887,10 +887,10 @@ async def _build_list_where_with_cursor( async def _fetch_list_rows( prisma_client: Any, resource_kind: str, - where: Dict[str, Any], + where: dict[str, Any], fetch_order: str, fetch_limit: int, -) -> Optional[List[Any]]: +) -> list[Any] | None: # created_at is not unique, so a second sort on the unique id column gives a # total order, keeping the limit+1 page boundary and cursor deterministic # across rows that share a created_at timestamp. @@ -915,11 +915,11 @@ async def _fetch_provider_scoped_list_rows( prisma_client: Any, resource_kind: str, provider: str, - where: Dict[str, Any], + where: dict[str, Any], fetch_order: str, raw_limit: int, fetch_limit: int, -) -> Tuple[List[Any], bool]: +) -> tuple[list[Any], bool]: """Fetch one page of list rows scoped to *provider* at the DB level. Both resource kinds carry a provider-distinguishing value that the query @@ -951,8 +951,8 @@ async def _fetch_provider_scoped_list_rows( return page, has_more -def _serialize_file_list_item(row: Any) -> Dict[str, Any]: - item: Dict[str, Any] = { +def _serialize_file_list_item(row: Any) -> dict[str, Any]: + item: dict[str, Any] = { "id": row.unified_file_id, "object": "file", "created_at": int(row.created_at.timestamp()) if row.created_at else None, @@ -964,8 +964,8 @@ def _serialize_file_list_item(row: Any) -> Dict[str, Any]: return item -def _serialize_batch_list_item(row: Any) -> Dict[str, Any]: - item: Dict[str, Any] = {} +def _serialize_batch_list_item(row: Any) -> dict[str, Any]: + item: dict[str, Any] = {} file_object = _parse_file_object(row.file_object) if isinstance(file_object, dict): item.update(file_object) @@ -974,7 +974,7 @@ def _serialize_batch_list_item(row: Any) -> Dict[str, Any]: return item -def _list_boundary_ids(rows: List[Any], resource_kind: str) -> Tuple[Optional[str], Optional[str]]: +def _list_boundary_ids(rows: list[Any], resource_kind: str) -> tuple[str | None, str | None]: if not rows: return None, None id_attr = "unified_file_id" if resource_kind == "files" else "unified_object_id" @@ -986,8 +986,8 @@ async def list_passthrough_ids_from_db( route: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: Any, - query_params: Optional[Dict[str, Any]] = None, -) -> Optional[Dict[str, Any]]: + query_params: dict[str, Any] | None = None, +) -> dict[str, Any] | None: """Query the DB for managed IDs the caller owns and return an OpenAI-style paginated list response. @@ -1069,7 +1069,7 @@ async def rewrite_path_ids( """ budget = _RawIdGuardBudget() segments = path.split("/") - new_segments: List[str] = [] + new_segments: list[str] = [] changed = False for seg in segments: decoded_seg = unquote(seg) @@ -1092,12 +1092,12 @@ async def rewrite_path_ids( async def rewrite_query_ids( - params: Optional[Dict[str, Any]], + params: dict[str, Any] | None, provider: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: Any, managed_files_hook: Any, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: """ Walk query param values and resolve any passthrough managed IDs. Returns *params* unchanged (same object) when nothing is resolved. @@ -1106,7 +1106,7 @@ async def rewrite_query_ids( return params budget = _RawIdGuardBudget() mutated = dict(params) - rewritten_keys: List[str] = [] + rewritten_keys: list[str] = [] for key, val in list(mutated.items()): if isinstance(val, str): if is_managed(val): @@ -1124,12 +1124,12 @@ async def rewrite_query_ids( async def rewrite_body_ids( - body: Optional[Dict[str, Any]], + body: dict[str, Any] | None, provider: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: Any, managed_files_hook: Any, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: """ Recursively walk a request body dict/list and resolve any passthrough managed IDs. Skips litellm internal keys (``litellm_*``). @@ -1144,7 +1144,7 @@ async def rewrite_body_ids( if depth >= _MAX_BODY_REWRITE_DEPTH: return node if isinstance(node, dict): - result: Dict[str, Any] = {} + result: dict[str, Any] = {} changed_inner = False for k, v in node.items(): # Skip litellm internal injection keys (e.g. litellm_logging_obj) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e97a0f38975..b957618d776 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -8,7 +8,7 @@ from base64 import b64encode from collections.abc import AsyncGenerator, Mapping from datetime import datetime from itertools import groupby -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, cast from urllib.parse import urlencode, urlparse import httpx @@ -92,17 +92,17 @@ router = APIRouter() pass_through_endpoint_logging = PassThroughEndpointLogging() # Global registry to track registered pass-through routes and prevent memory leaks -_registered_pass_through_routes: Dict[str, Dict[str, Union[str, bool, List[str], Dict[str, Any]]]] = {} +_registered_pass_through_routes: dict[str, dict[str, str | bool | list[str] | dict[str, Any]]] = {} -def get_response_body(response: httpx.Response) -> Optional[dict]: +def get_response_body(response: httpx.Response) -> dict | None: try: return response.json() except Exception: return None -async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optional[dict]: +async def set_env_variables_in_header(custom_headers: dict | None) -> dict | None: """ checks if any headers on config.yaml are defined as os.environ/COHERE_API_KEY etc @@ -127,7 +127,7 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"): _langfuse_secret_key = get_secret_str(_langfuse_secret_key) headers["Authorization"] = "Basic " + b64encode( - f"{_langfuse_public_key}:{_langfuse_secret_key}".encode("utf-8") + f"{_langfuse_public_key}:{_langfuse_secret_key}".encode() ).decode("ascii") else: # for all other headers @@ -294,8 +294,8 @@ async def chat_completion_pass_through_endpoint( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - {}".format(str(e))) - error_msg = f"{str(e)}" + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.completion(): Exception occured - {e!s}") + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -308,8 +308,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): @staticmethod def get_response_headers( headers: httpx.Headers, - litellm_call_id: Optional[str] = None, - custom_headers: Optional[dict] = None, + litellm_call_id: str | None = None, + custom_headers: dict | None = None, ) -> dict: # Exclude headers that uvicorn writes itself (server, date) and # encoding/length headers that don't survive re-serialization. @@ -364,8 +364,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): async_client: httpx.AsyncClient, url: str, headers: dict, - requested_query_params: Optional[dict] = None, - custom_body: Optional[dict] = None, + requested_query_params: dict | None = None, + custom_body: dict | None = None, ) -> httpx.Response: """ Make a non-streaming HTTP request @@ -395,8 +395,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): async_client: httpx.AsyncClient, url: httpx.URL, headers: dict, - requested_query_params: Optional[dict] = None, - _parsed_body: Optional[dict] = None, + requested_query_params: dict | None = None, + _parsed_body: dict | None = None, forward_multipart: bool = False, ) -> httpx.Response: """ @@ -445,8 +445,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): @staticmethod async def _build_request_files_from_upload_file( - upload_file: Union[UploadFile, StarletteUploadFile], - ) -> Tuple[Optional[str], bytes, Optional[str]]: + upload_file: UploadFile | StarletteUploadFile, + ) -> tuple[str | None, bytes, str | None]: """Build a request files dict from an UploadFile object""" file_content = await upload_file.read() return (upload_file.filename, file_content, upload_file.content_type) @@ -457,7 +457,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): async_client: httpx.AsyncClient, url: httpx.URL, headers: dict, - requested_query_params: Optional[dict] = None, + requested_query_params: dict | None = None, stream: bool = False, ) -> httpx.Response: """Process multipart/form-data requests, handling both files and form fields. @@ -529,8 +529,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): user_api_key_dict: UserAPIKeyAuth, passthrough_logging_payload: PassthroughStandardLoggingPayload, logging_obj: LiteLLMLoggingObj, - _parsed_body: Optional[dict] = None, - litellm_call_id: Optional[str] = None, + _parsed_body: dict | None = None, + litellm_call_id: str | None = None, ) -> dict: """ Filter out litellm params from the request body @@ -587,7 +587,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return kwargs @staticmethod - def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: Optional[bool]) -> str: + def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: bool | None) -> str: """ Helper function to construct the full target URL with subpath handling. @@ -608,8 +608,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): # Ensure base_target ends with / and subpath doesn't start with / if not base_target.endswith("/"): base_target = base_target + "/" - if subpath.startswith("/"): - subpath = subpath[1:] + subpath = subpath.removeprefix("/") # Resolve any '..' segments in the subpath so it cannot climb above # the base_target prefix that the operator configured. Preserve a @@ -654,8 +653,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): @staticmethod def _update_stream_param_based_on_request_body( parsed_body: dict, - stream: Optional[bool] = None, - ) -> Optional[bool]: + stream: bool | None = None, + ) -> bool | None: """ If stream is provided in the request body, use it. Otherwise, use the stream parameter passed to the `pass_through_request` function @@ -665,7 +664,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return stream -def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[dict]) -> None: +def _carry_guardrail_logging_info(request_data: dict, guardrail_data: dict | None) -> None: """Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``. Post-call guardrails run against a throwaway ``hook_data`` dict (its @@ -689,10 +688,10 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d def _build_passthrough_failure_request_payload( - parsed_body: Optional[dict], - kwargs: Optional[dict], - logging_obj: Optional[LiteLLMLoggingObj], - custom_llm_provider: Optional[str], + parsed_body: dict | None, + kwargs: dict | None, + logging_obj: LiteLLMLoggingObj | None, + custom_llm_provider: str | None, upstream_usage: UpstreamReportedUsage | None = None, ) -> dict: """Build the ``request_data`` dict passed to ``post_call_failure_hook``. @@ -777,16 +776,16 @@ async def pass_through_request( target: str, custom_headers: dict, user_api_key_dict: UserAPIKeyAuth, - custom_body: Optional[dict] = None, - forward_headers: Optional[bool] = False, - merge_query_params: Optional[bool] = False, - query_params: Optional[dict] = None, - default_query_params: Optional[dict] = None, - stream: Optional[bool] = None, - cost_per_request: Optional[float] = None, - custom_llm_provider: Optional[str] = None, - guardrails_config: Optional[dict] = None, - timeout: Optional[float] = None, + custom_body: dict | None = None, + forward_headers: bool | None = False, + merge_query_params: bool | None = False, + query_params: dict | None = None, + default_query_params: dict | None = None, + stream: bool | None = None, + cost_per_request: float | None = None, + custom_llm_provider: str | None = None, + guardrails_config: dict | None = None, + timeout: float | None = None, ): """ Pass through endpoint handler, makes the httpx request for pass-through endpoints and ensures logging hooks are called @@ -819,16 +818,16 @@ async def pass_through_request( # Initialize variables ######################################################### litellm_call_id = str(uuid.uuid4()) - url: Optional[httpx.URL] = None + url: httpx.URL | None = None # parsed request body - _parsed_body: Optional[dict] = None + _parsed_body: dict | None = None # kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload - kwargs: Optional[dict] = None - logging_obj: Optional[Logging] = None + kwargs: dict | None = None + logging_obj: Logging | None = None # the dict post-call guardrails wrote their logging info into; the failure # handler reuses it so a guardrail block still surfaces its span/logs - post_call_guardrail_data: Optional[dict] = None + post_call_guardrail_data: dict | None = None ######################################################### try: @@ -840,7 +839,7 @@ async def pass_through_request( forward_headers=forward_headers, ) - requested_query_params: Optional[dict] = query_params or dict(request.query_params) + requested_query_params: dict | None = query_params or dict(request.query_params) endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url)) @@ -850,7 +849,7 @@ async def pass_through_request( # Tolerate request objects without `state` (test fixtures) and only honor # values httpx accepts for `content=`. _request_state = getattr(request, "state", None) - state_raw_body: Optional[Union[str, bytes]] = ( + state_raw_body: str | bytes | None = ( getattr(_request_state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, None) if _request_state is not None else None @@ -1126,7 +1125,7 @@ async def pass_through_request( else: # SigV4-signed callers (Bedrock) supply the exact pre-signed bytes; # otherwise httpx encodes the parsed JSON dict as before. - body_kwargs: Dict[str, Any] = ( + body_kwargs: dict[str, Any] = ( {"content": state_raw_body} if state_raw_body is not None else {"json": _parsed_body} ) req = async_client.build_request( @@ -1297,7 +1296,7 @@ async def pass_through_request( # responses; response_body itself is parsed unconditionally so the # failure-hook log payload below still reflects upstream error bodies. _content_modified = False - response_body: Optional[dict] = get_response_body(response) + response_body: dict | None = get_response_body(response) failure_request_payload = _build_passthrough_failure_request_payload( parsed_body=_parsed_body, @@ -1503,7 +1502,7 @@ async def pass_through_request( ) else: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {}".format(str(e)) + f"litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {e!s}" ) ######################################################### @@ -1545,7 +1544,7 @@ async def pass_through_request( headers=custom_headers, ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -1584,7 +1583,7 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di async def _parse_request_data_by_content_type( request: Request, -) -> Tuple[Optional[Any], Optional[Any], Optional[Any], Optional[Any]]: +) -> tuple[Any | None, Any | None, Any | None, Any | None]: """ Parse request data based on content type. @@ -1644,19 +1643,19 @@ async def _parse_request_data_by_content_type( def create_pass_through_route( endpoint, target: str, - custom_headers: Optional[Mapping[str, Any]] = None, - _forward_headers: Optional[bool] = False, - _merge_query_params: Optional[bool] = False, - dependencies: Optional[List] = None, - include_subpath: Optional[bool] = False, - cost_per_request: Optional[float] = None, - custom_llm_provider: Optional[str] = None, - is_streaming_request: Optional[bool] = False, - query_params: Optional[dict] = None, - default_query_params: Optional[dict] = None, - guardrails: Optional[Dict[str, Any]] = None, - config_file_path: Optional[str] = None, - timeout: Optional[float] = None, + custom_headers: Mapping[str, Any] | None = None, + _forward_headers: bool | None = False, + _merge_query_params: bool | None = False, + dependencies: list | None = None, + include_subpath: bool | None = False, + cost_per_request: float | None = None, + custom_llm_provider: str | None = None, + is_streaming_request: bool | None = False, + query_params: dict | None = None, + default_query_params: dict | None = None, + guardrails: dict[str, Any] | None = None, + config_file_path: str | None = None, + timeout: float | None = None, ): # check if target is an adapter.py or a url from litellm._uuid import uuid @@ -1766,12 +1765,12 @@ def create_pass_through_route( final_query_params.update(query_params) # Programmatic callers set LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY on # request.state (see Bedrock proxy). Parsed JSON envelope otherwise. - state_custom_body: Optional[dict] = getattr( + state_custom_body: dict | None = getattr( request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, None, ) - final_custom_body: Optional[dict] = None + final_custom_body: dict | None = None if isinstance(state_custom_body, dict): final_custom_body = state_custom_body elif isinstance(custom_body_data, dict): @@ -1783,16 +1782,16 @@ def create_pass_through_route( target=full_target, custom_headers=headers_dict, user_api_key_dict=user_api_key_dict, - forward_headers=cast(Optional[bool], param_forward_headers), - merge_query_params=cast(Optional[bool], param_merge_query_params), + forward_headers=cast(bool | None, param_forward_headers), + merge_query_params=cast(bool | None, param_merge_query_params), query_params=final_query_params, - default_query_params=cast(Optional[dict], param_default_query_params), + default_query_params=cast(dict | None, param_default_query_params), stream=is_streaming_request or stream, custom_body=final_custom_body, - cost_per_request=cast(Optional[float], param_cost_per_request), + cost_per_request=cast(float | None, param_cost_per_request), custom_llm_provider=custom_llm_provider, - guardrails_config=cast(Optional[dict], param_guardrails), - timeout=cast(Optional[float], param_timeout), + guardrails_config=cast(dict | None, param_guardrails), + timeout=cast(float | None, param_timeout), ) finally: if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY): @@ -1807,10 +1806,10 @@ def create_pass_through_route( def create_websocket_passthrough_route( endpoint: str, target: str, - custom_headers: Optional[dict] = None, - _forward_headers: Optional[bool] = False, - dependencies: Optional[List] = None, - cost_per_request: Optional[float] = None, + custom_headers: dict | None = None, + _forward_headers: bool | None = False, + dependencies: list | None = None, + cost_per_request: float | None = None, ): """ Create a WebSocket passthrough route function. @@ -1860,9 +1859,9 @@ async def websocket_passthrough_request( target: str, custom_headers: dict, user_api_key_dict: UserAPIKeyAuth, - forward_headers: Optional[bool] = False, - endpoint: Optional[str] = None, - cost_per_request: Optional[float] = None, + forward_headers: bool | None = False, + endpoint: str | None = None, + cost_per_request: float | None = None, accept_websocket: bool = True, ): """ @@ -1932,7 +1931,7 @@ async def websocket_passthrough_request( # Create a dummy request object for WebSocket connections to maintain compatibility # with the existing _init_kwargs_for_pass_through_endpoint function class DummyRequest: - def __init__(self, url: str, method: str = "WEBSOCKET", headers: Optional[dict] = None): + def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None): self.url = url self.method = method self.headers = headers or {} @@ -2049,7 +2048,7 @@ async def websocket_passthrough_request( verbose_proxy_logger.debug( f"WebSocket passthrough ({endpoint}): Client message is not a valid setup message: {e}" ) - pass # Not a JSON message or doesn't contain setup data + # Not a JSON message or doesn't contain setup data await upstream_ws.send(text_data) elif bytes_data is not None: @@ -2122,7 +2121,6 @@ async def websocket_passthrough_request( except (ConnectionClosedOK, ConnectionClosedError) as e: verbose_proxy_logger.debug(f"Upstream WebSocket connection closed: {e}") - pass except asyncio.CancelledError: verbose_proxy_logger.debug("asyncio.CancelledError in forward_upstream_to_client") raise @@ -2302,7 +2300,7 @@ async def _relay_passthrough_response_bytes( url_route: str, start_time: datetime, logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, success_handler_kwargs: dict, ) -> AsyncGenerator[bytes, None]: """ @@ -2344,7 +2342,7 @@ async def _relay_passthrough_response_bytes( ) -def _extract_model_from_vertex_ai_setup(setup_response: dict) -> Optional[str]: +def _extract_model_from_vertex_ai_setup(setup_response: dict) -> str | None: """ Extract the model name from Vertex AI Live setup response. @@ -2383,7 +2381,7 @@ class SafeRouteAdder: """ @staticmethod - def _is_path_registered(app: FastAPI, path: str, methods: List[str]) -> bool: + def _is_path_registered(app: FastAPI, path: str, methods: list[str]) -> bool: """ Check if a path with any of the specified methods is already registered on the app. @@ -2411,8 +2409,8 @@ class SafeRouteAdder: app: FastAPI, path: str, endpoint: Any, - methods: List[str], - dependencies: Optional[List] = None, + methods: list[str], + dependencies: list | None = None, ) -> bool: """ Add an API route to the app only if it doesn't already exist. @@ -2455,18 +2453,18 @@ class InitPassThroughEndpointHelpers: app: FastAPI, path: str, target: str, - custom_headers: Optional[dict], - forward_headers: Optional[bool], - merge_query_params: Optional[bool], - dependencies: Optional[List], - cost_per_request: Optional[float], + custom_headers: dict | None, + forward_headers: bool | None, + merge_query_params: bool | None, + dependencies: list | None, + cost_per_request: float | None, endpoint_id: str, - guardrails: Optional[dict] = None, - methods: Optional[List[str]] = None, - default_query_params: Optional[dict] = None, - config_file_path: Optional[str] = None, + guardrails: dict | None = None, + methods: list[str] | None = None, + default_query_params: dict | None = None, + config_file_path: str | None = None, auth: bool = False, - timeout: Optional[float] = None, + timeout: float | None = None, ): """Add exact path route for pass-through endpoint""" # Default to all methods if none specified (backward compatibility) @@ -2538,18 +2536,18 @@ class InitPassThroughEndpointHelpers: app: FastAPI, path: str, target: str, - custom_headers: Optional[dict], - forward_headers: Optional[bool], - merge_query_params: Optional[bool], - dependencies: Optional[List], - cost_per_request: Optional[float], + custom_headers: dict | None, + forward_headers: bool | None, + merge_query_params: bool | None, + dependencies: list | None, + cost_per_request: float | None, endpoint_id: str, - guardrails: Optional[dict] = None, - methods: Optional[List[str]] = None, - default_query_params: Optional[dict] = None, - config_file_path: Optional[str] = None, + guardrails: dict | None = None, + methods: list[str] | None = None, + default_query_params: dict | None = None, + config_file_path: str | None = None, auth: bool = False, - timeout: Optional[float] = None, + timeout: float | None = None, ): """Add wildcard route for sub-paths""" # Default to all methods if none specified (backward compatibility) @@ -2644,7 +2642,7 @@ class InitPassThroughEndpointHelpers: _registered_pass_through_routes.clear() @staticmethod - def get_all_registered_pass_through_routes() -> List[str]: + def get_all_registered_pass_through_routes() -> list[str]: """Get all registered pass-through endpoints from the registry""" return list(_registered_pass_through_routes.keys()) @@ -2687,7 +2685,7 @@ class InitPassThroughEndpointHelpers: # Keys are in format: "{endpoint_id}:exact:{path}:{methods}" or "{endpoint_id}:subpath:{path}:{methods}" # For backward compatibility, also support old format: "{endpoint_id}:exact:{path}" or "{endpoint_id}:subpath:{path}" # Extract unique paths from keys for quick checking - for key in _registered_pass_through_routes.keys(): + for key in _registered_pass_through_routes: parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?] if len(parts) >= 3: route_type = parts[1] @@ -2701,10 +2699,10 @@ class InitPassThroughEndpointHelpers: return False @staticmethod - def get_registered_pass_through_route(route: str, method: Optional[str] = None) -> Optional[Dict[str, Any]]: + def get_registered_pass_through_route(route: str, method: str | None = None) -> dict[str, Any] | None: """Get passthrough params for a given route and optionally filter by HTTP method""" comparison_route = InitPassThroughEndpointHelpers._route_for_registry_lookup(route) - for key in _registered_pass_through_routes.keys(): + for key in _registered_pass_through_routes: parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?] if len(parts) >= 3: route_type = parts[1] @@ -2714,7 +2712,7 @@ class InitPassThroughEndpointHelpers: # but keep supporting test fixtures / older registry entries that # only encoded methods in the route key. methods_entry = _registered_pass_through_routes[key].get("methods", []) - route_methods: List[str] = methods_entry if isinstance(methods_entry, list) else [] + route_methods: list[str] = methods_entry if isinstance(methods_entry, list) else [] if not route_methods and len(parts) == 4: route_methods = parts[3].split(",") @@ -2735,21 +2733,21 @@ class InitPassThroughEndpointHelpers: def _get_combined_pass_through_endpoints( - pass_through_endpoints: Union[List[Dict], List[PassThroughGenericEndpoint]], - config_pass_through_endpoints: List[Dict], + pass_through_endpoints: list[dict] | list[PassThroughGenericEndpoint], + config_pass_through_endpoints: list[dict], ): """Get combined pass-through endpoints from db + config""" return pass_through_endpoints + config_pass_through_endpoints async def _register_pass_through_endpoint( - endpoint: Union[Dict[str, Any], PassThroughGenericEndpoint], + endpoint: dict[str, Any] | PassThroughGenericEndpoint, app: FastAPI, premium_user: bool, visited_endpoints: set[str], - config_file_path: Optional[str] = None, + config_file_path: str | None = None, ) -> None: - endpoint_data: Dict[str, Any] + endpoint_data: dict[str, Any] if isinstance(endpoint, PassThroughGenericEndpoint): endpoint_data = endpoint.model_dump() else: @@ -2840,8 +2838,8 @@ async def _register_pass_through_endpoint( async def initialize_pass_through_endpoints( - pass_through_endpoints: Union[List[Dict], List[PassThroughGenericEndpoint]], - config_file_path: Optional[str] = None, + pass_through_endpoints: list[dict] | list[PassThroughGenericEndpoint], + config_file_path: str | None = None, ): """ 1. Create a global list of pass-through endpoints (db + config) @@ -2870,7 +2868,7 @@ async def initialize_pass_through_endpoints( ) ## get combined pass-through endpoints from db + config - combined_pass_through_endpoints: List[Union[Dict, PassThroughGenericEndpoint]] + combined_pass_through_endpoints: list[dict | PassThroughGenericEndpoint] if config_passthrough_endpoints is not None: combined_pass_through_endpoints = _get_combined_pass_through_endpoints( # type: ignore @@ -2908,7 +2906,7 @@ async def initialize_pass_through_endpoints( _registered_pass_through_routes.pop(endpoint_key, None) -def _get_pass_through_endpoints_from_config() -> List[PassThroughGenericEndpoint]: +def _get_pass_through_endpoints_from_config() -> list[PassThroughGenericEndpoint]: """ Get pass-through endpoints defined in the config file. These are read-only and cannot be edited via the UI. @@ -2921,7 +2919,7 @@ def _get_pass_through_endpoints_from_config() -> List[PassThroughGenericEndpoint if config_passthrough_endpoints is None or len(config_passthrough_endpoints) == 0: return [] - returned_endpoints: List[PassThroughGenericEndpoint] = [] + returned_endpoints: list[PassThroughGenericEndpoint] = [] for endpoint in config_passthrough_endpoints: try: if isinstance(endpoint, dict): @@ -2944,9 +2942,9 @@ def _get_pass_through_endpoints_from_config() -> List[PassThroughGenericEndpoint async def _get_pass_through_endpoints_from_db( - endpoint_id: Optional[str] = None, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> List[PassThroughGenericEndpoint]: + endpoint_id: str | None = None, + user_api_key_dict: UserAPIKeyAuth | None = None, +) -> list[PassThroughGenericEndpoint]: from litellm.proxy._types import LitellmUserRoles from litellm.proxy.proxy_server import get_config_general_settings @@ -2959,11 +2957,11 @@ async def _get_pass_through_endpoints_from_db( except Exception: return [] - pass_through_endpoint_data: Optional[List] = response.field_value + pass_through_endpoint_data: list | None = response.field_value if pass_through_endpoint_data is None: return [] - returned_endpoints: List[PassThroughGenericEndpoint] = [] + returned_endpoints: list[PassThroughGenericEndpoint] = [] if endpoint_id is None: # Return all endpoints from DB, mark as not from config for endpoint in pass_through_endpoint_data: @@ -2992,9 +2990,9 @@ async def _get_pass_through_endpoints_from_db( async def _filter_endpoints_by_team_allowed_routes( team_id: str, - pass_through_endpoints: List[PassThroughGenericEndpoint], + pass_through_endpoints: list[PassThroughGenericEndpoint], prisma_client, -) -> List[PassThroughGenericEndpoint]: +) -> list[PassThroughGenericEndpoint]: """ Filter pass-through endpoints based on team's allowed_passthrough_routes metadata. @@ -3043,9 +3041,9 @@ async def _filter_endpoints_by_team_allowed_routes( response_model=PassThroughEndpointResponse, ) async def get_pass_through_endpoints( - endpoint_id: Optional[str] = None, + endpoint_id: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - team_id: Optional[str] = None, + team_id: str | None = None, ): """ GET configured pass through endpoint. @@ -3117,7 +3115,7 @@ async def update_pass_through_endpoints( detail={"error": "No pass-through endpoints found"}, ) - pass_through_endpoint_data: Optional[List] = response.field_value + pass_through_endpoint_data: list | None = response.field_value if pass_through_endpoint_data is None: raise HTTPException( status_code=404, @@ -3185,7 +3183,7 @@ async def update_pass_through_endpoints( await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict) # Re-register the route with updated headers - _custom_headers: Optional[dict] = updated_endpoint.headers or {} + _custom_headers: dict | None = updated_endpoint.headers or {} _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers) if updated_endpoint.include_subpath: @@ -3261,7 +3259,7 @@ async def create_pass_through_endpoints( if response.field_value is None: response.field_value = [data_dict] - elif isinstance(response.field_value, List): + elif isinstance(response.field_value, list): response.field_value.append(data_dict) ## Update db @@ -3276,7 +3274,7 @@ async def create_pass_through_endpoints( created_endpoint = PassThroughGenericEndpoint(**data_dict) # Register the new route - _custom_headers: Optional[dict] = created_endpoint.headers or {} + _custom_headers: dict | None = created_endpoint.headers or {} _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers) if created_endpoint.include_subpath: @@ -3346,7 +3344,7 @@ async def delete_pass_through_endpoints( response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None) ## Update field by removing endpoint - pass_through_endpoint_data: Optional[List] = response.field_value + pass_through_endpoint_data: list | None = response.field_value if response.field_value is None or pass_through_endpoint_data is None: raise HTTPException( status_code=400, @@ -3359,7 +3357,7 @@ async def delete_pass_through_endpoints( if found_endpoint is None: raise HTTPException( status_code=400, - detail={"error": "Endpoint with ID '{}' was not found in pass-through endpoint list.".format(endpoint_id)}, + detail={"error": f"Endpoint with ID '{endpoint_id}' was not found in pass-through endpoint list."}, ) # Find the index for deleting from the list @@ -3395,9 +3393,9 @@ async def delete_pass_through_endpoints( def _find_endpoint_by_id( - endpoints_data: List, + endpoints_data: list, endpoint_id: str, -) -> Optional[PassThroughGenericEndpoint]: +) -> PassThroughGenericEndpoint | None: """ Find an endpoint by ID. @@ -3409,7 +3407,7 @@ def _find_endpoint_by_id( Found endpoint or None if not found """ for endpoint in endpoints_data: - _endpoint: Optional[PassThroughGenericEndpoint] = None + _endpoint: PassThroughGenericEndpoint | None = None if isinstance(endpoint, dict): _endpoint = PassThroughGenericEndpoint(**endpoint) elif isinstance(endpoint, PassThroughGenericEndpoint): diff --git a/litellm/proxy/pass_through_endpoints/passthrough_endpoint_router.py b/litellm/proxy/pass_through_endpoints/passthrough_endpoint_router.py index 7f62a822dee..ed35e3d1b46 100644 --- a/litellm/proxy/pass_through_endpoints/passthrough_endpoint_router.py +++ b/litellm/proxy/pass_through_endpoints/passthrough_endpoint_router.py @@ -1,5 +1,3 @@ -from typing import Dict, Optional - import litellm from litellm._logging import verbose_router_logger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( @@ -16,15 +14,15 @@ class PassthroughEndpointRouter: """ def __init__(self): - self.credentials: Dict[str, str] = {} - self.deployment_key_to_vertex_credentials: Dict[str, VertexPassThroughCredentials] = {} - self.default_vertex_config: Optional[VertexPassThroughCredentials] = None + self.credentials: dict[str, str] = {} + self.deployment_key_to_vertex_credentials: dict[str, VertexPassThroughCredentials] = {} + self.default_vertex_config: VertexPassThroughCredentials | None = None def set_pass_through_credentials( self, custom_llm_provider: str, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, ): """ Set credentials for a pass-through endpoint. Used when a user adds a pass-through LLM endpoint on the UI. @@ -45,8 +43,8 @@ class PassthroughEndpointRouter: def get_credentials( self, custom_llm_provider: str, - region_name: Optional[str], - ) -> Optional[str]: + region_name: str | None, + ) -> str | None: credential_name = self._get_credential_name_for_provider( custom_llm_provider=custom_llm_provider, region_name=region_name, @@ -77,7 +75,7 @@ class PassthroughEndpointRouter: vertex_credentials=get_secret_str("DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"), ) - def set_default_vertex_config(self, config: Optional[dict] = None): + def set_default_vertex_config(self, config: dict | None = None): """Sets vertex configuration from provided config and/or environment variables Args: @@ -104,7 +102,7 @@ class PassthroughEndpointRouter: self, project_id: str, location: str, - vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, ): """ Add the vertex credentials for the given project-id, location @@ -124,7 +122,7 @@ class PassthroughEndpointRouter: ) self.deployment_key_to_vertex_credentials[deployment_key] = vertex_pass_through_credentials - def _get_deployment_key(self, project_id: Optional[str], location: Optional[str]) -> Optional[str]: + def _get_deployment_key(self, project_id: str | None, location: str | None) -> str | None: """ Get the deployment key for the given project-id, location """ @@ -132,13 +130,13 @@ class PassthroughEndpointRouter: return None return f"{project_id}-{location}" - def get_vector_store_credentials(self, vector_store_id: str) -> Optional[LiteLLM_ManagedVectorStore]: + def get_vector_store_credentials(self, vector_store_id: str) -> LiteLLM_ManagedVectorStore | None: """ Get the vector store credentials for the given vector store id """ if litellm.vector_store_registry is None: return None - vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( + vector_store_to_run: LiteLLM_ManagedVectorStore | None = ( litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( vector_store_id=vector_store_id ) @@ -146,8 +144,8 @@ class PassthroughEndpointRouter: return vector_store_to_run def get_vertex_credentials( - self, project_id: Optional[str], location: Optional[str] - ) -> Optional[VertexPassThroughCredentials]: + self, project_id: str | None, location: str | None + ) -> VertexPassThroughCredentials | None: """ Get the vertex credentials for the given project-id, location """ @@ -166,7 +164,7 @@ class PassthroughEndpointRouter: def _get_credential_name_for_provider( self, custom_llm_provider: str, - region_name: Optional[str], + region_name: str | None, ) -> str: if region_name is None: return f"{custom_llm_provider.upper()}_API_KEY" @@ -175,8 +173,8 @@ class PassthroughEndpointRouter: def _get_region_name_from_api_base( self, custom_llm_provider: str, - api_base: Optional[str], - ) -> Optional[str]: + api_base: str | None, + ) -> str | None: """ Get the region name from the API base. diff --git a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py index 662c921bd4b..45d62d40e6d 100644 --- a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py +++ b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py @@ -7,7 +7,7 @@ Handles guardrail execution for passthrough endpoints with: - Automatic inheritance from org/team/key levels when enabled """ -from typing import Any, Dict, List, Optional, Union +from typing import Any, Union from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( @@ -20,7 +20,7 @@ from litellm.proxy.pass_through_endpoints.jsonpath_extractor import JsonPathExtr # Type for raw guardrails config input (before normalization) # Can be a list of names or a dict with settings PassThroughGuardrailsConfigInput = Union[ - List[str], # Simple list: ["guard-1", "guard-2"] + list[str], # Simple list: ["guard-1", "guard-2"] PassThroughGuardrailsConfig, # Dict: {"guard-1": {"request_fields": [...]}} ] @@ -41,8 +41,8 @@ class PassthroughGuardrailHandler: @staticmethod def normalize_config( - guardrails_config: Optional[PassThroughGuardrailsConfigInput], - ) -> Optional[PassThroughGuardrailsConfig]: + guardrails_config: PassThroughGuardrailsConfigInput | None, + ) -> PassThroughGuardrailsConfig | None: """ Normalize guardrails config to dict format. @@ -70,7 +70,7 @@ class PassthroughGuardrailHandler: @staticmethod def is_enabled( - guardrails_config: Optional[PassThroughGuardrailsConfigInput], + guardrails_config: PassThroughGuardrailsConfigInput | None, ) -> bool: """ Check if guardrails are enabled for a passthrough endpoint. @@ -85,8 +85,8 @@ class PassthroughGuardrailHandler: @staticmethod def get_guardrail_names( - guardrails_config: Optional[PassThroughGuardrailsConfigInput], - ) -> List[str]: + guardrails_config: PassThroughGuardrailsConfigInput | None, + ) -> list[str]: """Get the list of guardrail names configured for a passthrough endpoint.""" normalized = PassthroughGuardrailHandler.normalize_config(guardrails_config) if normalized is None: @@ -95,9 +95,9 @@ class PassthroughGuardrailHandler: @staticmethod def get_settings( - guardrails_config: Optional[PassThroughGuardrailsConfigInput], + guardrails_config: PassThroughGuardrailsConfigInput | None, guardrail_name: str, - ) -> Optional[PassThroughGuardrailSettings]: + ) -> PassThroughGuardrailSettings | None: """Get settings for a specific guardrail from the passthrough config.""" normalized = PassthroughGuardrailHandler.normalize_config(guardrails_config) if normalized is None: @@ -115,7 +115,7 @@ class PassthroughGuardrailHandler: @staticmethod def prepare_input( request_data: dict, - guardrail_settings: Optional[PassThroughGuardrailSettings], + guardrail_settings: PassThroughGuardrailSettings | None, ) -> str: """ Prepare input text for guardrail execution based on field targeting settings. @@ -136,7 +136,7 @@ class PassthroughGuardrailHandler: @staticmethod def prepare_output( response_data: dict, - guardrail_settings: Optional[PassThroughGuardrailSettings], + guardrail_settings: PassThroughGuardrailSettings | None, ) -> str: """ Prepare output text for guardrail execution based on field targeting settings. @@ -158,7 +158,7 @@ class PassthroughGuardrailHandler: async def execute( request_data: dict, user_api_key_dict: UserAPIKeyAuth, - guardrails_config: Optional[PassThroughGuardrailsConfig], + guardrails_config: PassThroughGuardrailsConfig | None, event_type: str = "pre_call", ) -> dict: """ @@ -204,8 +204,8 @@ class PassthroughGuardrailHandler: @staticmethod def collect_guardrails( user_api_key_dict: UserAPIKeyAuth, - passthrough_guardrails_config: Optional[PassThroughGuardrailsConfigInput], - ) -> Optional[Dict[str, bool]]: + passthrough_guardrails_config: PassThroughGuardrailsConfigInput | None, + ) -> dict[str, bool] | None: """ Collect guardrails for a passthrough endpoint. @@ -243,7 +243,7 @@ class PassthroughGuardrailHandler: return None # Passthrough is enabled - collect guardrails - guardrails_to_run: Dict[str, bool] = {} + guardrails_to_run: dict[str, bool] = {} # Add passthrough-specific guardrails for guardrail_name in normalized_config.keys(): @@ -251,7 +251,7 @@ class PassthroughGuardrailHandler: verbose_proxy_logger.debug("Added passthrough-specific guardrail: %s", guardrail_name) # Add org/team/key level guardrails using shared helper - temp_data: Dict[str, Any] = {"metadata": {}} + temp_data: dict[str, Any] = {"metadata": {}} _add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, @@ -278,7 +278,7 @@ class PassthroughGuardrailHandler: data: dict, guardrail_name: str, is_request: bool = True, - ) -> Optional[str]: + ) -> str | None: """ Get the text to check for a guardrail, respecting field targeting settings. diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 4dc1e0e70dd..010cf8a7561 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -1,5 +1,4 @@ from datetime import datetime -from typing import List, Optional, Tuple import httpx @@ -33,14 +32,14 @@ class PassThroughStreamingHandler: @staticmethod async def chunk_processor( response: httpx.Response, - request_body: Optional[dict], + request_body: dict | None, litellm_logging_obj: LiteLLMLoggingObj, endpoint_type: EndpointType, start_time: datetime, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, ): - raw_bytes: List[bytes] = [] + raw_bytes: list[bytes] = [] logging_scheduled = False model_name = PassThroughStreamingHandler._extract_model_for_cost_injection( request_body=request_body, @@ -90,7 +89,7 @@ class PassThroughStreamingHandler: yield chunk except Exception as e: - verbose_proxy_logger.error(f"Error in chunk_processor: {str(e)}") + verbose_proxy_logger.error(f"Error in chunk_processor: {e!s}") raise finally: # GeneratorExit (raised on client disconnect) is not caught by @@ -116,7 +115,7 @@ class PassThroughStreamingHandler: ) ) except Exception as e: - verbose_proxy_logger.error(f"Error scheduling chunk_processor logging: {str(e)}") + verbose_proxy_logger.error(f"Error scheduling chunk_processor logging: {e!s}") @staticmethod async def _route_streaming_logging_to_handler( @@ -126,9 +125,9 @@ class PassThroughStreamingHandler: request_body: dict, endpoint_type: EndpointType, start_time: datetime, - raw_bytes: List[bytes], + raw_bytes: list[bytes], end_time: datetime, - model: Optional[str] = None, + model: str | None = None, ): """ Route the logging for the collected chunks to the appropriate handler @@ -166,7 +165,7 @@ class PassThroughStreamingHandler: **kwargs, ) except Exception as e: - verbose_proxy_logger.error(f"Error in _route_streaming_logging_to_handler: {str(e)}") + verbose_proxy_logger.error(f"Error in _route_streaming_logging_to_handler: {e!s}") @staticmethod def _build_passthrough_logging_result( @@ -176,10 +175,10 @@ class PassThroughStreamingHandler: request_body: dict, endpoint_type: EndpointType, start_time: datetime, - raw_bytes: List[bytes], + raw_bytes: list[bytes], end_time: datetime, - model: Optional[str], - ) -> Tuple[PassThroughEndpointLoggingResultValues, dict]: + model: str | None, + ) -> tuple[PassThroughEndpointLoggingResultValues, dict]: """ Synchronous, CPU-bound reconstruction of the standard logging payload from collected raw SSE bytes. Extracted from @@ -188,7 +187,7 @@ class PassThroughStreamingHandler: loop; an off-loop dispatch is a future change, not part of this PR. """ all_chunks = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(raw_bytes) - standard_logging_response_object: Optional[PassThroughEndpointLoggingResultValues] = None + standard_logging_response_object: PassThroughEndpointLoggingResultValues | None = None kwargs: dict = {} if endpoint_type == EndpointType.ANTHROPIC: anthropic_passthrough_logging_handler_result = ( @@ -245,11 +244,11 @@ class PassThroughStreamingHandler: @staticmethod def _extract_model_for_cost_injection( - request_body: Optional[dict], + request_body: dict | None, url_route: str, endpoint_type: EndpointType, litellm_logging_obj: LiteLLMLoggingObj, - ) -> Optional[str]: + ) -> str | None: """ Extract model name for cost injection from various sources. """ @@ -274,7 +273,7 @@ class PassThroughStreamingHandler: return None @staticmethod - def _convert_raw_bytes_to_str_lines(raw_bytes: List[bytes]) -> List[str]: + def _convert_raw_bytes_to_str_lines(raw_bytes: list[bytes]) -> list[str]: """ Converts a list of raw bytes into a list of string lines, similar to aiter_lines() diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 20ed84c8636..7447c489d94 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -1,6 +1,6 @@ import json from datetime import datetime -from typing import Any, Optional, Union +from typing import Any from urllib.parse import urlparse import httpx @@ -93,11 +93,9 @@ class PassThroughEndpointLogging: async def _handle_logging( self, logging_obj: LiteLLMLoggingObj, - standard_logging_response_object: Union[ - StandardPassThroughResponseObject, - PassThroughEndpointLoggingResultValues, - dict, - ], + standard_logging_response_object: StandardPassThroughResponseObject + | PassThroughEndpointLoggingResultValues + | dict, result: str, start_time: datetime, end_time: datetime, @@ -124,7 +122,7 @@ class PassThroughEndpointLogging: def normalize_llm_passthrough_logging_payload( self, httpx_response: httpx.Response, - response_body: Optional[dict], + response_body: dict | None, request_body: dict, logging_obj: LiteLLMLoggingObj, url_route: str, @@ -132,14 +130,14 @@ class PassThroughEndpointLogging: start_time: datetime, end_time: datetime, cache_hit: bool, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): return_dict = { "standard_logging_response_object": None, "kwargs": kwargs, } - standard_logging_response_object: Optional[Any] = None + standard_logging_response_object: Any | None = None if self.is_gemini_route(url_route, custom_llm_provider): gemini_passthrough_logging_handler_result = GeminiPassthroughLoggingHandler.gemini_passthrough_handler( @@ -268,7 +266,7 @@ class PassThroughEndpointLogging: async def pass_through_async_success_handler( self, httpx_response: httpx.Response, - response_body: Optional[dict], + response_body: dict | None, logging_obj: LiteLLMLoggingObj, url_route: str, result: str, @@ -277,10 +275,10 @@ class PassThroughEndpointLogging: cache_hit: bool, request_body: dict, passthrough_logging_payload: PassthroughStandardLoggingPayload, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): - standard_logging_response_object: Optional[PassThroughEndpointLoggingResultValues] = None + standard_logging_response_object: PassThroughEndpointLoggingResultValues | None = None logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload if self.is_assemblyai_route(url_route): if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True: @@ -358,9 +356,7 @@ class PassThroughEndpointLogging: def is_assemblyai_route(self, url_route: str): parsed_url = urlparse(url_route) - if parsed_url.hostname == "api.assemblyai.com": - return True - elif "/transcript" in parsed_url.path: + if parsed_url.hostname == "api.assemblyai.com" or "/transcript" in parsed_url.path: return True return False @@ -380,7 +376,7 @@ class PassThroughEndpointLogging: return True return False - def is_cursor_route(self, url_route: str, custom_llm_provider: Optional[str] = None): + def is_cursor_route(self, url_route: str, custom_llm_provider: str | None = None): """Check if the URL route is a Cursor Cloud Agents API route.""" if custom_llm_provider == "cursor": return True @@ -409,7 +405,7 @@ class PassThroughEndpointLogging: return _is_openai_compatible_url(url_route) - def is_gemini_route(self, url_route: str, custom_llm_provider: Optional[str] = None): + def is_gemini_route(self, url_route: str, custom_llm_provider: str | None = None): """Check if the URL route is a Gemini API route.""" for route in self.TRACKED_GEMINI_ROUTES: if route in url_route and custom_llm_provider == "gemini": diff --git a/litellm/proxy/plugin_routes.py b/litellm/proxy/plugin_routes.py index b6bd594096f..e8ea90d48b7 100644 --- a/litellm/proxy/plugin_routes.py +++ b/litellm/proxy/plugin_routes.py @@ -31,9 +31,9 @@ from collections.abc import Mapping from cryptography.fernet import Fernet, InvalidToken from fastapi import APIRouter, Depends, HTTPException, Request, Response +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import PluginConfig, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider router = APIRouter() diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index fe9ad3bef6d..ed0d98c6e6a 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -6,7 +6,7 @@ This allows the same policy to be attached to multiple scopes. """ from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_proxy_logger from litellm.repositories.table_repositories import PolicyAttachmentRepository @@ -41,11 +41,11 @@ class AttachmentRegistry: """ def __init__(self): - self._attachments: List[PolicyAttachment] = [] + self._attachments: list[PolicyAttachment] = [] self._config_attachments: tuple[PolicyAttachment, ...] = () self._initialized: bool = False - def load_attachments(self, attachments_config: List[Dict[str, Any]]) -> None: + def load_attachments(self, attachments_config: list[dict[str, Any]]) -> None: """ Load attachments from a configuration list. @@ -60,14 +60,14 @@ class AttachmentRegistry: self._attachments.append(attachment) verbose_proxy_logger.debug(f"Loaded attachment for policy: {attachment.policy}") except Exception as e: - verbose_proxy_logger.error(f"Error loading attachment: {str(e)}") - raise ValueError(f"Invalid attachment: {str(e)}") from e + verbose_proxy_logger.error(f"Error loading attachment: {e!s}") + raise ValueError(f"Invalid attachment: {e!s}") from e self._config_attachments = tuple(self._attachments) self._initialized = True verbose_proxy_logger.info(f"Loaded {len(self._attachments)} policy attachments") - def _parse_attachment(self, attachment_data: Dict[str, Any]) -> PolicyAttachment: + def _parse_attachment(self, attachment_data: dict[str, Any]) -> PolicyAttachment: """ Parse an attachment from raw configuration data. @@ -86,7 +86,7 @@ class AttachmentRegistry: tags=attachment_data.get("tags"), ) - def get_attached_policies(self, context: PolicyMatchContext) -> List[str]: + def get_attached_policies(self, context: PolicyMatchContext) -> list[str]: """ Get list of policy names attached to the given context. @@ -98,7 +98,7 @@ class AttachmentRegistry: """ return [r["policy_name"] for r in self.get_attached_policies_with_reasons(context)] - def get_attached_policies_with_reasons(self, context: PolicyMatchContext) -> List[Dict[str, Any]]: + def get_attached_policies_with_reasons(self, context: PolicyMatchContext) -> list[dict[str, Any]]: """ Get list of policy names and match reasons for the given context. @@ -107,7 +107,7 @@ class AttachmentRegistry: """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - results: List[Dict[str, Any]] = [] + results: list[dict[str, Any]] = [] seen_policies: set = set() for attachment in self._attachments: @@ -166,7 +166,7 @@ class AttachmentRegistry: attached = self.get_attached_policies(context) return policy_name in attached - def get_all_attachments(self) -> List[PolicyAttachment]: + def get_all_attachments(self) -> list[PolicyAttachment]: """ Get all loaded attachments. @@ -184,7 +184,7 @@ class AttachmentRegistry: """ return self._config_attachments - def get_attachments_for_policy(self, policy_name: str) -> List[PolicyAttachment]: + def get_attachments_for_policy(self, policy_name: str) -> list[PolicyAttachment]: """ Get all attachments for a specific policy. @@ -263,7 +263,7 @@ class AttachmentRegistry: self, attachment_request: PolicyAttachmentCreateRequest, prisma_client: "PrismaClient", - created_by: Optional[str] = None, + created_by: str | None = None, ) -> PolicyAttachmentDBResponse: """ Add a policy attachment to the database. @@ -318,13 +318,13 @@ class AttachmentRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error adding attachment to DB: {e}") - raise Exception(f"Error adding attachment to DB: {str(e)}") + raise Exception(f"Error adding attachment to DB: {e!s}") async def delete_attachment_from_db( self, attachment_id: str, prisma_client: "PrismaClient", - ) -> Dict[str, str]: + ) -> dict[str, str]: """ Delete a policy attachment from the database. @@ -354,13 +354,13 @@ class AttachmentRegistry: return {"message": f"Attachment {attachment_id} deleted successfully"} except Exception as e: verbose_proxy_logger.exception(f"Error deleting attachment from DB: {e}") - raise Exception(f"Error deleting attachment from DB: {str(e)}") + raise Exception(f"Error deleting attachment from DB: {e!s}") async def get_attachment_by_id_from_db( self, attachment_id: str, prisma_client: "PrismaClient", - ) -> Optional[PolicyAttachmentDBResponse]: + ) -> PolicyAttachmentDBResponse | None: """ Get a policy attachment by ID from the database. @@ -394,12 +394,12 @@ class AttachmentRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error getting attachment from DB: {e}") - raise Exception(f"Error getting attachment from DB: {str(e)}") + raise Exception(f"Error getting attachment from DB: {e!s}") async def get_all_attachments_from_db( self, prisma_client: "PrismaClient", - ) -> List[PolicyAttachmentDBResponse]: + ) -> list[PolicyAttachmentDBResponse]: """ Get all policy attachments from the database. @@ -432,7 +432,7 @@ class AttachmentRegistry: ] except Exception as e: verbose_proxy_logger.exception(f"Error getting attachments from DB: {e}") - raise Exception(f"Error getting attachments from DB: {str(e)}") + raise Exception(f"Error getting attachments from DB: {e!s}") async def sync_attachments_from_db( self, @@ -468,11 +468,11 @@ class AttachmentRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error syncing attachments from DB: {e}") - raise Exception(f"Error syncing attachments from DB: {str(e)}") + raise Exception(f"Error syncing attachments from DB: {e!s}") # Global singleton instance -_attachment_registry: Optional[AttachmentRegistry] = None +_attachment_registry: AttachmentRegistry | None = None def get_attachment_registry() -> AttachmentRegistry: diff --git a/litellm/proxy/policy_engine/condition_evaluator.py b/litellm/proxy/policy_engine/condition_evaluator.py index de27b6d6d23..02268fcd721 100644 --- a/litellm/proxy/policy_engine/condition_evaluator.py +++ b/litellm/proxy/policy_engine/condition_evaluator.py @@ -5,7 +5,6 @@ Supports model-based conditions with exact match or regex patterns. """ import re -from typing import List, Optional, Union from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import ( @@ -26,7 +25,7 @@ class ConditionEvaluator: @staticmethod def evaluate( - condition: Optional[PolicyCondition], + condition: PolicyCondition | None, context: PolicyMatchContext, ) -> bool: """ @@ -56,8 +55,8 @@ class ConditionEvaluator: @staticmethod def _evaluate_model_condition( - condition: Union[str, List[str]], - model: Optional[str], + condition: str | list[str], + model: str | None, ) -> bool: """ Evaluate a model condition. diff --git a/litellm/proxy/policy_engine/init_policies.py b/litellm/proxy/policy_engine/init_policies.py index a3b0d1a6cdb..9fb700770d2 100644 --- a/litellm/proxy/policy_engine/init_policies.py +++ b/litellm/proxy/policy_engine/init_policies.py @@ -6,7 +6,7 @@ Configuration structure: - policy_attachments: Define WHERE policies apply (teams, keys, models) """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry @@ -25,8 +25,8 @@ _reset_color_code = "\033[0m" def _print_policies_on_startup( - policies_config: Dict[str, Any], - policy_attachments_config: Optional[List[Dict[str, Any]]] = None, + policies_config: dict[str, Any], + policy_attachments_config: list[dict[str, Any]] | None = None, ) -> None: """ Print loaded policies to console on startup (similar to model list). @@ -96,8 +96,8 @@ def _print_policies_on_startup( async def init_policies( - policies_config: Dict[str, Any], - policy_attachments_config: Optional[List[Dict[str, Any]]] = None, + policies_config: dict[str, Any], + policy_attachments_config: list[dict[str, Any]] | None = None, prisma_client: Optional["PrismaClient"] = None, validate_db: bool = True, fail_on_error: bool = True, @@ -167,7 +167,7 @@ async def init_policies( policy_registry.load_policies(policies_config) verbose_proxy_logger.info(f"Successfully loaded {len(policies_config)} policies") except Exception as e: - verbose_proxy_logger.error(f"Failed to load policies: {str(e)}") + verbose_proxy_logger.error(f"Failed to load policies: {e!s}") raise # Load attachments if provided @@ -176,15 +176,15 @@ async def init_policies( attachment_registry.load_attachments(policy_attachments_config) verbose_proxy_logger.info(f"Successfully loaded {len(policy_attachments_config)} policy attachments") except Exception as e: - verbose_proxy_logger.error(f"Failed to load policy attachments: {str(e)}") + verbose_proxy_logger.error(f"Failed to load policy attachments: {e!s}") raise return validation_result def init_policies_sync( - policies_config: Dict[str, Any], - policy_attachments_config: Optional[List[Dict[str, Any]]] = None, + policies_config: dict[str, Any], + policy_attachments_config: list[dict[str, Any]] | None = None, fail_on_error: bool = True, ) -> None: """ @@ -217,7 +217,7 @@ def init_policies_sync( ) -def get_policies_summary() -> Dict[str, Any]: +def get_policies_summary() -> dict[str, Any]: """ Get a summary of loaded policies for debugging/display. @@ -234,7 +234,7 @@ def get_policies_summary() -> Dict[str, Any]: resolved = PolicyResolver.get_all_resolved_policies() - summary: Dict[str, Any] = { + summary: dict[str, Any] = { "initialized": True, "policy_count": len(resolved), "attachment_count": len(attachment_registry.get_all_attachments()), diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index 1f507bb4c54..983d2da124b 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -6,7 +6,7 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding. """ import time -from typing import Any, List, Literal, Optional +from typing import Any, Literal import litellm from litellm._logging import verbose_proxy_logger @@ -35,7 +35,7 @@ class PipelineExecutor: @staticmethod async def execute_steps( - steps: List[PipelineStep], + steps: list[PipelineStep], mode: str, data: dict, user_api_key_dict: Any, @@ -56,7 +56,7 @@ class PipelineExecutor: Returns: PipelineExecutionResult with terminal action and step results """ - step_results: List[PipelineStepResult] = [] + step_results: list[PipelineStepResult] = [] working_data = data.copy() if "metadata" in working_data: working_data["metadata"] = working_data["metadata"].copy() @@ -140,9 +140,9 @@ class PipelineExecutor: call_type: str, ) -> tuple[ Literal["pass", "fail", "error"], - Optional[dict], - Optional[str], - Optional[Exception], + dict | None, + str | None, + Exception | None, ]: """ Run a single pipeline step's guardrail. @@ -209,7 +209,7 @@ class PipelineExecutor: return ("error", None, str(e), e) @staticmethod - def find_guardrail_callback(guardrail_name: str) -> Optional[CustomGuardrail]: + def find_guardrail_callback(guardrail_name: str) -> CustomGuardrail | None: """Look up an initialized guardrail callback by name from litellm.callbacks.""" for callback in litellm.callbacks: if isinstance(callback, CustomGuardrail): diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index cff1378c676..1d0fd67cfc8 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -4,8 +4,6 @@ CRUD ENDPOINTS FOR POLICIES Provides REST API endpoints for managing policies and policy attachments. """ -from typing import Optional - from fastapi import APIRouter, Depends, HTTPException from litellm._logging import verbose_proxy_logger @@ -75,7 +73,7 @@ def _config_attachment_to_db_response(index: int, attachment: PolicyAttachment) dependencies=[Depends(user_api_key_auth)], response_model=PolicyListDBResponse, ) -async def list_policies(version_status: Optional[str] = None): +async def list_policies(version_status: str | None = None): """ List all policies from the database and config.yaml. Optionally filter by version_status. diff --git a/litellm/proxy/policy_engine/policy_matcher.py b/litellm/proxy/policy_engine/policy_matcher.py index acaa9629d83..48daaa732ba 100644 --- a/litellm/proxy/policy_engine/policy_matcher.py +++ b/litellm/proxy/policy_engine/policy_matcher.py @@ -7,8 +7,6 @@ apply to a given request based on team alias, key alias, and model. Policies are matched via policy_attachments which define WHERE each policy applies. """ -from typing import Dict, List, Optional - from litellm._logging import verbose_proxy_logger from litellm.proxy.auth.route_checks import RouteChecks from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext, PolicyScope @@ -26,7 +24,7 @@ class PolicyMatcher: """ @staticmethod - def matches_pattern(value: Optional[str], patterns: List[str]) -> bool: + def matches_pattern(value: str | None, patterns: list[str]) -> bool: """ Check if a value matches any of the given patterns. @@ -94,7 +92,7 @@ class PolicyMatcher: @staticmethod def get_matching_policies( context: PolicyMatchContext, - ) -> List[str]: + ) -> list[str]: """ Get list of policy names that match the given context via attachments. @@ -118,7 +116,7 @@ class PolicyMatcher: @staticmethod def get_matching_policies_from_registry( context: PolicyMatchContext, - ) -> List[str]: + ) -> list[str]: """ Get list of policy names that match the given context from the global registry. @@ -132,10 +130,10 @@ class PolicyMatcher: @staticmethod def get_policies_with_matching_conditions( - policy_names: List[str], + policy_names: list[str], context: PolicyMatchContext, - policies: Optional[Dict[str, Policy]] = None, - ) -> List[str]: + policies: dict[str, Policy] | None = None, + ) -> list[str]: """ Filter policies to only those whose conditions match the context. diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 01b88836387..32dfc44b8ba 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -187,8 +187,8 @@ class PolicyRegistry: self._policies[policy_name] = policy verbose_proxy_logger.debug(f"Loaded policy: {policy_name}") except Exception as e: - verbose_proxy_logger.error(f"Error loading policy '{policy_name}': {str(e)}") - raise ValueError(f"Invalid policy '{policy_name}': {str(e)}") from e + verbose_proxy_logger.error(f"Error loading policy '{policy_name}': {e!s}") + raise ValueError(f"Invalid policy '{policy_name}': {e!s}") from e self._config_policies = dict(self._policies) self._sources = {policy_name: "config" for policy_name in self._policies} @@ -310,7 +310,7 @@ class PolicyRegistry: self._sources = {} self._initialized = False - def get_source(self, policy_name: str) -> Optional[Literal["db", "config"]]: + def get_source(self, policy_name: str) -> Literal["db", "config"] | None: """ Return the provenance of an in-memory policy, or None if unknown. """ @@ -433,7 +433,7 @@ class PolicyRegistry: return _row_to_policy_db_response(created_policy) except Exception as e: verbose_proxy_logger.exception(f"Error adding policy to DB: {e}") - raise Exception(f"Error adding policy to DB: {str(e)}") + raise Exception(f"Error adding policy to DB: {e!s}") async def update_policy_in_db( self, @@ -497,7 +497,7 @@ class PolicyRegistry: return _row_to_policy_db_response(updated_policy) except Exception as e: verbose_proxy_logger.exception(f"Error updating policy in DB: {e}") - raise Exception(f"Error updating policy in DB: {str(e)}") + raise Exception(f"Error updating policy in DB: {e!s}") async def delete_policy_from_db( self, @@ -547,7 +547,7 @@ class PolicyRegistry: return result except Exception as e: verbose_proxy_logger.exception(f"Error deleting policy from DB: {e}") - raise Exception(f"Error deleting policy from DB: {str(e)}") + raise Exception(f"Error deleting policy from DB: {e!s}") async def get_policy_by_id_from_db( self, @@ -573,7 +573,7 @@ class PolicyRegistry: return _row_to_policy_db_response(policy) except Exception as e: verbose_proxy_logger.exception(f"Error getting policy from DB: {e}") - raise Exception(f"Error getting policy from DB: {str(e)}") + raise Exception(f"Error getting policy from DB: {e!s}") def get_policy_by_id_for_request(self, policy_id: str) -> tuple[str, Policy] | None: """ @@ -620,7 +620,7 @@ class PolicyRegistry: return [_row_to_policy_db_response(p) for p in policies] except Exception as e: verbose_proxy_logger.exception(f"Error getting policies from DB: {e}") - raise Exception(f"Error getting policies from DB: {str(e)}") + raise Exception(f"Error getting policies from DB: {e!s}") async def sync_policies_from_db( self, @@ -689,7 +689,7 @@ class PolicyRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}") - raise Exception(f"Error syncing policies from DB: {str(e)}") + raise Exception(f"Error syncing policies from DB: {e!s}") async def resolve_guardrails_from_db( self, @@ -742,7 +742,7 @@ class PolicyRegistry: return sorted(resolved_policy.guardrails) except Exception as e: verbose_proxy_logger.exception(f"Error resolving guardrails from DB: {e}") - raise Exception(f"Error resolving guardrails from DB: {str(e)}") + raise Exception(f"Error resolving guardrails from DB: {e!s}") async def get_versions_by_policy_name( self, @@ -772,7 +772,7 @@ class PolicyRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error getting versions: {e}") - raise Exception(f"Error getting versions: {str(e)}") + raise Exception(f"Error getting versions: {e!s}") async def create_new_version( self, @@ -858,7 +858,7 @@ class PolicyRegistry: return _row_to_policy_db_response(created) except Exception as e: verbose_proxy_logger.exception(f"Error creating new version: {e}") - raise Exception(f"Error creating new version: {str(e)}") + raise Exception(f"Error creating new version: {e!s}") async def update_version_status( self, @@ -963,7 +963,7 @@ class PolicyRegistry: return _row_to_policy_db_response(updated) except Exception as e: verbose_proxy_logger.exception(f"Error updating version status: {e}") - raise Exception(f"Error updating version status: {str(e)}") + raise Exception(f"Error updating version status: {e!s}") async def compare_versions( self, @@ -1016,7 +1016,7 @@ class PolicyRegistry: ) except Exception as e: verbose_proxy_logger.exception(f"Error comparing versions: {e}") - raise Exception(f"Error comparing versions: {str(e)}") + raise Exception(f"Error comparing versions: {e!s}") async def delete_all_versions( self, @@ -1047,7 +1047,7 @@ class PolicyRegistry: return {"message": message} except Exception as e: verbose_proxy_logger.exception(f"Error deleting all versions: {e}") - raise Exception(f"Error deleting all versions: {str(e)}") + raise Exception(f"Error deleting all versions: {e!s}") # Global singleton instance diff --git a/litellm/proxy/policy_engine/policy_resolver.py b/litellm/proxy/policy_engine/policy_resolver.py index f5520f4b4b4..65c8236a0b3 100644 --- a/litellm/proxy/policy_engine/policy_resolver.py +++ b/litellm/proxy/policy_engine/policy_resolver.py @@ -8,8 +8,6 @@ Handles: - Combining guardrails from multiple matching policies """ -from typing import Dict, List, Optional, Set, Tuple - from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import ( GuardrailPipeline, @@ -31,9 +29,9 @@ class PolicyResolver: @staticmethod def resolve_inheritance_chain( policy_name: str, - policies: Dict[str, Policy], - visited: Optional[Set[str]] = None, - ) -> List[str]: + policies: dict[str, Policy], + visited: set[str] | None = None, + ) -> list[str]: """ Get the inheritance chain for a policy (from root to policy). @@ -69,8 +67,8 @@ class PolicyResolver: @staticmethod def resolve_policy_guardrails( policy_name: str, - policies: Dict[str, Policy], - context: Optional[PolicyMatchContext] = None, + policies: dict[str, Policy], + context: PolicyMatchContext | None = None, ) -> ResolvedPolicy: """ Resolve the final guardrails for a single policy, including inheritance. @@ -93,7 +91,7 @@ class PolicyResolver: inheritance_chain = PolicyResolver.resolve_inheritance_chain(policy_name=policy_name, policies=policies) # Start with empty set of guardrails - guardrails: Set[str] = set() + guardrails: set[str] = set() # Apply each policy in the chain (from root to leaf) for chain_policy_name in inheritance_chain: @@ -129,9 +127,9 @@ class PolicyResolver: @staticmethod def resolve_guardrails_for_context( context: PolicyMatchContext, - policies: Optional[Dict[str, Policy]] = None, - policy_names: Optional[List[str]] = None, - ) -> List[str]: + policies: dict[str, Policy] | None = None, + policy_names: list[str] | None = None, + ) -> list[str]: """ Resolve the final list of guardrails for a request context. @@ -171,7 +169,7 @@ class PolicyResolver: return [] # Resolve each matching policy and combine guardrails - all_guardrails: Set[str] = set() + all_guardrails: set[str] = set() for policy_name in matching_policy_names: resolved = PolicyResolver.resolve_policy_guardrails( @@ -190,9 +188,9 @@ class PolicyResolver: @staticmethod def resolve_pipelines_for_context( context: PolicyMatchContext, - policies: Optional[Dict[str, Policy]] = None, - policy_names: Optional[List[str]] = None, - ) -> List[Tuple[str, GuardrailPipeline]]: + policies: dict[str, Policy] | None = None, + policy_names: list[str] | None = None, + ) -> list[tuple[str, GuardrailPipeline]]: """ Resolve pipelines from matching policies for a request context. @@ -223,7 +221,7 @@ class PolicyResolver: if not matching_policy_names: return [] - pipelines: List[Tuple[str, GuardrailPipeline]] = [] + pipelines: list[tuple[str, GuardrailPipeline]] = [] for policy_name in matching_policy_names: policy = policies.get(policy_name) if policy is None: @@ -238,14 +236,14 @@ class PolicyResolver: @staticmethod def get_pipeline_managed_guardrails( - pipelines: List[Tuple[str, GuardrailPipeline]], - ) -> Set[str]: + pipelines: list[tuple[str, GuardrailPipeline]], + ) -> set[str]: """ Get the set of guardrail names managed by pipelines. These guardrails should be excluded from normal independent execution. """ - managed: Set[str] = set() + managed: set[str] = set() for _policy_name, pipeline in pipelines: for step in pipeline.steps: managed.add(step.guardrail) @@ -253,9 +251,9 @@ class PolicyResolver: @staticmethod def get_all_resolved_policies( - policies: Optional[Dict[str, Policy]] = None, - context: Optional[PolicyMatchContext] = None, - ) -> Dict[str, ResolvedPolicy]: + policies: dict[str, Policy] | None = None, + context: PolicyMatchContext | None = None, + ) -> dict[str, ResolvedPolicy]: """ Resolve all policies and return their final guardrails. @@ -276,7 +274,7 @@ class PolicyResolver: return {} policies = registry.get_all_policies() - resolved: Dict[str, ResolvedPolicy] = {} + resolved: dict[str, ResolvedPolicy] = {} for policy_name in policies: resolved[policy_name] = PolicyResolver.resolve_policy_guardrails( diff --git a/litellm/proxy/policy_engine/policy_validator.py b/litellm/proxy/policy_engine/policy_validator.py index 4db5bc0435e..824f009c474 100644 --- a/litellm/proxy/policy_engine/policy_validator.py +++ b/litellm/proxy/policy_engine/policy_validator.py @@ -10,7 +10,7 @@ Validates: """ import asyncio -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.proxy.auth.route_checks import RouteChecks @@ -63,7 +63,7 @@ class PolicyValidator: """ return "*" in pattern or "?" in pattern - def get_available_guardrails(self) -> Set[str]: + def get_available_guardrails(self) -> set[str]: """ Get set of available guardrail names from the guardrail registry. @@ -78,7 +78,7 @@ class PolicyValidator: guardrails = IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails() return {g.get("guardrail_name", "") for g in guardrails if g.get("guardrail_name")} except Exception as e: - verbose_proxy_logger.warning(f"Could not get guardrails from registry: {str(e)}") + verbose_proxy_logger.warning(f"Could not get guardrails from registry: {e!s}") return set() async def check_team_alias_exists(self, team_alias: str) -> bool: @@ -100,7 +100,7 @@ class PolicyValidator: ) return team is not None except Exception as e: - verbose_proxy_logger.warning(f"Could not check team alias '{team_alias}': {str(e)}") + verbose_proxy_logger.warning(f"Could not check team alias '{team_alias}': {e!s}") return True # Assume valid on error async def check_key_alias_exists(self, key_alias: str) -> bool: @@ -122,7 +122,7 @@ class PolicyValidator: ) return key is not None except Exception as e: - verbose_proxy_logger.warning(f"Could not check key alias '{key_alias}': {str(e)}") + verbose_proxy_logger.warning(f"Could not check key alias '{key_alias}': {e!s}") return True # Assume valid on error def check_model_exists(self, model: str) -> bool: @@ -151,7 +151,7 @@ class PolicyValidator: return False except Exception as e: - verbose_proxy_logger.warning(f"Could not check model '{model}': {str(e)}") + verbose_proxy_logger.warning(f"Could not check model '{model}': {e!s}") return True # Assume valid on error @staticmethod @@ -222,10 +222,10 @@ class PolicyValidator: def _validate_inheritance_chain( self, policy_name: str, - policies: Dict[str, Policy], - visited: Optional[Set[str]] = None, + policies: dict[str, Policy], + visited: set[str] | None = None, max_depth: int = 100, - ) -> List[PolicyValidationError]: + ) -> list[PolicyValidationError]: """ Validate the inheritance chain for a policy. @@ -243,7 +243,7 @@ class PolicyValidator: Returns: List of validation errors """ - errors: List[PolicyValidationError] = [] + errors: list[PolicyValidationError] = [] # Prevent infinite recursion if max_depth <= 0: @@ -295,7 +295,7 @@ class PolicyValidator: async def validate_policies( self, - policies: Dict[str, Policy], + policies: dict[str, Policy], validate_db: bool = True, ) -> PolicyValidationResponse: """ @@ -308,8 +308,8 @@ class PolicyValidator: Returns: PolicyValidationResponse with errors and warnings """ - errors: List[PolicyValidationError] = [] - warnings: List[PolicyValidationError] = [] + errors: list[PolicyValidationError] = [] + warnings: list[PolicyValidationError] = [] # Get available guardrails available_guardrails = self.get_available_guardrails() @@ -363,10 +363,10 @@ class PolicyValidator: def _validate_pipeline( policy_name: str, policy: Policy, - available_guardrails: Set[str], - ) -> List[PolicyValidationError]: + available_guardrails: set[str], + ) -> list[PolicyValidationError]: """Validate a policy's pipeline configuration.""" - errors: List[PolicyValidationError] = [] + errors: list[PolicyValidationError] = [] pipeline = policy.pipeline if pipeline is None: return errors @@ -404,7 +404,7 @@ class PolicyValidator: async def validate_policy_config( self, - policy_config: Dict[str, Any], + policy_config: dict[str, Any], validate_db: bool = True, ) -> PolicyValidationResponse: """ @@ -422,8 +422,8 @@ class PolicyValidator: from litellm.proxy.policy_engine.policy_registry import PolicyRegistry # First, try to parse the policies - errors: List[PolicyValidationError] = [] - policies: Dict[str, Policy] = {} + errors: list[PolicyValidationError] = [] + policies: dict[str, Policy] = {} temp_registry = PolicyRegistry() @@ -436,7 +436,7 @@ class PolicyValidator: PolicyValidationError( policy_name=policy_name, error_type=PolicyValidationErrorType.INVALID_SYNTAX, - message=f"Failed to parse policy: {str(e)}", + message=f"Failed to parse policy: {e!s}", ) ) diff --git a/litellm/proxy/prompts/init_prompts.py b/litellm/proxy/prompts/init_prompts.py index a39f06b1242..67e961b2da8 100644 --- a/litellm/proxy/prompts/init_prompts.py +++ b/litellm/proxy/prompts/init_prompts.py @@ -2,20 +2,18 @@ Similar to init_guardrails.py, but for prompts. """ -from typing import Dict, List, Optional - from litellm._logging import verbose_proxy_logger def init_prompts( - all_prompts: List[Dict], - config_file_path: Optional[str] = None, + all_prompts: list[dict], + config_file_path: str | None = None, ): from litellm.types.prompts.init_prompts import PromptSpec from .prompt_registry import IN_MEMORY_PROMPT_REGISTRY - prompt_list: List[PromptSpec] = [] + prompt_list: list[PromptSpec] = [] for prompt in all_prompts: initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt( diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 91843e28283..c0c8ef2de54 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -4,7 +4,7 @@ CRUD ENDPOINTS FOR PROMPTS import tempfile from pathlib import Path -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast from fastapi import ( APIRouter, @@ -100,7 +100,7 @@ def get_version_number(prompt_id: str) -> int: return 1 -def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None) -> str: +def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) -> str: """ Construct a versioned prompt ID from a base prompt_id and version number. @@ -127,7 +127,7 @@ def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None) return f"{base_id}.v{version}" -def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any]) -> str: +def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: dict[str, Any]) -> str: """ Find the latest version of a prompt from available prompt IDs. @@ -152,7 +152,7 @@ def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any]) # Find all versions of this prompt matching_versions = [] - for stored_prompt_id in all_prompt_ids.keys(): + for stored_prompt_id in all_prompt_ids: if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id: version_num = get_version_number(prompt_id=stored_prompt_id) matching_versions.append((version_num, stored_prompt_id)) @@ -166,7 +166,7 @@ def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any]) return prompt_id -def get_latest_prompt_versions(prompts: List[PromptSpec]) -> List[PromptSpec]: +def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]: """ Filter a list of prompts to return only the latest version of each unique prompt. @@ -176,7 +176,7 @@ def get_latest_prompt_versions(prompts: List[PromptSpec]) -> List[PromptSpec]: Returns: List of PromptSpec objects with only the latest version of each prompt """ - latest_prompts: Dict[str, PromptSpec] = {} + latest_prompts: dict[str, PromptSpec] = {} for prompt in prompts: base_id = get_base_prompt_id(prompt_id=prompt.prompt_id) @@ -268,12 +268,12 @@ def create_versioned_prompt_spec(db_prompt) -> PromptSpec: class Prompt(BaseModel): prompt_id: str litellm_params: PromptLiteLLMParams - prompt_info: Optional[PromptInfo] = None + prompt_info: PromptInfo | None = None class PatchPromptRequest(BaseModel): - litellm_params: Optional[PromptLiteLLMParams] = None - prompt_info: Optional[PromptInfo] = None + litellm_params: PromptLiteLLMParams | None = None + prompt_info: PromptInfo | None = None @router.get( @@ -283,7 +283,7 @@ class PatchPromptRequest(BaseModel): response_model=ListPromptsResponse, ) async def list_prompts( - environment: Optional[str] = None, + environment: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -323,7 +323,7 @@ async def list_prompts( # check key metadata for prompts key_metadata = user_api_key_dict.metadata if key_metadata is not None: - prompts = cast(Optional[List[str]], key_metadata.get("prompts", None)) + prompts = cast(list[str] | None, key_metadata.get("prompts", None)) if prompts is not None: all_prompts = [ IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id] @@ -382,7 +382,7 @@ async def list_prompts( ) async def get_prompt_versions( prompt_id: str, - environment: Optional[str] = None, + environment: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -433,7 +433,7 @@ async def get_prompt_versions( # Query DB for versions versioned_prompts = [] if prisma_client is not None: - where_clause: Dict[str, Any] = {"prompt_id": base_prompt_id} + where_clause: dict[str, Any] = {"prompt_id": base_prompt_id} if environment: where_clause["environment"] = environment db_prompts = await PromptRepository(prisma_client).table.find_many( @@ -485,7 +485,7 @@ async def get_prompt_versions( return ListPromptsResponse(prompts=versioned_prompts) -def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Optional[PromptTemplateBase]: +def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> PromptTemplateBase | None: """Resolve the raw prompt template from dotprompt content or the in-memory registry.""" from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY @@ -540,7 +540,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Option ) async def get_prompt_info( prompt_id: str, - environment: Optional[str] = None, + environment: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -576,9 +576,9 @@ async def get_prompt_info( from litellm.proxy.proxy_server import prisma_client ## CHECK IF USER HAS ACCESS TO PROMPT - prompts: Optional[List[str]] = None + prompts: list[str] | None = None if user_api_key_dict.metadata is not None: - prompts = cast(Optional[List[str]], user_api_key_dict.metadata.get("prompts", None)) + prompts = cast(list[str] | None, user_api_key_dict.metadata.get("prompts", None)) if prompts is not None and prompt_id not in prompts: raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found") if user_api_key_dict.user_role is not None and ( @@ -595,7 +595,7 @@ async def get_prompt_info( base_prompt_id = get_base_prompt_id(prompt_id=prompt_id) # Query all environments this prompt exists in (lightweight: distinct on environment) - all_environments: List[str] = [] + all_environments: list[str] = [] if prisma_client is not None: all_prompt_rows = await PromptRepository(prisma_client).table.find_many( where={"prompt_id": base_prompt_id}, @@ -609,7 +609,7 @@ async def get_prompt_info( prompt_spec = None requested_version = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None if environment and prisma_client is not None: - where_clause: Dict[str, Any] = { + where_clause: dict[str, Any] = { "prompt_id": base_prompt_id, "environment": environment, } @@ -882,7 +882,7 @@ async def update_prompt( ) async def delete_prompt( prompt_id: str, - environment: Optional[str] = None, + environment: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -943,7 +943,7 @@ async def delete_prompt( base_prompt_id = get_base_prompt_id(prompt_id=prompt_id) # Build delete filter; scope to environment if provided - delete_where: Dict[str, Any] = {"prompt_id": base_prompt_id} + delete_where: dict[str, Any] = {"prompt_id": base_prompt_id} if environment: delete_where["environment"] = environment @@ -994,7 +994,7 @@ def _reload_prompt_in_registry(registry: Any, versioned_id: str, updated_prompt_ async def patch_prompt( prompt_id: str, request: PatchPromptRequest, - environment: Optional[str] = None, + environment: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -1040,7 +1040,7 @@ async def patch_prompt( requested_version = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None # Build query to find the exact row by composite unique key - find_where: Dict[str, Any] = { + find_where: dict[str, Any] = { "prompt_id": base_prompt_id, "environment": env, } @@ -1091,7 +1091,7 @@ async def patch_prompt( raise HTTPException(status_code=400, detail="litellm_params cannot be None") # Build update data dict - update_data: Dict[str, Any] = { + update_data: dict[str, Any] = { "litellm_params": updated_litellm_params.model_dump_json(), "prompt_info": updated_prompt_info.model_dump_json(), } @@ -1264,7 +1264,7 @@ async def test_prompt( async def convert_prompt_file_to_json( file: UploadFile = File(...), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Convert a .prompt file to JSON format. @@ -1304,7 +1304,7 @@ async def convert_prompt_file_to_json( } except Exception as e: - raise HTTPException(status_code=500, detail=f"Error converting prompt file: {str(e)}") + raise HTTPException(status_code=500, detail=f"Error converting prompt file: {e!s}") finally: # Clean up temp file diff --git a/litellm/proxy/prompts/prompt_registry.py b/litellm/proxy/prompts/prompt_registry.py index 8ec1225c5f2..9e00d4e7c63 100644 --- a/litellm/proxy/prompts/prompt_registry.py +++ b/litellm/proxy/prompts/prompt_registry.py @@ -2,7 +2,6 @@ import importlib import os from collections.abc import Callable from pathlib import Path -from typing import Dict, Optional from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_prompt_management import CustomPromptManagement @@ -25,7 +24,7 @@ def get_prompt_initializer_from_integrations(): Returns: Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions """ - discovered_initializers: Dict[str, Callable] = {} + discovered_initializers: dict[str, Callable] = {} try: # Get the path to the prompt_integrations directory @@ -91,12 +90,12 @@ class InMemoryPromptRegistry: """ def __init__(self): - self.IN_MEMORY_PROMPTS: Dict[str, PromptSpec] = {} + self.IN_MEMORY_PROMPTS: dict[str, PromptSpec] = {} """ Prompt id to Prompt object mapping """ - self.prompt_id_to_custom_prompt: Dict[str, Optional[CustomPromptManagement]] = {} + self.prompt_id_to_custom_prompt: dict[str, CustomPromptManagement | None] = {} """ Guardrail id to CustomGuardrail object mapping """ @@ -104,8 +103,8 @@ class InMemoryPromptRegistry: def initialize_prompt( self, prompt: PromptSpec, - config_file_path: Optional[str] = None, - ) -> Optional[PromptSpec]: + config_file_path: str | None = None, + ) -> PromptSpec | None: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -118,7 +117,7 @@ class InMemoryPromptRegistry: verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS") return self.IN_MEMORY_PROMPTS[prompt_id] - custom_prompt_callback: Optional[CustomPromptManagement] = None + custom_prompt_callback: CustomPromptManagement | None = None litellm_params_data = prompt.litellm_params verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) @@ -155,13 +154,13 @@ class InMemoryPromptRegistry: return parsed_prompt - def get_prompt_by_id(self, prompt_id: str) -> Optional[PromptSpec]: + def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None: """ Get a prompt by its ID from memory """ return self.IN_MEMORY_PROMPTS.get(prompt_id) - def get_prompt_callback_by_id(self, prompt_id: str) -> Optional[CustomPromptManagement]: + def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None: """ Get a prompt callback by its ID from memory """ diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 041e088ab9c..5842c49953d 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -8,7 +8,7 @@ import sys import urllib.parse as urlparse from collections.abc import Iterable from pathlib import Path -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any import click import httpx @@ -59,11 +59,11 @@ class LiteLLMDatabaseConnectionPool(Enum): def _build_db_connection_url_params( connection_limit: int, - pool_timeout: Optional[Union[int, float]], - connect_timeout: Optional[Union[int, float]] = None, - socket_timeout: Optional[Union[int, float]] = None, + pool_timeout: float | None, + connect_timeout: float | None = None, + socket_timeout: float | None = None, disable_prepared_statements: bool = False, - extra_params: Optional[dict] = None, + extra_params: dict | None = None, ) -> dict: """Build the Prisma DATABASE_URL query params controlling connection pool behavior. @@ -92,7 +92,7 @@ def _build_db_connection_url_params( return params -def append_query_params(url: Optional[str], params: dict) -> str: +def append_query_params(url: str | None, params: dict) -> str: from litellm._logging import verbose_proxy_logger verbose_proxy_logger.debug(f"url: {url}") @@ -127,7 +127,7 @@ class ProxyInitializationHelpers: host: str, port: int, model: str, - test: Union[bool, str], + test: bool | str, ): request_model = model or "gpt-3.5-turbo" click.echo(f"\nLiteLLM: Making a test ChatCompletions request to your proxy. Model={request_model}") @@ -176,9 +176,9 @@ class ProxyInitializationHelpers: def _get_default_unvicorn_init_args( host: str, port: int, - log_config: Optional[str] = None, - keepalive_timeout: Optional[int] = None, - timeout_worker_healthcheck: Optional[int] = None, + log_config: str | None = None, + keepalive_timeout: int | None = None, + timeout_worker_healthcheck: int | None = None, ) -> dict: """ Get the arguments for `uvicorn` worker @@ -217,7 +217,7 @@ class ProxyInitializationHelpers: @staticmethod def _apply_uvicorn_max_requests_jitter( uvicorn_args: dict, - max_requests_before_restart: Optional[int], + max_requests_before_restart: int | None, jitter: int, ) -> None: """ @@ -243,7 +243,7 @@ class ProxyInitializationHelpers: ) @staticmethod - def _get_reload_options(config_path: Optional[str]) -> dict: + def _get_reload_options(config_path: str | None) -> dict: """Build uvicorn reload kwargs so --reload also reacts to .env and YAML edits.""" cwd = os.path.abspath(os.getcwd()) reload_dirs = [cwd] @@ -264,7 +264,7 @@ class ProxyInitializationHelpers: } @staticmethod - def _patch_statreload_extra_paths(paths: Iterable[Optional[str]]) -> bool: + def _patch_statreload_extra_paths(paths: Iterable[str | None]) -> bool: """Make uvicorn's StatReload reloader notice non-Python dev files (the --config YAML and .env). @@ -305,7 +305,7 @@ class ProxyInitializationHelpers: return True @staticmethod - def _configure_dev_reload(uvicorn_args: dict, config_path: Optional[str]) -> None: + def _configure_dev_reload(uvicorn_args: dict, config_path: str | None) -> None: """Wire up --reload (dev only): watch *.py, the --config YAML, and .env, and signal reloaded workers to re-read .env with override so edits to existing keys actually take effect rather than staying masked by the @@ -329,7 +329,7 @@ class ProxyInitializationHelpers: port: int, ssl_certfile_path: str, ssl_keyfile_path: str, - ciphers: Optional[str] = None, + ciphers: str | None = None, ): """ Initialize litellm with `hypercorn` @@ -360,11 +360,11 @@ class ProxyInitializationHelpers: host: str, port: int, num_workers: int, - ssl_certfile_path: Optional[str], - ssl_keyfile_path: Optional[str], - max_requests_before_restart: Optional[int], - ciphers: Optional[str], - granian_runtime_threads: Optional[int] = None, + ssl_certfile_path: str | None, + ssl_keyfile_path: str | None, + max_requests_before_restart: int | None, + ciphers: str | None, + granian_runtime_threads: int | None = None, ) -> None: """ Run the proxy with Granian (Rust-backed ASGI server, HTTP/1 + HTTP/2). @@ -413,8 +413,8 @@ class ProxyInitializationHelpers: num_workers: int, ssl_certfile_path: str, ssl_keyfile_path: str, - max_requests_before_restart: Optional[int] = None, - max_requests_before_restart_jitter: Optional[int] = None, + max_requests_before_restart: int | None = None, + max_requests_before_restart_jitter: int | None = None, ): """ Run litellm with `gunicorn` @@ -545,7 +545,7 @@ class ProxyInitializationHelpers: @staticmethod def _maybe_setup_prometheus_multiproc_dir( num_workers: int, - litellm_settings: Optional[dict], + litellm_settings: dict | None, ) -> None: """ Auto-create PROMETHEUS_MULTIPROC_DIR when running with multiple workers @@ -899,8 +899,8 @@ def run_server( keepalive_timeout, timeout_worker_healthcheck, max_requests_before_restart, - max_requests_before_restart_jitter: Optional[int], - limit_concurrency: Optional[int], + max_requests_before_restart_jitter: int | None, + limit_concurrency: int | None, enforce_prisma_migration_check: bool, use_v2_migration_resolver: bool, reload: bool, @@ -1003,11 +1003,11 @@ def run_server( db_connection_pool_limit = 100 # Starts optional due to config fallback checks; guaranteed non-None before use. - db_connection_timeout: Optional[Union[int, float]] = 60 - db_connect_timeout: Optional[Union[int, float]] = None - db_socket_timeout: Optional[Union[int, float]] = None + db_connection_timeout: int | float | None = 60 + db_connect_timeout: int | float | None = None + db_socket_timeout: int | float | None = None db_disable_prepared_statements: bool = False - db_extra_connection_params: Optional[dict] = None + db_extra_connection_params: dict | None = None general_settings = {} ### GET DB TOKEN FOR IAM AUTH ### diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ec35f362f01..b310b8e1dbf 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -17,15 +17,12 @@ import traceback import warnings from collections.abc import AsyncGenerator, Callable, Mapping from datetime import datetime, timedelta, timezone +from types import UnionType from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Set, - Tuple, TypedDict, Union, cast, @@ -137,7 +134,7 @@ else: Span = Any OpenTelemetry = Any -REALTIME_REQUEST_SCOPE_TEMPLATE: Dict[str, Any] = { +REALTIME_REQUEST_SCOPE_TEMPLATE: dict[str, Any] = { "type": "http", "method": "POST", "path": "/v1/realtime", @@ -191,7 +188,7 @@ def generate_feedback_box(): print() # noqa: T201 print("\033[1;37m" + "#" + "-" * box_width + "#\033[0m") # noqa: T201 print("\033[1;37m" + "#" + " " * box_width + "#\033[0m") # noqa: T201 - print("\033[1;37m" + "# {:^59} #\033[0m".format(message)) # noqa: T201 + print("\033[1;37m" + f"# {message:^59} #\033[0m") # noqa: T201 print( # noqa: T201 "\033[1;37m" + "# {:^59} #\033[0m".format("https://github.com/BerriAI/litellm/issues/new") ) @@ -484,8 +481,8 @@ try: shutdown_billing_metrics_recorder as _shutdown_billing_metrics_recorder, ) - build_billing_metrics_recorder: Optional[Callable[..., Optional[BillingRecorder]]] = _build_billing_metrics_recorder - shutdown_billing_metrics_recorder: Optional[Callable[[], None]] = _shutdown_billing_metrics_recorder + build_billing_metrics_recorder: Callable[..., BillingRecorder | None] | None = _build_billing_metrics_recorder + shutdown_billing_metrics_recorder: Callable[[], None] | None = _shutdown_billing_metrics_recorder except ImportError: build_billing_metrics_recorder = None shutdown_billing_metrics_recorder = None @@ -658,7 +655,7 @@ from fastapi.staticfiles import StaticFiles enterprise_router = APIRouter() try: # when using litellm cli - import litellm.proxy.enterprise as enterprise + from litellm.proxy import enterprise except Exception: # when using litellm docker image try: @@ -673,7 +670,7 @@ try: from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig enterprise_router = _enterprise_router - enterprise_proxy_config: Optional[EnterpriseProxyConfig] = EnterpriseProxyConfig() + enterprise_proxy_config: EnterpriseProxyConfig | None = EnterpriseProxyConfig() except ImportError: enterprise_proxy_config = None ################### @@ -682,7 +679,7 @@ server_root_path = get_server_root_path() _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() premium_user_data: Optional["EnterpriseLicenseData"] = _license_check.airgapped_license_data -global_max_parallel_request_retries_env: Optional[str] = os.getenv("LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES") +global_max_parallel_request_retries_env: str | None = os.getenv("LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES") proxy_state = ProxyState() SENSITIVE_DATA_MASKER = SensitiveDataMasker() @@ -727,7 +724,7 @@ if global_max_parallel_request_retries_env is None: else: global_max_parallel_request_retries = int(global_max_parallel_request_retries_env) -global_max_parallel_request_retry_timeout_env: Optional[str] = os.getenv( +global_max_parallel_request_retry_timeout_env: str | None = os.getenv( "LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRY_TIMEOUT" ) if global_max_parallel_request_retry_timeout_env is None: @@ -838,7 +835,7 @@ async def _initialize_shared_aiohttp_session(): _build_aiohttp_keepalive_socket_factory, ) - connector_kwargs: Dict[str, Any] = { + connector_kwargs: dict[str, Any] = { "keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT, "ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE, } @@ -912,17 +909,15 @@ async def proxy_startup_event(app: FastAPI): raise ## CHECK PREMIUM USER - verbose_proxy_logger.debug( - "litellm.proxy.proxy_server.py::startup() - CHECKING PREMIUM USER - {}".format(premium_user) - ) + verbose_proxy_logger.debug(f"litellm.proxy.proxy_server.py::startup() - CHECKING PREMIUM USER - {premium_user}") if premium_user is False: premium_user = _license_check.is_premium() ## CHECK MASTER KEY IN ENVIRONMENT ## master_key = get_secret_str("LITELLM_MASTER_KEY") ### LOAD CONFIG ### - worker_config: Optional[Union[str, dict]] = get_secret("WORKER_CONFIG") # type: ignore - env_config_yaml: Optional[str] = get_secret_str("CONFIG_FILE_PATH") + worker_config: str | dict | None = get_secret("WORKER_CONFIG") # type: ignore + env_config_yaml: str | None = get_secret_str("CONFIG_FILE_PATH") verbose_proxy_logger.debug("worker_config: %s", _redact_worker_config_for_logging(worker_config)) # check if it's a valid file path if env_config_yaml is not None: @@ -934,16 +929,14 @@ async def proxy_startup_event(app: FastAPI): ) = await proxy_config.load_config(router=llm_router, config_file_path=env_config_yaml) elif worker_config is not None: if ( - isinstance(worker_config, str) - and os.path.isfile(worker_config) - and proxy_config.is_yaml(config_file_path=worker_config) - ): ( - llm_router, - llm_model_list, - general_settings, - ) = await proxy_config.load_config(router=llm_router, config_file_path=worker_config) - elif os.environ.get("LITELLM_CONFIG_BUCKET_NAME") is not None and isinstance(worker_config, str): + isinstance(worker_config, str) + and os.path.isfile(worker_config) + and proxy_config.is_yaml(config_file_path=worker_config) + ) + or os.environ.get("LITELLM_CONFIG_BUCKET_NAME") is not None + and isinstance(worker_config, str) + ): ( llm_router, llm_model_list, @@ -959,7 +952,7 @@ async def proxy_startup_event(app: FastAPI): # check if DATABASE_URL in environment - load from there if prisma_client is None: - _db_url: Optional[str] = get_secret("DATABASE_URL", None) # type: ignore + _db_url: str | None = get_secret("DATABASE_URL", None) # type: ignore prisma_client = await ProxyStartupEvent._setup_prisma_client( database_url=_db_url, proxy_logging_obj=proxy_logging_obj, @@ -1180,7 +1173,7 @@ _OPENAPI_HTTP_METHODS = { # the UI. Kept here at module scope to match the analogous descriptor # `is_secret` flags in litellm.proxy.config_resolvers and the # `_CACHE_SENSITIVE_FIELDS` constant in the cache endpoint file. -_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"} +_ALERTING_SENSITIVE_VARS: set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"} def _strip_operation_id_method_suffix(operation_id: str) -> str: @@ -1191,11 +1184,11 @@ def _strip_operation_id_method_suffix(operation_id: str) -> str: def ensure_unique_openapi_operation_ids( - openapi_schema: Dict[str, Any], - reserved_operation_ids: Optional[Set[str]] = None, -) -> Dict[str, Any]: + openapi_schema: dict[str, Any], + reserved_operation_ids: set[str] | None = None, +) -> dict[str, Any]: operation_entries = [] - operation_id_counts: Dict[str, int] = {} + operation_id_counts: dict[str, int] = {} for path_item in openapi_schema.get("paths", {}).values(): if not isinstance(path_item, dict): continue @@ -1209,7 +1202,7 @@ def ensure_unique_openapi_operation_ids( operation_id_counts[operation_id] = operation_id_counts.get(operation_id, 0) + 1 used_operation_ids = set(reserved_operation_ids or set()) - seen_operation_ids: Set[str] = set() + seen_operation_ids: set[str] = set() for method, operation, operation_id in operation_entries: should_rewrite = ( operation_id_counts[operation_id] > 1 @@ -1412,7 +1405,7 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) -def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Optional[Exception] = None) -> None: +def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Exception | None = None) -> None: parent_otel_span = getattr(request.state, "parent_otel_span", None) if parent_otel_span is None: return @@ -1499,8 +1492,8 @@ router = APIRouter() def _get_cors_config( - cors_origins_env: Optional[str] = None, - cors_credentials_env: Optional[str] = None, + cors_origins_env: str | None = None, + cors_credentials_env: str | None = None, ): """ Compute CORS allowed origins and credentials flag from environment variables. @@ -1970,7 +1963,7 @@ def mount_swagger_ui(): mount_swagger_ui() docs_url = _get_docs_url() -root_redirect_url: Optional[str] = os.getenv("ROOT_REDIRECT_URL") +root_redirect_url: str | None = os.getenv("ROOT_REDIRECT_URL") if docs_url != "/" and root_redirect_url is not None: @app.get("/", include_in_schema=False) @@ -1978,8 +1971,6 @@ if docs_url != "/" and root_redirect_url is not None: return RedirectResponse(url=root_redirect_url) # type: ignore[arg-type] -from typing import Dict - user_api_base = None user_model = None user_debug = False @@ -1989,19 +1980,19 @@ user_temperature = None user_telemetry = True user_config = None user_headers = None -user_config_file_path: Optional[str] = None +user_config_file_path: str | None = None local_logging = True # writes logs to a local api_log.json file for debugging experimental = False #### GLOBAL VARIABLES #### -llm_router: Optional[Router] = None -llm_model_list: Optional[list] = None +llm_router: Router | None = None +llm_model_list: list | None = None general_settings: dict = {} -config_passthrough_endpoints: Optional[List[Dict[str, Any]]] = None +config_passthrough_endpoints: list[dict[str, Any]] | None = None log_file = "api_log.json" worker_config = None -master_key: Optional[str] = None +master_key: str | None = None otel_logging = False -prisma_client: Optional[PrismaClient] = None +prisma_client: PrismaClient | None = None shared_aiohttp_session: Optional["ClientSession"] = None # Global shared session for connection reuse user_api_key_cache: UserApiKeyCache = UserApiKeyCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value @@ -2010,9 +2001,9 @@ spend_counter_cache = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_ cli_sso_session_cache = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS) model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=user_api_key_cache) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) -redis_usage_cache: Optional[RedisCache] = None # redis cache used for tracking spend, tpm/rpm limits -polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False -native_background_mode: List[str] = [] # Models that should use native provider background mode instead of polling +redis_usage_cache: RedisCache | None = None # redis cache used for tracking spend, tpm/rpm limits +polling_via_cache_enabled: Literal["all"] | list[str] | bool = False +native_background_mode: list[str] = [] # Models that should use native provider background mode instead of polling polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None user_custom_key_generate = None @@ -2029,13 +2020,13 @@ use_queue = False health_check_interval = None health_check_concurrency = None health_check_details = None -health_check_results: Dict[str, Union[int, List[Dict[str, Any]]]] = {} +health_check_results: dict[str, int | list[dict[str, Any]]] = {} background_health_check_loop_active = False background_health_check_cycle_seq = 0 -queue: List = [] +queue: list = [] litellm_proxy_budget_name = LITELLM_PROXY_BUDGET_NAME litellm_proxy_admin_name = LITELLM_PROXY_ADMIN_NAME -ui_access_mode: Union[Literal["admin", "all"], Dict] = "all" +ui_access_mode: Literal["admin", "all"] | dict = "all" proxy_budget_rescheduler_min_time = PROXY_BUDGET_RESCHEDULER_MIN_TIME proxy_budget_rescheduler_max_time = PROXY_BUDGET_RESCHEDULER_MAX_TIME proxy_batch_polling_interval = PROXY_BATCH_POLLING_INTERVAL @@ -2044,9 +2035,9 @@ proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS litellm_master_key_hash = None disable_spend_logs = False jwt_handler = JWTHandler() -prompt_injection_detection_obj: Optional[_OPTIONAL_PromptInjectionDetection] = None +prompt_injection_detection_obj: _OPTIONAL_PromptInjectionDetection | None = None store_model_in_db: bool = False -open_telemetry_logger: Optional[OpenTelemetry] = None +open_telemetry_logger: OpenTelemetry | None = None ### INITIALIZE GLOBAL LOGGING OBJECT ### proxy_logging_obj: ProxyLogging = ProxyLogging(user_api_key_cache=user_api_key_cache, premium_user=premium_user) ### REDIS QUEUE ### @@ -2063,7 +2054,7 @@ last_anthropic_beta_headers_reload = None ### DB WRITER ### -db_writer_client: Optional[AsyncHTTPHandler] = None +db_writer_client: AsyncHTTPHandler | None = None ### logger ### @@ -2072,7 +2063,7 @@ def _resolve_typed_dict_type(typ): from typing_extensions import _TypedDictMeta # type: ignore origin = get_origin(typ) - if origin is Union: # Check if it's a Union (like Optional) + if origin is Union or origin is UnionType: # Check if it's a Union (like Optional) for arg in get_args(typ): if isinstance(arg, _TypedDictMeta): return arg @@ -2081,13 +2072,13 @@ def _resolve_typed_dict_type(typ): return None -def _resolve_pydantic_type(typ) -> List: +def _resolve_pydantic_type(typ) -> list: """Resolve the actual TypedDict class from a potentially wrapped type.""" origin = get_origin(typ) typs = [] - if origin is Union: # Check if it's a Union (like Optional) + if origin is Union or origin is UnionType: # Check if it's a Union (like Optional) for arg in get_args(typ): - if arg is not None and not isinstance(arg, type(None)) and "NoneType" not in str(arg): + if arg is not None and "NoneType" not in str(arg): typs.append(arg) elif isinstance(typ, type) and isinstance(typ, BaseModel): return [typ] @@ -2367,14 +2358,14 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float) async def increment_spend_counters( - token: Optional[str], - team_id: Optional[str], - user_id: Optional[str], - response_cost: Optional[float], - org_id: Optional[str] = None, - budget_reservation: Optional[dict] = None, - end_user_id: Optional[str] = None, - tags: Optional[List[str]] = None, + token: str | None, + team_id: str | None, + user_id: str | None, + response_cost: float | None, + org_id: str | None = None, + budget_reservation: dict | None = None, + end_user_id: str | None = None, + tags: list[str] | None = None, ): """ Atomically increment spend counters for budget enforcement. @@ -2528,9 +2519,9 @@ async def increment_spend_counters( async def _reconcile_budget_reservation_for_counter_update( - budget_reservation: Optional[dict], - response_cost: Optional[float], -) -> Set[str]: + budget_reservation: dict | None, + response_cost: float | None, +) -> set[str]: if budget_reservation is None: return set() @@ -2563,10 +2554,10 @@ async def _reconcile_budget_reservation_for_counter_update( async def _increment_end_user_and_tag_spend_counters( - end_user_id: Optional[str], - tags: Optional[List[str]], + end_user_id: str | None, + tags: list[str] | None, response_cost: float, - reserved_counter_keys: Set[str], + reserved_counter_keys: set[str], ) -> None: if end_user_id is not None: await _init_and_increment_unreserved_spend_counter( @@ -2579,7 +2570,7 @@ async def _increment_end_user_and_tag_spend_counters( if tags is None: return - seen_tags: Set[str] = set() + seen_tags: set[str] = set() for tag_name in tags: if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags: continue @@ -2593,9 +2584,9 @@ async def _increment_end_user_and_tag_spend_counters( async def _increment_org_spend_counter( - org_id: Optional[str], + org_id: str | None, response_cost: float, - reserved_counter_keys: Set[str], + reserved_counter_keys: set[str], ) -> None: if org_id is None: return @@ -2610,9 +2601,9 @@ async def _increment_org_spend_counter( async def _init_and_increment_unreserved_spend_counter( counter_key: str, - source_cache_key: Union[str, List[str]], + source_cache_key: str | list[str], increment: float, - reserved_counter_keys: Set[str], + reserved_counter_keys: set[str], ) -> None: if counter_key in reserved_counter_keys: return @@ -2626,7 +2617,7 @@ async def _init_and_increment_unreserved_spend_counter( async def _init_and_increment_spend_counter( counter_key: str, - source_cache_key: Union[str, List[str]], + source_cache_key: str | list[str], increment: float, ): """ @@ -2656,7 +2647,7 @@ async def _init_and_increment_window_spend_counter( counter_key: str, entity_type: str, entity_id: str, - window_start: Optional[datetime], + window_start: datetime | None, increment: float, ): if window_start is None: @@ -2679,7 +2670,7 @@ async def _init_and_increment_window_spend_counter( async def _ensure_spend_counter_initialized( counter_key: str, - source_cache_key: Union[str, List[str]], + source_cache_key: str | list[str], ): is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) if is_warm is False: @@ -2698,7 +2689,7 @@ async def _ensure_spend_counter_initialized( async def _get_source_cache_base_spend( - source_cache_key: Union[str, List[str]], + source_cache_key: str | list[str], ) -> float: source_cache_keys = [source_cache_key] if isinstance(source_cache_key, str) else source_cache_key for cache_key in source_cache_keys: @@ -2799,13 +2790,13 @@ async def _invalidate_spend_counter(counter_key: str): async def update_cache( - token: Optional[str], - user_id: Optional[str], - end_user_id: Optional[str], - team_id: Optional[str], - response_cost: Optional[float], - parent_otel_span: Optional[Span], # type: ignore - tags: Optional[List[str]] = None, + token: str | None, + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + response_cost: float | None, + parent_otel_span: Span | None, # type: ignore + tags: list[str] | None = None, ): """ Use this to update the cache with new user spend. @@ -2813,7 +2804,7 @@ async def update_cache( Put any alerting logic in here. """ - values_to_update_in_cache: List[Tuple[Any, Any]] = [] + values_to_update_in_cache: list[tuple[Any, Any]] = [] ### UPDATE KEY SPEND ### async def _update_key_cache(token: str, response_cost: float): @@ -2925,7 +2916,7 @@ async def update_cache( if end_user_id is None or response_cost is None: return - _id = "end_user_id:{}".format(end_user_id) + _id = f"end_user_id:{end_user_id}" try: # Fetch the existing cost for the given user cached_end_user = await user_api_key_cache.async_get_cache(key=_id) @@ -2967,13 +2958,13 @@ async def update_cache( if team_id is None or response_cost is None: return - _id = "team_id:{}".format(team_id) + _id = f"team_id:{team_id}" try: cached_team = await user_api_key_cache.async_get_cache(key=_id) if cached_team is None: # do nothing if team not in api key cache return - existing_spend_obj: Optional[LiteLLM_TeamTableCachedObj] = CacheCodec.deserialize( + existing_spend_obj: LiteLLM_TeamTableCachedObj | None = CacheCodec.deserialize( cached_team, LiteLLM_TeamTableCachedObj ) if existing_spend_obj is None: @@ -3104,7 +3095,7 @@ def run_ollama_serve(): """) -def _get_process_rss_mb() -> Optional[float]: +def _get_process_rss_mb() -> float | None: """ Get process RSS memory in MB. On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes. @@ -3134,13 +3125,13 @@ def _is_unexpected_keyword_argument_type_error(exc: BaseException) -> bool: async def _run_direct_health_check_with_instrumentation( model_list: list, - details: Optional[bool], - max_concurrency: Optional[int], + details: bool | None, + max_concurrency: int | None, instrumentation_context: dict, ): """Call ``perform_health_check``, retrying with fewer kwargs on unexpected-kw TypeErrors.""" _hc_filter = health_check_filter_kwargs_from_general_settings(general_settings) - last_type_error: Optional[TypeError] = None + last_type_error: TypeError | None = None for extra_kwargs in ( { "instrumentation_context": instrumentation_context, @@ -3212,7 +3203,7 @@ def _get_endpoint_exception_status(endpoint: dict, exceptions: dict) -> int: def _write_health_state_to_router_cache( healthy_endpoints: list, unhealthy_endpoints: list, - exceptions_by_model_id: Optional[dict] = None, + exceptions_by_model_id: dict | None = None, ) -> None: """ Write deployment health states to the router's health state cache @@ -3498,7 +3489,7 @@ class StreamingCallbackError(Exception): # active), so the runtime gate cannot distinguish a YAML-sourced value # from a DB-sourced value. Scrubbing at the merge boundary closes that # gap without tracking source on every config dict entry. -_DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: Dict[str, Tuple[str, ...]] = { +_DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: dict[str, tuple[str, ...]] = { "litellm_settings": ("post_call_rules",), "general_settings": ( "custom_auth", @@ -3508,7 +3499,7 @@ _DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: Dict[str, Tuple[str, ...]] = { "custom_ui_sso_sign_in_handler", ), } -_DB_OVERLAY_REMOTE_MODULE_LIST_FIELDS: Dict[str, Tuple[str, ...]] = { +_DB_OVERLAY_REMOTE_MODULE_LIST_FIELDS: dict[str, tuple[str, ...]] = { "litellm_settings": ( "callbacks", "success_callback", @@ -3522,7 +3513,7 @@ def _is_remote_module_url(value: Any) -> bool: return isinstance(value, str) and (value.startswith("s3://") or value.startswith("gcs://")) -def _scrub_guardrail_inner(inner: Dict[str, Any]) -> None: +def _scrub_guardrail_inner(inner: dict[str, Any]) -> None: """Strip remote-URL entries from a guardrail's ``callbacks`` list and ``guardrail`` (v2 module-path) field. Mutates in place.""" cbs = inner.get("callbacks") @@ -3642,7 +3633,7 @@ def _scrub_db_overlay_remote_module_loads(section: str, db_value: Any) -> Any: return sanitized -def _normalize_user_url_validation(value: object) -> Optional[bool]: +def _normalize_user_url_validation(value: object) -> bool | None: if value is None: return None if isinstance(value, str): @@ -3832,10 +3823,10 @@ class ProxyConfig: """ def __init__(self) -> None: - self.config: Dict[str, Any] = {} - self._last_semantic_filter_config: Optional[Dict[str, Any]] = None - self._last_hashicorp_vault_config: Optional[Dict[str, Any]] = None - self.worker_registry: List["WorkerRegistryEntry"] = [] + self.config: dict[str, Any] = {} + self._last_semantic_filter_config: dict[str, Any] | None = None + self._last_hashicorp_vault_config: dict[str, Any] | None = None + self.worker_registry: list[WorkerRegistryEntry] = [] def is_yaml(self, config_file_path: str) -> bool: if not os.path.isfile(config_file_path): @@ -3852,9 +3843,9 @@ class ProxyConfig: with open(file_path, "r") as file: return yaml.safe_load(file) or {} except Exception as e: - raise Exception(f"Error loading yaml file {file_path}: {str(e)}") + raise Exception(f"Error loading yaml file {file_path}: {e!s}") - async def _get_config_from_file(self, config_file_path: Optional[str] = None) -> dict: + async def _get_config_from_file(self, config_file_path: str | None = None) -> dict: """ Given a config file path, load the config from the file. Args: @@ -4037,7 +4028,7 @@ class ProxyConfig: config[key] = get_secret(value) return config - def _get_team_config(self, team_id: str, all_teams_config: List[Dict]) -> Dict: + def _get_team_config(self, team_id: str, all_teams_config: list[dict]) -> dict: team_config: dict = {} for team in all_teams_config: if "team_id" not in team: @@ -4160,7 +4151,7 @@ class ProxyConfig: llm_router.cache_responses = True verbose_proxy_logger.debug("Set router.cache_responses=True after initializing cache") - async def get_config(self, config_file_path: Optional[str] = None) -> dict: + async def get_config(self, config_file_path: str | None = None) -> dict: """ Load config file Supports reading from: @@ -4226,13 +4217,11 @@ class ProxyConfig: return copy.deepcopy(self.config) except Exception as e: verbose_proxy_logger.debug( - "ProxyConfig:get_config_state(): Error returning copy of config state. self.config={}\nError: {}".format( - self.config, e - ) + f"ProxyConfig:get_config_state(): Error returning copy of config state. self.config={self.config}\nError: {e}" ) return {} - def load_credential_list(self, config: dict) -> List[CredentialItem]: + def load_credential_list(self, config: dict) -> list[CredentialItem]: """ Load the credential list from the database """ @@ -4242,7 +4231,7 @@ class ProxyConfig: credential_list = [CredentialItem(**cred) for cred in credential_list_dict] return credential_list - def parse_search_tools(self, config: dict) -> Optional[List[SearchToolTypedDict]]: + def parse_search_tools(self, config: dict) -> list[SearchToolTypedDict] | None: """ Parse and validate search tools from config. Loads environment variables and casts to SearchToolTypedDict. @@ -4263,7 +4252,7 @@ class ProxyConfig: if not search_tools_raw: return None - search_tools_parsed: List[SearchToolTypedDict] = [] + search_tools_parsed: list[SearchToolTypedDict] = [] print( # noqa: T201 "\033[32mLiteLLM: Proxy initialized with Search Tools:\033[0m" @@ -4292,14 +4281,14 @@ class ProxyConfig: search_tool_typed: SearchToolTypedDict = SearchToolTypedDict(**search_tool) # type: ignore search_tools_parsed.append(search_tool_typed) except Exception as e: - verbose_proxy_logger.error(f"Error parsing search tool {search_tool_name}: {str(e)}") + verbose_proxy_logger.error(f"Error parsing search tool {search_tool_name}: {e!s}") continue return search_tools_parsed if search_tools_parsed else None # Environment variable keys that must not be overridden via config because # they can alter process execution, library loading, or network routing. - _BLOCKED_ENV_KEYS: Set[str] = { + _BLOCKED_ENV_KEYS: set[str] = { "PATH", "LD_PRELOAD", "LD_LIBRARY_PATH", @@ -4333,7 +4322,7 @@ class ProxyConfig: # ``` ######################################################### if isinstance(value, str) and value.startswith("os.environ/"): - resolved_secret_string: Optional[str] = get_secret_str(secret_name=value) + resolved_secret_string: str | None = get_secret_str(secret_name=value) if resolved_secret_string is not None: os.environ[key] = resolved_secret_string else: @@ -4350,9 +4339,8 @@ class ProxyConfig: if "LITELLM_LICENSE" in environment_variables: _license_check.license_str = os.getenv("LITELLM_LICENSE", None) premium_user = _license_check.is_premium() - return - async def load_config(self, router: Optional[litellm.Router], config_file_path: str): + async def load_config(self, router: litellm.Router | None, config_file_path: str): """ Load config values into proxy global state """ @@ -4969,7 +4957,7 @@ class ProxyConfig: run_ollama_serve() ## ASSISTANT SETTINGS - assistants_config: Optional[AssistantsTypedDict] = None + assistants_config: AssistantsTypedDict | None = None assistant_settings = config.get("assistant_settings", None) if assistant_settings: for k, v in assistant_settings["litellm_params"].items(): @@ -4980,7 +4968,7 @@ class ProxyConfig: assistants_config = AssistantsTypedDict(**assistant_settings) # type: ignore ## SEARCH TOOLS SETTINGS - search_tools: Optional[List[SearchToolTypedDict]] = self.parse_search_tools(config) + search_tools: list[SearchToolTypedDict] | None = self.parse_search_tools(config) ## SANDBOX TOOLS SETTINGS from litellm.sandbox.sandbox_tools import register_sandbox_tools @@ -5042,7 +5030,7 @@ class ProxyConfig: router._update_redis_cache(cache=redis_usage_cache) # Guardrail settings - guardrails_v2: Optional[List[Dict]] = None + guardrails_v2: list[dict] | None = None if config is not None: guardrails_v2 = config.get("guardrails", None) @@ -5061,7 +5049,7 @@ class ProxyConfig: ) ## Prompt settings - prompts: Optional[List[Dict]] = None + prompts: list[dict] | None = None if config is not None: prompts = config.get("prompts", None) if prompts: @@ -5078,7 +5066,7 @@ class ProxyConfig: return router, router.get_model_list(), general_settings - async def _init_non_llm_configs(self, config: dict, config_file_path: Optional[str] = None): + async def _init_non_llm_configs(self, config: dict, config_file_path: str | None = None): """ Initialize non-LLM configs eg. MCP tools, vector stores, etc. """ @@ -5132,7 +5120,7 @@ class ProxyConfig: async def _init_policy_engine( self, - config: Optional[dict], + config: dict | None, prisma_client: Optional["PrismaClient"], llm_router: Optional["Router"], ): @@ -5208,12 +5196,11 @@ class ProxyConfig: ) if _logger is not None: litellm.logging_callback_manager.add_litellm_callback(_logger) - pass def initialize_secret_manager( self, - key_management_system: Optional[str], - config_file_path: Optional[str] = None, + key_management_system: str | None, + config_file_path: str | None = None, ): """ Initialize the relevant secret manager if `key_management_system` is provided @@ -5275,7 +5262,7 @@ class ProxyConfig: Return model info w/ id """ - _id: Optional[str] = getattr(model, "model_id", None) + _id: str | None = getattr(model, "model_id", None) if _id is not None: model.model_info["id"] = _id model.model_info["db_model"] = True @@ -5448,7 +5435,7 @@ class ProxyConfig: async def _update_llm_router( self, - new_models: Optional[Json], + new_models: Json | None, proxy_logging_obj: ProxyLogging, ) -> frozenset[str] | None: global llm_router, llm_model_list, master_key, general_settings @@ -5507,7 +5494,7 @@ class ProxyConfig: self._add_deployment(db_models=models_list) except Exception as e: - verbose_proxy_logger.exception(f"Error adding/deleting model to llm_router: {str(e)}") + verbose_proxy_logger.exception(f"Error adding/deleting model to llm_router: {e!s}") if llm_router is not None: llm_model_list = llm_router.get_model_list() @@ -5532,7 +5519,7 @@ class ProxyConfig: def _add_callback_from_db_to_in_memory_litellm_callbacks( self, callback: str, - event_types: List[Literal["success", "failure"]], + event_types: list[Literal["success", "failure"]], existing_callbacks: list, ) -> None: """ @@ -5587,7 +5574,7 @@ class ProxyConfig: existing_callbacks=litellm.callbacks, ) - def _encrypt_env_variables(self, environment_variables: dict, new_encryption_key: Optional[str] = None) -> dict: + def _encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Encrypts a dictionary of environment variables and returns them. """ @@ -5628,9 +5615,7 @@ class ProxyConfig: decrypted_variables[k] = decrypted_value return decrypted_variables - def _encrypt_env_variables_for_db( - self, environment_variables: dict, new_encryption_key: Optional[str] = None - ) -> dict: + def _encrypt_env_variables_for_db(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Idempotently encrypt environment variables for a DB write. @@ -5652,7 +5637,7 @@ class ProxyConfig: ) @staticmethod - def _parse_router_settings_value(value: Any) -> Optional[dict]: + def _parse_router_settings_value(value: Any) -> dict | None: """ Parse a router_settings value that may be a dict or a JSON/YAML string. @@ -5661,7 +5646,7 @@ class ProxyConfig: if value is None: return None - parsed: Optional[dict] = None + parsed: dict | None = None if isinstance(value, dict): parsed = value elif isinstance(value, str): @@ -5682,9 +5667,9 @@ class ProxyConfig: async def _get_hierarchical_router_settings( self, user_api_key_dict: Optional["UserAPIKeyAuth"], - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, proxy_logging_obj: Optional["ProxyLogging"] = None, - ) -> Optional[dict]: + ) -> dict | None: """ Get router_settings in priority order: Key > Team @@ -5728,8 +5713,8 @@ class ProxyConfig: async def _add_router_settings_from_db_config( self, config_data: dict, - llm_router: Optional[Router], - prisma_client: Optional[PrismaClient], + llm_router: Router | None, + prisma_client: PrismaClient | None, ) -> None: """ Adds router settings from DB config to litellm proxy @@ -5898,7 +5883,7 @@ class ProxyConfig: except ValueError: verbose_proxy_logger.error("Invalid maximum_spend_logs_retention_interval value") - async def _update_general_settings(self, db_general_settings: Optional[Json]): + async def _update_general_settings(self, db_general_settings: Json | None): """ Pull from DB, read general settings value """ @@ -6069,7 +6054,7 @@ class ProxyConfig: self, prisma_client: PrismaClient, config: dict, - store_model_in_db: Optional[bool], + store_model_in_db: bool | None, ): if store_model_in_db is not True: verbose_proxy_logger.info("'store_model_in_db' is not True, skipping db updates") @@ -6107,7 +6092,7 @@ class ProxyConfig: return config - def _should_load_db_object(self, object_type: Union[str, SupportedDBObjectType]) -> bool: + def _should_load_db_object(self, object_type: str | SupportedDBObjectType) -> bool: """ Check if an object type should be loaded from the database based on general_settings.supported_db_objects. @@ -6139,7 +6124,7 @@ class ProxyConfig: # Check if the object type is in the list (supports both str and enum values) return any(str(obj) == object_type_str for obj in supported_db_objects) - async def _get_models_from_db(self, prisma_client: PrismaClient) -> Optional[list]: + async def _get_models_from_db(self, prisma_client: PrismaClient) -> list | None: """ Fetch all model deployments from the DB. @@ -6153,7 +6138,7 @@ class ProxyConfig: return new_models except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy_server.py::add_deployment() - Error getting new models from DB - {}".format(str(e)) + f"litellm.proxy_server.py::add_deployment() - Error getting new models from DB - {e!s}" ) return None @@ -6210,9 +6195,7 @@ class ProxyConfig: await self._init_non_llm_objects_in_db(prisma_client=prisma_client) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:add_deployment - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.py::ProxyConfig:add_deployment - {e!s}") return still_desired_ids @@ -6356,7 +6339,7 @@ class ProxyConfig: self._decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_sso_settings_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_sso_settings_in_db - {e!s}" ) async def _init_hashicorp_vault_config_override(self, prisma_client: PrismaClient): @@ -6514,7 +6497,7 @@ class ProxyConfig: ) except Exception as e: - verbose_proxy_logger.exception(f"Error in _check_and_reload_model_cost_map: {str(e)}") + verbose_proxy_logger.exception(f"Error in _check_and_reload_model_cost_map: {e!s}") async def _check_and_reload_anthropic_beta_headers(self, prisma_client: PrismaClient): """ @@ -6611,7 +6594,7 @@ class ProxyConfig: ) except Exception as e: - verbose_proxy_logger.exception(f"Error in _check_and_reload_anthropic_beta_headers: {str(e)}") + verbose_proxy_logger.exception(f"Error in _check_and_reload_anthropic_beta_headers: {e!s}") def _get_prompt_spec_for_db_prompt(self, db_prompt): """ @@ -6640,9 +6623,7 @@ class ProxyConfig: prompt_spec = self._get_prompt_spec_for_db_prompt(db_prompt=prompt) IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=prompt_spec) except Exception as e: - verbose_proxy_logger.debug( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - {}".format(str(e)) - ) + verbose_proxy_logger.debug(f"litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - {e!s}") async def _init_guardrails_in_db(self, prisma_client: PrismaClient): from litellm.proxy.guardrails.guardrail_registry import ( @@ -6652,7 +6633,7 @@ class ProxyConfig: ) try: - guardrails_in_db: List[Guardrail] = await GuardrailRegistry.get_all_guardrails_from_db( + guardrails_in_db: list[Guardrail] = await GuardrailRegistry.get_all_guardrails_from_db( prisma_client=prisma_client ) verbose_proxy_logger.debug("guardrails from the DB %s", str(guardrails_in_db)) @@ -6669,9 +6650,7 @@ class ProxyConfig: # pod. Config-loaded entries are never touched. IN_MEMORY_GUARDRAIL_HANDLER.reconcile_db_guardrails(db_guardrail_ids=db_guardrail_ids) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - {e!s}") async def _init_policies_in_db(self, prisma_client: PrismaClient): """ @@ -6695,9 +6674,7 @@ class ProxyConfig: verbose_proxy_logger.debug("Successfully synced policies and attachments from DB") except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_policies_in_db - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.py::ProxyConfig:_init_policies_in_db - {e!s}") async def _init_tool_policy_in_db(self, prisma_client: PrismaClient): """ @@ -6712,7 +6689,7 @@ class ProxyConfig: verbose_proxy_logger.debug("Successfully synced tool policy from DB") except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_tool_policy_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_tool_policy_in_db - {e!s}" ) async def _init_vector_stores_in_db(self, prisma_client: PrismaClient): @@ -6731,7 +6708,7 @@ class ProxyConfig: litellm.vector_store_registry.add_vector_store_to_registry(vector_store=vector_store) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - {e!s}" ) async def _init_vector_store_indexes_in_db(self, prisma_client: PrismaClient): @@ -6755,7 +6732,7 @@ class ProxyConfig: litellm.vector_store_index_registry.upsert_vector_store_index(vector_store_index=vector_store_index) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - {e!s}" ) async def _init_mcp_servers_in_db(self): @@ -6780,7 +6757,7 @@ class ProxyConfig: await backfill_null_oauth2_flows(prisma_client) except Exception as e: # noqa: BLE001 verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {e!s}" ) try: @@ -6788,16 +6765,14 @@ class ProxyConfig: await backfill_discovery_stamped_issuers(prisma_client) except Exception as e: # noqa: BLE001 verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format( - str(e) - ) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {e!s}" ) try: await global_mcp_server_manager.reload_servers_from_database() except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - {e!s}" ) async def init_mcp_servers_from_db(self) -> None: @@ -6826,7 +6801,7 @@ class ProxyConfig: await global_mcp_server_manager.reload_servers_from_database() except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {e!s}" ) async def _init_agents_in_db(self, prisma_client: PrismaClient): @@ -6838,9 +6813,7 @@ class ProxyConfig: db_agents = await AGENT_REGISTRY.get_all_agents_from_db(prisma_client=prisma_client) AGENT_REGISTRY.load_agents_from_db_and_config(db_agents=db_agents) except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_agents_in_db - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.py::ProxyConfig:_init_agents_in_db - {e!s}") async def _init_search_tools_in_db(self, prisma_client: PrismaClient): """ @@ -6881,7 +6854,7 @@ class ProxyConfig: except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_search_tools_in_db - {}".format(str(e)) + f"litellm.proxy.proxy_server.py::ProxyConfig:_init_search_tools_in_db - {e!s}" ) @staticmethod @@ -6906,7 +6879,7 @@ class ProxyConfig: await initialize_pass_through_endpoints_in_db() - def decrypt_credentials(self, credential: Union[dict, BaseModel]) -> CredentialItem: + def decrypt_credentials(self, credential: dict | BaseModel) -> CredentialItem: if isinstance(credential, dict): credential_object = CredentialItem(**credential) elif isinstance(credential, BaseModel): @@ -6919,7 +6892,7 @@ class ProxyConfig: credential_object.credential_values = decrypted_credential_values return credential_object - async def delete_credentials(self, db_credentials: List[CredentialItem]): + async def delete_credentials(self, db_credentials: list[CredentialItem]): """ Create all-up list of db credentials + local credentials Compare to the litellm.credential_list @@ -6948,7 +6921,7 @@ class ProxyConfig: CredentialAccessor.upsert_credentials(credentials) # upsert credentials that are in the all-up list except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy_server.py::get_credentials() - Error getting credentials from DB - {}".format(str(e)) + f"litellm.proxy_server.py::get_credentials() - Error getting credentials from DB - {e!s}" ) return [] @@ -7128,14 +7101,14 @@ async def async_assistants_data_generator(response, user_api_key_dict: UserAPIKe try: yield f"data: {c}\n\n" except Exception as e: - yield f"data: {str(e)}\n\n" + yield f"data: {e!s}\n\n" # Streaming is done, yield the [DONE] chunk done_message = "[DONE]" yield f"data: {done_message}\n\n" except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.async_assistants_data_generator(): Exception occured - {}".format(str(e)) + f"litellm.proxy.proxy_server.async_assistants_data_generator(): Exception occured - {e!s}" ) await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, @@ -7297,7 +7270,7 @@ def _restamp_streaming_chunk_model( def _fast_serialize_simple_model_response_stream( chunk: ModelResponseStream, -) -> Optional[bytes]: +) -> bytes | None: """ Serialize the common OpenAI text streaming chunk without the full Pydantic serializer. Fall back for richer chunks so tool calls, logprobs, usage, and @@ -7372,7 +7345,7 @@ def _fast_serialize_simple_model_response_stream( return orjson.dumps(payload) -def _serialize_streaming_chunk(chunk: BaseModel) -> Union[str, bytes]: +def _serialize_streaming_chunk(chunk: BaseModel) -> str | bytes: if isinstance(chunk, ModelResponseStream): serialized_chunk = _fast_serialize_simple_model_response_stream(chunk) if serialized_chunk is not None: @@ -7406,7 +7379,7 @@ async def _apply_streaming_chunk_hooks( user_api_key_dict: UserAPIKeyAuth, request_data: dict, str_so_far: str, -) -> Tuple[Any, str]: +) -> tuple[Any, str]: chunk = await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=chunk, @@ -7421,7 +7394,7 @@ async def _apply_streaming_chunk_hooks( return chunk, str_so_far -def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]: +def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes: if isinstance(chunk, bytes): return b"data: " + chunk + b"\n\n" return f"data: {chunk}\n\n" @@ -7453,7 +7426,7 @@ async def async_data_generator( stream_completed = False client_disconnected = False try: - error_message: Optional[str] = None + error_message: str | None = None requested_model_from_client = _get_client_requested_model_for_streaming(request_data=request_data) ( fallback_was_attempted, @@ -7576,7 +7549,7 @@ async def async_data_generator( try: yield _format_streaming_sse_chunk(chunk=chunk) except Exception as e: - yield f"data: {str(e)}\n\n" + yield f"data: {e!s}\n\n" if pending_fallback_event: yield _format_fallback_metadata_sse_event( @@ -7614,9 +7587,7 @@ async def async_data_generator( client_disconnected = True raise except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {e!s}") await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, @@ -7714,9 +7685,9 @@ class ProxyStartupEvent: @classmethod def _initialize_startup_logging( cls, - llm_router: Optional[Router], + llm_router: Router | None, proxy_logging_obj: ProxyLogging, - redis_usage_cache: Optional[RedisCache], + redis_usage_cache: RedisCache | None, ): """Initialize logging and alerting on startup""" ## COST TRACKING ## @@ -7758,7 +7729,7 @@ class ProxyStartupEvent: @staticmethod def _validate_redis_transaction_buffer_config( general_settings: dict, - redis_usage_cache: Optional[RedisCache], + redis_usage_cache: RedisCache | None, ): """ Validates that when use_redis_transaction_buffer is enabled, @@ -7766,9 +7737,7 @@ class ProxyStartupEvent: """ from litellm.secret_managers.main import str_to_bool - _use_redis_transaction_buffer: Optional[Union[bool, str]] = general_settings.get( - "use_redis_transaction_buffer", False - ) + _use_redis_transaction_buffer: bool | str | None = general_settings.get("use_redis_transaction_buffer", False) if isinstance(_use_redis_transaction_buffer, str): _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) @@ -7790,7 +7759,7 @@ class ProxyStartupEvent: @staticmethod async def _init_coordination_redis_from_db( litellm_settings: Mapping[str, object], - llm_router: Optional[Router], + llm_router: Router | None, ) -> RedisCache | None: """ Applies a coordination_redis block saved to the database, which the admin @@ -7853,8 +7822,8 @@ class ProxyStartupEvent: @classmethod async def _initialize_semantic_tool_filter( cls, - llm_router: Optional[Router], - litellm_settings: Dict[str, Any], + llm_router: Router | None, + litellm_settings: dict[str, Any], ): """Initialize MCP semantic tool filter if configured""" from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook @@ -7886,7 +7855,7 @@ class ProxyStartupEvent: def _initialize_jwt_auth( cls, general_settings: dict, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, ): """Initialize JWT auth on startup""" @@ -8302,7 +8271,6 @@ class ProxyStartupEvent: verbose_proxy_logger.debug( "Checking batch cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." ) - pass ### CHECK RESPONSES COST ### if llm_router is not None and PROXY_BATCH_POLLING_ENABLED: @@ -8332,7 +8300,6 @@ class ProxyStartupEvent: verbose_proxy_logger.debug( "Checking responses cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." ) - pass # MEMORY LEAK FIX: Start scheduler with paused=False to avoid backlog processing # Do NOT reset job times to "now" as this can trigger the memory leak @@ -8429,7 +8396,7 @@ class ProxyStartupEvent: LITELLM_KEY_ROTATION_ENABLED, ) - key_rotation_enabled: Optional[bool] = str_to_bool(LITELLM_KEY_ROTATION_ENABLED) + key_rotation_enabled: bool | None = str_to_bool(LITELLM_KEY_ROTATION_ENABLED) verbose_proxy_logger.debug(f"key_rotation_enabled: {key_rotation_enabled}") if key_rotation_enabled is True: @@ -8482,7 +8449,7 @@ class ProxyStartupEvent: LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS, ) - expired_ui_session_key_cleanup_enabled: Optional[bool] = str_to_bool( + expired_ui_session_key_cleanup_enabled: bool | None = str_to_bool( LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED ) verbose_proxy_logger.debug(f"expired_ui_session_key_cleanup_enabled: {expired_ui_session_key_cleanup_enabled}") @@ -8582,16 +8549,16 @@ class ProxyStartupEvent: @classmethod async def _setup_prisma_client( cls, - database_url: Optional[str], + database_url: str | None, proxy_logging_obj: ProxyLogging, user_api_key_cache: UserApiKeyCache, - ) -> Optional[PrismaClient]: + ) -> PrismaClient | None: """ - Sets up prisma client - Adds necessary views to proxy """ try: - prisma_client: Optional[PrismaClient] = None + prisma_client: PrismaClient | None = None if database_url is not None: try: prisma_client = PrismaClient(database_url=database_url, proxy_logging_obj=proxy_logging_obj) @@ -8749,14 +8716,14 @@ class ProxyStartupEvent: ) # if project requires model list async def model_list( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - return_wildcard_routes: Optional[bool] = False, - team_id: Optional[str] = None, - include_model_access_groups: Optional[bool] = False, - only_model_access_groups: Optional[bool] = False, - include_metadata: Optional[bool] = False, - fallback_type: Optional[str] = None, - scope: Optional[str] = None, - healthy_only: Optional[bool] = False, + return_wildcard_routes: bool | None = False, + team_id: str | None = None, + include_model_access_groups: bool | None = False, + only_model_access_groups: bool | None = False, + include_metadata: bool | None = False, + fallback_type: str | None = None, + scope: str | None = None, + healthy_only: bool | None = False, ): """ Use `/model/info` - to get detailed model information, example - pricing, mode, etc. @@ -8814,7 +8781,7 @@ async def model_list( # Opt-in: also hide models whose deployments are all unhealthy per background # health checks. Empty when health state is unavailable or stale (fail open). - unhealthy_names: Set[str] = set() + unhealthy_names: set[str] = set() if healthy_only and llm_router is not None: unhealthy_names = await llm_router.async_get_fully_unhealthy_model_names() if not unhealthy_names: @@ -8933,8 +8900,8 @@ async def model_list( async def model_info( model_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - team_id: Optional[str] = None, - healthy_only: Optional[bool] = False, + team_id: str | None = None, + healthy_only: bool | None = False, ): """ Retrieve information about a specific model accessible to your API key. @@ -9018,7 +8985,7 @@ async def model_info( ) -def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage": +def _blocked_response_usage(original_response: Any | None) -> "litellm.Usage": """ Token usage for a synthetic guardrail-blocked response. @@ -9058,7 +9025,7 @@ def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage" async def chat_completion( request: Request, fastapi_response: Response, - model: Optional[str] = None, + model: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -9223,7 +9190,7 @@ async def chat_completion( async def completion( request: Request, fastapi_response: Response, - model: Optional[str] = None, + model: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -9365,8 +9332,8 @@ async def completion( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - {}".format(str(e))) - error_msg = f"{str(e)}" + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.completion(): Exception occured - {e!s}") + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -9403,7 +9370,7 @@ async def completion( async def embeddings( request: Request, fastapi_response: Response, - model: Optional[str] = None, + model: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -9535,7 +9502,7 @@ async def moderations( ``` """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly body = await request.body() @@ -9604,9 +9571,7 @@ async def moderations( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.moderations(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.moderations(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e)), @@ -9615,7 +9580,7 @@ async def moderations( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -9657,7 +9622,7 @@ async def audio_speech( https://platform.openai.com/docs/api-reference/audio/createSpeech """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly body = await request.body() @@ -9752,7 +9717,7 @@ async def audio_speech( original_exception=e, request_data=data, ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.audio_speech(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.audio_speech(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) raise e @@ -9779,7 +9744,7 @@ async def audio_transcriptions( https://platform.openai.com/docs/api-reference/audio/createTranscription?lang=curl """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly form_data = await get_form_data(request) @@ -9894,9 +9859,7 @@ async def audio_transcriptions( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.audio_transcription(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.audio_transcription(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e.detail)), @@ -9905,7 +9868,7 @@ async def audio_transcriptions( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -9925,15 +9888,15 @@ async def audio_transcriptions( @app.websocket("/vertex_ai/live") async def vertex_ai_live_passthrough_endpoint( websocket: WebSocket, - model: Optional[str] = fastapi.Query( + model: str | None = fastapi.Query( None, description="Optional model name, used to determine Vertex region for global models.", ), - vertex_project: Optional[str] = fastapi.Query( + vertex_project: str | None = fastapi.Query( None, description="Override the Vertex AI project id used for the upstream connection.", ), - vertex_location: Optional[str] = fastapi.Query( + vertex_location: str | None = fastapi.Query( None, description="Override the Vertex AI region (for example, 'us-central1').", ), @@ -9961,12 +9924,12 @@ async def vertex_ai_live_passthrough_endpoint( @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) -def _realtime_query_params_template(model: Optional[str], intent: Optional[str]) -> Tuple[Tuple[str, str], ...]: +def _realtime_query_params_template(model: str | None, intent: str | None) -> tuple[tuple[str, str], ...]: """ Build a hashable representation of the realtime query params so we can cache the repetitive model/intent combinations. """ - params: List[Tuple[str, str]] = [] + params: list[tuple[str, str]] = [] if model is not None: params.append(("model", model)) if intent is not None: @@ -9979,9 +9942,9 @@ def _realtime_query_params_template(model: Optional[str], intent: Optional[str]) @app.websocket("/realtime") async def realtime_websocket_endpoint( websocket: WebSocket, - model: Optional[str] = fastapi.Query(None, description="The model to use for the websocket connection."), - intent: Optional[str] = fastapi.Query(None, description="The intent of the websocket connection."), - guardrails: Optional[str] = fastapi.Query( + model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."), + intent: str | None = fastapi.Query(None, description="The intent of the websocket connection."), + guardrails: str | None = fastapi.Query( None, description="Comma-separated list of guardrail names to apply to this request.", ), @@ -10017,7 +9980,7 @@ async def realtime_websocket_endpoint( # Only use explicit parameters, not all query params query_params = cast(RealtimeQueryParams, dict(_realtime_query_params_template(model, intent))) - data: Dict[str, Any] = { + data: dict[str, Any] = { "model": route_model, "websocket": websocket, "query_params": query_params, # Only explicit params @@ -10133,7 +10096,7 @@ async def get_assistants( API Reference docs - https://platform.openai.com/docs/api-reference/assistants/listAssistants """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly await request.body() @@ -10182,7 +10145,7 @@ async def get_assistants( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.get_assistants(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.get_assistants(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10192,7 +10155,7 @@ async def get_assistants( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10273,9 +10236,7 @@ async def create_assistant( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.create_assistant(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.create_assistant(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10285,7 +10246,7 @@ async def create_assistant( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10316,7 +10277,7 @@ async def delete_assistant( API Reference docs - https://platform.openai.com/docs/api-reference/assistants/createAssistant """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly @@ -10364,9 +10325,7 @@ async def delete_assistant( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.delete_assistant(): Exception occured - {}".format(str(e)) - ) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.delete_assistant(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10376,7 +10335,7 @@ async def delete_assistant( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10406,7 +10365,7 @@ async def create_threads( API Reference - https://platform.openai.com/docs/api-reference/threads/createThread """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly await request.body() @@ -10455,7 +10414,7 @@ async def create_threads( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.create_threads(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.create_threads(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10465,7 +10424,7 @@ async def create_threads( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10496,7 +10455,7 @@ async def get_thread( API Reference - https://platform.openai.com/docs/api-reference/threads/getThread """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Include original request and headers in the data data = await add_litellm_data_to_request( @@ -10542,7 +10501,7 @@ async def get_thread( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.get_thread(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.get_thread(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10552,7 +10511,7 @@ async def get_thread( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10583,7 +10542,7 @@ async def add_messages( API Reference - https://platform.openai.com/docs/api-reference/messages/createMessage """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Use orjson to parse JSON data, orjson speeds up requests significantly body = await request.body() @@ -10633,7 +10592,7 @@ async def add_messages( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.add_messages(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.add_messages(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10643,7 +10602,7 @@ async def add_messages( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10674,7 +10633,7 @@ async def get_messages( API Reference - https://platform.openai.com/docs/api-reference/messages/listMessages """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: # Include original request and headers in the data data = await add_litellm_data_to_request( @@ -10720,7 +10679,7 @@ async def get_messages( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.get_messages(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.get_messages(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10730,7 +10689,7 @@ async def get_messages( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10761,7 +10720,7 @@ async def run_thread( API Reference: https://platform.openai.com/docs/api-reference/runs/createRun """ global proxy_logging_obj - data: Dict = {} + data: dict = {} try: body = await request.body() data = orjson.loads(body) @@ -10821,7 +10780,7 @@ async def run_thread( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.run_thread(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.run_thread(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( @@ -10831,7 +10790,7 @@ async def run_thread( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -10865,7 +10824,7 @@ from litellm.repositories.user_repository import UserRepository def _get_provider_token_counter( deployment: dict, model_to_use: str -) -> Tuple[Optional[BaseTokenCounter], Optional[str], Optional[str]]: +) -> tuple[BaseTokenCounter | None, str | None, str | None]: """ Auto-route to the correct provider's token counter based on model/deployment. Uses the existing get_provider_model_info infrastructure with switch-case pattern. @@ -10876,8 +10835,8 @@ def _get_provider_token_counter( from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider full_model = deployment.get("litellm_params", {}).get("model", "") - model: Optional[str] = None - custom_llm_provider: Optional[str] = None + model: str | None = None + custom_llm_provider: str | None = None try: # Use existing LiteLLM logic to determine provider @@ -10920,14 +10879,14 @@ def _get_provider_token_counter( async def _try_provider_token_count( provider_counter: "BaseTokenCounter", - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, model_to_use: str, - messages: Optional[list], - contents: Optional[list], - deployment: Optional[Dict[str, Any]], + messages: list | None, + contents: list | None, + deployment: dict[str, Any] | None, request_model: str, - tools: Optional[list] = None, - system: Optional[str] = None, + tools: list | None = None, + system: str | None = None, ) -> Optional["TokenCountResponse"]: """Attempt provider-specific token counting. Returns result on success, None to fall through to local counting.""" if not provider_counter.should_use_token_counting_api(custom_llm_provider=custom_llm_provider): @@ -10998,9 +10957,9 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) if prompt is None and messages is None and contents is None: raise HTTPException(status_code=400, detail="prompt or messages or contents must be provided") - deployment: Optional[Dict[str, Any]] = None + deployment: dict[str, Any] | None = None litellm_model_name = None - model_info: Optional[ModelMapInfo] = None + model_info: ModelMapInfo | None = None if llm_router is not None: # get 1 deployment corresponding to the model try: @@ -11012,7 +10971,6 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) verbose_proxy_logger.exception( "litellm.proxy.proxy_server.token_counter(): Exception occured while getting deployment" ) - pass if deployment is not None: litellm_model_name = deployment.get("litellm_params", {}).get("model") model_info = deployment.get("model_info", {}) @@ -11026,8 +10984,8 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) ) # use litellm model name, if it's not avalable then fallback to request.model # Try provider-specific token counting first - only for non-direct requests (from provider endpoints) - provider_counter: Optional[BaseTokenCounter] = None - custom_llm_provider: Optional[str] = None + provider_counter: BaseTokenCounter | None = None + custom_llm_provider: str | None = None if call_endpoint is True and deployment is not None: # Auto-route to the correct provider based on model provider_counter, _model, custom_llm_provider = _get_provider_token_counter(deployment, model_to_use) @@ -11059,10 +11017,10 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) ) # Default LiteLLM token counting - custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None + custom_tokenizer: CustomHuggingfaceTokenizer | None = None if model_info is not None: custom_tokenizer = cast( - Optional[CustomHuggingfaceTokenizer], + CustomHuggingfaceTokenizer | None, model_info.get("custom_tokenizer", None), ) _tokenizer_used = litellm.utils._select_tokenizer(model=model_to_use, custom_tokenizer=custom_tokenizer) @@ -11107,7 +11065,7 @@ async def supported_openai_params(model: str): ) } except Exception: - raise HTTPException(status_code=400, detail={"error": "Could not map model={}".format(model)}) + raise HTTPException(status_code=400, detail={"error": f"Could not map model={model}"}) @router.post( @@ -11133,10 +11091,10 @@ async def transform_request(request: TransformRequestBody): async def _check_if_model_is_user_added( - models: List[Dict], + models: list[dict], user_api_key_dict: UserAPIKeyAuth, - prisma_client: Optional[PrismaClient], -) -> List[Dict]: + prisma_client: PrismaClient | None, +) -> list[dict]: """ Check if model is in db @@ -11161,29 +11119,29 @@ async def _check_if_model_is_user_added( return filtered_models -def _check_if_model_is_team_model(models: List[DeploymentTypedDict], user_row: LiteLLM_UserTable) -> List[Dict]: +def _check_if_model_is_team_model(models: list[DeploymentTypedDict], user_row: LiteLLM_UserTable) -> list[dict]: """ Check if model is a team model Check if user is a member of the team that the model belongs to """ - user_team_models: List[Dict] = [] + user_team_models: list[dict] = [] for model in models: model_team_id = model.get("model_info", {}).get("team_id", None) if model_team_id is not None: if model_team_id in user_row.teams: - user_team_models.append(cast(Dict, model)) + user_team_models.append(cast(dict, model)) return user_team_models async def non_admin_all_models( - all_models: List[Dict], + all_models: list[dict], llm_router: Router, user_api_key_dict: UserAPIKeyAuth, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, ): """ Check if model is in db @@ -11225,13 +11183,13 @@ async def non_admin_all_models( def _add_team_models_to_all_models( - team_db_objects_typed: List[LiteLLM_TeamTable], + team_db_objects_typed: list[LiteLLM_TeamTable], llm_router: Router, -) -> Dict[str, Set[str]]: +) -> dict[str, set[str]]: """ Add team models to all models """ - team_models: Dict[str, Set[str]] = {} + team_models: dict[str, set[str]] = {} for team_object in team_db_objects_typed: if ( @@ -11247,9 +11205,7 @@ def _add_team_models_to_all_models( # if team model id set, check if team id in user_teams team_model_id = model.get("model_info", {}).get("team_id", None) can_add_model = False - if team_model_id is None: - can_add_model = True - elif team_model_id in team_object.team_id: + if team_model_id is None or team_model_id in team_object.team_id: can_add_model = True if can_add_model: @@ -11266,11 +11222,11 @@ def _add_team_models_to_all_models( async def _add_access_group_models_to_team_models( - team_db_objects_typed: List[LiteLLM_TeamTable], + team_db_objects_typed: list[LiteLLM_TeamTable], llm_router: Router, prisma_client: PrismaClient, - team_models: Dict[str, Set[str]], -) -> Dict[str, Set[str]]: + team_models: dict[str, set[str]], +) -> dict[str, set[str]]: """ Resolve models reachable via team access groups and merge them into team_models. @@ -11281,8 +11237,8 @@ async def _add_access_group_models_to_team_models( (not directly in team.models) are included in the UI model listing. """ # First pass: identify eligible teams and collect all distinct access group IDs - eligible_teams: List[LiteLLM_TeamTable] = [] - all_access_group_ids: Set[str] = set() + eligible_teams: list[LiteLLM_TeamTable] = [] + all_access_group_ids: set[str] = set() for team_object in team_db_objects_typed: if not team_object.access_group_ids: @@ -11303,13 +11259,13 @@ async def _add_access_group_models_to_team_models( access_group_rows = await AccessGroupRepository(prisma_client).table.find_many( where={"access_group_id": {"in": list(all_access_group_ids)}} ) - ag_model_map: Dict[str, List[str]] = { + ag_model_map: dict[str, list[str]] = { row.access_group_id: row.access_model_names or [] for row in access_group_rows } # Second pass: resolve deployments for each eligible team for team_object in eligible_teams: - model_names: Set[str] = set() + model_names: set[str] = set() for ag_id in team_object.access_group_ids or []: model_names.update(ag_model_map.get(ag_id, [])) @@ -11325,10 +11281,10 @@ async def _add_access_group_models_to_team_models( async def get_all_team_models( - user_teams: Union[List[str], Literal["*"]], + user_teams: list[str] | Literal["*"], prisma_client: PrismaClient, llm_router: Router, -) -> Dict[str, List[str]]: +) -> dict[str, list[str]]: """ Get all models across all teams user is in. @@ -11337,7 +11293,7 @@ async def get_all_team_models( 3. Return {"model_id": ["team_id1", "team_id2"]} """ - team_db_objects_typed: List[LiteLLM_TeamTable] = [] + team_db_objects_typed: list[LiteLLM_TeamTable] = [] if user_teams == "*": team_db_objects = await TeamRepository(prisma_client).table.find_many() @@ -11365,7 +11321,7 @@ async def get_all_team_models( ) # convert set to list - returned_team_models: Dict[str, List[str]] = {} + returned_team_models: dict[str, list[str]] = {} for model_id, team_ids in team_models.items(): returned_team_models[model_id] = list(team_ids) @@ -11375,7 +11331,7 @@ async def get_all_team_models( def get_direct_access_models( user_db_object: LiteLLM_UserTable, llm_router: Router, -) -> List[str]: +) -> list[str]: """ Get all models that user has direct access to. @@ -11393,7 +11349,7 @@ def get_direct_access_models( ] -def _filter_models_to_user_accessible(all_models: List[Dict]) -> List[Dict]: +def _filter_models_to_user_accessible(all_models: list[dict]) -> list[dict]: """Keep only deployments the caller can use via direct access or team membership.""" return [ _model @@ -11407,14 +11363,14 @@ async def _populate_team_access_on_models( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, llm_router: Router, - all_models: List[Dict], -) -> List[Dict]: + all_models: list[dict], +) -> list[dict]: """ Populate `model_info.access_via_team_ids` and `model_info.direct_access` without filtering the model list. """ - user_teams: Optional[Union[List[str], Literal["*"]]] = None - direct_access_models: List[str] = [] + user_teams: list[str] | Literal["*"] | None = None + direct_access_models: list[str] = [] if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: user_teams = "*" direct_access_models = llm_router.get_model_ids(exclude_team_models=True) # has access to all models @@ -11462,8 +11418,8 @@ async def get_all_team_and_direct_access_models( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, llm_router: Router, - all_models: List[Dict], -) -> List[Dict]: + all_models: list[dict], +) -> list[dict]: """ Get all models across all teams user is in. """ @@ -11477,8 +11433,8 @@ async def get_all_team_and_direct_access_models( def _enrich_model_info_with_litellm_data( - model: Dict[str, Any], debug: bool = False, llm_router: Optional[Router] = None -) -> Dict[str, Any]: + model: dict[str, Any], debug: bool = False, llm_router: Router | None = None +) -> dict[str, Any]: """ Enrich a model dictionary with litellm model info (pricing, context window, etc.) and remove sensitive information. @@ -11539,9 +11495,9 @@ def _enrich_model_info_with_litellm_data( async def _get_caller_byok_team_scope( - user_api_key_dict: Optional[UserAPIKeyAuth], - prisma_client: Optional[Any], -) -> Optional[Set[str]]: + user_api_key_dict: UserAPIKeyAuth | None, + prisma_client: Any | None, +) -> set[str] | None: """ Return the team IDs whose BYOK rows the caller is allowed to see via `/v2/model/info` search results. @@ -11575,7 +11531,7 @@ async def _get_caller_byok_team_scope( return key_team_scope | set(user_row.teams or []) -def _byok_row_outside_caller_teams(model_info_dict: Dict[str, Any], allowed_team_ids: Optional[Set[str]]) -> bool: +def _byok_row_outside_caller_teams(model_info_dict: dict[str, Any], allowed_team_ids: set[str] | None) -> bool: """Whether a team BYOK row belongs to a team the caller is not a member of. `team_id` is only set on team BYOK rows; non-team rows fall through @@ -11600,13 +11556,13 @@ async def _fetch_db_models_for_search( prisma_client: Any, proxy_config: Any, search_lower: str, - db_model_ids_in_router: Set[str], + db_model_ids_in_router: set[str], router_models_count: int, page: int, size: int, - sort_by: Optional[str], - is_byok_outside_caller_teams: Callable[[Dict[str, Any]], bool], -) -> Tuple[List[Dict[str, Any]], int]: + sort_by: str | None, + is_byok_outside_caller_teams: Callable[[dict[str, Any]], bool], +) -> tuple[list[dict[str, Any]], int]: """ Run the bounded DB query that backs `/v2/model/info?search=`. Returns `(decrypted_models, total_count)` where `total_count` is the cheap @@ -11622,7 +11578,7 @@ async def _fetch_db_models_for_search( filter for `team_public_model_name` instead and keep the DB cost bounded by `search`. """ - db_where_condition: Dict[str, Any] = {"model_name": {"contains": search_lower, "mode": "insensitive"}} + db_where_condition: dict[str, Any] = {"model_name": {"contains": search_lower, "mode": "insensitive"}} if db_model_ids_in_router: db_where_condition["model_id"] = {"not": {"in": list(db_model_ids_in_router)}} @@ -11651,7 +11607,7 @@ async def _fetch_db_models_for_search( if not is_byok_outside_caller_teams(m.model_info if isinstance(m.model_info, dict) else {}) ] - decrypted: List[Dict[str, Any]] = [] + decrypted: list[dict[str, Any]] = [] for db_model in matching_db_rows: decrypted_models = proxy_config.decrypt_model_list_from_db([db_model]) if decrypted_models: @@ -11661,15 +11617,15 @@ async def _fetch_db_models_for_search( async def _apply_search_filter_to_models( - all_models: List[Dict[str, Any]], + all_models: list[dict[str, Any]], search: str, - prisma_client: Optional[Any], + prisma_client: Any | None, proxy_config: Any, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, page: int = 1, size: int = 50, - sort_by: Optional[str] = None, -) -> Tuple[List[Dict[str, Any]], Optional[int]]: + sort_by: str | None = None, +) -> tuple[list[dict[str, Any]], int | None]: """ Apply search filter to models, querying database for additional matching models. @@ -11703,10 +11659,10 @@ async def _apply_search_filter_to_models( prisma_client=prisma_client, ) - def _is_byok_outside_caller_teams(model_info_dict: Dict[str, Any]) -> bool: + def _is_byok_outside_caller_teams(model_info_dict: dict[str, Any]) -> bool: return _byok_row_outside_caller_teams(model_info_dict, allowed_team_ids) - def _model_matches_search(m: Dict[str, Any]) -> bool: + def _model_matches_search(m: dict[str, Any]) -> bool: # Team BYOK models persist an internal `model_name` # (e.g. `model_name_{team_id}_{uuid}`) and expose the user-facing # name via `model_info.team_public_model_name`. Match both so the @@ -11745,7 +11701,7 @@ async def _apply_search_filter_to_models( router_models_count = config_models_count + db_models_in_router_count # Query database for additional models with search term - db_models: List[Dict[str, Any]] = [] + db_models: list[dict[str, Any]] = [] if prisma_client is not None: try: db_models, db_models_total_count = await _fetch_db_models_for_search( @@ -11761,7 +11717,7 @@ async def _apply_search_filter_to_models( ) search_total_count = router_models_count + db_models_total_count except Exception as e: - verbose_proxy_logger.exception(f"Error querying database models with search: {str(e)}") + verbose_proxy_logger.exception(f"Error querying database models with search: {e!s}") search_total_count = router_models_count else: search_total_count = router_models_count @@ -11769,7 +11725,7 @@ async def _apply_search_filter_to_models( return filtered_router_models + db_models, search_total_count -def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]: +def _normalize_datetime_for_sorting(dt: Any) -> datetime | None: """ Normalize a datetime value to a timezone-aware UTC datetime for sorting. @@ -11812,10 +11768,10 @@ def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]: def _sort_models( - all_models: List[Dict[str, Any]], - sort_by: Optional[str], + all_models: list[dict[str, Any]], + sort_by: str | None, sort_order: str = "asc", -) -> List[Dict[str, Any]]: +) -> list[dict[str, Any]]: """ Sort models by the specified field and order. @@ -11838,7 +11794,7 @@ def _sort_models( reverse = sort_order.lower() == "desc" - def get_sort_key(model: Dict[str, Any]) -> Any: + def get_sort_key(model: dict[str, Any]) -> Any: model_info = model.get("model_info", {}) if sort_by == "model_name": @@ -11896,7 +11852,7 @@ def _sort_models( sorted_models = sorted(all_models, key=get_sort_key, reverse=reverse) return sorted_models except Exception as e: - verbose_proxy_logger.exception(f"Error sorting models by {sort_by}: {str(e)}") + verbose_proxy_logger.exception(f"Error sorting models by {sort_by}: {e!s}") return all_models @@ -11917,12 +11873,12 @@ def _is_auto_router_model(model: Mapping[str, object]) -> bool: def _paginate_models_response( - all_models: List[Dict[str, Any]], + all_models: list[dict[str, Any]], page: int, size: int, - total_count: Optional[int], - search: Optional[str], -) -> Dict[str, Any]: + total_count: int | None, + search: str | None, +) -> dict[str, Any]: """ Paginate models and return response dictionary. @@ -11956,9 +11912,9 @@ def _paginate_models_response( } -def _team_models_resolve_to_names(team_models: List[str], access_groups: Dict[str, Any]) -> List[str]: +def _team_models_resolve_to_names(team_models: list[str], access_groups: dict[str, Any]) -> list[str]: """Expand team model entries (including access group names) to concrete model names.""" - resolved: List[str] = [] + resolved: list[str] = [] for name in team_models: if name in access_groups: resolved.extend(access_groups[name]) @@ -11967,7 +11923,7 @@ def _team_models_resolve_to_names(team_models: List[str], access_groups: Dict[st return resolved -async def _load_team_object_for_model_filter(team_id: str, prisma_client: PrismaClient) -> Optional[LiteLLM_TeamTable]: +async def _load_team_object_for_model_filter(team_id: str, prisma_client: PrismaClient) -> LiteLLM_TeamTable | None: """Load team row from DB; returns None if missing or on error.""" try: team_db_object = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) @@ -11976,7 +11932,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma return None return LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) except Exception as e: - verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") + verbose_proxy_logger.exception(f"Error fetching team {team_id}: {e!s}") return None @@ -11985,9 +11941,9 @@ async def _gather_team_accessible_model_ids( team_id: str, prisma_client: PrismaClient, llm_router: Router, -) -> Set[str]: +) -> set[str]: """Collect model IDs the team can use from router config and DB.""" - team_accessible_model_ids: Set[str] = set() + team_accessible_model_ids: set[str] = set() access_groups = llm_router.get_model_access_groups() if llm_router else {} if not team_object.models or SpecialModelNames.all_proxy_models.value in team_object.models: @@ -12001,7 +11957,7 @@ async def _gather_team_accessible_model_ids( if team_model_id is None or team_model_id == team_id: team_accessible_model_ids.add(model_id) else: - resolved_model_names: Set[str] = set() + resolved_model_names: set[str] = set() for model_name in team_object.models: if model_name in access_groups: resolved_model_names.update(access_groups[model_name]) @@ -12026,7 +11982,7 @@ async def _gather_team_accessible_model_ids( if db_model.model_id: team_accessible_model_ids.add(db_model.model_id) except Exception as e: - verbose_proxy_logger.debug(f"Error querying database models for team {team_id}: {str(e)}") + verbose_proxy_logger.debug(f"Error querying database models for team {team_id}: {e!s}") return team_accessible_model_ids @@ -12072,12 +12028,12 @@ async def _authorize_team_id_query( async def _filter_models_by_team_id( - all_models: List[Dict[str, Any]], + all_models: list[dict[str, Any]], team_id: str, prisma_client: PrismaClient, llm_router: Router, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> List[Dict[str, Any]]: + user_api_key_dict: UserAPIKeyAuth | None = None, +) -> list[dict[str, Any]]: """ Filter models by team ID. Returns models where: - team_id matches the model's BYOK team_id, OR @@ -12140,11 +12096,11 @@ async def _filter_models_by_team_id( async def _find_model_by_id( model_id: str, - search: Optional[str], + search: str | None, llm_router, prisma_client, proxy_config, -) -> tuple[list, Optional[int]]: +) -> tuple[list, int | None]: """Find a model by its ID and optionally filter by search term.""" found_model = None @@ -12164,7 +12120,7 @@ async def _find_model_by_id( if decrypted_models: found_model = decrypted_models[0] except Exception as e: - verbose_proxy_logger.exception(f"Error querying database for modelId {model_id}: {str(e)}") + verbose_proxy_logger.exception(f"Error querying database for modelId {model_id}: {e!s}") # If model found, verify search filter if provided if found_model is not None: @@ -12177,7 +12133,7 @@ async def _find_model_by_id( # Set all_models to the found model or empty list all_models = [found_model] if found_model is not None else [] - search_total_count: Optional[int] = len(all_models) + search_total_count: int | None = len(all_models) return all_models, search_total_count @@ -12188,25 +12144,25 @@ async def _find_model_by_id( ) async def model_info_v2( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - model: Optional[str] = fastapi.Query(None, description="Specify the model name (optional)"), - user_models_only: Optional[bool] = fastapi.Query(False, description="Only return models added by this user"), - include_team_models: Optional[bool] = fastapi.Query( + model: str | None = fastapi.Query(None, description="Specify the model name (optional)"), + user_models_only: bool | None = fastapi.Query(False, description="Only return models added by this user"), + include_team_models: bool | None = fastapi.Query( False, description="Return all models across all teams user is in." ), - debug: Optional[bool] = False, + debug: bool | None = False, page: int = Query(1, description="Page number", ge=1), size: int = Query(50, description="Page size", ge=1), - search: Optional[str] = fastapi.Query(None, description="Search model names (case-insensitive partial match)"), - modelId: Optional[str] = fastapi.Query(None, description="Search for a specific model by its unique ID"), - teamId: Optional[str] = fastapi.Query( + search: str | None = fastapi.Query(None, description="Search model names (case-insensitive partial match)"), + modelId: str | None = fastapi.Query(None, description="Search for a specific model by its unique ID"), + teamId: str | None = fastapi.Query( None, description="Filter models by team ID. Returns models with direct_access=True or teamId in access_via_team_ids", ), - sortBy: Optional[str] = fastapi.Query( + sortBy: str | None = fastapi.Query( None, description="Field to sort by. Options: model_name, created_at, updated_at, costs, status", ), - sortOrder: Optional[str] = fastapi.Query( + sortOrder: str | None = fastapi.Query( "asc", description="Sort order. Options: asc, desc", ), @@ -12415,9 +12371,9 @@ async def model_info_v2( ) async def model_streaming_metrics( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - _selected_model_group: Optional[str] = None, - startTime: Optional[datetime] = None, - endTime: Optional[datetime] = None, + _selected_model_group: str | None = None, + startTime: datetime | None = None, + endTime: datetime | None = None, ): global prisma_client, llm_router if prisma_client is None: @@ -12520,7 +12476,7 @@ async def model_streaming_metrics( """ # convert daily entries to list of dicts - response: List[dict] = [] + response: list[dict] = [] # sort daily entries by date _daily_entries = dict(sorted(_daily_entries.items(), key=lambda item: item[0])) @@ -12545,11 +12501,11 @@ async def model_streaming_metrics( ) async def model_metrics( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - _selected_model_group: Optional[str] = "gpt-4-32k", - startTime: Optional[datetime] = None, - endTime: Optional[datetime] = None, - api_key: Optional[str] = None, - customer: Optional[str] = None, + _selected_model_group: str | None = "gpt-4-32k", + startTime: datetime | None = None, + endTime: datetime | None = None, + api_key: str | None = None, + customer: str | None = None, ): global prisma_client, llm_router if prisma_client is None: @@ -12635,7 +12591,7 @@ async def model_metrics( """ # convert daily entries to list of dicts - response: List[dict] = [] + response: list[dict] = [] # sort daily entries by date _daily_entries = dict(sorted(_daily_entries.items(), key=lambda item: item[0])) @@ -12660,11 +12616,11 @@ async def model_metrics( ) async def model_metrics_slow_responses( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - _selected_model_group: Optional[str] = "gpt-4-32k", - startTime: Optional[datetime] = None, - endTime: Optional[datetime] = None, - api_key: Optional[str] = None, - customer: Optional[str] = None, + _selected_model_group: str | None = "gpt-4-32k", + startTime: datetime | None = None, + endTime: datetime | None = None, + api_key: str | None = None, + customer: str | None = None, ): global prisma_client, llm_router, proxy_logging_obj if prisma_client is None: @@ -12749,11 +12705,11 @@ ORDER BY ) async def model_metrics_exceptions( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - _selected_model_group: Optional[str] = None, - startTime: Optional[datetime] = None, - endTime: Optional[datetime] = None, - api_key: Optional[str] = None, - customer: Optional[str] = None, + _selected_model_group: str | None = None, + startTime: datetime | None = None, + endTime: datetime | None = None, + api_key: str | None = None, + customer: str | None = None, ): global prisma_client, llm_router if prisma_client is None: @@ -12795,7 +12751,7 @@ async def model_metrics_exceptions( LIMIT 200; """ db_response = await prisma_client.db.query_raw(sql_query, startTime, endTime, _selected_model_group, api_key) - response: List[dict] = [] + response: list[dict] = [] exception_types = set() """ @@ -12826,7 +12782,7 @@ async def model_metrics_exceptions( return {"data": response, "exception_types": list(exception_types)} -def _deployment_matches_allowed_model_names(model: Dict[str, Any], allowed_model_names: Set[str]) -> bool: +def _deployment_matches_allowed_model_names(model: dict[str, Any], allowed_model_names: set[str]) -> bool: """Match a router deployment against allowed public model names. Team-scoped rows store an internal routing key in ``model_name``; callers @@ -12845,7 +12801,7 @@ def _deployment_matches_allowed_model_names(model: Dict[str, Any], allowed_model def _get_v1_model_info_allowed_model_names( user_api_key_dict: UserAPIKeyAuth, llm_router: Router, -) -> Optional[Set[str]]: +) -> set[str] | None: """Return key/team allowlisted public model names, or None if unrestricted.""" model_access_groups = llm_router.get_model_access_groups() proxy_model_list = llm_router.get_model_names() @@ -12875,9 +12831,9 @@ def _get_v1_model_info_allowed_model_names( def _filter_v1_model_info_deployments( - all_models: List[dict], - allowed_model_names: Optional[Set[str]], -) -> List[dict]: + all_models: list[dict], + allowed_model_names: set[str] | None, +) -> list[dict]: if allowed_model_names is None: return all_models return [model for model in all_models if _deployment_matches_allowed_model_names(model, allowed_model_names)] @@ -12961,12 +12917,12 @@ def _get_proxy_model_info(model: dict) -> dict: ) async def model_info_v1( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - litellm_model_id: Optional[str] = None, - include_team_models: Optional[bool] = fastapi.Query( + litellm_model_id: str | None = None, + include_team_models: bool | None = fastapi.Query( False, description="When true, filter to deployments the caller can use via direct access or team membership.", ), - teamId: Optional[str] = fastapi.Query( + teamId: str | None = fastapi.Query( None, description="Filter models by team ID. Returns models with direct_access=True or teamId in access_via_team_ids", ), @@ -13019,7 +12975,7 @@ async def model_info_v1( if user_model is not None: # user is trying to get specific model from litellm router try: - model_info: Dict = cast(Dict, litellm.get_model_info(model=user_model)) + model_info: dict = cast(dict, litellm.get_model_info(model=user_model)) except Exception: model_info = {} _deployment_info = Deployment( @@ -13067,7 +13023,7 @@ async def model_info_v1( detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"}, ) _deployment_info_dict = _get_proxy_model_info(model=deployment_info.model_dump(exclude_none=True)) - single_model_list: List[dict] = [_deployment_info_dict] + single_model_list: list[dict] = [_deployment_info_dict] if prisma_client is not None: single_model_list = await _populate_team_access_on_models( user_api_key_dict=user_api_key_dict, @@ -13091,7 +13047,7 @@ async def model_info_v1( # expanded model names from get_complete_model_list(). Team-scoped rows # use internal routing keys (model_name_{team_id}_{uuid}) and were omitted # when v1 resolved models only via public model_name strings. - all_models: List[dict] = copy.deepcopy(llm_router.model_list) + all_models: list[dict] = copy.deepcopy(llm_router.model_list) alias_models = copy.deepcopy(llm_router.get_model_list_from_model_alias()) all_models.extend(alias_models) @@ -13150,9 +13106,9 @@ async def model_info_v1( def _get_model_group_info( - llm_router: Router, all_models_str: List[str], model_group: Optional[str] -) -> List[ModelGroupInfoProxy]: - model_groups: List[ModelGroupInfoProxy] = [] + llm_router: Router, all_models_str: list[str], model_group: str | None +) -> list[ModelGroupInfoProxy]: + model_groups: list[ModelGroupInfoProxy] = [] unique_models = [] for model in all_models_str: @@ -13190,7 +13146,7 @@ def _get_model_group_info( ) async def model_group_info( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - model_group: Optional[str] = None, + model_group: str | None = None, ): """ Get information about all the deployments on litellm proxy, including config.yaml descriptions (except api key and api base) @@ -13364,7 +13320,7 @@ async def model_group_info( return_wildcard_routes=False, user_api_key_cache=user_api_key_cache, ) - model_groups: List[ModelGroupInfoProxy] = _get_model_group_info( + model_groups: list[ModelGroupInfoProxy] = _get_model_group_info( llm_router=llm_router, all_models_str=all_models_str, model_group=model_group ) @@ -13443,12 +13399,7 @@ async def alerting_settings( if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) ## get general settings from db @@ -13459,7 +13410,7 @@ async def alerting_settings( if db_general_settings is not None and db_general_settings.param_value is not None: db_general_settings_dict = dict(db_general_settings.param_value) alerting_args_dict: dict = db_general_settings_dict.get("alerting_args", {}) # type: ignore - alerting_values: Optional[list] = db_general_settings_dict.get("alerting") # type: ignore + alerting_values: list | None = db_general_settings_dict.get("alerting") # type: ignore else: alerting_args_dict = {} alerting_values = None @@ -13500,7 +13451,7 @@ async def alerting_settings( for field_name, field_info in SlackAlertingArgs.model_fields.items(): if field_name in allowed_args: - _stored_in_db: Optional[bool] = None + _stored_in_db: bool | None = None if field_name in alerting_args_dict: _stored_in_db = True else: @@ -13529,7 +13480,7 @@ async def alerting_settings( async def async_queue_request( request: Request, fastapi_response: Response, - model: Optional[str] = None, + model: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): global general_settings, user_debug, proxy_logging_obj @@ -13619,7 +13570,7 @@ async def async_queue_request( ) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -13785,7 +13736,7 @@ async def login_v2(request: Request): json_response.set_cookie(key="token", value=jwt_token) return json_response except Exception as e: - verbose_proxy_logger.exception("litellm.proxy.proxy_server.login_v2(): Exception occurred - {}".format(str(e))) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.login_v2(): Exception occurred - {e!s}") if isinstance(e, ProxyException): raise e elif isinstance(e, HTTPException): @@ -13796,7 +13747,7 @@ async def login_v2(request: Request): code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=error_msg, type=ProxyErrorTypes.auth_error, @@ -13862,7 +13813,7 @@ async def login_v3(request: Request): status_code=status.HTTP_200_OK, ) except Exception as e: - verbose_proxy_logger.exception("litellm.proxy.proxy_server.login_v3(): Exception occurred - {}".format(str(e))) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.login_v3(): Exception occurred - {e!s}") if isinstance(e, ProxyException): raise e elif isinstance(e, HTTPException): @@ -13873,7 +13824,7 @@ async def login_v3(request: Request): code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=error_msg, type=ProxyErrorTypes.auth_error, @@ -13935,9 +13886,7 @@ async def login_v3_exchange(request: Request): except ProxyException: raise except Exception as e: - verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.login_v3_exchange(): Exception occurred - {}".format(str(e)) - ) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.login_v3_exchange(): Exception occurred - {e!s}") raise ProxyException( message=str(e), type=ProxyErrorTypes.auth_error, @@ -14032,7 +13981,7 @@ async def onboarding(invite_link: str, request: Request): algorithm="HS256", ) - litellm_dashboard_ui += "?token={}&user_email={}".format(jwt_token, user_email) + litellm_dashboard_ui += f"?token={jwt_token}&user_email={user_email}" return { "login_url": litellm_dashboard_ui, "token": jwt_token, @@ -14186,9 +14135,7 @@ async def claim_onboarding_link(data: InvitationClaim, request: Request): raise HTTPException( status_code=401, detail={ - "error": "Invalid invitation link. The user id submitted does not match the user id this link is attached to. Got={}, Expected={}".format( - data.user_id, invite_obj.user_id - ) + "error": f"Invalid invitation link. The user id submitted does not match the user id this link is attached to. Got={data.user_id}, Expected={invite_obj.user_id}" }, ) @@ -14430,10 +14377,7 @@ async def new_invitation(data: InvitationNew, user_api_key_dict: UserAPIKeyAuth raise HTTPException( status_code=400, detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) + "error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}" }, ) @@ -14491,12 +14435,7 @@ async def invitation_info(invitation_id: str, user_api_key_dict: UserAPIKeyAuth if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) response = await InvitationLinkRepository(prisma_client).table.find_unique(where={"id": invitation_id}) @@ -14543,7 +14482,7 @@ async def invitation_update( if user_api_key_dict.user_id is None: raise HTTPException( status_code=500, - detail={"error": "Unable to identify user id. Received={}".format(user_api_key_dict.user_id)}, + detail={"error": f"Unable to identify user id. Received={user_api_key_dict.user_id}"}, ) current_time = litellm.utils.get_utc_datetime() @@ -14608,12 +14547,7 @@ async def invitation_delete( if not is_proxy_admin and not is_other_admin: raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) # Org admins can only delete invitations they created @@ -14779,11 +14713,11 @@ async def update_config( return {"message": "Config updated successfully"} except Exception as e: - verbose_proxy_logger.error("litellm.proxy.proxy_server.update_config(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.update_config(): Exception occured - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -14882,7 +14816,7 @@ async def update_config_general_settings( if data.field_name not in ConfigGeneralSettings.model_fields: raise HTTPException( status_code=400, - detail={"error": "Invalid field={} passed in.".format(data.field_name)}, + detail={"error": f"Invalid field={data.field_name} passed in."}, ) try: @@ -14890,11 +14824,7 @@ async def update_config_general_settings( except Exception: raise HTTPException( status_code=400, - detail={ - "error": "Invalid type of field value={} passed in.".format( - type(data.field_value), - ) - }, + detail={"error": f"Invalid type of field value={type(data.field_value)} passed in."}, ) ## get general settings from db @@ -14971,7 +14901,7 @@ def _redact_secret_values_in_obj(value: JsonValue, depth: int = 0) -> JsonValue: return value -def _redact_config_param_value_for_logging(param_name: Optional[str], param_value: JsonValue) -> JsonValue: +def _redact_config_param_value_for_logging(param_name: str | None, param_value: JsonValue) -> JsonValue: if param_name == "environment_variables" and isinstance(param_value, dict): return {key: "REDACTED" for key in param_value} if isinstance(param_value, (dict, list)): @@ -14989,7 +14919,7 @@ def _redact_general_setting_value(field_name: str, value: JsonValue, is_full_adm return value -def _dump_redacted_config(value: Optional[JsonValue], *, redact_all_values: bool = False) -> Optional[str]: +def _dump_redacted_config(value: JsonValue | None, *, redact_all_values: bool = False) -> str | None: # `default=str` matches the sibling audit-log serializers in # team_endpoints.py and the LiteLLM_AuditLogs validator, so a YAML-loaded # value with a non-JSON-native leaf (datetime, custom object) cannot turn @@ -15004,8 +14934,8 @@ def _dump_redacted_config(value: Optional[JsonValue], *, redact_all_values: bool async def create_config_audit_log( param_name: str, action: AUDIT_ACTIONS, - before_value: Optional[JsonValue], - after_value: Optional[JsonValue], + before_value: JsonValue | None, + after_value: JsonValue | None, user_api_key_dict: UserAPIKeyAuth, table_name: LitellmTableNames = LitellmTableNames.CONFIG_TABLE_NAME, ) -> None: @@ -15041,7 +14971,7 @@ _EXTRA_SECRET_CALLBACK_ENV_VARS = frozenset( ) -def _redact_callback_env_vars(env_vars: dict[str, Optional[str]]) -> dict[str, Optional[str]]: +def _redact_callback_env_vars(env_vars: dict[str, str | None]) -> dict[str, str | None]: """Return a copy of ``env_vars`` with values for keys classified as sensitive by ``is_sensitive_callback_key`` replaced with ``"REDACTED"``. ``None`` values pass through unchanged. @@ -15108,7 +15038,7 @@ async def get_config_general_settings( if field_name not in ConfigGeneralSettings.model_fields: raise HTTPException( status_code=400, - detail={"error": "Invalid field={} passed in.".format(field_name)}, + detail={"error": f"Invalid field={field_name} passed in."}, ) ## get general settings from db @@ -15120,7 +15050,7 @@ async def get_config_general_settings( if db_general_settings is None or db_general_settings.param_value is None: raise HTTPException( status_code=400, - detail={"error": "Field name={} not in DB".format(field_name)}, + detail={"error": f"Field name={field_name} not in DB"}, ) else: general_settings = dict(db_general_settings.param_value) @@ -15140,7 +15070,7 @@ async def get_config_general_settings( else: raise HTTPException( status_code=400, - detail={"error": "Field name={} not in DB".format(field_name)}, + detail={"error": f"Field name={field_name} not in DB"}, ) @@ -15274,7 +15204,7 @@ async def _reset_general_settings_ui_litellm_field(field_name: str, user_api_key async def get_config_list( config_type: Literal["general_settings"], user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> List[ConfigList]: +) -> list[ConfigList]: """ List the available fields + current values for a given type of setting (currently just 'general_settings'user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),) """ @@ -15295,12 +15225,7 @@ async def get_config_list( if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) is_full_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN @@ -15354,7 +15279,7 @@ async def get_config_list( nested_fields = [ FieldDetail( field_name=sub_field, - field_type=sub_field_type.__name__, + field_type=getattr(sub_field_type, "__name__", str(sub_field_type)), field_description="", # Add custom logic if descriptions are available field_default_value=_redact_general_setting_value( sub_field, @@ -15431,7 +15356,7 @@ async def get_config_list( for litellm_field_name, spec in _GENERAL_SETTINGS_UI_LITELLM_FIELDS.items(): current_value: GeneralSettingsUILiteLLMValue = getattr(litellm, litellm_field_name, None) default_value = _general_settings_ui_litellm_default(spec) - stored_in_db_litellm: Optional[bool] + stored_in_db_litellm: bool | None if litellm_field_name in db_litellm_settings: stored_in_db_litellm = True elif current_value != default_value: @@ -15484,12 +15409,7 @@ async def delete_config_general_settings( if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) if data.field_name in _GENERAL_SETTINGS_UI_LITELLM_FIELDS: @@ -15498,7 +15418,7 @@ async def delete_config_general_settings( if data.field_name not in ConfigGeneralSettings.model_fields: raise HTTPException( status_code=400, - detail={"error": "Invalid field={} passed in.".format(data.field_name)}, + detail={"error": f"Invalid field={data.field_name} passed in."}, ) ## get general settings from db @@ -15510,7 +15430,7 @@ async def delete_config_general_settings( if db_general_settings is None or db_general_settings.param_value is None: raise HTTPException( status_code=400, - detail={"error": "Field name={} not in config".format(data.field_name)}, + detail={"error": f"Field name={data.field_name} not in config"}, ) else: general_settings = dict(db_general_settings.param_value) @@ -15565,12 +15485,7 @@ async def delete_callback( if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: raise HTTPException( status_code=400, - detail={ - "error": "{}, your role={}".format( - CommonProxyErrors.not_allowed_access.value, - user_api_key_dict.user_role, - ) - }, + detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) if store_model_in_db is not True: @@ -15626,7 +15541,7 @@ async def delete_callback( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"litellm.proxy.proxy_server.delete_callback(): Exception occurred - {str(e)}") + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.delete_callback(): Exception occurred - {e!s}") verbose_proxy_logger.debug(traceback.format_exc()) raise ProxyException( message="Error deleting callback: " + str(e), @@ -15750,10 +15665,10 @@ async def get_config( "available_callbacks": all_available_callbacks, } except Exception as e: - verbose_proxy_logger.exception("litellm.proxy.proxy_server.get_config(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"litellm.proxy.proxy_server.get_config(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"Authentication Error({str(e)})"), + message=getattr(e, "detail", f"Authentication Error({e!s})"), type=ProxyErrorTypes.auth_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -15868,8 +15783,8 @@ async def reload_model_cost_map( "timestamp": current_time.isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to reload model cost map: {str(e)}") - raise HTTPException(status_code=500, detail=f"Failed to reload model cost map: {str(e)}") + verbose_proxy_logger.exception(f"Failed to reload model cost map: {e!s}") + raise HTTPException(status_code=500, detail=f"Failed to reload model cost map: {e!s}") @router.post( @@ -15925,10 +15840,10 @@ async def schedule_model_cost_map_reload( "timestamp": datetime.utcnow().isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to schedule model cost map reload: {str(e)}") + verbose_proxy_logger.exception(f"Failed to schedule model cost map reload: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to schedule model cost map reload: {str(e)}", + detail=f"Failed to schedule model cost map reload: {e!s}", ) @@ -15970,8 +15885,8 @@ async def cancel_model_cost_map_reload( "timestamp": datetime.utcnow().isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to cancel model cost map reload: {str(e)}") - raise HTTPException(status_code=500, detail=f"Failed to cancel model cost map reload: {str(e)}") + verbose_proxy_logger.exception(f"Failed to cancel model cost map reload: {e!s}") + raise HTTPException(status_code=500, detail=f"Failed to cancel model cost map reload: {e!s}") @router.get( @@ -16057,10 +15972,10 @@ async def get_model_cost_map_reload_status( "next_run": next_run, } except Exception as e: - verbose_proxy_logger.exception(f"Failed to get model cost map reload status: {str(e)}") + verbose_proxy_logger.exception(f"Failed to get model cost map reload status: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to get model cost map reload status: {str(e)}", + detail=f"Failed to get model cost map reload status: {e!s}", ) @@ -16105,10 +16020,10 @@ async def get_model_cost_map_source( "model_count": model_count, } except Exception as e: - verbose_proxy_logger.exception(f"Failed to get model cost map source info: {str(e)}") + verbose_proxy_logger.exception(f"Failed to get model cost map source info: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to get model cost map source info: {str(e)}", + detail=f"Failed to get model cost map source info: {e!s}", ) @@ -16184,8 +16099,8 @@ async def reload_anthropic_beta_headers( "timestamp": current_time.isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to reload anthropic beta headers: {str(e)}") - raise HTTPException(status_code=500, detail=f"Failed to reload anthropic beta headers: {str(e)}") + verbose_proxy_logger.exception(f"Failed to reload anthropic beta headers: {e!s}") + raise HTTPException(status_code=500, detail=f"Failed to reload anthropic beta headers: {e!s}") @router.post( @@ -16241,10 +16156,10 @@ async def schedule_anthropic_beta_headers_reload( "timestamp": datetime.utcnow().isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to schedule anthropic beta headers reload: {str(e)}") + verbose_proxy_logger.exception(f"Failed to schedule anthropic beta headers reload: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to schedule anthropic beta headers reload: {str(e)}", + detail=f"Failed to schedule anthropic beta headers reload: {e!s}", ) @@ -16286,10 +16201,10 @@ async def cancel_anthropic_beta_headers_reload( "timestamp": datetime.utcnow().isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to cancel anthropic beta headers reload: {str(e)}") + verbose_proxy_logger.exception(f"Failed to cancel anthropic beta headers reload: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to cancel anthropic beta headers reload: {str(e)}", + detail=f"Failed to cancel anthropic beta headers reload: {e!s}", ) @@ -16378,10 +16293,10 @@ async def get_anthropic_beta_headers_reload_status( "next_run": next_run, } except Exception as e: - verbose_proxy_logger.exception(f"Failed to get anthropic beta headers reload status: {str(e)}") + verbose_proxy_logger.exception(f"Failed to get anthropic beta headers reload status: {e!s}") raise HTTPException( status_code=500, - detail=f"Failed to get anthropic beta headers reload status: {str(e)}", + detail=f"Failed to get anthropic beta headers reload status: {e!s}", ) @@ -16672,7 +16587,7 @@ async def _mcp_forward_as_path(path_segment: str, request: Request): return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) -async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: Optional[str]) -> List[str]: +async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: str | None) -> list[str]: """Validate a comma-separated ``/{name1,name2,...}/mcp`` segment. For each token, check (in order) whether it is a registered MCP server @@ -16700,7 +16615,7 @@ async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: Optional[str]) -> ) seen: set = set() - deduped: List[str] = [] + deduped: list[str] = [] for raw in csv_segment.split(","): token = raw.strip() if not token or token in seen: @@ -16710,7 +16625,7 @@ async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: Optional[str]) -> if len(deduped) >= DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS: break - resolved: List[str] = [] + resolved: list[str] = [] for token in deduped: if global_mcp_server_manager.get_mcp_server_by_name(token, client_ip=client_ip): resolved.append(token) diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 76d211d2c8f..6b8227ee94f 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -2,7 +2,7 @@ import json import os import re from importlib.resources import files -from typing import Any, Dict, List, Optional +from typing import Any from fastapi import APIRouter, HTTPException, Request @@ -39,7 +39,7 @@ router = APIRouter() # /public/endpoints — helpers # --------------------------------------------------------------------------- -_ENDPOINT_METADATA: Dict[str, Dict[str, str]] = { +_ENDPOINT_METADATA: dict[str, dict[str, str]] = { "chat_completions": {"label": "Chat Completions", "endpoint": "/chat/completions"}, "messages": {"label": "Messages", "endpoint": "/messages"}, "responses": {"label": "Responses", "endpoint": "/responses"}, @@ -101,33 +101,33 @@ _ENDPOINT_METADATA: Dict[str, Dict[str, str]] = { _SLUG_SUFFIX_RE = re.compile(r"\s*\(`[^`]+`\)\s*$") # Loaded once on first request; never invalidated (local file, no TTL needed). -_cached_endpoints: Optional[SupportedEndpointsResponse] = None +_cached_endpoints: SupportedEndpointsResponse | None = None def _clean_display_name(raw: str) -> str: return _SLUG_SUFFIX_RE.sub("", raw).strip() -def _build_endpoints(raw: Dict[str, Any]) -> List[Dict[str, Any]]: +def _build_endpoints(raw: dict[str, Any]) -> list[dict[str, Any]]: """Transform raw provider_endpoints_support_backup.json into the response shape.""" - providers: Dict[str, Any] = raw.get("providers", {}) + providers: dict[str, Any] = raw.get("providers", {}) # Collect endpoint keys in insertion order (union across all providers). seen: set = set() - all_keys: List[str] = [] + all_keys: list[str] = [] for provider_data in providers.values(): for key in provider_data.get("endpoints", {}): if key not in seen: seen.add(key) all_keys.append(key) - result: List[Dict[str, Any]] = [] + result: list[dict[str, Any]] = [] for key in all_keys: meta = _ENDPOINT_METADATA.get(key) label = meta["label"] if meta else key.replace("_", " ").title() path = meta["endpoint"] if meta else "/" + key.replace("_", "/") - supporting: List[Dict[str, str]] = [ + supporting: list[dict[str, str]] = [ { "slug": slug, "display_name": _clean_display_name(pd.get("display_name", slug)), @@ -140,7 +140,7 @@ def _build_endpoints(raw: Dict[str, Any]) -> List[Dict[str, Any]]: return result -def _load_endpoints() -> List[Dict[str, Any]]: +def _load_endpoints() -> list[dict[str, Any]]: raw = json.loads(files("litellm").joinpath("provider_endpoints_support_backup.json").read_text(encoding="utf-8")) return _build_endpoints(raw) @@ -151,7 +151,7 @@ def _load_endpoints() -> List[Dict[str, Any]]: @router.get( "/public/model_hub", tags=["public", "model management"], - response_model=List[ModelGroupInfoProxy], + response_model=list[ModelGroupInfoProxy], ) async def public_model_hub(): import litellm @@ -167,7 +167,7 @@ async def public_model_hub(): if llm_router is None: raise HTTPException(status_code=400, detail=CommonProxyErrors.no_llm_router.value) - model_groups: List[ModelGroupInfoProxy] = [] + model_groups: list[ModelGroupInfoProxy] = [] if litellm.public_model_groups is not None: model_groups = _get_model_group_info( llm_router=llm_router, @@ -203,7 +203,7 @@ async def public_model_hub(): @router.get( "/public/agent_hub", tags=["[beta] Agents", "public"], - response_model=List[AgentCard], + response_model=list[AgentCard], ) async def get_agents(request: Request): import litellm @@ -227,7 +227,7 @@ async def get_agents(request: Request): @router.get( "/public/mcp_hub", tags=["[beta] MCP", "public"], - response_model=List[MCPPublicServer], + response_model=list[MCPPublicServer], ) async def get_mcp_servers(): from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -314,9 +314,9 @@ async def public_model_hub_info(): @router.get( "/public/providers", tags=["public", "providers"], - response_model=List[str], + response_model=list[str], ) -async def get_supported_providers() -> List[str]: +async def get_supported_providers() -> list[str]: """ Return a sorted list of all providers supported by LiteLLM. """ @@ -327,9 +327,9 @@ async def get_supported_providers() -> List[str]: @router.get( "/public/providers/fields", tags=["public", "providers"], - response_model=List[ProviderCreateInfo], + response_model=list[ProviderCreateInfo], ) -async def get_provider_fields() -> List[ProviderCreateInfo]: +async def get_provider_fields() -> list[ProviderCreateInfo]: """ Return provider metadata required by the dashboard create-model flow. """ @@ -364,7 +364,7 @@ async def get_litellm_model_cost_map(): except Exception as e: raise HTTPException( status_code=500, - detail=f"Internal Server Error ({str(e)})", + detail=f"Internal Server Error ({e!s})", ) @@ -411,9 +411,9 @@ async def get_supported_endpoints() -> SupportedEndpointsResponse: @router.get( "/public/agents/fields", tags=["public", "[beta] Agents"], - response_model=List[AgentCreateInfo], + response_model=list[AgentCreateInfo], ) -async def get_agent_fields() -> List[AgentCreateInfo]: +async def get_agent_fields() -> list[AgentCreateInfo]: """ Return agent type metadata required by the dashboard create-agent flow. diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 27ffc49901b..a4ef30fb3c4 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -7,7 +7,7 @@ Provides: """ import base64 -from typing import Any, Dict, Optional, Tuple +from typing import Any import orjson from fastapi import APIRouter, Depends, HTTPException, Request, Response, status @@ -107,9 +107,9 @@ async def _authorize_nested_vector_store_ids( def _build_file_metadata_entry( response: Any, - file_data: Optional[Tuple[str, bytes, str]] = None, - file_url: Optional[str] = None, -) -> Dict[str, Any]: + file_data: tuple[str, bytes, str] | None = None, + file_url: str | None = None, +) -> dict[str, Any]: """ Build a file metadata entry for storing in vector_store_metadata. @@ -159,11 +159,11 @@ def _build_file_metadata_entry( async def _save_vector_store_to_db_from_rag_ingest( response: Any, - ingest_options: Dict[str, Any], + ingest_options: dict[str, Any], prisma_client, user_api_key_dict: UserAPIKeyAuth, - file_data: Optional[Tuple[str, bytes, str]] = None, - file_url: Optional[str] = None, + file_data: tuple[str, bytes, str] | None = None, + file_url: str | None = None, ) -> None: """ Helper function to save a newly created vector store from RAG ingest to the database. @@ -283,7 +283,7 @@ async def _save_vector_store_to_db_from_rag_ingest( async def parse_rag_ingest_request( request: Request, -) -> Tuple[Dict[str, Any], Optional[Tuple[str, bytes, str]], Optional[str], Optional[str]]: +) -> tuple[dict[str, Any], tuple[str, bytes, str] | None, str | None, str | None]: """ Parse RAG ingest request. @@ -300,7 +300,7 @@ async def parse_rag_ingest_request( file_data = None file_url = None file_id = None - ingest_options: Dict[str, Any] = {} + ingest_options: dict[str, Any] = {} if "multipart/form-data" in content_type: # Form upload @@ -485,7 +485,7 @@ async def rag_ingest( raise HTTPException(status_code=400, detail={"error": str(e)}) # Add litellm data - request_data: Dict[str, Any] = {} + request_data: dict[str, Any] = {} request_data = await add_litellm_data_to_request( data=request_data, request=request, @@ -653,7 +653,7 @@ async def rag_query( ) # Add litellm data - request_data: Dict[str, Any] = {} + request_data: dict[str, Any] = {} request_data = await add_litellm_data_to_request( data=request_data, request=request, diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 5a443deba83..62c947b58a0 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -2,7 +2,7 @@ import json import time -from typing import Any, Dict, Optional +from typing import Any import httpx from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -32,7 +32,7 @@ _DEFAULT_TRANSCRIPTION_MODEL = "gpt-realtime-whisper" _ALLOWED_SESSION_TYPES = ("realtime", "transcription") -def _coerce_realtime_session_type(session_type: Optional[str]) -> str: +def _coerce_realtime_session_type(session_type: str | None) -> str: if session_type in _ALLOWED_SESSION_TYPES: return session_type return "realtime" @@ -115,11 +115,11 @@ def _set_transcription_model_on_session( async def _prepare_client_secret_session( req: RealtimeClientSecretRequest, user_api_key_dict: UserAPIKeyAuth, - llm_model_list: Optional[list], + llm_model_list: list | None, llm_router: Any, -) -> tuple[str, Optional[dict], str]: +) -> tuple[str, dict | None, str]: session_type = _coerce_realtime_session_type(req.session.type if req.session else None) - session_data: Optional[dict] = req.session.model_dump(exclude_none=True) if req.session else None + session_data: dict | None = req.session.model_dump(exclude_none=True) if req.session else None if session_data is not None: session_data["type"] = session_type @@ -162,16 +162,16 @@ async def _prepare_client_secret_session( def _encode_realtime_token_payload( ephemeral_key: str, model_id: str, - user_id: Optional[str], - team_id: Optional[str], - expires_at: Optional[int], + user_id: str | None, + team_id: str | None, + expires_at: int | None, session_type: str = "realtime", ) -> str: """ Encode metadata with the upstream ephemeral key so /realtime/calls can route without requiring model as a query param. """ - payload: Dict[str, Any] = { + payload: dict[str, Any] = { "v": _REALTIME_TOKEN_VERSION, "ephemeral_key": ephemeral_key, "model_id": model_id, @@ -185,7 +185,7 @@ def _encode_realtime_token_payload( def _decode_realtime_token_payload( decrypted_value: str, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: """ Decode realtime token payload; returns None for legacy/raw ephemeral tokens. """ @@ -228,8 +228,8 @@ async def create_realtime_client_secret( from litellm.proxy.proxy_server import ( add_litellm_data_to_request, general_settings, - llm_router, llm_model_list, + llm_router, proxy_config, proxy_logging_obj, route_request, @@ -341,7 +341,7 @@ async def create_realtime_client_secret( encrypted_token: str = encrypt_value_helper(token_payload) upstream_json["value"] = encrypted_token - session_obj: Optional[dict] = upstream_json.get("session") + session_obj: dict | None = upstream_json.get("session") if isinstance(session_obj, dict): cs = session_obj.get("client_secret") if isinstance(cs, dict) and "value" in cs: @@ -380,7 +380,7 @@ async def proxy_realtime_calls( # Auth: the Bearer token is the encrypted ephemeral key issued by # /realtime/client_secrets, not a standard proxy API key. - auth_header: Optional[str] = request.headers.get("Authorization") + auth_header: str | None = request.headers.get("Authorization") if not auth_header or not auth_header.startswith("Bearer "): return Response( content=json.dumps({"error": "Missing or invalid Authorization header"}), @@ -540,8 +540,8 @@ async def create_realtime_transcription_session( from litellm.proxy.proxy_server import ( add_litellm_data_to_request, general_settings, - llm_router, llm_model_list, + llm_router, proxy_config, proxy_logging_obj, route_request, diff --git a/litellm/proxy/rerank_endpoints/endpoints.py b/litellm/proxy/rerank_endpoints/endpoints.py index a1b8d2a4821..69a5a9861d2 100644 --- a/litellm/proxy/rerank_endpoints/endpoints.py +++ b/litellm/proxy/rerank_endpoints/endpoints.py @@ -103,7 +103,7 @@ async def rerank( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.error("litellm.proxy.proxy_server.rerank(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.error(f"litellm.proxy.proxy_server.rerank(): Exception occured - {e!s}") if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e)), @@ -112,7 +112,7 @@ async def rerank( code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) else: - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index e594111f324..f17d546b88b 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -2,7 +2,7 @@ import asyncio import json import time from collections.abc import AsyncIterator -from typing import Any, Dict, Optional, cast +from typing import Any, cast from uuid import uuid4 import fastapi @@ -221,7 +221,7 @@ async def responses_api( ) managed_files_obj = cast( - Optional[_PROXY_LiteLLMManagedFiles], + _PROXY_LiteLLMManagedFiles | None, proxy_logging_obj.get_proxy_hook("managed_files"), ) @@ -251,7 +251,7 @@ async def responses_api( ) except Exception as e: verbose_proxy_logger.error( - f"Failed to store background response in managed objects table: {str(e)}" + f"Failed to store background response in managed objects table: {e!s}" ) return response @@ -922,7 +922,7 @@ async def cancel_response( async def _read_ws_model_from_first_frame( websocket: WebSocket, -) -> Optional[tuple]: +) -> tuple | None: """Read the first WS frame and return (model, raw_message), or None on error. Sends an appropriate error frame and closes the socket before returning None. @@ -990,7 +990,7 @@ async def _read_ws_model_from_first_frame( return model, first_message -def _extract_model_from_first_ws_event(first_event: Any) -> Optional[str]: +def _extract_model_from_first_ws_event(first_event: Any) -> str | None: """Extract model from a response.create WS event, handling flat and nested formats. Flat: {"type": "response.create", "model": "gpt-4o", ...} @@ -1006,7 +1006,7 @@ async def _enforce_responses_ws_first_frame_model_auth( request: Request, model: str, user_api_key_dict: UserAPIKeyAuth, - llm_router: Optional[Any], + llm_router: Any | None, ) -> None: from litellm.proxy.auth.user_api_key_auth import ( _enforce_key_and_fallback_model_access, @@ -1049,7 +1049,7 @@ async def _enforce_responses_ws_first_frame_model_auth( @router.websocket("/responses") async def responses_websocket_endpoint( websocket: WebSocket, - model: Optional[str] = fastapi.Query(None, description="The model to use for the responses WebSocket session."), + model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."), user_api_key_dict=Depends(user_api_key_auth_websocket), ): """ @@ -1088,14 +1088,14 @@ async def responses_websocket_endpoint( accept_kwargs["subprotocol"] = requested_protocols[0] await websocket.accept(**accept_kwargs) - first_message: Optional[str] = None + first_message: str | None = None if not model: result = await _read_ws_model_from_first_frame(websocket) if result is None: return model, first_message = result - data: Dict[str, Any] = { + data: dict[str, Any] = { "model": model, "websocket": websocket, } @@ -1104,7 +1104,7 @@ async def responses_websocket_endpoint( # Construct a synthetic Request for pre-call processing headers_list = list(websocket.scope.get("headers") or []) - scope: Dict[str, Any] = { + scope: dict[str, Any] = { "type": "http", "method": "POST", "path": "/v1/responses", diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 67d8c4021d8..84dcc5718e7 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -10,7 +10,7 @@ https://platform.openai.com/docs/api-reference/responses-streaming import asyncio import json -from typing import Any, Optional, cast +from typing import Any, cast from fastapi import Request, Response @@ -117,7 +117,7 @@ async def background_streaming_task( UPDATE_INTERVAL = 0.150 # 150ms batching interval # Track the terminal event from the stream (may not be "completed") - terminal_status: Optional[ResponsesAPIStatus] = ( + terminal_status: ResponsesAPIStatus | None = ( None # Will be set by response.completed/failed/incomplete/cancelled ) terminal_error = None @@ -294,7 +294,6 @@ async def background_streaming_task( except json.JSONDecodeError as e: verbose_proxy_logger.warning(f"Failed to parse streaming chunk: {e}") - pass # Final flush to ensure all accumulated state is saved await flush_state_if_needed(force=True) @@ -329,7 +328,7 @@ async def background_streaming_task( ) except Exception as e: - verbose_proxy_logger.error(f"Error in background streaming task for {polling_id}: {str(e)}") + verbose_proxy_logger.error(f"Error in background streaming task for {polling_id}: {e!s}") import traceback verbose_proxy_logger.error(traceback.format_exc()) diff --git a/litellm/proxy/response_polling/polling_handler.py b/litellm/proxy/response_polling/polling_handler.py index 44f3cfb32e4..4f2ad70cc7d 100644 --- a/litellm/proxy/response_polling/polling_handler.py +++ b/litellm/proxy/response_polling/polling_handler.py @@ -4,7 +4,7 @@ Response Polling Handler for Background Responses with Cache import json from datetime import datetime, timezone -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid4 @@ -18,7 +18,7 @@ class ResponsePollingHandler: CACHE_KEY_PREFIX = "litellm:polling:response:" POLLING_ID_PREFIX = "litellm_poll_" # Clear prefix to identify polling IDs - def __init__(self, redis_cache: Optional[RedisCache] = None, ttl: int = 3600): + def __init__(self, redis_cache: RedisCache | None = None, ttl: int = 3600): self.redis_cache = redis_cache self.ttl = ttl # Time-to-live for cache entries (default: 1 hour) @@ -40,7 +40,7 @@ class ResponsePollingHandler: async def create_initial_state( self, polling_id: str, - request_data: Dict[str, Any], + request_data: dict[str, Any], ) -> ResponsesAPIResponse: """ Create initial state in Redis for a polling request @@ -84,26 +84,26 @@ class ResponsePollingHandler: async def update_state( self, polling_id: str, - status: Optional[ResponsesAPIStatus] = None, - usage: Optional[Dict] = None, - error: Optional[Dict] = None, - incomplete_details: Optional[Dict] = None, - reasoning: Optional[Dict] = None, - tool_choice: Optional[Any] = None, - tools: Optional[list] = None, - output: Optional[list] = None, + status: ResponsesAPIStatus | None = None, + usage: dict | None = None, + error: dict | None = None, + incomplete_details: dict | None = None, + reasoning: dict | None = None, + tool_choice: Any | None = None, + tools: list | None = None, + output: list | None = None, # Additional ResponsesAPIResponse fields - model: Optional[str] = None, - instructions: Optional[str] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - max_output_tokens: Optional[int] = None, - previous_response_id: Optional[str] = None, - text: Optional[Dict] = None, - truncation: Optional[str] = None, - parallel_tool_calls: Optional[bool] = None, - user: Optional[str] = None, - store: Optional[bool] = None, + model: str | None = None, + instructions: str | None = None, + temperature: float | None = None, + top_p: float | None = None, + max_output_tokens: int | None = None, + previous_response_id: str | None = None, + text: dict | None = None, + truncation: str | None = None, + parallel_tool_calls: bool | None = None, + user: str | None = None, + store: bool | None = None, ) -> None: """ Update the polling state in Redis @@ -212,7 +212,7 @@ class ResponsePollingHandler: f"Updated polling state for {polling_id}: status={state['status']}, output_items={output_count}" ) - async def get_state(self, polling_id: str) -> Optional[Dict[str, Any]]: + async def get_state(self, polling_id: str) -> dict[str, Any] | None: """Get current polling state from Redis""" if not self.redis_cache: return None @@ -254,7 +254,7 @@ def should_use_polling_for_request( redis_cache, # RedisCache or None model: str, llm_router, # Router instance or None - native_background_mode: Optional[List[str]] = None, # List of models that should use native background mode + native_background_mode: list[str] | None = None, # List of models that should use native background mode ) -> bool: """ Determine if polling via cache should be used for a request. diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 11525b98ecc..5e775d1cbf7 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,6 +1,6 @@ import asyncio from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal import httpx from fastapi import HTTPException, status @@ -55,7 +55,7 @@ def _is_a2a_agent_model(model_name: Any) -> bool: return isinstance(model_name, str) and model_name.startswith("a2a/") -def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, team_id: Optional[str]) -> None: +def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, team_id: str | None) -> None: if not isinstance(model_name, str) or not model_name: return if not isinstance(llm_router, litellm.Router): @@ -206,12 +206,12 @@ def raise_if_mock_testing_params_disallowed(data: Mapping[str, object], *, allow def mock_testing_params_allowed() -> bool: """Read the opt-in from the running proxy's ``general_settings``.""" - import litellm.proxy.proxy_server as proxy_server + from litellm.proxy import proxy_server return proxy_server.general_settings.get(MOCK_TESTING_CONFIG_KEY, False) is True -def get_team_id_from_data(data: dict) -> Optional[str]: +def get_team_id_from_data(data: dict) -> str | None: """ Get the team id from the data's metadata or litellm_metadata params. """ @@ -226,7 +226,7 @@ def get_team_id_from_data(data: dict) -> Optional[str]: return None -_shared_session_lock: Optional[asyncio.Lock] = None +_shared_session_lock: asyncio.Lock | None = None def _get_shared_session_lock() -> asyncio.Lock: @@ -255,8 +255,8 @@ async def add_shared_session_to_data(data: dict) -> None: data: Dictionary to add the shared session to """ try: - import litellm.proxy.proxy_server as proxy_server from litellm._logging import verbose_proxy_logger + from litellm.proxy import proxy_server session = proxy_server.shared_aiohttp_session @@ -315,8 +315,8 @@ async def add_shared_session_to_data(data: dict) -> None: async def route_request( data: dict, - llm_router: Optional[LitellmRouter], - user_model: Optional[str], + llm_router: LitellmRouter | None, + user_model: str | None, route_type: Literal[ "acompletion", "atext_completion", @@ -414,7 +414,7 @@ async def route_request( "acancel_run", "adelete_run", ], - user_api_key_dict: Optional[UserAPIKeyAuth] = None, + user_api_key_dict: UserAPIKeyAuth | None = None, ): """ Common helper to route the request @@ -587,18 +587,18 @@ async def route_request( return getattr(llm_router, f"{route_type}")(**data) elif ( - is_proxy_admin_without_team - and data["model"] not in router_model_names - and data["model"] in llm_router.team_public_model_names + ( + is_proxy_admin_without_team + and data["model"] not in router_model_names + and data["model"] in llm_router.team_public_model_names + ) + or data["model"] in router_model_names + or llm_router.has_model_id(data["model"]) + or llm_router.model_group_alias is not None + and data["model"] in llm_router.model_group_alias ): return getattr(llm_router, f"{route_type}")(**data) - elif data["model"] in router_model_names or llm_router.has_model_id(data["model"]): - return getattr(llm_router, f"{route_type}")(**data) - - elif llm_router.model_group_alias is not None and data["model"] in llm_router.model_group_alias: - return getattr(llm_router, f"{route_type}")(**data) - elif data["model"] not in router_model_names: # Check wildcards before checking deployment_names # Priority: 1. Exact model_name match, 2. Wildcard match, 3. deployment_names match @@ -666,9 +666,7 @@ async def route_request( return result # Fall through to raise exception below if result is None - elif user_model is not None: - return getattr(litellm, f"{route_type}")(**data) - elif route_type == "allm_passthrough_route": + elif user_model is not None or route_type == "allm_passthrough_route": return getattr(litellm, f"{route_type}")(**data) # if no route found then it's a bad request diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 7b07d259f48..7c3a924b3b5 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -40,7 +40,7 @@ async def search( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - search_tool_name: Optional[str] = None, + search_tool_name: str | None = None, ): """ Search endpoint for performing web searches. @@ -170,7 +170,7 @@ async def search( team_object=team_object, ) except Exception as e: - verbose_proxy_logger.error(f"Search tool authorization failed for {search_tool_name_value}: {str(e)}") + verbose_proxy_logger.error(f"Search tool authorization failed for {search_tool_name_value}: {e!s}") raise if llm_router is not None and hasattr(llm_router, "search_tools"): diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 848623bead8..6b0bbacd131 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -3,7 +3,7 @@ CRUD ENDPOINTS FOR SEARCH TOOLS """ from datetime import datetime -from typing import Any, Dict, List, Optional, Union +from typing import Any from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel @@ -29,7 +29,7 @@ router = APIRouter() SEARCH_TOOL_REGISTRY = SearchToolRegistry() -def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, None]: +def _convert_datetime_to_str(value: datetime | str | None) -> str | None: """ Convert datetime object to ISO format string. @@ -47,9 +47,9 @@ def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, No async def _filter_visible_search_tools( - search_tools: List[SearchToolInfoResponse], + search_tools: list[SearchToolInfoResponse], user_api_key_dict: UserAPIKeyAuth, -) -> List[SearchToolInfoResponse]: +) -> list[SearchToolInfoResponse]: """ Drop search tools the caller is not authorized to invoke, applying the same key/team object_permission allowlists enforced on /search. Admins see all tools. @@ -70,7 +70,7 @@ async def _filter_visible_search_tools( user_api_key_cache, ) - team_object: Optional[LiteLLM_TeamTable] = None + team_object: LiteLLM_TeamTable | None = None if user_api_key_dict.team_id: team_object = await get_team_object( team_id=user_api_key_dict.team_id, @@ -80,7 +80,7 @@ async def _filter_visible_search_tools( proxy_logging_obj=proxy_logging_obj, ) - visible: List[SearchToolInfoResponse] = [] + visible: list[SearchToolInfoResponse] = [] for tool in search_tools: tool_name = tool.get("search_tool_name") if tool_name and await can_user_view_search_tool( @@ -151,7 +151,7 @@ async def list_search_tools( db_tool_names = {tool.get("search_tool_name") for tool in search_tools_from_db} - search_tool_configs: List[SearchToolInfoResponse] = [] + search_tool_configs: list[SearchToolInfoResponse] = [] config_search_tools = [] @@ -503,7 +503,7 @@ async def get_search_tool_info(search_tool_id: str): class TestSearchToolConnectionRequest(BaseModel): - litellm_params: Dict[str, Any] + litellm_params: dict[str, Any] @router.post( diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index 5934d686b5d..be4a588660c 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -3,7 +3,6 @@ Search Tool Registry for managing search tool configurations. """ from datetime import datetime, timezone -from typing import List, Optional from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -79,8 +78,8 @@ class SearchToolRegistry: return search_tool_dict except Exception as e: - verbose_proxy_logger.exception(f"Error adding search tool to DB: {str(e)}") - raise Exception(f"Error adding search tool to DB: {str(e)}") + verbose_proxy_logger.exception(f"Error adding search tool to DB: {e!s}") + raise Exception(f"Error adding search tool to DB: {e!s}") async def delete_search_tool_from_db(self, search_tool_id: str, prisma_client: PrismaClient): """ @@ -110,8 +109,8 @@ class SearchToolRegistry: "search_tool_name": existing_tool.search_tool_name, } except Exception as e: - verbose_proxy_logger.exception(f"Error deleting search tool from DB: {str(e)}") - raise Exception(f"Error deleting search tool from DB: {str(e)}") + verbose_proxy_logger.exception(f"Error deleting search tool from DB: {e!s}") + raise Exception(f"Error deleting search tool from DB: {e!s}") async def update_search_tool_in_db(self, search_tool_id: str, search_tool: SearchTool, prisma_client: PrismaClient): """ @@ -144,13 +143,13 @@ class SearchToolRegistry: # Convert to dict with ISO formatted datetimes return self._convert_prisma_to_dict(updated_search_tool) except Exception as e: - verbose_proxy_logger.exception(f"Error updating search tool in DB: {str(e)}") - raise Exception(f"Error updating search tool in DB: {str(e)}") + verbose_proxy_logger.exception(f"Error updating search tool in DB: {e!s}") + raise Exception(f"Error updating search tool in DB: {e!s}") @staticmethod async def get_all_search_tools_from_db( prisma_client: PrismaClient, - ) -> List[SearchTool]: + ) -> list[SearchTool]: """ Get all search tools from the database. @@ -169,7 +168,7 @@ class SearchToolRegistry: reason="get_all_search_tools_from_db_lookup_failure", ) - search_tools: List[SearchTool] = [] + search_tools: list[SearchTool] = [] for search_tool in search_tools_from_db: # Convert Prisma result to dict with ISO formatted datetimes search_tool_dict = SearchToolRegistry._convert_prisma_to_dict(search_tool) @@ -177,12 +176,12 @@ class SearchToolRegistry: return search_tools except Exception as e: - verbose_proxy_logger.exception(f"Error getting search tools from DB: {str(e)}") - raise Exception(f"Error getting search tools from DB: {str(e)}") + verbose_proxy_logger.exception(f"Error getting search tools from DB: {e!s}") + raise Exception(f"Error getting search tools from DB: {e!s}") async def get_search_tool_by_id_from_db( self, search_tool_id: str, prisma_client: PrismaClient - ) -> Optional[SearchTool]: + ) -> SearchTool | None: """ Get a search tool by its ID from the database. @@ -205,12 +204,12 @@ class SearchToolRegistry: search_tool_dict = self._convert_prisma_to_dict(search_tool) return SearchTool(**search_tool_dict) # type: ignore except Exception as e: - verbose_proxy_logger.exception(f"Error getting search tool from DB: {str(e)}") - raise Exception(f"Error getting search tool from DB: {str(e)}") + verbose_proxy_logger.exception(f"Error getting search tool from DB: {e!s}") + raise Exception(f"Error getting search tool from DB: {e!s}") async def get_search_tool_by_name_from_db( self, search_tool_name: str, prisma_client: PrismaClient - ) -> Optional[SearchTool]: + ) -> SearchTool | None: """ Get a search tool by its name from the database. @@ -233,5 +232,5 @@ class SearchToolRegistry: search_tool_dict = self._convert_prisma_to_dict(search_tool) return SearchTool(**search_tool_dict) # type: ignore except Exception as e: - verbose_proxy_logger.exception(f"Error getting search tool from DB: {str(e)}") - raise Exception(f"Error getting search tool from DB: {str(e)}") + verbose_proxy_logger.exception(f"Error getting search tool from DB: {e!s}") + raise Exception(f"Error getting search tool from DB: {e!s}") diff --git a/litellm/proxy/shutdown/graceful_shutdown_manager.py b/litellm/proxy/shutdown/graceful_shutdown_manager.py index 3b5e07d6d69..7449564fdf5 100644 --- a/litellm/proxy/shutdown/graceful_shutdown_manager.py +++ b/litellm/proxy/shutdown/graceful_shutdown_manager.py @@ -21,7 +21,6 @@ import asyncio import os import time from collections.abc import Callable -from typing import Optional from litellm._logging import verbose_proxy_logger from litellm.proxy.middleware.in_flight_requests_middleware import ( @@ -41,7 +40,7 @@ class GracefulShutdownManager: """ _is_shutting_down: bool = False - _shutdown_started_at: Optional[float] = None + _shutdown_started_at: float | None = None _drain_performed: bool = False @classmethod @@ -87,9 +86,9 @@ class GracefulShutdownManager: @classmethod async def wait_for_drain( cls, - timeout: Optional[float] = None, + timeout: float | None = None, exclude_self: bool = False, - count_fn: Optional[Callable[[], int]] = None, + count_fn: Callable[[], int] | None = None, poll_interval: float = _DRAIN_POLL_INTERVAL, log_interval: float = _DRAIN_LOG_INTERVAL, ) -> int: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 1623979ad8c..50ea349b4b7 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -5,7 +5,7 @@ import json from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import Any, Dict, List, NoReturn, Optional, cast +from typing import Any, NoReturn, cast from fastapi import HTTPException, status @@ -35,9 +35,9 @@ class _BudgetCounter: fallback_spend: float entity_type: str entity_id: str - source_cache_key: Optional[str] = None - spend_log_entity_id: Optional[str] = None - window_start: Optional[datetime] = None + source_cache_key: str | None = None + spend_log_entity_id: str | None = None + window_start: datetime | None = None _COUNTER_ENTITY_TYPES: Mapping[str, str] = { @@ -78,7 +78,7 @@ def _raise_reservation_unavailable(counter_key: str) -> NoReturn: ) -def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: +def get_reserved_counter_keys(budget_reservation: dict | None) -> set: if not budget_reservation: return set() entries = budget_reservation.get("entries") or [] @@ -87,7 +87,7 @@ def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: } -def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: Optional[UserAPIKeyAuth]) -> bool: +def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: UserAPIKeyAuth | None) -> bool: """ Whether an over-budget key's own ``max_budget`` reservation should be released rather than blocked, because the key opted into throttling: the @@ -103,7 +103,7 @@ def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: async def _apply_over_budget_reservation_policy( counter: _BudgetCounter, - valid_token: Optional[UserAPIKeyAuth], + valid_token: UserAPIKeyAuth | None, entry: dict[str, Any], applied_entries: list[dict[str, Any]], reservation_cost: float, @@ -147,17 +147,17 @@ async def _apply_over_budget_reservation_policy( async def reserve_budget_for_request( request_body: dict, route: str, - llm_router: Optional[Router], - valid_token: Optional[UserAPIKeyAuth], - team_object: Optional[LiteLLM_TeamTable], - user_object: Optional[LiteLLM_UserTable], - prisma_client: Optional[PrismaClient], + llm_router: Router | None, + valid_token: UserAPIKeyAuth | None, + team_object: LiteLLM_TeamTable | None, + user_object: LiteLLM_UserTable | None, + prisma_client: PrismaClient | None, user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, - end_user_id: Optional[str] = None, - end_user_object: Optional[Any] = None, + end_user_id: str | None = None, + end_user_object: Any | None = None, fail_closed_budget_enforcement: bool = False, -) -> Optional[dict]: +) -> dict | None: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): return None if route in {"/models", "/v1/models", "/utils/token_counter"}: @@ -179,7 +179,7 @@ async def reserve_budget_for_request( if not counters: return None - current_spend_by_counter_key: Dict[str, float] = {} + current_spend_by_counter_key: dict[str, float] = {} reservation_cost = estimate_request_max_cost( request_body=request_body, route=route, @@ -191,7 +191,7 @@ async def reserve_budget_for_request( if reservation_cost is None or reservation_cost <= 0: return None - applied_entries: List[Dict[str, Any]] = [] + applied_entries: list[dict[str, Any]] = [] try: for counter in counters: entry = _counter_to_reservation_entry( @@ -252,8 +252,8 @@ async def reserve_budget_for_request( async def reconcile_budget_reservation( - budget_reservation: Optional[dict], - actual_cost: Optional[float], + budget_reservation: dict | None, + actual_cost: float | None, finalize: bool = True, ) -> None: if not budget_reservation or budget_reservation.get("finalized") is True: @@ -270,7 +270,7 @@ async def reconcile_budget_reservation( budget_reservation["finalized"] = True -async def release_budget_reservation(budget_reservation: Optional[dict]) -> None: +async def release_budget_reservation(budget_reservation: dict | None) -> None: await reconcile_budget_reservation( budget_reservation=budget_reservation, actual_cost=0.0, @@ -311,7 +311,7 @@ async def release_budget_reservation_on_cancel( async def invalidate_budget_reservation_counters( - budget_reservation: Optional[dict], + budget_reservation: dict | None, ) -> None: if budget_reservation is None: return @@ -325,15 +325,15 @@ async def invalidate_budget_reservation_counters( async def _get_budget_counters( request_body: dict, valid_token: UserAPIKeyAuth, - team_object: Optional[LiteLLM_TeamTable], - user_object: Optional[LiteLLM_UserTable], - prisma_client: Optional[PrismaClient], + team_object: LiteLLM_TeamTable | None, + user_object: LiteLLM_UserTable | None, + prisma_client: PrismaClient | None, user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, - end_user_id: Optional[str] = None, - end_user_object: Optional[Any] = None, -) -> List[_BudgetCounter]: - counters: List[_BudgetCounter] = [] + end_user_id: str | None = None, + end_user_object: Any | None = None, +) -> list[_BudgetCounter]: + counters: list[_BudgetCounter] = [] if valid_token.token is not None: if valid_token.max_budget is not None and valid_token.max_budget > 0: @@ -437,9 +437,9 @@ async def _get_budget_counters( async def _get_end_user_budget_counter( valid_token: UserAPIKeyAuth, - end_user_id: Optional[str], - end_user_object: Optional[Any], -) -> Optional[_BudgetCounter]: + end_user_id: str | None, + end_user_object: Any | None, +) -> _BudgetCounter | None: end_user_id = end_user_id or valid_token.end_user_id if end_user_id is None: return None @@ -468,10 +468,10 @@ async def _get_end_user_budget_counter( async def _get_tag_budget_counters( request_body: dict, - prisma_client: Optional[PrismaClient], + prisma_client: PrismaClient | None, user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, -) -> List[_BudgetCounter]: +) -> list[_BudgetCounter]: from litellm.proxy.auth.auth_checks import get_tag_objects_batch from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body @@ -486,7 +486,7 @@ async def _get_tag_budget_counters( proxy_logging_obj=proxy_logging_obj, ) - counters: List[_BudgetCounter] = [] + counters: list[_BudgetCounter] = [] for tag_name in tag_names: tag_object = tag_objects.get(tag_name) if tag_object is None: @@ -508,7 +508,7 @@ async def _get_tag_budget_counters( return counters -def _dedupe_tags(tags: List[str]) -> List[str]: +def _dedupe_tags(tags: list[str]) -> list[str]: seen = set() deduped_tags = [] for tag in tags: @@ -521,22 +521,22 @@ def _dedupe_tags(tags: List[str]) -> List[str]: async def _get_team_member_budget_counter( valid_token: UserAPIKeyAuth, - team_object: Optional[LiteLLM_TeamTable], - user_object: Optional[LiteLLM_UserTable], + team_object: LiteLLM_TeamTable | None, + user_object: LiteLLM_UserTable | None, user_api_key_cache: DualCache, -) -> Optional[_BudgetCounter]: +) -> _BudgetCounter | None: if team_object is None or team_object.team_id is None or user_object is None or valid_token.user_id is None: return None membership_cache_key = f"team_membership:{valid_token.user_id}:{team_object.team_id}" cached_team_membership = await user_api_key_cache.async_get_cache(key=membership_cache_key) - team_membership: Optional[LiteLLM_TeamMembership] = None + team_membership: LiteLLM_TeamMembership | None = None if isinstance(cached_team_membership, LiteLLM_TeamMembership): team_membership = cached_team_membership elif isinstance(cached_team_membership, dict): team_membership = LiteLLM_TeamMembership(**cached_team_membership) - team_member_budget: Optional[float] = None + team_member_budget: float | None = None if team_membership is not None and team_membership.litellm_budget_table is not None: team_member_budget = team_membership.litellm_budget_table.max_budget else: @@ -563,10 +563,10 @@ async def _get_team_member_budget_counter( async def _get_org_budget_counter( valid_token: UserAPIKeyAuth, - team_object: Optional[LiteLLM_TeamTable], + team_object: LiteLLM_TeamTable | None, user_api_key_cache: DualCache, -) -> Optional[_BudgetCounter]: - org_id: Optional[str] = None +) -> _BudgetCounter | None: + org_id: str | None = None if valid_token.org_id is not None: org_id = valid_token.org_id elif team_object is not None and team_object.organization_id is not None: @@ -603,10 +603,10 @@ def _get_budget_limit_counters( entity_prefix: str, entity_type: str, entity_id: str, - budget_limits: Optional[Sequence[Any]], + budget_limits: Sequence[Any] | None, fallback_spend: float, -) -> List[_BudgetCounter]: - counters: List[_BudgetCounter] = [] +) -> list[_BudgetCounter]: + counters: list[_BudgetCounter] = [] if not budget_limits: return counters @@ -656,7 +656,7 @@ def _coerce_window(window: Any) -> dict: async def _reserve_counter( counter: _BudgetCounter, reservation_cost: float, -) -> Optional[float]: +) -> float | None: from litellm.proxy.proxy_server import ( _ensure_spend_counter_initialized, _ensure_window_spend_counter_initialized, @@ -725,7 +725,7 @@ async def _get_current_counter_value(counter: _BudgetCounter) -> float: async def _set_reserved_entries_actual_cost( - entries: List[dict], + entries: list[dict], actual_cost: float, default_reserved_cost: float, reseed_on_inconsistent: bool = True, @@ -806,7 +806,7 @@ async def _counter_can_apply_adjustment( async def _release_applied_entries_best_effort( - entries: List[dict], + entries: list[dict], default_reserved_cost: float, ) -> None: for entry in entries: @@ -832,7 +832,7 @@ async def _release_applied_entries_best_effort( async def _resize_applied_reservation( - entries: List[dict], + entries: list[dict], current_reserved_cost: float, new_reserved_cost: float, ) -> None: @@ -850,7 +850,7 @@ async def _resize_applied_reservation( def _counter_to_reservation_entry( counter: _BudgetCounter, reserved_cost: float, -) -> Dict[str, Any]: +) -> dict[str, Any]: return { "counter_key": counter.counter_key, "entity_type": counter.entity_type, @@ -867,7 +867,7 @@ def _get_entry_reserved_cost(entry: dict, default_reserved_cost: float) -> float return default_reserved_cost -def get_budget_window_start(window: Any) -> Optional[datetime]: +def get_budget_window_start(window: Any) -> datetime | None: window_dict = _coerce_window(window) budget_duration = window_dict.get("budget_duration") if budget_duration is None: @@ -885,7 +885,7 @@ def get_budget_window_start(window: Any) -> Optional[datetime]: return reset_at - timedelta(seconds=duration_seconds) -def _coerce_datetime(value: Any) -> Optional[datetime]: +def _coerce_datetime(value: Any) -> datetime | None: if value is None: return None if isinstance(value, datetime): @@ -901,8 +901,8 @@ def _coerce_datetime(value: Any) -> Optional[datetime]: def estimate_request_max_cost( request_body: dict, route: str, - llm_router: Optional[Router], -) -> Optional[float]: + llm_router: Router | None, +) -> float | None: model = get_model_from_request(request_body, route, llm_router=llm_router) if model is None: return None @@ -920,7 +920,7 @@ def estimate_request_max_cost( estimates = [estimate for estimate in estimates if estimate is not None] if not estimates: return None - return max(cast(List[float], estimates)) + return max(cast(list[float], estimates)) def estimate_request_input_cost( @@ -978,8 +978,8 @@ def _input_cost_for_cost_info( request_body: dict, route: str, model: str, - model_info: Dict[str, Any], -) -> Optional[float]: + model_info: dict[str, Any], +) -> float | None: input_tokens = _estimate_input_tokens( request_body=request_body, route=route, @@ -1003,8 +1003,8 @@ def _estimate_request_max_cost_for_model( request_body: dict, route: str, model: str, - llm_router: Optional[Router], -) -> Optional[float]: + llm_router: Router | None, +) -> float | None: estimates = [ _max_cost_for_cost_info( request_body=request_body, @@ -1022,8 +1022,8 @@ def _max_cost_for_cost_info( request_body: dict, route: str, model: str, - model_info: Dict[str, Any], -) -> Optional[float]: + model_info: dict[str, Any], +) -> float | None: image_cost = _estimate_image_generation_cost( request_body=request_body, model_info=model_info, @@ -1081,8 +1081,8 @@ def _max_cost_for_cost_info( def _estimate_image_generation_cost( request_body: dict, - model_info: Dict[str, Any], -) -> Optional[float]: + model_info: dict[str, Any], +) -> float | None: """ Reserve `n × per-image cost` for image-generation requests so concurrent requests against a depleted budget cannot all slip past the admission gate @@ -1119,8 +1119,8 @@ def _estimate_image_generation_cost( def _get_model_cost_info( model: str, - llm_router: Optional[Router], -) -> Optional[Dict[str, Any]]: + llm_router: Router | None, +) -> dict[str, Any] | None: if llm_router is not None: model_group_info = llm_router.get_model_group_info(model_group=model) if model_group_info is not None: @@ -1130,8 +1130,8 @@ def _get_model_cost_info( def _get_model_cost_infos( model: str, - llm_router: Optional[Router], -) -> List[Dict[str, Any]]: + llm_router: Router | None, +) -> list[dict[str, Any]]: """Cost-info candidates to estimate a request against for one model group. Reservation runs before routing, so the deployment that will serve the request @@ -1157,9 +1157,9 @@ def _get_model_cost_infos( def _deployment_tiered_pricing_table( - deployment: Dict[str, Any], + deployment: dict[str, Any], llm_router: Router, -) -> Optional[List[dict]]: +) -> list[dict] | None: model_id = deployment.get("model_info", {}).get("id") backend_model = deployment.get("litellm_params", {}).get("model") if not isinstance(model_id, str) or not isinstance(backend_model, str): @@ -1175,8 +1175,8 @@ def _deployment_tiered_pricing_table( def _get_deployment_tiered_pricing_tables( model: str, - llm_router: Optional[Router], -) -> List[List[dict]]: + llm_router: Router | None, +) -> list[list[dict]]: if llm_router is None: return [] deployments = llm_router.get_model_list(model_name=model) or [] @@ -1191,8 +1191,8 @@ def _estimate_input_tokens( request_body: dict, route: str, model: str, - model_info: Dict[str, Any], -) -> Optional[int]: + model_info: dict[str, Any], +) -> int | None: try: if "messages" in request_body: return litellm.token_counter( @@ -1228,12 +1228,12 @@ DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK = 16384 def _estimate_output_tokens( request_body: dict, route: str, - model_info: Dict[str, Any], -) -> Optional[int]: + model_info: dict[str, Any], +) -> int | None: if _is_input_only_route(route=route): return 0 - requested: Optional[int] = None + requested: int | None = None for key in ("max_completion_tokens", "max_tokens", "max_output_tokens"): requested = _to_int(request_body.get(key)) if requested is not None: @@ -1293,7 +1293,7 @@ def _is_input_only_route(route: str) -> bool: ) -def _to_float(value: Any) -> Optional[float]: +def _to_float(value: Any) -> float | None: if value is None: return None try: @@ -1302,7 +1302,7 @@ def _to_float(value: Any) -> Optional[float]: return None -def _to_int(value: Any) -> Optional[int]: +def _to_int(value: Any) -> int | None: if value is None: return None try: diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 827814a9716..37a53d06b0b 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -161,10 +161,10 @@ async def get_cloudzero_settings( # Re-raise HTTPExceptions as-is raise e except Exception as e: - verbose_proxy_logger.error(f"Error retrieving CloudZero settings: {str(e)}") + verbose_proxy_logger.error(f"Error retrieving CloudZero settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to retrieve CloudZero settings: {str(e)}"}, + detail={"error": f"Failed to retrieve CloudZero settings: {e!s}"}, ) @@ -238,10 +238,10 @@ async def update_cloudzero_settings( ) raise e except Exception as e: - verbose_proxy_logger.error(f"Error updating CloudZero settings: {str(e)}") + verbose_proxy_logger.error(f"Error updating CloudZero settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to update CloudZero settings: {str(e)}"}, + detail={"error": f"Failed to update CloudZero settings: {e!s}"}, ) @@ -275,7 +275,7 @@ async def is_cloudzero_setup_in_db() -> bool: return cloudzero_config is not None and cloudzero_config.param_value is not None except Exception as e: - verbose_proxy_logger.error(f"Error checking CloudZero status: {str(e)}") + verbose_proxy_logger.error(f"Error checking CloudZero status: {e!s}") return False @@ -317,7 +317,7 @@ async def is_cloudzero_setup() -> bool: return False except Exception as e: - verbose_proxy_logger.error(f"Error checking CloudZero setup: {str(e)}") + verbose_proxy_logger.error(f"Error checking CloudZero setup: {e!s}") return False @@ -364,10 +364,10 @@ async def init_cloudzero_settings( return CloudZeroInitResponse(message="CloudZero settings initialized successfully", status="success") except Exception as e: - verbose_proxy_logger.error(f"Error initializing CloudZero settings: {str(e)}") + verbose_proxy_logger.error(f"Error initializing CloudZero settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to initialize CloudZero settings: {str(e)}"}, + detail={"error": f"Failed to initialize CloudZero settings: {e!s}"}, ) @@ -422,10 +422,10 @@ async def cloudzero_dry_run_export( ) except Exception as e: - verbose_proxy_logger.error(f"Error performing CloudZero dry run export: {str(e)}") + verbose_proxy_logger.error(f"Error performing CloudZero dry run export: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to perform CloudZero dry run export: {str(e)}"}, + detail={"error": f"Failed to perform CloudZero dry run export: {e!s}"}, ) @@ -487,10 +487,10 @@ async def cloudzero_export( ) except Exception as e: - verbose_proxy_logger.error(f"Error performing CloudZero export: {str(e)}") + verbose_proxy_logger.error(f"Error performing CloudZero export: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to perform CloudZero export: {str(e)}"}, + detail={"error": f"Failed to perform CloudZero export: {e!s}"}, ) @@ -550,8 +550,8 @@ async def delete_cloudzero_settings( except HTTPException as e: raise e except Exception as e: - verbose_proxy_logger.error(f"Error deleting CloudZero settings: {str(e)}") + verbose_proxy_logger.error(f"Error deleting CloudZero settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to delete CloudZero settings: {str(e)}"}, + detail={"error": f"Failed to delete CloudZero settings: {e!s}"}, ) diff --git a/litellm/proxy/spend_tracking/spend_log_error_logger.py b/litellm/proxy/spend_tracking/spend_log_error_logger.py index e2987481c8b..41d908d2c26 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, Optional +from typing import Any from litellm._logging import verbose_proxy_logger from litellm.secret_managers.main import str_to_bool @@ -59,7 +59,7 @@ def should_suppress_spend_log_tracebacks() -> bool: def spend_log_error( message: str, *args: Any, - exc: Optional[BaseException] = None, + exc: BaseException | None = None, ) -> None: """Log a spend-tracking error, with the traceback gated on the env var. diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index cbaca95d688..d2c3b0d9391 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -7,14 +7,11 @@ from datetime import datetime, timedelta, timezone from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, NamedTuple, Protocol, TypedDict, TypeVar, - Union, ) import fastapi @@ -391,7 +388,7 @@ async def spend_user_fn( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def view_spend_tags( @@ -443,7 +440,7 @@ async def view_spend_tags( except Exception as e: if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"/spend/tags Error({str(e)})"), + message=getattr(e, "detail", f"/spend/tags Error({e!s})"), type="internal_error", param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -493,7 +490,7 @@ async def get_global_activity_internal_user( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, include_in_schema=False, ) @@ -635,7 +632,7 @@ async def get_global_activity_model_internal_user( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, include_in_schema=False, ) @@ -787,7 +784,7 @@ async def get_global_activity_model( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, include_in_schema=False, ) @@ -934,7 +931,7 @@ async def get_global_activity_exceptions_per_deployment( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, include_in_schema=False, ) @@ -1043,7 +1040,7 @@ async def get_global_activity_exceptions( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def get_global_spend_provider( @@ -1171,7 +1168,7 @@ async def get_global_spend_provider( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def get_global_spend_report( @@ -1464,7 +1461,7 @@ async def get_global_spend_report( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def global_get_all_tag_names(): @@ -1495,7 +1492,7 @@ async def global_get_all_tag_names(): except Exception as e: if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"/spend/all_tag_names Error({str(e)})"), + message=getattr(e, "detail", f"/spend/all_tag_names Error({e!s})"), type="internal_error", param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -1515,7 +1512,7 @@ async def global_get_all_tag_names(): tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def global_view_spend_tags( @@ -1651,7 +1648,7 @@ async def _get_spend_report_for_time_range( return response, spend_per_tag except Exception as e: - verbose_proxy_logger.error("Exception in _get_daily_spend_reports {}".format(str(e))) + verbose_proxy_logger.error(f"Exception in _get_daily_spend_reports {e!s}") @router.post( @@ -1801,7 +1798,7 @@ async def calculate_spend(request: SpendCalculateRequest): param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), ) - error_msg = f"{str(e)}" + error_msg = f"{e!s}" raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), @@ -1815,7 +1812,7 @@ async def calculate_spend(request: SpendCalculateRequest): tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": Dict[str, Any]}, + 200: {"model": dict[str, Any]}, }, ) @router.get( @@ -1824,7 +1821,7 @@ async def calculate_spend(request: SpendCalculateRequest): dependencies=[Depends(user_api_key_auth)], include_in_schema=False, responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def ui_view_spend_logs( @@ -2068,7 +2065,7 @@ async def ui_view_spend_logs( user_api_key_dict=user_api_key_dict, request_id=request_id, ) - permitted_team_ids: List[str] | None = None + permitted_team_ids: list[str] | None = None if not is_request_id_lookup and not is_admin_view: if team_id is not None: can_view_team = await _can_team_member_view_log( @@ -2079,7 +2076,7 @@ async def ui_view_spend_logs( if not can_view_team: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Not authorized to view team spend for team_id={}".format(team_id)}, + detail={"error": f"Not authorized to view team spend for team_id={team_id}"}, ) where_conditions["team_id"] = team_id where_conditions.pop("user", None) @@ -2111,7 +2108,7 @@ async def ui_view_spend_logs( # Build raw SQL to fetch paginated data WITHOUT heavy columns # (messages, response, proxy_server_request can be hundreds of KB per row). # These are only needed in the detail endpoint /spend/logs/ui/{request_id}. - sql_conditions: List[str] = [] + sql_conditions: list[str] = [] sql_params: list[object] = [] p = 1 # parameter index counter @@ -2269,15 +2266,15 @@ async def ui_view_spend_logs( class RequestResponsePayload(NamedTuple): - messages: Union[str, list, dict] | None - response: Union[str, list, dict] | None - proxy_server_request: Union[str, dict] | None + messages: str | list | dict | None + response: str | list | dict | None + proxy_server_request: str | dict | None _EMPTY_SPEND_LOG_VALUES = frozenset({"", "{}", "[]", "null"}) -def _spend_log_field_has_content(value: Union[str, list, dict] | None) -> bool: +def _spend_log_field_has_content(value: str | list | dict | None) -> bool: if value is None: return False if isinstance(value, str): @@ -2309,7 +2306,7 @@ def _hydrate_spend_log_metadata(rows: Sequence[Mapping[str, object]]) -> None: def _cold_storage_object_key_from_metadata( - metadata: Union[str, dict] | None, + metadata: str | dict | None, ) -> str | None: if isinstance(metadata, str): try: @@ -2462,7 +2459,7 @@ async def ui_view_request_response_for_request_id( tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def view_spend_logs( @@ -2670,7 +2667,7 @@ async def view_spend_logs( except Exception as e: if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "detail", f"/spend/logs Error({str(e)})"), + message=getattr(e, "detail", f"/spend/logs Error({e!s})"), type="internal_error", param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -2792,7 +2789,7 @@ async def global_spend_refresh(): } except Exception as e: - verbose_proxy_logger.exception("Failed to refresh materialized view - {}".format(str(e))) + verbose_proxy_logger.exception(f"Failed to refresh materialized view - {e!s}") return { "message": "Failed to refresh materialized view", "status": "failure", @@ -2833,7 +2830,7 @@ async def global_spend_for_internal_user( return response except Exception as e: - verbose_proxy_logger.error(f"/global/spend/logs Error: {str(e)}") + verbose_proxy_logger.error(f"/global/spend/logs Error: {e!s}") raise e @@ -3377,7 +3374,7 @@ async def provider_budgets() -> ProviderBudgetResponse: if router_budget_logger is None: raise ValueError("No router budget logger found") - provider_budget_response_dict: Dict[str, ProviderBudgetResponseObject] = {} + provider_budget_response_dict: dict[str, ProviderBudgetResponseObject] = {} for _provider, _budget_info in provider_budget_config.items(): _provider_spend = await router_budget_logger._get_current_provider_spend(_provider) or 0.0 _provider_budget_ttl = await router_budget_logger._get_current_provider_budget_reset_at(_provider) @@ -3390,7 +3387,7 @@ async def provider_budgets() -> ProviderBudgetResponse: provider_budget_response_dict[_provider] = provider_budget_response_object return ProviderBudgetResponse(providers=provider_budget_response_dict) except Exception as e: - verbose_proxy_logger.exception("/provider/budgets: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception(f"/provider/budgets: Exception occured - {e!s}") raise handle_exception_on_proxy(e) @@ -3425,7 +3422,7 @@ async def ui_get_spend_by_tags( # tags_str is a list of strings csv of tags # tags_str = tag1,tag2,tag3 # convert to list if it's not None - tags_list: List[str] | None = None + tags_list: list[str] | None = None if tags_str is not None and len(tags_str) > 0: tags_list = tags_str.split(",") @@ -3508,7 +3505,7 @@ async def ui_get_spend_by_tags( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, + 200: {"model": list[LiteLLM_SpendLogs]}, }, ) async def ui_view_session_spend_logs( @@ -3682,7 +3679,7 @@ async def _build_ui_spend_logs_response( counts = await _count_logs_per_session(prisma_client, session_ids) count_map = {r["session_id"]: r["_count"]["session_id"] for r in counts if r.get("session_id")} - session_spend_map: dict[str, dict[str, Union[int, float]]] = {} + session_spend_map: dict[str, dict[str, int | float]] = {} if enrich_session_counts and session_ids: from prisma.errors import PrismaError @@ -3732,7 +3729,7 @@ async def _build_ui_spend_logs_response( ) if enrich_session_counts: - enriched: List[dict] = [] + enriched: list[dict] = [] for row in data: row_dict = dict(row) if isinstance(row, dict) else row.model_dump() sid = row_dict.get("session_id") @@ -3872,14 +3869,14 @@ async def _assert_user_can_view_request_id( raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Not authorized to view spend log for request_id={}".format(request_id)}, + detail={"error": f"Not authorized to view spend log for request_id={request_id}"}, ) async def _get_permitted_team_ids_for_spend_logs( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, -) -> List[str]: +) -> list[str]: """ Return team IDs where the user is either a team admin or has the ``/spend/logs`` permission, allowing them to view team-wide spend logs. @@ -3904,12 +3901,10 @@ async def _get_permitted_team_ids_for_spend_logs( team_rows = await _find_team_rows(prisma_client, user_obj.teams) - permitted: List[str] = [] + permitted: list[str] = [] for team_row in team_rows: team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) - if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): - permitted.append(team_obj.team_id) - elif _team_member_has_permission( + if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) or _team_member_has_permission( user_api_key_dict=user_api_key_dict, team_obj=team_obj, permission=KeyManagementRoutes.SPEND_LOGS.value, diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index a6a67d57582..0dca38b47f2 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -3,10 +3,9 @@ import json import os import re import secrets -from datetime import datetime +from datetime import datetime, timezone from datetime import datetime as dt -from datetime import timezone -from typing import Any, List, Literal, Optional, cast +from typing import Any, Literal, cast from pydantic import BaseModel @@ -15,11 +14,11 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import ( LITELLM_TRUNCATED_PAYLOAD_FIELD, LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE, + REDACTED_BY_LITELM_STRING, ) from litellm.constants import ( MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB, ) -from litellm.constants import REDACTED_BY_LITELM_STRING from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, reconstruct_model_name, @@ -62,7 +61,7 @@ def _hash_api_key_for_spend_log(api_key: str) -> str: return stripped -def _is_master_key(api_key: Optional[str], _master_key: Optional[str]) -> bool: +def _is_master_key(api_key: str | None, _master_key: str | None) -> bool: """ Raw-only constant-time master-key comparison. The hashed form is never considered equivalent — only the raw master-key string matches. @@ -73,18 +72,18 @@ def _is_master_key(api_key: Optional[str], _master_key: Optional[str]) -> bool: def _get_spend_logs_metadata( - metadata: Optional[dict], - applied_guardrails: Optional[List[str]] = None, - batch_models: Optional[List[str]] = None, - mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] = None, - vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]] = None, - guardrail_information: Optional[List[StandardLoggingGuardrailInformation]] = None, - usage_object: Optional[dict] = None, - model_map_information: Optional[StandardLoggingModelInformation] = None, - cold_storage_object_key: Optional[str] = None, - litellm_overhead_time_ms: Optional[float] = None, - cost_breakdown: Optional[CostBreakdown] = None, - litellm_call_id: Optional[str] = None, + metadata: dict | None, + applied_guardrails: list[str] | None = None, + batch_models: list[str] | None = None, + mcp_tool_call_metadata: StandardLoggingMCPToolCall | None = None, + vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None = None, + guardrail_information: list[StandardLoggingGuardrailInformation] | None = None, + usage_object: dict | None = None, + model_map_information: StandardLoggingModelInformation | None = None, + cold_storage_object_key: str | None = None, + litellm_overhead_time_ms: float | None = None, + cost_breakdown: CostBreakdown | None = None, + litellm_call_id: str | None = None, ) -> SpendLogsMetadata: if metadata is None: return SpendLogsMetadata( @@ -172,12 +171,12 @@ def generate_hash_from_response(response_obj: Any) -> str: return hashlib.md5(str(response_obj).encode()).hexdigest() -def get_spend_logs_id(call_type: str, response_obj: dict, kwargs: dict) -> Optional[str]: +def get_spend_logs_id(call_type: str, response_obj: dict, kwargs: dict) -> str | None: if call_type == "aretrieve_batch" or call_type == "acreate_file": # Generate a hash from the response object - id: Optional[str] = generate_hash_from_response(response_obj) + id: str | None = generate_hash_from_response(response_obj) else: - id = cast(Optional[str], response_obj.get("id")) or cast(Optional[str], kwargs.get("litellm_call_id")) + id = cast(str | None, response_obj.get("id")) or cast(str | None, kwargs.get("litellm_call_id")) return id @@ -279,7 +278,7 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs usage = _combined_usage.model_dump() id = get_spend_logs_id(call_type or "acompletion", response_obj_dict, kwargs) - standard_logging_payload = cast(Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None)) + standard_logging_payload = cast(StandardLoggingPayload | None, kwargs.get("standard_logging_object", None)) end_user_id = get_end_user_id_for_cost_tracking(litellm_params) @@ -363,7 +362,7 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs standard_logging_payload.get("cost_breakdown", None) if standard_logging_payload is not None else None ), litellm_call_id=cast( - Optional[str], + str | None, kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"), ), ) @@ -403,12 +402,12 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs id = f"{id}_cache_hit{time.time()}" # SpendLogs does not allow duplicate request_id mcp_namespaced_tool_name = None - mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] = clean_metadata.get("mcp_tool_call_metadata") + mcp_tool_call_metadata: StandardLoggingMCPToolCall | None = clean_metadata.get("mcp_tool_call_metadata") if mcp_tool_call_metadata is not None: mcp_namespaced_tool_name = mcp_tool_call_metadata.get("namespaced_tool_name", None) # Extract agent_id for A2A requests (set directly on model_call_details) - agent_id: Optional[str] = kwargs.get("agent_id") or metadata.get("agent_id") + agent_id: str | None = kwargs.get("agent_id") or metadata.get("agent_id") custom_llm_provider = kwargs.get("custom_llm_provider") raw_model = cast(str, kwargs.get("model") or "") model_name = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {}) @@ -476,7 +475,7 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs def _get_session_id_for_spend_log( kwargs: dict, - standard_logging_payload: Optional[StandardLoggingPayload], + standard_logging_payload: StandardLoggingPayload | None, ) -> str: """ Get the session id for the spend log. @@ -497,7 +496,7 @@ def _get_session_id_for_spend_log( return str(uuid.uuid4()) -def _get_request_duration_ms(start_time: datetime, end_time: datetime) -> Optional[int]: +def _get_request_duration_ms(start_time: datetime, end_time: datetime) -> int | None: """Compute request duration in milliseconds from start and end times.""" try: return int((end_time - start_time).total_seconds() * 1000) @@ -514,7 +513,7 @@ def _ensure_datetime_utc(timestamp: datetime) -> datetime: async def get_spend_by_team( start_date: dt, end_date: dt, - team_id: Optional[str], + team_id: str | None, prisma_client: PrismaClient, ): sql_query = """ @@ -655,8 +654,8 @@ async def get_spend_by_team_and_customer( def _get_messages_for_spend_logs_payload( - standard_logging_payload: Optional[StandardLoggingPayload], - metadata: Optional[dict] = None, + standard_logging_payload: StandardLoggingPayload | None, + metadata: dict | None = None, ) -> str: if _should_store_prompts_and_responses_in_spend_logs(): if standard_logging_payload is not None: @@ -676,8 +675,8 @@ _SENSITIVE_REQUEST_BODY_KEYS = frozenset({"secret_fields"}) def _sanitize_request_body_for_spend_logs_payload( request_body: dict, - visited: Optional[set] = None, - max_string_length_prompt_in_db: Optional[int] = None, + visited: set | None = None, + max_string_length_prompt_in_db: int | None = None, ) -> dict: """ Recursively sanitize request body to prevent logging large base64 strings or other large values. @@ -856,7 +855,7 @@ def _redact_prompt_leaks_in_error_string(text: str) -> str: if not text: return text redaction = f'"{REDACTED_BY_LITELM_STRING}"' - out: List[str] = [] + out: list[str] = [] n = len(text) pos = 0 while pos < n: @@ -887,8 +886,8 @@ def _redact_prompt_leaks_in_error_string(text: str) -> str: def _sanitize_guardrail_information_for_spend_logs( - guardrail_information: Optional[List[StandardLoggingGuardrailInformation]], -) -> Optional[List[StandardLoggingGuardrailInformation]]: + guardrail_information: list[StandardLoggingGuardrailInformation] | None, +) -> list[StandardLoggingGuardrailInformation] | None: """ When ``store_prompts_in_spend_logs`` is False, redact prompt-carrying fields (``guardrail_request``, ``guardrail_response``, ``match_details``, @@ -963,8 +962,8 @@ def _redact_prompt_fields_in_guardrail_entry( def _sanitize_error_information_for_spend_logs( - error_information: Optional[StandardLoggingPayloadErrorInformation], -) -> Optional[StandardLoggingPayloadErrorInformation]: + error_information: StandardLoggingPayloadErrorInformation | None, +) -> StandardLoggingPayloadErrorInformation | None: """ Sanitize ``error_information`` before it lands in ``LiteLLM_SpendLogs.metadata``. @@ -997,7 +996,7 @@ def _sanitize_error_information_for_spend_logs( return cast(StandardLoggingPayloadErrorInformation, sanitized) -def _convert_to_json_serializable_dict(obj: Any, visited: Optional[set] = None, max_depth: int = 20) -> Any: +def _convert_to_json_serializable_dict(obj: Any, visited: set | None = None, max_depth: int = 20) -> Any: """ Convert object to JSON-serializable dict, handling Pydantic models safely. @@ -1054,7 +1053,7 @@ def _convert_to_json_serializable_dict(obj: Any, visited: Optional[set] = None, def _get_proxy_server_request_for_spend_logs_payload( metadata: dict, litellm_params: dict, - kwargs: Optional[dict] = None, + kwargs: dict | None = None, ) -> str: """ Only store if _should_store_prompts_and_responses_in_spend_logs() is True @@ -1062,7 +1061,7 @@ def _get_proxy_server_request_for_spend_logs_payload( If turn_off_message_logging is enabled, redact messages in the request body. """ if _should_store_prompts_and_responses_in_spend_logs(): - _proxy_server_request = cast(Optional[dict], litellm_params.get("proxy_server_request", {})) + _proxy_server_request = cast(dict | None, litellm_params.get("proxy_server_request", {})) if _proxy_server_request is not None: _request_body = _proxy_server_request.get("body", {}) or {} @@ -1102,8 +1101,8 @@ def _get_proxy_server_request_for_spend_logs_payload( def _get_vector_store_request_for_spend_logs_payload( - vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]], -) -> Optional[List[StandardLoggingVectorStoreRequest]]: + vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None, +) -> list[StandardLoggingVectorStoreRequest] | None: """ If user does not want to store prompts and responses, then remove the content from the vector store request metadata """ @@ -1126,8 +1125,8 @@ def _get_vector_store_request_for_spend_logs_payload( def _get_response_for_spend_logs_payload( - payload: Optional[StandardLoggingPayload], - kwargs: Optional[dict] = None, + payload: StandardLoggingPayload | None, + kwargs: dict | None = None, ) -> str: if payload is None: return "{}" @@ -1206,7 +1205,7 @@ def _get_status_for_spend_log( It's only a failure if metadata.get("status") is "failure" """ - _status: Optional[str] = metadata.get("status", None) + _status: str | None = metadata.get("status", None) if _status == "failure": return "failure" return "success" diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index 52d46e0d267..195731c3ed1 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -166,10 +166,10 @@ async def get_vantage_settings( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error retrieving Vantage settings: {str(e)}") + verbose_proxy_logger.error(f"Error retrieving Vantage settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to retrieve Vantage settings: {str(e)}"}, + detail={"error": f"Failed to retrieve Vantage settings: {e!s}"}, ) @@ -235,10 +235,10 @@ async def update_vantage_settings( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error updating Vantage settings: {str(e)}") + verbose_proxy_logger.error(f"Error updating Vantage settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to update Vantage settings: {str(e)}"}, + detail={"error": f"Failed to update Vantage settings: {e!s}"}, ) @@ -257,7 +257,7 @@ async def is_vantage_setup_in_db() -> bool: return vantage_config is not None and vantage_config.param_value is not None except Exception as e: - verbose_proxy_logger.error(f"Error checking Vantage status: {str(e)}") + verbose_proxy_logger.error(f"Error checking Vantage status: {e!s}") return False @@ -280,7 +280,7 @@ async def is_vantage_setup() -> bool: return True return False except Exception as e: - verbose_proxy_logger.error(f"Error checking Vantage setup: {str(e)}") + verbose_proxy_logger.error(f"Error checking Vantage setup: {e!s}") return False @@ -324,10 +324,10 @@ async def init_vantage_settings( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error initializing Vantage settings: {str(e)}") + verbose_proxy_logger.error(f"Error initializing Vantage settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to initialize Vantage settings: {str(e)}"}, + detail={"error": f"Failed to initialize Vantage settings: {e!s}"}, ) @@ -415,10 +415,10 @@ async def vantage_dry_run_export( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error performing Vantage dry run export: {str(e)}") + verbose_proxy_logger.error(f"Error performing Vantage dry run export: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to perform Vantage dry run export: {str(e)}"}, + detail={"error": f"Failed to perform Vantage dry run export: {e!s}"}, ) @@ -488,10 +488,10 @@ async def vantage_export( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error performing Vantage export: {str(e)}") + verbose_proxy_logger.error(f"Error performing Vantage export: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to perform Vantage export: {str(e)}"}, + detail={"error": f"Failed to perform Vantage export: {e!s}"}, ) @@ -548,8 +548,8 @@ async def delete_vantage_settings( except HTTPException: raise except Exception as e: - verbose_proxy_logger.error(f"Error deleting Vantage settings: {str(e)}") + verbose_proxy_logger.error(f"Error deleting Vantage settings: {e!s}") raise HTTPException( status_code=500, - detail={"error": f"Failed to delete Vantage settings: {str(e)}"}, + detail={"error": f"Failed to delete Vantage settings: {e!s}"}, ) diff --git a/litellm/proxy/types_utils/utils.py b/litellm/proxy/types_utils/utils.py index c879315d477..e9fb18b258e 100644 --- a/litellm/proxy/types_utils/utils.py +++ b/litellm/proxy/types_utils/utils.py @@ -3,10 +3,10 @@ import importlib import importlib.util import os from collections.abc import Callable -from typing import Any, Literal, Optional, get_type_hints +from typing import Any, Literal, get_type_hints -def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any: +def get_instance_fn(value: str, config_file_path: str | None = None) -> Any: module_name = value instance_name = None try: @@ -65,7 +65,7 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any: raise e -def _load_instance_from_remote_storage(remote_url: str, config_file_path: Optional[str] = None) -> Any: +def _load_instance_from_remote_storage(remote_url: str, config_file_path: str | None = None) -> Any: """ Load custom logger instance from S3 or GCS URL. @@ -176,7 +176,7 @@ def _load_instance_from_remote_storage(remote_url: str, config_file_path: Option return instance except Exception as e: - raise ImportError(f"Failed to load custom logger from {remote_url}: {str(e)}") from e + raise ImportError(f"Failed to load custom logger from {remote_url}: {e!s}") from e async def _download_gcs_file_wrapper(bucket_name: str, object_key: str, local_file_path: str) -> bool: @@ -190,13 +190,13 @@ async def _download_gcs_file_wrapper(bucket_name: str, object_key: str, local_fi except Exception as e: from litellm._logging import verbose_proxy_logger - verbose_proxy_logger.error(f"Error downloading from GCS: {str(e)}") + verbose_proxy_logger.error(f"Error downloading from GCS: {e!s}") return False def validate_custom_validate_return_type( - fn: Optional[Callable[..., Any]], -) -> Optional[Callable[..., Literal[True]]]: + fn: Callable[..., Any] | None, +) -> Callable[..., Literal[True]] | None: if fn is None: return None diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 8f3f8ad1bfc..08b4b12c1d1 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -4,7 +4,7 @@ import json import os from collections import Counter from collections.abc import Mapping -from typing import Any, Dict, List, Optional, Set, Tuple, Type, Union +from typing import Any from urllib.parse import urlparse from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile @@ -81,13 +81,13 @@ class UIThemeConfig(BaseModel): """Configuration for UI theme customization""" # Logo configuration - logo_url: Optional[str] = Field( + logo_url: str | None = Field( default=None, description="URL or path to custom logo image. Can be a local file path or HTTP/HTTPS URL", ) # Favicon configuration - favicon_url: Optional[str] = Field( + favicon_url: str | None = Field( default=None, description="URL to custom favicon image. Must be an HTTP/HTTPS URL to a .ico, .png, or .svg file", ) @@ -96,37 +96,31 @@ class UIThemeConfig(BaseModel): class SettingsResponse(BaseModel): """Base response model for settings with values and schema information""" - values: Dict[str, Any] + values: dict[str, Any] """The current configuration values""" - field_schema: Dict[str, Any] + field_schema: dict[str, Any] """Schema information including descriptions and property types for UI display""" class SSOSettingsResponse(SettingsResponse): """Response model for SSO settings""" - provenance: Dict[str, str] = Field(default_factory=dict) + provenance: dict[str, str] = Field(default_factory=dict) """Per-field source of each value: 'db', 'env', 'default', or 'unset'.""" class InternalUserSettingsResponse(SettingsResponse): """Response model for internal user settings""" - pass - class DefaultTeamSettingsResponse(SettingsResponse): """Response model for default team settings""" - pass - class UIThemeSettingsResponse(SettingsResponse): """Response model for UI theme settings""" - pass - class UISettings(BaseModel): """Configuration for UI-specific flags""" @@ -143,7 +137,7 @@ class UISettings(BaseModel): description="Prevents Team Admins from deleting users from the teams they manage. Useful for SCIM provisioning where team membership is defined externally.", ) - enabled_ui_pages_internal_users: Optional[List[str]] = Field( + enabled_ui_pages_internal_users: list[str] | None = Field( default=None, description="List of page keys that internal users (non-admins) can see in the UI sidebar. If not set, all pages are visible based on role permissions.", ) @@ -224,8 +218,6 @@ class UISettings(BaseModel): class UISettingsResponse(SettingsResponse): """Response model for UI settings""" - pass - # Allowlist of UI settings that can be stored ALLOWED_UI_SETTINGS_FIELDS = { @@ -269,16 +261,16 @@ _RUNTIME_GENERAL_SETTINGS_FLAGS = [ # include generics like ``Optional[int]`` / ``List[str]`` that are not # instances of ``type`` — so tightening this to ``type`` would reject # valid inputs. -_EXTRA_UI_SETTINGS_FIELDS: Dict[str, Tuple[Any, FieldInfo]] = {} +_EXTRA_UI_SETTINGS_FIELDS: dict[str, tuple[Any, FieldInfo]] = {} # Settings OSS knows about as enterprise-gated. If a caller sends one of # these keys and no extension package has registered it, the PATCH # endpoint returns 403 instead of silently dropping the value, so the # client gets a clear signal that the feature requires LiteLLM Enterprise. -_ENTERPRISE_ONLY_UI_SETTINGS: Set[str] = {"enable_projects_ui"} +_ENTERPRISE_ONLY_UI_SETTINGS: set[str] = {"enable_projects_ui"} # Memoized effective class; invalidated on registration. -_EFFECTIVE_UI_SETTINGS_CLASS: Optional[Type[UISettings]] = None +_EFFECTIVE_UI_SETTINGS_CLASS: type[UISettings] | None = None def register_extra_ui_setting(name: str, annotation: Any, field: FieldInfo) -> None: @@ -295,7 +287,7 @@ def register_extra_ui_setting(name: str, annotation: Any, field: FieldInfo) -> N _EFFECTIVE_UI_SETTINGS_CLASS = None -def _get_effective_ui_settings_class() -> Type[UISettings]: +def _get_effective_ui_settings_class() -> type[UISettings]: """Return UISettings with any extension-registered fields merged in. Memoized — pydantic ``create_model`` runs metaclass + schema work @@ -346,8 +338,6 @@ class MCPSemanticFilterSettings(BaseModel): class MCPSemanticFilterSettingsResponse(SettingsResponse): """Response model for MCP semantic filter settings""" - pass - @router.get( "/get/allowed_ips", @@ -382,7 +372,7 @@ async def add_allowed_ip( if prisma_client is None: raise Exception("No DB Connected") - _allowed_ips: List = general_settings.get("allowed_ips", []) + _allowed_ips: list = general_settings.get("allowed_ips", []) if ip_address.ip not in _allowed_ips: _allowed_ips.append(ip_address.ip) general_settings["allowed_ips"] = _allowed_ips @@ -441,7 +431,7 @@ async def delete_allowed_ip( proxy_config, ) - _allowed_ips: List = general_settings.get("allowed_ips", []) + _allowed_ips: list = general_settings.get("allowed_ips", []) if ip_address.ip in _allowed_ips: _allowed_ips.remove(ip_address.ip) general_settings["allowed_ips"] = _allowed_ips @@ -645,7 +635,7 @@ async def _validate_default_teams_exist(teams: list[str] | list[NewUserRequestTe ) -async def update_default_team_member_budget(teams: List[NewUserRequestTeam], user_api_key_dict: UserAPIKeyAuth): +async def update_default_team_member_budget(teams: list[NewUserRequestTeam], user_api_key_dict: UserAPIKeyAuth): """ 1. Update the max member budget for the team """ @@ -673,7 +663,7 @@ async def update_default_team_member_budget(teams: List[NewUserRequestTeam], use async def _update_litellm_setting( - settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings], + settings: DefaultInternalUserParams | DefaultTeamSSOParams | MCPSemanticFilterSettings, settings_key: str, success_message: str, user_api_key_dict: UserAPIKeyAuth, @@ -887,7 +877,7 @@ async def update_sso_settings( # create_config_audit_log's secret-name redaction to mask the # *_client_secret fields before the audit row is written. existing_sso_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) - before_sso_data: Optional[Dict[str, Any]] = None + before_sso_data: dict[str, Any] | None = None if existing_sso_record and existing_sso_record.sso_settings: stored = existing_sso_record.sso_settings if isinstance(stored, str): @@ -974,7 +964,7 @@ async def update_sso_settings( except Exception as e: raise HTTPException( status_code=500, - detail={"error": f"Error updating environment_variables: {str(e)}"}, + detail={"error": f"Error updating environment_variables: {e!s}"}, ) return { @@ -1016,7 +1006,7 @@ async def get_ui_theme_settings(): return result -def _validate_public_image_url(value: Optional[str], field_name: str) -> None: +def _validate_public_image_url(value: str | None, field_name: str) -> None: """ Reject anything that isn't a plain http(s) URL with a host. This value is later served via the unauthenticated /get_image endpoint, so local paths @@ -1192,7 +1182,7 @@ UI_SETTINGS_CACHE_KEY = "ui_settings:settings_dict" UI_SETTINGS_CACHE_TTL = 600 # 10 minutes -async def get_ui_settings_cached() -> Dict[str, Any]: +async def get_ui_settings_cached() -> dict[str, Any]: """ Return the persisted UI settings dict, using DualCache for reads. @@ -1211,7 +1201,7 @@ async def get_ui_settings_cached() -> Dict[str, Any]: return {} db_record = await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}) - ui_settings: Dict[str, Any] = {} + ui_settings: dict[str, Any] = {} if db_record and db_record.ui_settings: raw = db_record.ui_settings ui_settings = json.loads(raw) if isinstance(raw, str) else dict(raw) @@ -1243,7 +1233,7 @@ async def get_ui_settings(): detail={"error": "Database not connected. Please connect a database."}, ) - ui_settings: Dict[str, Any] = {} + ui_settings: dict[str, Any] = {} db_record = await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}) @@ -1271,7 +1261,7 @@ async def get_ui_settings(): await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) # Build config-like object for schema helper - config: Dict[str, Any] = {"litellm_settings": {"ui_settings": ui_settings}} + config: dict[str, Any] = {"litellm_settings": {"ui_settings": ui_settings}} return await _get_settings_with_schema( settings_key="ui_settings", @@ -1286,7 +1276,7 @@ async def get_ui_settings(): dependencies=[Depends(user_api_key_auth)], ) async def update_ui_settings( - settings_body: Dict[str, Any] = Body(...), + settings_body: dict[str, Any] = Body(...), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 86eee5c8d7a..40db863de5d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -19,11 +19,8 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, - Dict, - List, Literal, Optional, - Tuple, Union, cast, overload, @@ -196,7 +193,7 @@ def print_verbose(print_statement): """ import traceback - verbose_proxy_logger.debug("{}\n{}".format(print_statement, traceback.format_exc())) + verbose_proxy_logger.debug(f"{print_statement}\n{traceback.format_exc()}") if litellm.set_verbose: print(f"LiteLLM Proxy: {_redact_string(str(print_statement))}") # noqa: T201 @@ -235,7 +232,7 @@ class InternalUsageCache: async def async_get_cache( self, key, - litellm_parent_otel_span: Union[Span, None], + litellm_parent_otel_span: Span | None, local_only: bool = False, **kwargs, ) -> Any: @@ -250,7 +247,7 @@ class InternalUsageCache: self, key, value, - litellm_parent_otel_span: Union[Span, None], + litellm_parent_otel_span: Span | None, local_only: bool = False, **kwargs, ) -> None: @@ -264,8 +261,8 @@ class InternalUsageCache: async def async_batch_set_cache( self, - cache_list: List, - litellm_parent_otel_span: Union[Span, None], + cache_list: list, + litellm_parent_otel_span: Span | None, local_only: bool = False, **kwargs, ) -> None: @@ -279,7 +276,7 @@ class InternalUsageCache: async def async_batch_get_cache( self, keys: list, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, local_only: bool = False, ): return await self.dual_cache.async_batch_get_cache( @@ -292,7 +289,7 @@ class InternalUsageCache: self, key, value: float, - litellm_parent_otel_span: Union[Span, None], + litellm_parent_otel_span: Span | None, local_only: bool = False, **kwargs, ): @@ -334,7 +331,7 @@ class InternalUsageCache: ### LOGGING ### # Cache for inspect.signature checks — avoids repeated introspection per request -_CALLBACK_ACCEPTS_CALL_INFO: Dict[int, bool] = {} +_CALLBACK_ACCEPTS_CALL_INFO: dict[int, bool] = {} def _accepts_litellm_call_info(cb: CustomLogger) -> bool: @@ -394,11 +391,11 @@ class _CallbackCapabilities: # Tuple[(resolved_callback, "override" | "apply_guardrail"), ...] # Ordered the same as ``litellm.callbacks``; used to build the streaming # iterator chain without re-scanning per request. - iterator_overrides: Tuple[Tuple[Any, str], ...] = field(default_factory=tuple) + iterator_overrides: tuple[tuple[Any, str], ...] = field(default_factory=tuple) # Resolved CustomLogger callbacks in original order. Pre-resolving once # avoids the per-request ``get_custom_logger_compatible_class`` walk for # every string entry in ``litellm.callbacks``. - resolved_callbacks: Tuple[Any, ...] = field(default_factory=tuple) + resolved_callbacks: tuple[Any, ...] = field(default_factory=tuple) class ProxyLogging: @@ -424,16 +421,16 @@ class ProxyLogging: self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) self.max_budget_limiter = _PROXY_MaxBudgetLimiter() self.cache_control_check = _PROXY_CacheControlCheck() - self.alerting: Optional[List] = None + self.alerting: list | None = None self.alerting_threshold: float = 300 # default to 5 min. threshold - self.alert_types: List[AlertType] = DEFAULT_ALERT_TYPES - self.alert_to_webhook_url: Optional[dict] = None + self.alert_types: list[AlertType] = DEFAULT_ALERT_TYPES + self.alert_to_webhook_url: dict | None = None self.slack_alerting_instance: SlackAlerting = SlackAlerting( alerting_threshold=self.alerting_threshold, alerting=self.alerting, internal_usage_cache=self.internal_usage_cache.dual_cache, ) - self.email_logging_instance: Optional[Any] = None + self.email_logging_instance: Any | None = None if BaseEmailLogger is not None: email_logger_class = _get_email_logger_class() if email_logger_class is not None: @@ -444,7 +441,7 @@ class ProxyLogging: self.premium_user = premium_user self.service_logging_obj = ServiceLogging() self.db_spend_update_writer = DBSpendUpdateWriter() - self.proxy_hook_mapping: Dict[str, CustomLogger] = {} + self.proxy_hook_mapping: dict[str, CustomLogger] = {} # Guard flags to prevent duplicate background tasks self.daily_report_started: bool = False @@ -452,8 +449,8 @@ class ProxyLogging: def startup_event( self, - llm_router: Optional[Router], - redis_usage_cache: Optional[RedisCache], + llm_router: Router | None, + redis_usage_cache: RedisCache | None, ): """Initialize logging and alerting on proxy startup""" ## UPDATE SLACK ALERTING ## @@ -490,13 +487,13 @@ class ProxyLogging: def update_values( self, - alerting: Optional[List] = None, - alerting_threshold: Optional[float] = None, - redis_cache: Optional[RedisCache] = None, - alert_types: Optional[List[AlertType]] = None, - alerting_args: Optional[dict] = None, - alert_to_webhook_url: Optional[dict] = None, - alert_type_config: Optional[dict] = None, + alerting: list | None = None, + alerting_threshold: float | None = None, + redis_cache: RedisCache | None = None, + alert_types: list[AlertType] | None = None, + alerting_args: dict | None = None, + alert_to_webhook_url: dict | None = None, + alert_type_config: dict | None = None, ): updated_slack_alerting: bool = False if alerting is not None: @@ -542,7 +539,7 @@ class ProxyLogging: self.db_spend_update_writer.redis_update_buffer.redis_cache = redis_cache self.db_spend_update_writer.pod_lock_manager.redis_cache = redis_cache - def _add_proxy_hooks(self, llm_router: Optional[Router] = None): + def _add_proxy_hooks(self, llm_router: Router | None = None): """ Add proxy hooks to litellm.callbacks """ @@ -551,7 +548,7 @@ class ProxyLogging: for hook in PROXY_HOOKS: proxy_hook = get_proxy_hook(hook) expected_args = inspect.getfullargspec(proxy_hook).args - passed_in_args: Dict[str, Any] = {} + passed_in_args: dict[str, Any] = {} if "internal_usage_cache" in expected_args: passed_in_args["internal_usage_cache"] = self.internal_usage_cache if "prisma_client" in expected_args: @@ -561,20 +558,20 @@ class ProxyLogging: self.proxy_hook_mapping[hook] = proxy_hook_obj - def get_proxy_hook(self, hook: str) -> Optional[CustomLogger]: + def get_proxy_hook(self, hook: str) -> CustomLogger | None: """ Get a proxy hook from the proxy_hook_mapping """ return self.proxy_hook_mapping.get(hook) - def _init_litellm_callbacks(self, llm_router: Optional[Router] = None): + def _init_litellm_callbacks(self, llm_router: Router | None = None): self._add_proxy_hooks(llm_router) litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) # type: ignore # Track string callbacks and their initialized instances so we can # replace them in-place, preventing duplicates (string + instance) in # litellm.callbacks which caused double-counting of metrics. - string_callbacks_to_replace: Dict[int, CustomLogger] = {} + string_callbacks_to_replace: dict[int, CustomLogger] = {} for idx, callback in enumerate(litellm.callbacks): if isinstance(callback, str): @@ -620,7 +617,7 @@ class ProxyLogging: alerting_threshold += 100 await self.internal_usage_cache.async_set_cache( - key="request_status:{}".format(litellm_call_id), + key=f"request_status:{litellm_call_id}", value=status, local_only=True, ttl=alerting_threshold, @@ -676,7 +673,7 @@ class ProxyLogging: return synthetic_data - def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> Optional[Any]: + def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> Any | None: """ Convert LLM guardrail result back to MCP response format. """ @@ -736,7 +733,7 @@ class ProxyLogging: return None - def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> Optional[dict]: + def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> dict | None: """ Extract modified/masked arguments from the guardrail response content. """ @@ -773,7 +770,7 @@ class ProxyLogging: verbose_proxy_logger.error(f"Error extracting modified arguments: {e}") return None - def _parse_arguments_manually(self, args_text: str, original_args: dict) -> Optional[dict]: + def _parse_arguments_manually(self, args_text: str, original_args: dict) -> dict | None: """ Try to manually parse arguments when JSON parsing fails. This is a fallback for cases where the guardrail modifies the format. @@ -802,7 +799,7 @@ class ProxyLogging: verbose_proxy_logger.error(f"Error in manual argument parsing: {e}") return None - def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> Optional[Any]: + def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> Any | None: """ Convert LLM guardrail result back to MCP during call response format. """ @@ -839,7 +836,7 @@ class ProxyLogging: return None - def get_combined_callback_list(self, dynamic_success_callbacks: Optional[List], global_callbacks: List) -> List: + def get_combined_callback_list(self, dynamic_success_callbacks: list | None, global_callbacks: list) -> list: if dynamic_success_callbacks is None: return list(global_callbacks) return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks)) @@ -848,7 +845,7 @@ class ProxyLogging: self, response: MCPPreCallResponseObject, original_request: MCPPreCallRequestObject, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Parse the response from the pre_mcp_tool_call_hook @@ -881,7 +878,7 @@ class ProxyLogging: hidden_params=HiddenParams(), ) - def _convert_mcp_hook_response_to_kwargs(self, response_data: Optional[dict], original_kwargs: dict) -> dict: + def _convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: """ Helper function to convert pre_call_hook response back to kwargs for MCP usage. @@ -949,9 +946,9 @@ class ProxyLogging: callback: "CustomGuardrail", hook_type: str, data: dict, - user_api_key_dict: Optional[UserAPIKeyAuth], + user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, - response: Optional[Any] = None, + response: Any | None = None, ) -> Any: """ Execute a single guardrail's hook. @@ -1004,9 +1001,9 @@ class ProxyLogging: guardrail_name: str, hook_type: str, data: dict, - user_api_key_dict: Optional[UserAPIKeyAuth], + user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, - response: Optional[Any] = None, + response: Any | None = None, ) -> Any: """ Execute a guardrail using the router's load balancing. @@ -1047,10 +1044,10 @@ class ProxyLogging: self, callback: CustomGuardrail, data: dict, - user_api_key_dict: Optional[UserAPIKeyAuth], + user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, event_type: GuardrailEventHooks, - ) -> Optional[dict]: + ) -> dict | None: """ Process a guardrail callback during pre-call hook. @@ -1165,7 +1162,7 @@ class ProxyLogging: custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id) prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id) - litellm_prompt_id: Optional[str] = None + litellm_prompt_id: str | None = None if prompt_spec is not None: litellm_prompt_id = prompt_spec.litellm_params.prompt_id data.pop("prompt_id", None) @@ -1345,9 +1342,9 @@ class ProxyLogging: async def pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, - data: Optional[dict], + data: dict | None, call_type: CallTypesLiteral, - ) -> Optional[dict]: + ) -> dict | None: """ Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body. @@ -1414,7 +1411,7 @@ class ProxyLogging: and not (cb.guardrail_name and cb.guardrail_name in pipeline_managed) ) - deferred_route_exc: Optional[SensitiveDataRouteException] = None + deferred_route_exc: SensitiveDataRouteException | None = None for _callback in caps.resolved_callbacks: start_time = time.time() try: @@ -1546,9 +1543,9 @@ class ProxyLogging: async def _handle_sensitive_data_route_exception( self, exc: SensitiveDataRouteException, - data: Optional[dict], - user_api_key_dict: Optional[UserAPIKeyAuth], - ) -> Optional[dict]: + data: dict | None, + user_api_key_dict: UserAPIKeyAuth | None, + ) -> dict | None: """ Handle SensitiveDataRouteException by rerouting the current request to the target model and, when sticky_session_routing is enabled, persisting @@ -1599,7 +1596,7 @@ class ProxyLogging: guardrail_name: str, latency_seconds: float, status: str, - error_type: Optional[str], + error_type: str | None, hook_type: str, ) -> None: for prom_callback in litellm.callbacks: @@ -1624,7 +1621,7 @@ class ProxyLogging: guardrail_name = getattr(callback, "guardrail_name", None) or type(callback).__name__ start_time = time.perf_counter() status = "success" - error_type: Optional[str] = None + error_type: str | None = None try: return await coro except SensitiveDataRouteException: @@ -1666,7 +1663,7 @@ class ProxyLogging: # Cache for callback-capability detection. Keyed on a signature of # litellm.callbacks (length + each item's id) so we recompute when the # callback list mutates (add/remove) without iterating every request. - _callback_capabilities_cache: ClassVar[Dict[Tuple[int, Tuple[int, ...]], "_CallbackCapabilities"]] = {} + _callback_capabilities_cache: ClassVar[dict[tuple[int, tuple[int, ...]], "_CallbackCapabilities"]] = {} @staticmethod def _callback_capabilities() -> "_CallbackCapabilities": @@ -1691,8 +1688,8 @@ class ProxyLogging: has_streaming_chunk_override = False has_guardrail = False has_pre_call_override = False - iterator_overrides: List[Tuple[Any, str]] = [] # (callback, kind) - resolved_callbacks: List[Any] = [] + iterator_overrides: list[tuple[Any, str]] = [] # (callback, kind) + resolved_callbacks: list[Any] = [] for callback in callbacks: if isinstance(callback, str): @@ -1812,7 +1809,7 @@ class ProxyLogging: async def during_call_hook( self, data: dict, - user_api_key_dict: Optional[UserAPIKeyAuth], + user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, ): """ @@ -1954,7 +1951,7 @@ class ProxyLogging: message: str, level: Literal["Low", "Medium", "High"], alert_type: AlertType, - request_data: Optional[dict] = None, + request_data: dict | None = None, ): """ Alerting based on thresholds: - https://github.com/BerriAI/litellm/issues/1298 @@ -1989,7 +1986,7 @@ class ProxyLogging: if _url is not None: extra_kwargs["🪢 Langfuse Trace"] = _url - formatted_message += "\n\n🪢 Langfuse Trace: {}".format(_url) + formatted_message += f"\n\n🪢 Langfuse Trace: {_url}" if ( "metadata" in request_data and request_data["metadata"].get("alerting_metadata", None) is not None @@ -2058,10 +2055,10 @@ class ProxyLogging: request_data: dict, original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, - error_type: Optional[ProxyErrorTypes] = None, - route: Optional[str] = None, - traceback_str: Optional[str] = None, - ) -> Optional[HTTPException]: + error_type: ProxyErrorTypes | None = None, + route: str | None = None, + traceback_str: str | None = None, + ) -> HTTPException | None: """ Allows users to raise custom exceptions/log when a call fails, without having to deal with parsing Request body. Callbacks can return or raise HTTPException to transform error responses sent to clients. @@ -2145,11 +2142,11 @@ class ProxyLogging: request_data.pop("litellm_logging_obj", None) # Track the first HTTPException returned or raised by any callback - transformed_exception: Optional[HTTPException] = None + transformed_exception: HTTPException | None = None for callback in litellm.callbacks: try: - _callback: Optional[CustomLogger] = None + _callback: CustomLogger | None = None if isinstance(callback, str): _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( cast(_custom_logger_compatible_callbacks_literal, callback) @@ -2184,8 +2181,8 @@ class ProxyLogging: def _is_proxy_only_llm_api_error( self, original_exception: Exception, - error_type: Optional[ProxyErrorTypes] = None, - route: Optional[str] = None, + error_type: ProxyErrorTypes | None = None, + route: str | None = None, ) -> bool: """ Return True if the error is a Proxy Only LLM API Error @@ -2218,15 +2215,15 @@ class ProxyLogging: self, request_data: dict, user_api_key_dict: UserAPIKeyAuth, - route: Optional[str] = None, - original_exception: Optional[Exception] = None, + route: str | None = None, + original_exception: Exception | None = None, ): """ Handle logging for proxy only errors by calling `litellm_logging_obj.async_failure_handler` Is triggered when self._is_proxy_only_error() returns True """ - litellm_logging_obj: Optional[Logging] = request_data.get("litellm_logging_obj", None) + litellm_logging_obj: Logging | None = request_data.get("litellm_logging_obj", None) if litellm_logging_obj is None: from litellm._uuid import uuid @@ -2264,8 +2261,8 @@ class ProxyLogging: litellm_params=_litellm_params, ) - input: Union[list, str, dict] = "" - normalized_call_type: Optional[str] = None + input: list | str | dict = "" + normalized_call_type: str | None = None if "messages" in request_data and isinstance(request_data["messages"], list): input = request_data["messages"] litellm_logging_obj.model_call_details["messages"] = input @@ -2334,11 +2331,11 @@ class ProxyLogging: from litellm.proxy.proxy_server import llm_router from litellm.types.guardrails import GuardrailEventHooks - guardrail_callbacks: List[CustomGuardrail] = [] - other_callbacks: List[CustomLogger] = [] + guardrail_callbacks: list[CustomGuardrail] = [] + other_callbacks: list[CustomLogger] = [] try: for callback in litellm.callbacks: - _callback: Optional[CustomLogger] = None + _callback: CustomLogger | None = None if isinstance(callback, str): _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( cast(_custom_logger_compatible_callbacks_literal, callback) @@ -2376,7 +2373,7 @@ class ProxyLogging: ): continue - guardrail_response: Optional[Any] = None + guardrail_response: Any | None = None if "apply_guardrail" in type(callback).__dict__: data["guardrail_to_apply"] = callback @@ -2542,8 +2539,8 @@ class ProxyLogging: data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any, - request_headers: Optional[Dict[str, str]] = None, - ) -> Dict[str, str]: + request_headers: dict[str, str] | None = None, + ) -> dict[str, str]: """ Calls async_post_call_response_headers_hook on all CustomLogger callbacks. Merges all returned header dicts (later callbacks override earlier ones). @@ -2551,7 +2548,7 @@ class ProxyLogging: Returns: Dict[str, str]: Merged headers from all callbacks. """ - merged_headers: Dict[str, str] = {} + merged_headers: dict[str, str] = {} # Outer call sites in common_request_processing.py already gate this # call with ``has_post_call_response_headers_callbacks()``. The # cached detection makes the redundant interior guard cheap, but the @@ -2565,7 +2562,7 @@ class ProxyLogging: litellm_call_info = self._build_litellm_call_info(data=data, response=response) for callback in litellm.callbacks: - _callback: Optional[CustomLogger] = None + _callback: CustomLogger | None = None if isinstance(callback, str): _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( cast(_custom_logger_compatible_callbacks_literal, callback) @@ -2597,7 +2594,7 @@ class ProxyLogging: return merged_headers @staticmethod - def _build_litellm_call_info(data: dict, response: Any) -> Dict[str, Any]: + def _build_litellm_call_info(data: dict, response: Any) -> dict[str, Any]: """ Build a normalized dict of routing metadata from response._hidden_params and data, abstracting away the metadata vs litellm_metadata split. @@ -2626,9 +2623,9 @@ class ProxyLogging: async def async_post_call_streaming_hook( self, data: dict, - response: Union[ModelResponse, EmbeddingResponse, ImageResponse, ModelResponseStream], + response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream, user_api_key_dict: UserAPIKeyAuth, - str_so_far: Optional[str] = None, + str_so_far: str | None = None, ): """ Allow user to modify outgoing streaming data -> per chunk @@ -2648,7 +2645,7 @@ class ProxyLogging: from litellm.proxy.proxy_server import llm_router - response_str: Optional[str] = None + response_str: str | None = None if isinstance(response, (ModelResponse, ModelResponseStream)): response_str = litellm.get_response_string(response_obj=response) elif isinstance(response, dict) and self.is_a2a_streaming_response(response): @@ -2658,12 +2655,12 @@ class ProxyLogging: if response_str is not None: # Cache model-level guardrails check per-request to avoid repeated # dict lookups + llm_router.get_deployment() per callback per chunk. - _cached_guardrail_data: Optional[dict] = None + _cached_guardrail_data: dict | None = None _guardrail_data_computed = False for callback in litellm.callbacks: try: - _callback: Optional[CustomLogger] = None + _callback: CustomLogger | None = None if isinstance(callback, CustomGuardrail): # Main - V2 Guardrails implementation from litellm.types.guardrails import GuardrailEventHooks @@ -2831,7 +2828,7 @@ class ProxyLogging: return await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) - def _init_response_taking_too_long_task(self, data: Optional[dict] = None): + def _init_response_taking_too_long_task(self, data: dict | None = None): """ Initialize the response taking too long task if user is using slack alerting @@ -2876,7 +2873,7 @@ _DEPRECATED_KEY_CACHE_TTL_SECONDS = 60 async def _lookup_deprecated_key( db: Any, hashed_token: str, -) -> Optional[str]: +) -> str | None: """ Check if a token exists in the deprecated keys table and is still within its grace period. @@ -2942,11 +2939,11 @@ def _config_cache_key(param_name: str) -> str: return f"litellm_config:param:{param_name}" -def _pack_config_row(row: Any) -> Dict[str, Any]: +def _pack_config_row(row: Any) -> dict[str, Any]: return {"param_name": row.param_name, "param_value": row.param_value} -def _unpack_config_row(cached: Any) -> Optional[_ConfigRow]: +def _unpack_config_row(cached: Any) -> _ConfigRow | None: if cached is None or cached == _CONFIG_CACHE_MISS: return None if isinstance(cached, dict): @@ -2954,7 +2951,7 @@ def _unpack_config_row(cached: Any) -> Optional[_ConfigRow]: return None -async def get_config_param(prisma_client: Any, param_name: str) -> Optional[Any]: +async def get_config_param(prisma_client: Any, param_name: str) -> Any | None: """Cached read of a LiteLLM_Config row; returns row, _ConfigRow shim, or None.""" cache_key = _config_cache_key(param_name) cached = await litellm_config_cache.async_get_cache(cache_key) @@ -2972,7 +2969,7 @@ async def invalidate_config_param(param_name: str) -> None: await litellm_config_cache.async_delete_cache(_config_cache_key(param_name)) -async def prefetch_config_params(prisma_client: Any, param_names: List[str]) -> None: +async def prefetch_config_params(prisma_client: Any, param_names: list[str]) -> None: """Batch-load LiteLLM_Config rows into the cache with one find_many.""" if not param_names: return @@ -2996,20 +2993,20 @@ async def prefetch_config_params(prisma_client: Any, param_names: List[str]) -> class PrismaClient: - spend_log_transactions: List = [] + spend_log_transactions: list = [] _spend_log_transactions_lock = asyncio.Lock() - tool_usage_transactions: List["ToolUsageTransaction"] = [] + tool_usage_transactions: list["ToolUsageTransaction"] = [] _tool_usage_transactions_lock = asyncio.Lock() def __init__( self, database_url: str, proxy_logging_obj: ProxyLogging, - http_client: Optional[Any] = None, + http_client: Any | None = None, ): ## init logging object self.proxy_logging_obj = proxy_logging_obj - self.iam_token_db_auth: Optional[bool] = str_to_bool(os.getenv("IAM_TOKEN_DB_AUTH")) + self.iam_token_db_auth: bool | None = str_to_bool(os.getenv("IAM_TOKEN_DB_AUTH")) verbose_proxy_logger.debug("Creating Prisma Client..") try: from prisma import Prisma # type: ignore @@ -3042,7 +3039,7 @@ class PrismaClient: # reader endpoint and writes stay on the writer. Falls back to the # writer-only wrapper when the env var is unset, preserving existing # single-DB deployments. - self.db: Union[PrismaWrapper, RoutingPrismaWrapper] + self.db: PrismaWrapper | RoutingPrismaWrapper if read_replica_url: try: # If IAM auth is enabled, the reader refreshes its own token on @@ -3070,7 +3067,7 @@ class PrismaClient: ) read_replica_url = reader_iam_endpoint.build_url(reader_token) os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url - reader_kwargs: Dict[str, Any] = {"datasource": {"url": read_replica_url}} + reader_kwargs: dict[str, Any] = {"datasource": {"url": read_replica_url}} if http_client is not None: reader_prisma = Prisma(http=http_client, **reader_kwargs) else: @@ -3106,7 +3103,7 @@ class PrismaClient: else: self.db = writer_wrapper # Client to connect to Prisma db self._db_reconnect_lock = asyncio.Lock() - self._db_health_watchdog_task: Optional[asyncio.Task] = None + self._db_health_watchdog_task: asyncio.Task | None = None self._db_last_reconnect_attempt_ts: float = 0.0 self._db_reconnect_cooldown_seconds: int = max(1, int(os.getenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "15"))) self._db_health_watchdog_interval_seconds: int = max( @@ -3135,7 +3132,7 @@ class PrismaClient: self._engine_pid: int = 0 self._watching_engine: bool = False self._engine_confirmed_dead: bool = False - self._engine_wait_thread: Optional[threading.Thread] = None + self._engine_wait_thread: threading.Thread | None = None verbose_proxy_logger.debug("Success - Created Prisma Client") @property @@ -3153,7 +3150,7 @@ class PrismaClient: """ return cast("TransactionManager", self.db.tx()) # cast-ok: wrappers delegate tx via __getattr__ (untyped) - def get_request_status(self, payload: Union[dict, SpendLogsPayload]) -> Literal["success", "failure"]: + def get_request_status(self, payload: dict | SpendLogsPayload) -> Literal["success", "failure"]: """ Determine if a request was successful or failed based on payload metadata. @@ -3165,9 +3162,9 @@ class PrismaClient: """ try: # Get metadata and convert to dict if it's a JSON string - payload_metadata: Union[Dict, SpendLogsMetadata, str] = payload.get("metadata", {}) + payload_metadata: dict | SpendLogsMetadata | str = payload.get("metadata", {}) if isinstance(payload_metadata, str): - payload_metadata_json: Union[Dict, SpendLogsMetadata] = cast(Dict, json.loads(payload_metadata)) + payload_metadata_json: dict | SpendLogsMetadata = cast(dict, json.loads(payload_metadata)) else: payload_metadata_json = payload_metadata @@ -3278,9 +3275,7 @@ class PrismaClient: missing_views = expected_views_set - ret_view_names_set verbose_proxy_logger.warning( - "\n\n\033[93mNot all views exist in db, needed for UI 'Usage' tab. Missing={}.\nRun 'create_views.py' from https://github.com/BerriAI/litellm/tree/main/db_scripts to create missing views.\033[0m\n".format( - missing_views - ) + f"\n\n\033[93mNot all views exist in db, needed for UI 'Usage' tab. Missing={missing_views}.\nRun 'create_views.py' from https://github.com/BerriAI/litellm/tree/main/db_scripts to create missing views.\033[0m\n" ) except Exception: @@ -3338,9 +3333,9 @@ class PrismaClient: reason=f"prisma_get_generic_data_{table_name}_lookup_failure", ) except Exception as e: - error_msg = f"LiteLLM Prisma Client Exception get_generic_data: {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception get_generic_data: {e!s}" verbose_proxy_logger.error(error_msg) - error_msg = error_msg + "\nException Type: {}".format(type(e)) + error_msg = error_msg + f"\nException Type: {type(e)}" error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() _duration = end_time - start_time @@ -3355,7 +3350,7 @@ class PrismaClient: raise e - async def _query_first_with_cached_plan_fallback(self, sql_query: str, *args) -> Optional[dict]: + async def _query_first_with_cached_plan_fallback(self, sql_query: str, *args) -> dict | None: """ Execute a query, recovering once from PostgreSQL's "cached plan must not change result type" error. @@ -3411,38 +3406,29 @@ class PrismaClient: @log_db_metrics async def get_data( self, - token: Optional[Union[str, list]] = None, - user_id: Optional[str] = None, - user_id_list: Optional[list] = None, - team_id: Optional[str] = None, - team_id_list: Optional[list] = None, - key_val: Optional[dict] = None, - table_name: Optional[ - Literal[ - "user", - "key", - "config", - "spend", - "enduser", - "budget", - "team", - "user_notification", - "combined_view", - ] - ] = None, + token: str | list | None = None, + user_id: str | None = None, + user_id_list: list | None = None, + team_id: str | None = None, + team_id_list: list | None = None, + key_val: dict | None = None, + table_name: Literal[ + "user", "key", "config", "spend", "enduser", "budget", "team", "user_notification", "combined_view" + ] + | None = None, query_type: Literal["find_unique", "find_all"] = "find_unique", - expires: Optional[datetime] = None, - reset_at: Optional[datetime] = None, - offset: Optional[int] = None, # pagination, what row number to start from - limit: Optional[int] = None, # pagination, number of rows to getch when find_all==True - parent_otel_span: Optional[Span] = None, - proxy_logging_obj: Optional[ProxyLogging] = None, - budget_id_list: Optional[List[str]] = None, + expires: datetime | None = None, + reset_at: datetime | None = None, + offset: int | None = None, # pagination, what row number to start from + limit: int | None = None, # pagination, number of rows to getch when find_all==True + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, + budget_id_list: list[str] | None = None, check_deprecated: bool = True, ): args_passed_in = locals() start_time = time.time() - hashed_token: Optional[str] = None + hashed_token: str | None = None try: response: Any = None if (token is not None and table_name is None) or (table_name is not None and table_name == "key"): @@ -3773,7 +3759,7 @@ class PrismaClient: if response["team_blocked"] is None: response["team_blocked"] = False - team_member: Optional[Member] = None + team_member: Member | None = None if response["team_members_with_roles"] is not None and response["user_id"] is not None: ## find the team member corresponding to user id """ @@ -3933,7 +3919,7 @@ class PrismaClient: tasks.append(updated_table_row) await asyncio.gather(*tasks) # invalidate cache so other pods see writes from save_config - for k in data.keys(): + for k in data: await invalidate_config_param(k) verbose_proxy_logger.info("Data Inserted into Config Table") elif table_name == "spend": @@ -3962,7 +3948,7 @@ class PrismaClient: except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception in insert_data: {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception in insert_data: {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -3987,15 +3973,15 @@ class PrismaClient: ) async def update_data( self, - token: Optional[str] = None, + token: str | None = None, data: dict = {}, - data_list: Optional[List] = None, - user_id: Optional[str] = None, - team_id: Optional[str] = None, + data_list: list | None = None, + user_id: str | None = None, + team_id: str | None = None, query_type: Literal["update", "update_many"] = "update", - table_name: Optional[Literal["user", "key", "config", "spend", "team", "enduser", "budget"]] = None, - update_key_values: Optional[dict] = None, - update_key_values_custom_query: Optional[dict] = None, + table_name: Literal["user", "key", "config", "spend", "team", "enduser", "budget"] | None = None, + update_key_values: dict | None = None, + update_key_values_custom_query: dict | None = None, ): """ Update existing data @@ -4211,7 +4197,7 @@ class PrismaClient: except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception - update_data: {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception - update_data: {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -4236,10 +4222,10 @@ class PrismaClient: ) async def delete_data( self, - tokens: Optional[List] = None, - team_id_list: Optional[List] = None, - table_name: Optional[Literal["user", "key", "config", "spend", "team"]] = None, - user_id: Optional[str] = None, + tokens: list | None = None, + team_id_list: list | None = None, + table_name: Literal["user", "key", "config", "spend", "team"] | None = None, + user_id: str | None = None, ): """ Allow user to delete a key(s) @@ -4248,7 +4234,7 @@ class PrismaClient: """ start_time = time.time() try: - if tokens is not None and isinstance(tokens, List): + if tokens is not None and isinstance(tokens, list): hashed_tokens = [] for token in tokens: if isinstance(token, str) and token.startswith("sk-"): @@ -4267,17 +4253,17 @@ class PrismaClient: ) verbose_proxy_logger.debug("deleted_tokens: %s", deleted_tokens) return {"deleted_keys": deleted_tokens} - elif table_name == "team" and team_id_list is not None and isinstance(team_id_list, List): + elif table_name == "team" and team_id_list is not None and isinstance(team_id_list, list): # admin only endpoint -> `/team/delete` await TeamRepository(self).table.delete_many(where={"team_id": {"in": team_id_list}}) return {"deleted_teams": team_id_list} - elif table_name == "key" and team_id_list is not None and isinstance(team_id_list, List): + elif table_name == "key" and team_id_list is not None and isinstance(team_id_list, list): # admin only endpoint -> `/team/delete` await VerificationTokenRepository(self).table.delete_many(where={"team_id": {"in": team_id_list}}) except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception - delete_data: {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception - delete_data: {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -4310,7 +4296,7 @@ class PrismaClient: except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception connect(): {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception connect(): {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -4340,7 +4326,7 @@ class PrismaClient: except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception disconnect(): {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception disconnect(): {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -4715,7 +4701,7 @@ class PrismaClient: self._cleanup_engine_watcher() asyncio.create_task(self._start_engine_watcher()) - async def _run_reconnect_cycle(self, timeout_seconds: Optional[float] = None) -> None: + async def _run_reconnect_cycle(self, timeout_seconds: float | None = None) -> None: """ Run a reconnect cycle with a single overall timeout budget. @@ -4825,7 +4811,7 @@ class PrismaClient: self, force: bool, reason: str, - timeout_seconds: Optional[float], + timeout_seconds: float | None, ) -> bool: now = time.time() if force is False and now - self._db_last_reconnect_attempt_ts < self._db_reconnect_cooldown_seconds: @@ -4873,8 +4859,8 @@ class PrismaClient: self, reason: str, force: bool = False, - timeout_seconds: Optional[float] = None, - lock_timeout_seconds: Optional[float] = None, + timeout_seconds: float | None = None, + lock_timeout_seconds: float | None = None, ) -> bool: """ Attempt to reconnect the Prisma client in a singleflight manner. @@ -5029,7 +5015,7 @@ class PrismaClient: except Exception as e: import traceback - error_msg = f"LiteLLM Prisma Client Exception disconnect(): {str(e)}" + error_msg = f"LiteLLM Prisma Client Exception disconnect(): {e!s}" print_verbose(error_msg) error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() @@ -5093,7 +5079,7 @@ class PrismaClient: ) # Health Check Database Methods - def _validate_response_time(self, response_time_ms: Optional[float]) -> Optional[float]: + def _validate_response_time(self, response_time_ms: float | None) -> float | None: """Validate and clean response time value""" if response_time_ms is None: return None @@ -5104,7 +5090,7 @@ class PrismaClient: verbose_proxy_logger.warning(f"Invalid response_time_ms value: {response_time_ms}") return None - def _clean_details(self, details: Optional[dict]) -> Optional[dict]: + def _clean_details(self, details: dict | None) -> dict | None: """Clean and validate details JSON""" if not isinstance(details, dict): return None @@ -5120,11 +5106,11 @@ class PrismaClient: status: str, healthy_count: int = 0, unhealthy_count: int = 0, - error_message: Optional[str] = None, - response_time_ms: Optional[float] = None, - details: Optional[dict] = None, - checked_by: Optional[str] = None, - model_id: Optional[str] = None, + error_message: str | None = None, + response_time_ms: float | None = None, + details: dict | None = None, + checked_by: str | None = None, + model_id: str | None = None, ): """Save health check result to database""" try: @@ -5157,10 +5143,10 @@ class PrismaClient: async def get_health_check_history( self, - model_name: Optional[str] = None, + model_name: str | None = None, limit: int = 100, offset: int = 0, - status_filter: Optional[str] = None, + status_filter: str | None = None, ): """ Get health check history with optional filtering @@ -5221,7 +5207,6 @@ async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): if hasattr(user_row, "model_dump_json") and callable(getattr(user_row, "model_dump_json")): cache_value = user_row.model_dump_json() cache.set_cache(key=cache_key, value=cache_value, ttl=600) # store for 10 minutes - return def _should_use_smtp_ssl(smtp_port: int) -> bool: @@ -5240,9 +5225,9 @@ def _create_smtp_connection(smtp_host: str, smtp_port: int) -> smtplib.SMTP: async def send_email( - receiver_email: Optional[str] = None, - subject: Optional[str] = None, - html: Optional[str] = None, + receiver_email: str | None = None, + subject: str | None = None, + html: str | None = None, ): """ smtp_host, @@ -5393,7 +5378,7 @@ class ProxyUpdateSpend: n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, - end_user_list_transactions: Dict[str, float], + end_user_list_transactions: dict[str, float], ): for i in range(n_retry_times + 1): start_time = time.time() @@ -5431,9 +5416,9 @@ class ProxyUpdateSpend: async def update_spend_logs( n_retry_times: int, prisma_client: PrismaClient, - db_writer_client: Optional[AsyncHTTPHandler], + db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, - logs_to_process: Optional[List[Dict[str, Any]]] = None, + logs_to_process: list[dict[str, Any]] | None = None, ): BATCH_SIZE = 1000 # Preferred size of each batch to write to the database MAX_LOGS_PER_INTERVAL = 10000 # Maximum number of logs to flush in a single interval @@ -5458,7 +5443,7 @@ class ProxyUpdateSpend: if len(logs_to_process) > 0 and base_url is not None and db_writer_client is not None: if not base_url.endswith("/"): base_url += "/" - verbose_proxy_logger.debug("base_url: {}".format(base_url)) + verbose_proxy_logger.debug(f"base_url: {base_url}") json_data = json.dumps(logs_to_process) response = await db_writer_client.post( url=base_url + "spend/update", @@ -5526,7 +5511,7 @@ class ProxyUpdateSpend: async def update_spend( prisma_client: PrismaClient, - db_writer_client: Optional[AsyncHTTPHandler], + db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, ): """ @@ -5554,7 +5539,7 @@ async def update_spend( # Check queue size with lock protection async with prisma_client._spend_log_transactions_lock: queue_size = len(prisma_client.spend_log_transactions) - verbose_proxy_logger.debug("Spend Logs transactions: {}".format(queue_size)) + verbose_proxy_logger.debug(f"Spend Logs transactions: {queue_size}") async with prisma_client._tool_usage_transactions_lock: tool_usage_queue_size = len(prisma_client.tool_usage_transactions) @@ -5619,7 +5604,7 @@ async def update_daily_tag_spend( async def update_spend_logs_job( prisma_client: PrismaClient, - db_writer_client: Optional[AsyncHTTPHandler], + db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, ): """ @@ -5692,7 +5677,7 @@ async def update_spend_logs_job( async def _monitor_spend_logs_queue( prisma_client: PrismaClient, - db_writer_client: Optional[AsyncHTTPHandler], + db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, ): """ @@ -5828,7 +5813,7 @@ def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_ """ import traceback - error_msg = f"[Non-Blocking]LiteLLM Prisma Client Exception - update spend logs: {str(e)}" + error_msg = f"[Non-Blocking]LiteLLM Prisma Client Exception - update spend logs: {e!s}" error_traceback = error_msg + "\n" + traceback.format_exc() end_time = time.time() _duration = end_time - start_time @@ -5849,7 +5834,7 @@ def _get_month_end_date(today: date) -> date: return date(today.year, today.month + 1, 1) - timedelta(days=1) -def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: Optional[float]): +def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None): if soft_budget_limit is None: # If there's no limit, we can't exceed it. return False @@ -5876,7 +5861,7 @@ def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: Opti return False -def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: Optional[float]) -> Optional[tuple]: +def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: if soft_budget_limit is None: return None @@ -5928,7 +5913,7 @@ def _to_ns(dt): def _check_and_merge_model_level_guardrails( data: dict, - llm_router: Optional[Router], + llm_router: Router | None, trust_client_model_info: bool = True, ) -> dict: """ @@ -5962,7 +5947,7 @@ def _check_and_merge_model_level_guardrails( # Medium on #29654). team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id") - model_level_guardrails: Optional[list] = None + model_level_guardrails: list | None = None if model_id is not None: deployment = llm_router.get_deployment(model_id=model_id) if deployment is None: @@ -6059,7 +6044,7 @@ def get_error_message_str(e: Exception) -> str: return error_message -def _get_redoc_url() -> Optional[str]: +def _get_redoc_url() -> str | None: """ Get the Redoc URL from the environment variables. @@ -6076,7 +6061,7 @@ def _get_redoc_url() -> Optional[str]: return "/redoc" -def _get_docs_url() -> Optional[str]: +def _get_docs_url() -> str | None: """ Get the docs (Swagger UI) URL from the environment variables. @@ -6093,7 +6078,7 @@ def _get_docs_url() -> Optional[str]: return "/" -def _get_openapi_url() -> Optional[str]: +def _get_openapi_url() -> str | None: """ Get the OpenAPI JSON URL from the environment variables. @@ -6120,7 +6105,7 @@ def handle_exception_on_proxy(e: Exception) -> ProxyException: if isinstance(e, HTTPException): return ProxyException( - message=getattr(e, "detail", f"error({str(e)})"), + message=getattr(e, "detail", f"error({e!s})"), type=ProxyErrorTypes.internal_server_error, param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), @@ -6136,7 +6121,7 @@ def handle_exception_on_proxy(e: Exception) -> ProxyException: ) -def _premium_user_check(feature: Optional[str] = None): +def _premium_user_check(feature: str | None = None): """ Raises an HTTPException if the user is not a premium user """ @@ -6156,7 +6141,7 @@ def _premium_user_check(feature: Optional[str] = None): ) -def is_known_model(model: Optional[str], llm_router: Optional[Router]) -> bool: +def is_known_model(model: str | None, llm_router: Router | None) -> bool: """ Returns True if the model is in the llm_router model names """ @@ -6206,7 +6191,7 @@ def join_paths(base_path: str, route: str) -> str: return final_path -def get_custom_url(request_base_url: str, route: Optional[str] = None) -> str: +def get_custom_url(request_base_url: str, route: str | None = None) -> str: # Use environment variable value, otherwise use URL from request server_base_url = get_proxy_base_url() if server_base_url is not None: @@ -6226,7 +6211,7 @@ def get_custom_url(request_base_url: str, route: Optional[str] = None) -> str: return join_paths(base_url, server_root_path) -def get_proxy_base_url() -> Optional[str]: +def get_proxy_base_url() -> str | None: """ Get the proxy base url from the environment variables. """ @@ -6243,7 +6228,7 @@ def get_server_root_path() -> str: return os.getenv("SERVER_ROOT_PATH", "") -def normalize_route_for_root_path(route: str) -> Optional[str]: +def normalize_route_for_root_path(route: str) -> str | None: """Strip SERVER_ROOT_PATH prefix. Returns de-prefixed route, or None if route is not under root path.""" root_path = get_server_root_path() if root_path and root_path != "/": @@ -6283,7 +6268,7 @@ def is_valid_api_key(key: str) -> bool: return False -def construct_database_url_from_env_vars() -> Optional[str]: +def construct_database_url_from_env_vars() -> str | None: """ Construct a DATABASE_URL from individual environment variables. Returns: @@ -6324,15 +6309,15 @@ async def get_available_models_for_user( user_api_key_dict: "UserAPIKeyAuth", llm_router: Optional["Router"], general_settings: dict, - user_model: Optional[str], + user_model: str | None, prisma_client: Optional["PrismaClient"] = None, proxy_logging_obj: Optional["ProxyLogging"] = None, - team_id: Optional[str] = None, + team_id: str | None = None, include_model_access_groups: bool = False, only_model_access_groups: bool = False, return_wildcard_routes: bool = False, user_api_key_cache: Optional["UserApiKeyCache"] = None, -) -> List[str]: +) -> list[str]: """ Get the list of models available to a user based on their API key and team permissions. @@ -6376,7 +6361,7 @@ async def get_available_models_for_user( ) # Get team models - team_models: List[str] = user_api_key_dict.team_models + team_models: list[str] = user_api_key_dict.team_models # If specific team_id is provided, validate and get team models if team_id and prisma_client and proxy_logging_obj and user_api_key_cache: @@ -6421,7 +6406,7 @@ def create_model_info_response( model_id: str, provider: str, include_metadata: bool = False, - fallback_type: Optional[str] = None, + fallback_type: str | None = None, llm_router: Optional["Router"] = None, get_model_info: Callable[[str], ModelInfo] = litellm.get_model_info, ) -> ModelInfoResponse: @@ -6494,7 +6479,7 @@ def create_model_info_response( def validate_model_access( model_id: str, - available_models: List[str], + available_models: list[str], ) -> None: """ Validate that a model is accessible to the user. @@ -6523,11 +6508,11 @@ def validate_model_access( if model_id not in available_models: raise HTTPException( status_code=404, - detail="The model `{}` does not exist or is not accessible".format(model_id), + detail=f"The model `{model_id}` does not exist or is not accessible", ) -_PRESERVED_NONE_FIELDS: List[tuple[str, str]] = [ +_PRESERVED_NONE_FIELDS: list[tuple[str, str]] = [ ("message", "content"), # null when tool_calls present (issue #6677) ("message", "role"), # always required by OpenAI spec ("delta", "content"), # null in streaming chunks @@ -6536,9 +6521,9 @@ _PRESERVED_NONE_FIELDS: List[tuple[str, str]] = [ def model_dump_with_preserved_fields( obj: Any, - preserve_fields: Optional[List[str]] = None, + preserve_fields: list[str] | None = None, exclude_unset: bool = True, -) -> Dict[str, Any]: +) -> dict[str, Any]: """ Serialize a Pydantic model to a dictionary while preserving specific fields even if they are None. diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 4e7890e6ed8..6cc053012f8 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Optional +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -27,10 +27,10 @@ router = APIRouter() async def _update_request_data_with_litellm_managed_vector_store_registry( - data: Dict, + data: dict, vector_store_id: str, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> Dict: + user_api_key_dict: UserAPIKeyAuth | None = None, +) -> dict: """ Update the request data with the litellm managed vector store registry. @@ -42,7 +42,7 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( Raises: HTTPException: If user doesn't have access to the vector store """ - vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = await get_litellm_managed_vector_store( + vector_store_to_run: LiteLLM_ManagedVectorStore | None = await get_litellm_managed_vector_store( vector_store_id=vector_store_id ) if vector_store_to_run is not None: @@ -341,10 +341,10 @@ async def vector_store_retrieve( async def vector_store_list( request: Request, fastapi_response: Response, - after: Optional[str] = None, - before: Optional[str] = None, - limit: Optional[int] = 20, - order: Optional[str] = "desc", + after: str | None = None, + before: str | None = None, + limit: int | None = 20, + order: str | None = "desc", user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 84054f1398d..811597f3821 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -10,7 +10,7 @@ All /vector_store management endpoints import copy import json -from typing import Any, Dict, List, Optional +from typing import Any from fastapi import APIRouter, Depends, HTTPException @@ -58,7 +58,7 @@ _REDACT_LITELLM_PARAMS_MAX_DEPTH = 10 # management responses), so the cache doesn't widen the disclosure surface. _EMBEDDING_CONFIG_CACHE_TTL = 60 _EMBEDDING_CONFIG_CACHE_MAX_SIZE = 256 -_embedding_config_cache: Optional[InMemoryCache] = None +_embedding_config_cache: InMemoryCache | None = None def _get_embedding_config_cache() -> InMemoryCache: @@ -103,7 +103,7 @@ def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> An return json.dumps(_redact_sensitive_litellm_params(parsed, _depth + 1)) if not isinstance(litellm_params, dict): return litellm_params - out: Dict[str, Any] = {} + out: dict[str, Any] = {} for k, v in litellm_params.items(): if _LITELLM_PARAMS_MASKER.is_sensitive_key(k): out[k] = REDACTED_BY_LITELM_STRING @@ -141,7 +141,7 @@ async def _fetch_and_authorize_vector_store( return typed -def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> Optional[Dict[str, Any]]: +def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> dict[str, Any] | None: """ Resolve embedding config from router's config-defined models. @@ -177,7 +177,7 @@ def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> O litellm_params = deployment.litellm_params # Build embedding config from model params - embedding_config: Dict[str, Any] = {} + embedding_config: dict[str, Any] = {} # Extract api_key api_key = getattr(litellm_params, "api_key", None) @@ -211,13 +211,13 @@ def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> O ) return embedding_config except Exception as e: - verbose_proxy_logger.debug(f"Error resolving embedding config from router for model {model_name}: {str(e)}") + verbose_proxy_logger.debug(f"Error resolving embedding config from router for model {model_name}: {e!s}") continue return None -async def _resolve_embedding_config_from_db(embedding_model: str, prisma_client) -> Optional[Dict[str, Any]]: +async def _resolve_embedding_config_from_db(embedding_model: str, prisma_client) -> dict[str, Any] | None: """ Resolve embedding config from database model configuration. @@ -299,13 +299,13 @@ async def _resolve_embedding_config_from_db(embedding_model: str, prisma_client) ) return embedding_config except Exception as e: - verbose_proxy_logger.debug(f"Error resolving embedding config for model {model_name}: {str(e)}") + verbose_proxy_logger.debug(f"Error resolving embedding config for model {model_name}: {e!s}") continue return None -async def _resolve_embedding_config(embedding_model: str, prisma_client, llm_router=None) -> Optional[Dict[str, Any]]: +async def _resolve_embedding_config(embedding_model: str, prisma_client, llm_router=None) -> dict[str, Any] | None: """ Resolve embedding config from either router (config-defined) or database models. @@ -387,13 +387,13 @@ async def create_vector_store_in_db( vector_store_id: str, custom_llm_provider: str, prisma_client, - vector_store_name: Optional[str] = None, - vector_store_description: Optional[str] = None, - vector_store_metadata: Optional[Dict] = None, - litellm_params: Optional[Dict] = None, - litellm_credential_name: Optional[str] = None, - team_id: Optional[str] = None, - user_id: Optional[str] = None, + vector_store_name: str | None = None, + vector_store_description: str | None = None, + vector_store_metadata: dict | None = None, + litellm_params: dict | None = None, + litellm_credential_name: str | None = None, + team_id: str | None = None, + user_id: str | None = None, ) -> LiteLLM_ManagedVectorStore: """ Helper function to create a vector store in the database. @@ -425,7 +425,7 @@ async def create_vector_store_in_db( ) # Prepare data for database - data_to_create: Dict[str, Any] = { + data_to_create: dict[str, Any] = { "vector_store_id": vector_store_id, "custom_llm_provider": custom_llm_provider, } @@ -512,7 +512,7 @@ async def new_vector_store( # Extract and validate metadata metadata = vector_store.get("vector_store_metadata") - validated_metadata: Optional[Dict] = None + validated_metadata: dict | None = None if metadata is not None and isinstance(metadata, dict): validated_metadata = metadata @@ -542,7 +542,7 @@ async def new_vector_store( "vector_store": response_vs, } except Exception as e: - verbose_proxy_logger.exception(f"Error creating vector store: {str(e)}") + verbose_proxy_logger.exception(f"Error creating vector store: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -576,7 +576,7 @@ async def list_vector_stores( from litellm.proxy.proxy_server import prisma_client - vector_store_map: Dict[str, LiteLLM_ManagedVectorStore] = {} + vector_store_map: dict[str, LiteLLM_ManagedVectorStore] = {} db_vector_store_ids: set = set() try: @@ -594,7 +594,7 @@ async def list_vector_stores( if litellm.vector_store_registry is not None: in_memory_vector_stores = copy.deepcopy(litellm.vector_store_registry.vector_stores) - vector_stores_to_delete_from_memory: List[str] = [] + vector_stores_to_delete_from_memory: list[str] = [] for vector_store in in_memory_vector_stores: vector_store_id = vector_store.get("vector_store_id", None) @@ -647,7 +647,7 @@ async def list_vector_stores( return response except Exception as e: - verbose_proxy_logger.exception(f"Error listing vector stores: {str(e)}") + verbose_proxy_logger.exception(f"Error listing vector stores: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -727,7 +727,7 @@ async def delete_vector_store( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception(f"Error deleting vector store: {str(e)}") + verbose_proxy_logger.exception(f"Error deleting vector store: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -764,7 +764,7 @@ async def get_vector_store_info( vector_store_metadata = vector_store.get("vector_store_metadata") # Parse metadata if it's a JSON string - parsed_metadata: Optional[dict] = None + parsed_metadata: dict | None = None if isinstance(vector_store_metadata, str): parsed_metadata = json.loads(vector_store_metadata) elif isinstance(vector_store_metadata, dict): @@ -799,7 +799,7 @@ async def get_vector_store_info( # the catch-all below would otherwise rewrite them as 500. raise except Exception as e: - verbose_proxy_logger.exception(f"Error getting vector store info: {str(e)}") + verbose_proxy_logger.exception(f"Error getting vector store info: {e!s}") raise HTTPException(status_code=500, detail=str(e)) @@ -888,5 +888,5 @@ async def update_vector_store( # as 500 with the original status code embedded in the detail. raise except Exception as e: - verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}") + verbose_proxy_logger.exception(f"Error updating vector store: {e!s}") raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 05149a0f6c0..099a94812eb 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -1,6 +1,6 @@ import json import re -from typing import Any, Dict, Literal, Optional +from typing import Any, Literal from fastapi import HTTPException, Request @@ -52,7 +52,7 @@ def assert_proxy_admin_for_vector_store_index_management( ) -def _suffix_after_index_name(request_path: str, index_name: str) -> Optional[str]: +def _suffix_after_index_name(request_path: str, index_name: str) -> str | None: """Return the path suffix after ``/indexes/{index_name}``, or None if absent.""" match = re.search(rf"/indexes/{re.escape(index_name)}(?=$|[/?])", request_path) if match is None: @@ -94,7 +94,7 @@ def _is_vector_store_index_lifecycle_request( def _object_permission_allows_vector_store( - object_permission: Optional[LiteLLM_ObjectPermissionTable], + object_permission: LiteLLM_ObjectPermissionTable | None, vector_store_id: str, ) -> bool: """Returns True if an object permission explicitly allowlists the vector store.""" @@ -107,8 +107,8 @@ def _object_permission_allows_vector_store( async def _get_object_permission_for_id( - object_permission_id: Optional[str], -) -> Optional[LiteLLM_ObjectPermissionTable]: + object_permission_id: str | None, +) -> LiteLLM_ObjectPermissionTable | None: """Load an object permission record by id, using the shared cache/DB helper.""" if not object_permission_id: return None @@ -172,7 +172,7 @@ async def can_user_access_vector_store( if _object_permission_allows_vector_store(key_object_permission, vector_store_id): return True - team_object_permission: Optional[LiteLLM_ObjectPermissionTable] = user_api_key_dict.team_object_permission + team_object_permission: LiteLLM_ObjectPermissionTable | None = user_api_key_dict.team_object_permission if team_object_permission is None: team_object_permission = await _get_object_permission_for_id(user_api_key_dict.team_object_permission_id) if _object_permission_allows_vector_store(team_object_permission, vector_store_id): @@ -186,7 +186,7 @@ async def can_user_access_vector_store( async def get_litellm_managed_vector_store( vector_store_id: str, -) -> Optional[LiteLLM_ManagedVectorStore]: +) -> LiteLLM_ManagedVectorStore | None: """ Resolve a LiteLLM-managed vector store from the registry or shared cache. @@ -262,7 +262,7 @@ async def assert_user_can_access_vector_store_id( vector_store_id: str, user_api_key_dict: UserAPIKeyAuth, detail: str = "Access denied: You do not have permission to access this vector store", -) -> Optional[LiteLLM_ManagedVectorStore]: +) -> LiteLLM_ManagedVectorStore | None: """ Resolve a managed vector store id and enforce ownership if it exists. @@ -291,8 +291,8 @@ def _does_endpoint_match(endpoint_path: str, request_path: str) -> bool: def check_vector_store_permission( index_name: str, permission: str, - key_metadata: Optional[Dict[str, Any]], - team_metadata: Optional[Dict[str, Any]], + key_metadata: dict[str, Any] | None, + team_metadata: dict[str, Any] | None, ) -> bool: """ Check if a specific permission is allowed for a given vector store index. @@ -343,7 +343,7 @@ def is_allowed_to_call_vector_store_endpoint( index_name: str, request: Request, user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Literal[True]]: +) -> Literal[True] | None: """ Check if the user is allowed to call the vector store endpoint. @@ -432,7 +432,7 @@ def is_allowed_to_call_vector_store_files_endpoint( vector_store_id: str, request: Request, user_api_key_dict: UserAPIKeyAuth, -) -> Optional[Literal[True]]: +) -> Literal[True] | None: if ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -453,7 +453,7 @@ def is_allowed_to_call_vector_store_files_endpoint( request_route = get_request_route(request) - permission_type: Optional[str] = None + permission_type: str | None = None for endpoint in provider_vector_store_endpoints.get("read", ()): if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route): permission_type = "read" diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 44935fc57c9..06bcc524ea7 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Dict, Optional +from typing import TYPE_CHECKING, Optional from fastapi import APIRouter, Depends, Request, Response from fastapi.responses import ORJSONResponse @@ -31,11 +31,11 @@ router = APIRouter() def _update_request_data_with_managed_file_id( - data: Dict, + data: dict, file_id: str, request: Request, llm_router: Optional["Router"] = None, -) -> tuple[Dict, Optional[str]]: +) -> tuple[dict, str | None]: """ Update request data with model routing information from managed file ID. @@ -188,7 +188,7 @@ async def _authorize_model_routing_hint( *, model: str, llm_router: Optional["Router"], - user_api_key_dict: Optional[UserAPIKeyAuth], + user_api_key_dict: UserAPIKeyAuth | None, ) -> None: if user_api_key_dict is None: return @@ -215,11 +215,11 @@ async def _authorize_model_routing_hint( async def _update_request_data_with_model_routing_hint( - data: Dict, + data: dict, request: Request, llm_router: Optional["Router"] = None, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> Dict: + user_api_key_dict: UserAPIKeyAuth | None = None, +) -> dict: if data.get("api_key") is not None or data.get("api_base") is not None: return data @@ -323,12 +323,12 @@ async def _update_request_data_with_model_routing_hint( def _update_request_data_with_litellm_managed_vector_store_registry( - data: Dict, + data: dict, vector_store_id: str, llm_router: Optional["Router"] = None, - managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None, + managed_vector_store: LiteLLM_ManagedVectorStore | None = None, should_lookup_registry: bool = True, -) -> Dict: +) -> dict: """ Update request data with model routing information from managed vector store. @@ -416,9 +416,9 @@ def _update_request_data_with_litellm_managed_vector_store_registry( async def _resolve_provider( *, - data: Dict, + data: dict, request: Request, -) -> Optional[LlmProviders]: +) -> LlmProviders | None: provider = ( data.get("custom_llm_provider") or get_custom_llm_provider_from_request_headers(request=request) @@ -439,7 +439,7 @@ async def _resolve_provider( def _maybe_check_permissions( *, - provider: Optional[LlmProviders], + provider: LlmProviders | None, vector_store_id: str, request: Request, user_api_key_dict: UserAPIKeyAuth, @@ -593,7 +593,7 @@ async def vector_store_file_list( ) query_params = dict(request.query_params) - data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id} + data: dict[str, str | None] = {"vector_store_id": vector_store_id} data.update(query_params) data["vector_store_id"] = vector_store_id managed_vector_store = await assert_user_can_access_vector_store_id( @@ -689,7 +689,7 @@ async def vector_store_file_retrieve( version, ) - data: Dict[str, str] = { + data: dict[str, str] = { "vector_store_id": vector_store_id, "file_id": file_id, } @@ -791,7 +791,7 @@ async def vector_store_file_content( version, ) - data: Dict[str, str] = { + data: dict[str, str] = { "vector_store_id": vector_store_id, "file_id": file_id, } @@ -995,7 +995,7 @@ async def vector_store_file_delete( version, ) - data: Dict[str, str] = { + data: dict[str, str] = { "vector_store_id": vector_store_id, "file_id": file_id, } diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index 79a188817bf..767e526804c 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -11,7 +11,6 @@ Logging Pass-Through Endpoints import base64 import os from base64 import b64encode -from typing import Optional from urllib.parse import unquote import httpx @@ -61,7 +60,7 @@ def _normalize_langfuse_base_url(base_target_url: str) -> str: except Exception as e: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": f"Invalid Langfuse host: {str(e)}"}, + detail={"error": f"Invalid Langfuse host: {e!s}"}, ) if base_url.scheme not in ("http", "https") or not base_url.host: @@ -104,8 +103,8 @@ def _validate_langfuse_proxy_path(endpoint: str) -> str: def _get_langfuse_proxy_credentials( *, dynamic_host_supplied: bool, - dynamic_langfuse_public_key: Optional[str], - dynamic_langfuse_secret_key: Optional[str], + dynamic_langfuse_public_key: str | None, + dynamic_langfuse_secret_key: str | None, ): if dynamic_host_supplied: if not dynamic_langfuse_public_key or not dynamic_langfuse_secret_key: @@ -138,7 +137,7 @@ def _build_langfuse_proxy_target( except SSRFError as e: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": f"Invalid Langfuse host: {str(e)}"}, + detail={"error": f"Invalid Langfuse host: {e!s}"}, ) custom_headers["Host"] = host_header return target_url, custom_headers @@ -173,15 +172,15 @@ async def langfuse_proxy_route( decoded_str = decoded_bytes.decode("utf-8") api_key = decoded_str.split(":")[1] # assume api key is passed in as secret key - user_api_key_dict = await user_api_key_auth(request=request, api_key="Bearer {}".format(api_key)) + user_api_key_dict = await user_api_key_auth(request=request, api_key=f"Bearer {api_key}") - callback_settings_obj: Optional[TeamCallbackMetadata] = _get_dynamic_logging_metadata( + callback_settings_obj: TeamCallbackMetadata | None = _get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) - dynamic_langfuse_public_key: Optional[str] = None - dynamic_langfuse_secret_key: Optional[str] = None - dynamic_langfuse_host: Optional[str] = None + dynamic_langfuse_public_key: str | None = None + dynamic_langfuse_secret_key: str | None = None + dynamic_langfuse_host: str | None = None if callback_settings_obj is not None and callback_settings_obj.callback_vars is not None: for k, v in callback_settings_obj.callback_vars.items(): if k == "langfuse_public_key": @@ -206,7 +205,7 @@ async def langfuse_proxy_route( dynamic_host_supplied=dynamic_host_supplied, ) - langfuse_combined_key = "Basic " + b64encode(f"{langfuse_public_key}:{langfuse_secret_key}".encode("utf-8")).decode( + langfuse_combined_key = "Basic " + b64encode(f"{langfuse_public_key}:{langfuse_secret_key}".encode()).decode( "ascii" ) target_headers["Authorization"] = langfuse_combined_key diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 5809967cff5..d45b2fc29a6 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -1,6 +1,6 @@ #### Video Endpoints ##### -from typing import Any, Dict, Optional +from typing import Any import orjson from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile @@ -44,7 +44,7 @@ router = APIRouter() async def video_generation( request: Request, fastapi_response: Response, - input_reference: Optional[UploadFile] = File(None), + input_reference: UploadFile | None = File(None), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -160,7 +160,7 @@ async def video_list( # Read query parameters query_params = dict(request.query_params) - data: Dict[str, Any] = {"query_params": query_params} + data: dict[str, Any] = {"query_params": query_params} # Extract custom_llm_provider from headers, query params, or body custom_llm_provider = ( @@ -245,7 +245,7 @@ async def video_status( ) # Create data with video_id - data: Dict[str, Any] = {"video_id": video_id} + data: dict[str, Any] = {"video_id": video_id} decoded = decode_video_id_with_provider(video_id) provider_from_id = decoded.get("custom_llm_provider") @@ -344,7 +344,7 @@ async def video_content( ) # Create data with video_id - data: Dict[str, Any] = {"video_id": video_id} + data: dict[str, Any] = {"video_id": video_id} decoded = decode_video_id_with_provider(video_id) provider_from_id = decoded.get("custom_llm_provider") @@ -654,7 +654,7 @@ async def video_get_character( ) original_requested_character_id = character_id - data: Dict[str, Any] = {"character_id": character_id} + data: dict[str, Any] = {"character_id": character_id} decoded = decode_character_id_with_provider(character_id) provider_from_id = decoded.get("custom_llm_provider") diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index 412a0e87d88..f891425322e 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -1,11 +1,11 @@ -from typing import Any, Dict, Optional +from typing import Any import orjson from litellm.types.videos.utils import encode_character_id_with_provider -def extract_model_from_target_model_names(target_model_names: Any) -> Optional[str]: +def extract_model_from_target_model_names(target_model_names: Any) -> str | None: if isinstance(target_model_names, str): target_model_names = [m.strip() for m in target_model_names.split(",") if m.strip()] elif not isinstance(target_model_names, list): @@ -13,7 +13,7 @@ def extract_model_from_target_model_names(target_model_names: Any) -> Optional[s return target_model_names[0] if target_model_names else None -def get_custom_provider_from_data(data: Dict[str, Any]) -> Optional[str]: +def get_custom_provider_from_data(data: dict[str, Any]) -> str | None: custom_llm_provider = data.get("custom_llm_provider") if custom_llm_provider: return custom_llm_provider @@ -35,7 +35,7 @@ def get_custom_provider_from_data(data: Dict[str, Any]) -> Optional[str]: return None -def encode_character_id_in_response(response: Any, custom_llm_provider: str, model_id: Optional[str]) -> Any: +def encode_character_id_in_response(response: Any, custom_llm_provider: str, model_id: str | None) -> Any: if isinstance(response, dict) and response.get("id"): response["id"] = encode_character_id_with_provider( character_id=response["id"], diff --git a/litellm/proxy_auth/__init__.py b/litellm/proxy_auth/__init__.py index 27624a94fb9..c7f29867f3b 100644 --- a/litellm/proxy_auth/__init__.py +++ b/litellm/proxy_auth/__init__.py @@ -15,16 +15,16 @@ Usage: from .credentials import ( AccessToken, - TokenCredential, AzureADCredential, GenericOAuth2Credential, ProxyAuthHandler, + TokenCredential, ) __all__ = [ "AccessToken", - "TokenCredential", "AzureADCredential", "GenericOAuth2Credential", "ProxyAuthHandler", + "TokenCredential", ] diff --git a/litellm/proxy_auth/credentials.py b/litellm/proxy_auth/credentials.py index 5383e17e793..e2a28889099 100644 --- a/litellm/proxy_auth/credentials.py +++ b/litellm/proxy_auth/credentials.py @@ -7,7 +7,7 @@ It follows the same TokenCredential protocol used by Azure SDK. import time from dataclasses import dataclass -from typing import Any, Optional, Protocol, runtime_checkable +from typing import Any, Protocol, runtime_checkable @dataclass @@ -71,7 +71,7 @@ class AzureADCredential: cred = AzureADCredential(credential=azure_cred) """ - def __init__(self, credential: Optional[Any] = None): + def __init__(self, credential: Any | None = None): """ Initialize with an optional Azure credential. @@ -137,7 +137,7 @@ class GenericOAuth2Credential: self.client_id = client_id self.client_secret = client_secret self.token_url = token_url - self._cached_token: Optional[AccessToken] = None + self._cached_token: AccessToken | None = None def get_token(self, scope: str) -> AccessToken: """ @@ -214,7 +214,7 @@ class ProxyAuthHandler: """ self.credential = credential self.scope = scope - self._cached_token: Optional[AccessToken] = None + self._cached_token: AccessToken | None = None def get_token(self) -> AccessToken: """ diff --git a/litellm/rag/__init__.py b/litellm/rag/__init__.py index e387cf837eb..610cf04c0a3 100644 --- a/litellm/rag/__init__.py +++ b/litellm/rag/__init__.py @@ -7,7 +7,7 @@ Upload -> (OCR) -> Chunk -> Embed -> Vector Store from litellm.rag.main import aingest, aquery, ingest, query -__all__ = ["ingest", "aingest", "query", "aquery"] +__all__ = ["aingest", "aquery", "ingest", "query"] # Expose at litellm.rag level for convenience diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 527b76e42ff..76e6a4c574c 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -14,17 +14,17 @@ from __future__ import annotations import base64 from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_logger from litellm._uuid import uuid4 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, httpxSpecialProvider, ) -from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.rag.ingestion.file_parsers import extract_text_from_pdf from litellm.rag.text_splitters import RecursiveCharacterTextSplitter from litellm.types.rag import RAGIngestOptions, RAGIngestResponse @@ -47,7 +47,7 @@ class BaseRAGIngestion(ABC): def __init__( self, ingest_options: RAGIngestOptions, - router: Optional["Router"] = None, + router: Router | None = None, ): self.ingest_options = ingest_options self.router = router @@ -55,12 +55,12 @@ class BaseRAGIngestion(ABC): # Extract configs from options self.ocr_config = ingest_options.get("ocr") - self.chunking_strategy: Dict[str, Any] = cast( - Dict[str, Any], + self.chunking_strategy: dict[str, Any] = cast( + dict[str, Any], ingest_options.get("chunking_strategy") or {"type": "auto"}, ) self.embedding_config = ingest_options.get("embedding") - self.vector_store_config: Dict[str, Any] = cast(Dict[str, Any], ingest_options.get("vector_store") or {}) + self.vector_store_config: dict[str, Any] = cast(dict[str, Any], ingest_options.get("vector_store") or {}) self.ingest_name = ingest_options.get("name") # Load credentials from litellm_credential_name if provided in vector_store config @@ -100,10 +100,10 @@ class BaseRAGIngestion(ABC): async def upload( self, - file_data: Optional[Tuple[str, bytes, str]] = None, - file_url: Optional[str] = None, - file_id: Optional[str] = None, - ) -> Tuple[Optional[str], Optional[bytes], Optional[str], Optional[str]]: + file_data: tuple[str, bytes, str] | None = None, + file_url: str | None = None, + file_id: str | None = None, + ) -> tuple[str | None, bytes | None, str | None, str | None]: """ Upload / prepare file for ingestion. @@ -135,9 +135,9 @@ class BaseRAGIngestion(ABC): async def ocr( self, - file_content: Optional[bytes], - content_type: Optional[str], - ) -> Optional[str]: + file_content: bytes | None, + content_type: str | None, + ) -> str | None: """ Perform OCR on file content to extract text. @@ -187,10 +187,10 @@ class BaseRAGIngestion(ABC): def chunk( self, - text: Optional[str], - file_content: Optional[bytes], + text: str | None, + file_content: bytes | None, ocr_was_used: bool, - ) -> List[str]: + ) -> list[str]: """ Split text into chunks using RecursiveCharacterTextSplitter. @@ -203,7 +203,7 @@ class BaseRAGIngestion(ABC): List of text chunks """ # Get text to chunk - text_to_chunk: Optional[str] = None + text_to_chunk: str | None = None if text: text_to_chunk = text elif file_content and not ocr_was_used: @@ -235,7 +235,7 @@ class BaseRAGIngestion(ABC): separators = splitter_args.get("separators", None) # Build splitter kwargs - splitter_kwargs: Dict[str, Any] = { + splitter_kwargs: dict[str, Any] = { "chunk_size": chunk_size, "chunk_overlap": chunk_overlap, } @@ -247,8 +247,8 @@ class BaseRAGIngestion(ABC): async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Generate embeddings for text chunks. @@ -273,13 +273,13 @@ class BaseRAGIngestion(ABC): @abstractmethod async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store content in vector store. @@ -296,13 +296,12 @@ class BaseRAGIngestion(ABC): Returns: Tuple of (vector_store_id, file_id) """ - pass async def ingest( self, - file_data: Optional[Tuple[str, bytes, str]] = None, - file_url: Optional[str] = None, - file_id: Optional[str] = None, + file_data: tuple[str, bytes, str] | None = None, + file_url: str | None = None, + file_id: str | None = None, ) -> RAGIngestResponse: """ Execute the full ingestion pipeline. diff --git a/litellm/rag/ingestion/bedrock_ingestion.py b/litellm/rag/ingestion/bedrock_ingestion.py index e4d636b7d30..10dc3af4319 100644 --- a/litellm/rag/ingestion/bedrock_ingestion.py +++ b/litellm/rag/ingestion/bedrock_ingestion.py @@ -14,7 +14,7 @@ from __future__ import annotations import asyncio import json import uuid -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM @@ -25,7 +25,7 @@ if TYPE_CHECKING: from litellm.types.rag import RAGIngestOptions -def _get_str_or_none(value: Any) -> Optional[str]: +def _get_str_or_none(value: Any) -> str | None: """Cast config value to Optional[str].""" return str(value) if value is not None else None @@ -85,8 +85,8 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router) BaseAWSLLM.__init__(self) @@ -99,7 +99,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): # Optional config self._data_source_id = self.vector_store_config.get("data_source_id") self._s3_bucket = self.vector_store_config.get("s3_bucket") - self._s3_prefix: Optional[str] = ( + self._s3_prefix: str | None = ( str(self.vector_store_config.get("s3_prefix")) if self.vector_store_config.get("s3_prefix") else None ) self.embedding_model = self.vector_store_config.get("embedding_model") or "amazon.titan-embed-text-v2:0" @@ -114,13 +114,13 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): ) # Will be set during initialization - self.data_source_id: Optional[str] = None - self.s3_bucket: Optional[str] = None + self.data_source_id: str | None = None + self.s3_bucket: str | None = None self.s3_prefix: str = self._s3_prefix or "data/" self._config_initialized = False # Track resources we create (for cleanup if needed) - self._created_resources: Dict[str, Any] = {} + self._created_resources: dict[str, Any] = {} async def _ensure_config_initialized(self): """Lazily initialize KB config - either detect from existing or create new.""" @@ -229,7 +229,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): verbose_logger.debug(f"Creating S3 bucket: {bucket_name}") - create_params: Dict[str, Any] = {"Bucket": bucket_name} + create_params: dict[str, Any] = {"Bucket": bucket_name} if self.aws_region_name != "us-east-1": create_params["CreateBucketConfiguration"] = {"LocationConstraint": self.aws_region_name} @@ -239,7 +239,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): verbose_logger.info(f"Created S3 bucket: {bucket_name}") return bucket_name - async def _create_opensearch_collection(self, unique_id: str, account_id: str, caller_arn: str) -> Tuple[str, str]: + async def _create_opensearch_collection(self, unique_id: str, account_id: str, caller_arn: str) -> tuple[str, str]: """Create OpenSearch Serverless collection for vector storage.""" oss = self._get_boto3_client("opensearchserverless") collection_name = f"litellm-kb-{unique_id}" @@ -602,8 +602,8 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Bedrock handles embedding internally - skip this step. @@ -614,13 +614,13 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store content in Bedrock Knowledge Base. diff --git a/litellm/rag/ingestion/file_parsers/pdf_parser.py b/litellm/rag/ingestion/file_parsers/pdf_parser.py index a992e957f37..2b4e07b224f 100644 --- a/litellm/rag/ingestion/file_parsers/pdf_parser.py +++ b/litellm/rag/ingestion/file_parsers/pdf_parser.py @@ -4,12 +4,10 @@ PDF text extraction utilities. Provides text extraction from PDF files using pypdf or PyPDF2. """ -from typing import Optional - from litellm._logging import verbose_logger -def extract_text_from_pdf(file_content: bytes) -> Optional[str]: +def extract_text_from_pdf(file_content: bytes) -> str | None: """ Extract text from PDF using pypdf if available. diff --git a/litellm/rag/ingestion/gemini_ingestion.py b/litellm/rag/ingestion/gemini_ingestion.py index 3f1d46bbb11..5722936b742 100644 --- a/litellm/rag/ingestion/gemini_ingestion.py +++ b/litellm/rag/ingestion/gemini_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, Dict, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( @@ -35,16 +35,16 @@ class GeminiRAGIngestion(BaseRAGIngestion): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): super().__init__(ingest_options=ingest_options, router=router) self.model_info = GeminiModelInfo() async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Gemini handles embedding internally - skip this step. @@ -56,13 +56,13 @@ class GeminiRAGIngestion(BaseRAGIngestion): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store content in Gemini File Search store. @@ -83,11 +83,11 @@ class GeminiRAGIngestion(BaseRAGIngestion): """ vector_store_id = self.vector_store_config.get("vector_store_id") - vector_store_config = cast(Dict[str, Any], self.vector_store_config) + vector_store_config = cast(dict[str, Any], self.vector_store_config) # Get API credentials - api_key = cast(Optional[str], vector_store_config.get("api_key")) or GeminiModelInfo.get_api_key() - api_base = cast(Optional[str], vector_store_config.get("api_base")) or GeminiModelInfo.get_api_base() + api_key = cast(str | None, vector_store_config.get("api_key")) or GeminiModelInfo.get_api_key() + api_base = cast(str | None, vector_store_config.get("api_base")) or GeminiModelInfo.get_api_base() if not api_key: raise ValueError("GEMINI_API_KEY or GOOGLE_API_KEY is required for Gemini File Search") @@ -172,7 +172,7 @@ class GeminiRAGIngestion(BaseRAGIngestion): vector_store_id: str, filename: str, file_content: bytes, - content_type: Optional[str], + content_type: str | None, ) -> str: """ Upload a file to Gemini File Search store using resumable upload. @@ -228,7 +228,7 @@ class GeminiRAGIngestion(BaseRAGIngestion): url = f"{api_base}/upload/v1beta/{vector_store_id}:uploadToFileSearchStore" # Build request body with chunking config and metadata if provided - request_body: Dict[str, Any] = {"displayName": filename} + request_body: dict[str, Any] = {"displayName": filename} # Add chunking configuration if provided chunking_strategy = self.chunking_strategy @@ -244,7 +244,7 @@ class GeminiRAGIngestion(BaseRAGIngestion): # Add custom metadata if provided in vector_store_config custom_metadata = cast( - Optional[List[Dict[str, Any]]], + list[dict[str, Any]] | None, self.vector_store_config.get("custom_metadata"), ) if custom_metadata: diff --git a/litellm/rag/ingestion/openai_ingestion.py b/litellm/rag/ingestion/openai_ingestion.py index ca5575a7e30..864a4f21290 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, List, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion @@ -33,15 +33,15 @@ class OpenAIRAGIngestion(BaseRAGIngestion): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): super().__init__(ingest_options=ingest_options, router=router) async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ OpenAI handles embedding internally - skip this step. @@ -53,13 +53,13 @@ class OpenAIRAGIngestion(BaseRAGIngestion): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store content in OpenAI vector store. diff --git a/litellm/rag/ingestion/s3_vectors_ingestion.py b/litellm/rag/ingestion/s3_vectors_ingestion.py index 3ec623657ee..6abd0737ba6 100644 --- a/litellm/rag/ingestion/s3_vectors_ingestion.py +++ b/litellm/rag/ingestion/s3_vectors_ingestion.py @@ -17,7 +17,7 @@ from __future__ import annotations import hashlib import uuid -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -59,8 +59,8 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router) BaseAWSLLM.__init__(self) @@ -127,7 +127,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): return S3_VECTORS_DEFAULT_DIMENSION - def _get_dimension_from_config(self) -> Optional[int]: + def _get_dimension_from_config(self) -> int | None: """ Get vector dimension from config if explicitly provided. @@ -163,8 +163,8 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): self, method: str, url: str, - data: Optional[str] = None, - headers: Optional[Dict[str, str]] = None, + data: str | None = None, + headers: dict[str, str] | None = None, ) -> Any: """ Helper to sign and execute AWS API requests using httpx + SigV4. @@ -332,7 +332,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): verbose_logger.exception(f"Error creating vector index: {e}") raise - async def _put_vectors(self, vectors: List[Dict[str, Any]]): + async def _put_vectors(self, vectors: list[dict[str, Any]]): """ Call PutVectors API to store vectors in S3 Vectors. @@ -364,8 +364,8 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Generate embeddings using LiteLLM's embedding API. @@ -384,7 +384,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): verbose_logger.debug(f"Generating embeddings for {len(chunks)} chunks using {embedding_model}") # Convert to list to ensure type compatibility - input_chunks: List[str] = list(chunks) + input_chunks: list[str] = list(chunks) if self.router: response = await self.router.aembedding(model=embedding_model, input=input_chunks) @@ -395,13 +395,13 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store vectors in S3 Vectors using PutVectors API. @@ -441,7 +441,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): vectors = [] for i, (chunk, embedding) in enumerate(zip(chunks, embeddings)): # Build metadata dict - metadata: Dict[str, str] = { + metadata: dict[str, str] = { "source_text": chunk, # Non-filterable (for reference) "chunk_index": str(i), # Filterable } @@ -464,7 +464,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): vector_store_id = f"{self.vector_bucket_name}:{self.index_name}" return vector_store_id, filename - async def query_vector_store(self, vector_store_id: str, query: str, top_k: int = 5) -> Optional[Dict[str, Any]]: + async def query_vector_store(self, vector_store_id: str, query: str, top_k: int = 5) -> dict[str, Any] | None: """ Query S3 Vectors using QueryVectors API. diff --git a/litellm/rag/ingestion/vertex_ai_ingestion.py b/litellm/rag/ingestion/vertex_ai_ingestion.py index 34cd1a88a61..ababca1a955 100644 --- a/litellm/rag/ingestion/vertex_ai_ingestion.py +++ b/litellm/rag/ingestion/vertex_ai_ingestion.py @@ -10,7 +10,7 @@ Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-refer from __future__ import annotations import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( @@ -40,8 +40,8 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): def __init__( self, - ingest_options: "RAGIngestOptions", - router: Optional["Router"] = None, + ingest_options: RAGIngestOptions, + router: Router | None = None, ): BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router) VertexBase.__init__(self) @@ -56,8 +56,8 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): async def embed( self, - chunks: List[str], - ) -> Optional[List[List[float]]]: + chunks: list[str], + ) -> list[list[float]] | None: """ Vertex AI RAG Engine handles embedding internally - skip this step. @@ -69,13 +69,13 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): async def store( self, - file_content: Optional[bytes], - filename: Optional[str], - content_type: Optional[str], - chunks: List[str], - embeddings: Optional[List[List[float]]], + file_content: bytes | None, + filename: str | None, + content_type: str | None, + chunks: list[str], + embeddings: list[list[float]] | None, existing_file_id: str | None = None, - ) -> Tuple[Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None]: """ Store content in Vertex AI RAG corpus. @@ -120,7 +120,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): async def _create_rag_corpus( self, display_name: str, - description: Optional[str] = None, + description: str | None = None, ) -> str: """ Create a Vertex AI RAG corpus. @@ -148,7 +148,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): url = f"{base_url}/v1beta1/projects/{self.project_id}/locations/{self.location}/ragCorpora" # Build request body with camelCase keys (Vertex AI API format) - request_body: Dict[str, Any] = { + request_body: dict[str, Any] = { "displayName": display_name, } @@ -282,7 +282,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): rag_corpus_id: str, filename: str, file_content: bytes, - content_type: Optional[str], + content_type: str | None, ) -> str: """ Upload a file to Vertex AI RAG corpus using multipart upload. @@ -308,7 +308,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): url = f"{base_url}/upload/v1beta1/{rag_corpus_id}/ragFiles:upload" # Build metadata for the file with snake_case keys (as per upload API docs) - metadata: Dict[str, Any] = { + metadata: dict[str, Any] = { "rag_file": { "display_name": filename, } @@ -390,7 +390,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): async def _import_files_from_gcs( self, rag_corpus_id: str, - gcs_uris: List[str], + gcs_uris: list[str], ) -> str: """ Import files from Google Cloud Storage into RAG corpus. @@ -414,7 +414,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase): url = f"{base_url}/v1beta1/{rag_corpus_id}/ragFiles:import" # Build request body with camelCase keys (Vertex AI API format) - request_body: Dict[str, Any] = {"importRagFilesConfig": {"gcsSource": {"uris": gcs_uris}}} + request_body: dict[str, Any] = {"importRagFilesConfig": {"gcsSource": {"uris": gcs_uris}}} # Add chunking configuration if provided chunking_strategy = self.chunking_strategy diff --git a/litellm/rag/main.py b/litellm/rag/main.py index 80d6efe9ddb..64b55b4320f 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -7,7 +7,7 @@ Upload -> (OCR) -> Chunk -> Embed -> Vector Store from __future__ import annotations -__all__ = ["ingest", "aingest", "query", "aquery"] +__all__ = ["aingest", "aquery", "ingest", "query"] import asyncio import contextvars @@ -17,12 +17,6 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - Dict, - List, - Optional, - Tuple, - Type, - Union, ) import httpx @@ -50,7 +44,7 @@ if TYPE_CHECKING: # Registry of provider-specific ingestion classes -INGESTION_REGISTRY: Dict[str, Type[BaseRAGIngestion]] = { +INGESTION_REGISTRY: dict[str, type[BaseRAGIngestion]] = { "openai": OpenAIRAGIngestion, "bedrock": BedrockRAGIngestion, "gemini": GeminiRAGIngestion, @@ -59,7 +53,7 @@ INGESTION_REGISTRY: Dict[str, Type[BaseRAGIngestion]] = { } -def get_ingestion_class(provider: str) -> Type[BaseRAGIngestion]: +def get_ingestion_class(provider: str) -> type[BaseRAGIngestion]: """ Get the ingestion class for a given provider. @@ -81,10 +75,10 @@ def get_ingestion_class(provider: str) -> Type[BaseRAGIngestion]: async def _execute_ingest_pipeline( ingest_options: RAGIngestOptions, - file_data: Optional[Tuple[str, bytes, str]] = None, - file_url: Optional[str] = None, - file_id: Optional[str] = None, - router: Optional["Router"] = None, + file_data: tuple[str, bytes, str] | None = None, + file_url: str | None = None, + file_id: str | None = None, + router: Router | None = None, ) -> RAGIngestResponse: """ Execute the RAG ingest pipeline using provider-specific implementation. @@ -125,12 +119,12 @@ async def _execute_ingest_pipeline( @client async def aingest( - ingest_options: Dict[str, Any], - file_data: Optional[Tuple[str, bytes, str]] = None, - file: Optional[Dict[str, str]] = None, - file_url: Optional[str] = None, - file_id: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + ingest_options: dict[str, Any], + file_data: tuple[str, bytes, str] | None = None, + file: dict[str, str] | None = None, + file_url: str | None = None, + file_id: str | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, ) -> RAGIngestResponse: """ @@ -213,9 +207,9 @@ def _suppressed_sub_call_billing() -> Iterator[None]: async def _execute_query_pipeline( model: str, - messages: List[Any], - retrieval_config: Dict[str, Any], - rerank: Optional[Dict[str, Any]] = None, + messages: list[Any], + retrieval_config: dict[str, Any], + rerank: dict[str, Any] | None = None, stream: bool = False, **kwargs, ) -> ModelResponse: @@ -224,7 +218,7 @@ async def _execute_query_pipeline( """ # Extract router from kwargs - use it for completion if available # to properly resolve virtual model names - router: Optional["Router"] = kwargs.pop("router", None) + router: Router | None = kwargs.pop("router", None) # 1. Extract query from last user message query_text = RAGQuery.extract_query_from_messages(messages) @@ -320,9 +314,9 @@ async def _execute_query_pipeline( @client async def aquery( model: str, - messages: List[Any], - retrieval_config: Dict[str, Any], - rerank: Optional[Dict[str, Any]] = None, + messages: list[Any], + retrieval_config: dict[str, Any], + rerank: dict[str, Any] | None = None, stream: bool = False, **kwargs, ) -> ModelResponse: @@ -367,12 +361,12 @@ async def aquery( @client def query( model: str, - messages: List[Any], - retrieval_config: Dict[str, Any], - rerank: Optional[Dict[str, Any]] = None, + messages: list[Any], + retrieval_config: dict[str, Any], + rerank: dict[str, Any] | None = None, stream: bool = False, **kwargs, -) -> Union[ModelResponse, Coroutine[Any, Any, ModelResponse]]: +) -> ModelResponse | Coroutine[Any, Any, ModelResponse]: """ Query a RAG pipeline. """ @@ -412,14 +406,14 @@ def query( @client def ingest( - ingest_options: Dict[str, Any], - file_data: Optional[Tuple[str, bytes, str]] = None, - file: Optional[Dict[str, str]] = None, - file_url: Optional[str] = None, - file_id: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + ingest_options: dict[str, Any], + file_data: tuple[str, bytes, str] | None = None, + file: dict[str, str] | None = None, + file_url: str | None = None, + file_id: str | None = None, + timeout: float | httpx.Timeout | None = None, **kwargs, -) -> Union[RAGIngestResponse, Coroutine[Any, Any, RAGIngestResponse]]: +) -> RAGIngestResponse | Coroutine[Any, Any, RAGIngestResponse]: """ Ingest a document into a vector store. @@ -448,7 +442,7 @@ def ingest( local_vars = locals() try: _is_async = kwargs.pop("aingest", False) is True - router: Optional["Router"] = kwargs.get("router") + router: Router | None = kwargs.get("router") # Convert file dict to file_data tuple if provided if file is not None and file_data is None: diff --git a/litellm/rag/rag_query.py b/litellm/rag/rag_query.py index 65cd8a3572e..c1933f84f98 100644 --- a/litellm/rag/rag_query.py +++ b/litellm/rag/rag_query.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, Union +from typing import Any from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage from litellm.types.utils import ModelResponse @@ -12,7 +12,7 @@ class RAGQuery: CONTENT_PREFIX_STRING = "Context:\n\n" @staticmethod - def extract_query_from_messages(messages: List[AllMessageValues]) -> Optional[str]: + def extract_query_from_messages(messages: list[AllMessageValues]) -> str | None: """ Extract the query from the last user message. """ @@ -36,7 +36,7 @@ class RAGQuery: return None @staticmethod - def build_context_message(context_chunks: List[Any]) -> ChatCompletionUserMessage: + def build_context_message(context_chunks: list[Any]) -> ChatCompletionUserMessage: """ Process search results and build a context message. """ @@ -44,10 +44,10 @@ class RAGQuery: for chunk in context_chunks: if isinstance(chunk, dict): - result_content: Optional[List[VectorStoreResultContent]] = chunk.get("content") + result_content: list[VectorStoreResultContent] | None = chunk.get("content") if result_content: for content_item in result_content: - content_text: Optional[str] = content_item.get("text") + content_text: str | None = content_item.get("text") if content_text: context_content += content_text + "\n\n" elif "text" in chunk: # Fallback for simple dict with text @@ -64,7 +64,7 @@ class RAGQuery: def add_search_results_to_response( response: ModelResponse, search_results: VectorStoreSearchResponse, - rerank_results: Optional[Any] = None, + rerank_results: Any | None = None, ) -> ModelResponse: """ Add search results to the response choices. @@ -88,9 +88,9 @@ class RAGQuery: @staticmethod def extract_documents_from_search( search_response: Any, - ) -> List[Union[str, Dict[str, Any]]]: + ) -> list[str | dict[str, Any]]: """Extract text documents from vector store search response.""" - documents: List[Union[str, Dict[str, Any]]] = [] + documents: list[str | dict[str, Any]] = [] for result in search_response.get("data", []): content_list = result.get("content", []) for content in content_list: @@ -99,7 +99,7 @@ class RAGQuery: return documents @staticmethod - def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> List[Any]: + def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> list[Any]: """Get the original search results corresponding to the top reranked results.""" top_chunks = [] original_results = search_response.get("data", []) diff --git a/litellm/rag/text_splitters/recursive_character_text_splitter.py b/litellm/rag/text_splitters/recursive_character_text_splitter.py index 48041172369..db929468916 100644 --- a/litellm/rag/text_splitters/recursive_character_text_splitter.py +++ b/litellm/rag/text_splitters/recursive_character_text_splitter.py @@ -4,8 +4,6 @@ RecursiveCharacterTextSplitter for RAG ingestion. A simple implementation that splits text recursively by different separators. """ -from typing import List, Optional - from litellm.constants import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE @@ -21,17 +19,17 @@ class RecursiveCharacterTextSplitter: self, chunk_size: int = DEFAULT_CHUNK_SIZE, chunk_overlap: int = DEFAULT_CHUNK_OVERLAP, - separators: Optional[List[str]] = None, + separators: list[str] | None = None, ): self.chunk_size = chunk_size self.chunk_overlap = chunk_overlap self.separators = separators or ["\n\n", "\n", " ", ""] - def split_text(self, text: str) -> List[str]: + def split_text(self, text: str) -> list[str]: """Split text into chunks.""" return self._split_text(text, self.separators) - def _split_text(self, text: str, separators: List[str], depth: int = 0) -> List[str]: + def _split_text(self, text: str, separators: list[str], depth: int = 0) -> list[str]: """Recursively split text using separators.""" from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -39,11 +37,11 @@ class RecursiveCharacterTextSplitter: # Max depth reached, return text as-is split into chunk_size pieces return [text[i : i + self.chunk_size] for i in range(0, len(text), self.chunk_size)] - final_chunks: List[str] = [] + final_chunks: list[str] = [] # Get the appropriate separator separator = separators[-1] - new_separators: List[str] = [] + new_separators: list[str] = [] for i, sep in enumerate(separators): if sep == "": @@ -61,7 +59,7 @@ class RecursiveCharacterTextSplitter: splits = list(text) # Merge splits into chunks - good_splits: List[str] = [] + good_splits: list[str] = [] for split in splits: if len(split) < self.chunk_size: good_splits.append(split) @@ -87,10 +85,10 @@ class RecursiveCharacterTextSplitter: return final_chunks - def _merge_splits(self, splits: List[str], separator: str) -> List[str]: + def _merge_splits(self, splits: list[str], separator: str) -> list[str]: """Merge splits into chunks respecting chunk_size and chunk_overlap.""" - chunks: List[str] = [] - current_chunk: List[str] = [] + chunks: list[str] = [] + current_chunk: list[str] = [] current_length = 0 for split in splits: @@ -119,9 +117,9 @@ class RecursiveCharacterTextSplitter: return chunks - def _force_split(self, text: str) -> List[str]: + def _force_split(self, text: str) -> list[str]: """Force split text by chunk_size when no separator works.""" - chunks: List[str] = [] + chunks: list[str] = [] start = 0 while start < len(text): diff --git a/litellm/rag/utils.py b/litellm/rag/utils.py index 49b8037de57..c195a882d5b 100644 --- a/litellm/rag/utils.py +++ b/litellm/rag/utils.py @@ -4,13 +4,13 @@ RAG utility functions. Provides provider configuration utilities similar to ProviderConfigManager. """ -from typing import TYPE_CHECKING, Type +from typing import TYPE_CHECKING if TYPE_CHECKING: from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion -def get_rag_ingestion_class(custom_llm_provider: str) -> Type["BaseRAGIngestion"]: +def get_rag_ingestion_class(custom_llm_provider: str) -> type["BaseRAGIngestion"]: """ Get the appropriate RAG ingestion class for a provider. diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 5ecf4d91ff6..0e55097f92a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -1,7 +1,7 @@ """Abstraction function for OpenAI's realtime API""" import os -from typing import Any, Dict, Optional, cast +from typing import Any, cast import litellm from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, request_timeout @@ -57,8 +57,8 @@ def _build_litellm_metadata(kwargs: dict) -> dict: def _get_realtime_http_provider_config( custom_llm_provider: str, - dynamic_api_base: Optional[str], - dynamic_api_key: Optional[str], + dynamic_api_base: str | None, + dynamic_api_key: str | None, litellm_params: GenericLiteLLMParams, ) -> tuple[Any, str, str]: """ @@ -72,7 +72,7 @@ def _get_realtime_http_provider_config( BaseRealtimeHTTPConfig, ) - provider_config: Optional[BaseRealtimeHTTPConfig] = None + provider_config: BaseRealtimeHTTPConfig | None = None if custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_http_config( model="", @@ -97,10 +97,10 @@ def _get_realtime_http_provider_config( @wrapper_client async def acreate_realtime_client_secret( - model: Optional[str] = None, - session: Optional[Dict[str, Any]] = None, - expires_after: Optional[Dict[str, Any]] = None, - timeout: Optional[float] = None, + model: str | None = None, + session: dict[str, Any] | None = None, + expires_after: dict[str, Any] | None = None, + timeout: float | None = None, **kwargs, ): req = RealtimeClientSecretRequest( @@ -158,9 +158,9 @@ async def acreate_realtime_client_secret( @wrapper_client async def acreate_realtime_transcription_session( - model: Optional[str] = None, - transcription_session: Optional[Dict[str, Any]] = None, - timeout: Optional[float] = None, + model: str | None = None, + transcription_session: dict[str, Any] | None = None, + timeout: float | None = None, **kwargs, ): """ @@ -232,9 +232,9 @@ async def acreate_realtime_transcription_session( async def arealtime_calls( openai_ephemeral_key: str, sdp_body: bytes, - model: Optional[str] = None, - session: Optional[Dict[str, Any]] = None, - timeout: Optional[float] = None, + model: str | None = None, + session: dict[str, Any] | None = None, + timeout: float | None = None, **kwargs, ): model_name = model or "gpt-4o-realtime-preview" @@ -285,13 +285,13 @@ async def arealtime_calls( async def _arealtime( model: str, websocket: Any, # fastapi websocket - api_base: Optional[str] = None, - api_key: Optional[str] = None, - api_version: Optional[str] = None, - azure_ad_token: Optional[str] = None, - client: Optional[Any] = None, - timeout: Optional[float] = None, - query_params: Optional[RealtimeQueryParams] = None, + api_base: str | None = None, + api_key: str | None = None, + api_version: str | None = None, + azure_ad_token: str | None = None, + client: Any | None = None, + timeout: float | None = None, + query_params: RealtimeQueryParams | None = None, **kwargs, ): """ @@ -299,8 +299,8 @@ async def _arealtime( For PROXY use only. """ - headers = cast(Optional[dict], kwargs.get("headers")) - extra_headers = cast(Optional[dict], kwargs.get("extra_headers")) + headers = cast(dict | None, kwargs.get("headers")) + extra_headers = cast(dict | None, kwargs.get("extra_headers")) if headers is None: headers = {} if extra_headers is not None: @@ -334,7 +334,7 @@ async def _arealtime( custom_llm_provider=_custom_llm_provider, ) - provider_config: Optional[BaseRealtimeConfig] = None + provider_config: BaseRealtimeConfig | None = None if _custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_config( model=model, @@ -511,11 +511,11 @@ async def _arealtime( async def _realtime_health_check( model: str, custom_llm_provider: str, - api_key: Optional[str], - api_base: Optional[str] = None, - api_version: Optional[str] = None, - realtime_protocol: Optional[str] = None, - model_params: Optional[dict] = None, + api_key: str | None, + api_base: str | None = None, + api_version: str | None = None, + realtime_protocol: str | None = None, + model_params: dict | None = None, ): """ Health check for realtime API - tries connection to the realtime API websocket @@ -535,7 +535,7 @@ async def _realtime_health_check( """ import websockets - url: Optional[str] = None + url: str | None = None if custom_llm_provider == "azure": url = azure_realtime._construct_url( api_base=api_base or "", diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 29c953e06cf..1fc3d8dadaf 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -68,62 +68,62 @@ from litellm.repositories.verification_token_repository import ( ) __all__ = [ - "PrismaTableRepository", - "PolicyRepository", - "AgentsRepository", - "GuardrailsRepository", - "MCPServerRepository", - "ManagedObjectRepository", - "OrganizationMembershipRepository", - "SpendLogsRepository", - "ClaudeCodePluginRepository", - "TeamMembershipRepository", - "EndUserRepository", - "ManagedVectorStoresRepository", - "MCPUserCredentialsRepository", - "PromptRepository", - "TagRepository", - "InvitationLinkRepository", - "JWTKeyMappingRepository", - "ManagedFileRepository", - "MemoryRepository", - "SearchToolsRepository", - "ConfigOverridesRepository", - "MCPToolsetRepository", - "ToolRepository", - "DeletedVerificationTokenRepository", - "WorkflowRunRepository", - "ModelTableRepository", "AccessGroupRepository", - "SSOConfigRepository", - "UISettingsRepository", - "DailyGuardrailMetricsRepository", - "PolicyAttachmentRepository", - "DeletedTeamRepository", - "SkillsRepository", - "CacheConfigRepository", - "ManagedVectorStoreIndexRepository", - "WorkflowMessageRepository", - "DailyTagSpendRepository", - "DailyToolSpendRepository", - "SpendLogToolIndexRepository", - "SpendLogGuardrailIndexRepository", - "UserNotificationsRepository", - "HealthCheckRepository", - "DeprecatedVerificationTokenRepository", - "WorkflowEventRepository", - "DailyPolicyMetricsRepository", - "AdaptiveRouterStateRepository", - "AuditLogRepository", "AdaptiveRouterSessionRepository", + "AdaptiveRouterStateRepository", + "AgentsRepository", + "AuditLogRepository", "BudgetRepository", + "CacheConfigRepository", + "ClaudeCodePluginRepository", + "ConfigOverridesRepository", "ConfigRepository", "CredentialsRepository", + "DailyGuardrailMetricsRepository", + "DailyPolicyMetricsRepository", + "DailyTagSpendRepository", + "DailyToolSpendRepository", + "DeletedTeamRepository", + "DeletedVerificationTokenRepository", + "DeprecatedVerificationTokenRepository", + "EndUserRepository", + "GuardrailsRepository", + "HealthCheckRepository", + "InvitationLinkRepository", + "JWTKeyMappingRepository", + "MCPServerRepository", + "MCPToolsetRepository", + "MCPUserCredentialsRepository", + "ManagedFileRepository", + "ManagedObjectRepository", + "ManagedVectorStoreIndexRepository", + "ManagedVectorStoresRepository", + "MemoryRepository", "ModelRepository", + "ModelTableRepository", "ObjectPermissionRepository", + "OrganizationMembershipRepository", "OrganizationRepository", + "PolicyAttachmentRepository", + "PolicyRepository", + "PrismaTableRepository", "ProjectRepository", + "PromptRepository", + "SSOConfigRepository", + "SearchToolsRepository", + "SkillsRepository", + "SpendLogGuardrailIndexRepository", + "SpendLogToolIndexRepository", + "SpendLogsRepository", + "TagRepository", + "TeamMembershipRepository", "TeamRepository", + "ToolRepository", + "UISettingsRepository", + "UserNotificationsRepository", "UserRepository", "VerificationTokenRepository", + "WorkflowEventRepository", + "WorkflowMessageRepository", + "WorkflowRunRepository", ] diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 755e4595c01..7f7333d6d2b 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -4,7 +4,7 @@ Base repository class with common functionality. from abc import ABC, abstractmethod from collections.abc import Iterable, Mapping, Sequence -from typing import Any, Dict, Generic, List, Optional, Protocol, Tuple, Type, TypeVar, Union, runtime_checkable +from typing import Any, Generic, Protocol, TypeVar, Union, runtime_checkable from pydantic import BaseModel @@ -13,19 +13,19 @@ T = TypeVar("T", bound=BaseModel) @runtime_checkable class SupportsModelDump(Protocol): - def model_dump(self) -> Dict[str, object]: ... + def model_dump(self) -> dict[str, object]: ... @runtime_checkable class SupportsDict(Protocol): - def dict(self) -> Dict[str, object]: ... + def dict(self) -> dict[str, object]: ... DbRecord = Union[ Mapping[str, object], SupportsModelDump, SupportsDict, - Sequence[Tuple[str, object]], + Sequence[tuple[str, object]], ] @@ -60,34 +60,34 @@ class BaseRepository(ABC, Generic[T]): @property @abstractmethod - def model_class(self) -> Type[T]: + def model_class(self) -> type[T]: """Return the domain model class for this repository.""" ... - def _to_model(self, record: Optional[DbRecord]) -> Optional[T]: + def _to_model(self, record: DbRecord | None) -> T | None: """Convert a database record to a domain model.""" if record is None: return None return self.model_class.model_validate(record_to_dict(record)) - def _to_model_list(self, records: Iterable[Optional[DbRecord]]) -> List[T]: + def _to_model_list(self, records: Iterable[DbRecord | None]) -> list[T]: """Convert a list of database records to domain models.""" return [model for record in records if record is not None and (model := self._to_model(record)) is not None] - async def find_by_id(self, id_value: str, id_field: str = "id") -> Optional[T]: + async def find_by_id(self, id_value: str, id_field: str = "id") -> T | None: """Find a record by its primary key.""" record = await self.table.find_unique(where={id_field: id_value}) return self._to_model(record) async def find_many( self, - where: Optional[Dict[str, Any]] = None, - skip: Optional[int] = None, - take: Optional[int] = None, - order: Optional[Dict[str, str]] = None, - ) -> List[T]: + where: dict[str, Any] | None = None, + skip: int | None = None, + take: int | None = None, + order: dict[str, str] | None = None, + ) -> list[T]: """Find multiple records matching the criteria.""" - kwargs: Dict[str, Any] = {} + kwargs: dict[str, Any] = {} if where: kwargs["where"] = where if skip is not None: @@ -100,24 +100,24 @@ class BaseRepository(ABC, Generic[T]): records = await self.table.find_many(**kwargs) return self._to_model_list(records) - async def create(self, data: Dict[str, Any]) -> T: + async def create(self, data: dict[str, Any]) -> T: """Create a new record.""" record = await self.table.create(data=data) model = self._to_model(record) assert model is not None return model - async def update(self, id_value: str, data: Dict[str, Any], id_field: str = "id") -> Optional[T]: + async def update(self, id_value: str, data: dict[str, Any], id_field: str = "id") -> T | None: """Update an existing record.""" record = await self.table.update(where={id_field: id_value}, data=data) return self._to_model(record) - async def delete(self, id_value: str, id_field: str = "id") -> Optional[T]: + async def delete(self, id_value: str, id_field: str = "id") -> T | None: """Delete a record by its primary key.""" record = await self.table.delete(where={id_field: id_value}) return self._to_model(record) - async def count(self, where: Optional[Dict[str, Any]] = None) -> int: + async def count(self, where: dict[str, Any] | None = None) -> int: """Count records matching the criteria.""" return await self.table.count(where=where) diff --git a/litellm/repositories/budget_repository.py b/litellm/repositories/budget_repository.py index aa2676fa72b..9ae4afc0317 100644 --- a/litellm/repositories/budget_repository.py +++ b/litellm/repositories/budget_repository.py @@ -2,7 +2,7 @@ Budget repository for database operations on LiteLLM_BudgetTable. """ -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm.models.budget import LiteLLM_BudgetTable from litellm.repositories.base_repository import BaseRepository @@ -16,26 +16,26 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): return self.prisma_client.db.litellm_budgettable @property - def model_class(self) -> Type[LiteLLM_BudgetTable]: + def model_class(self) -> type[LiteLLM_BudgetTable]: return LiteLLM_BudgetTable - async def find_by_id(self, budget_id: str, id_field: str = "budget_id") -> Optional[LiteLLM_BudgetTable]: + async def find_by_id(self, budget_id: str, id_field: str = "budget_id") -> LiteLLM_BudgetTable | None: return await super().find_by_id(budget_id, id_field) async def create_budget( self, created_by: str, - max_budget: Optional[float] = None, - soft_budget: Optional[float] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - model_max_budget: Optional[Dict[str, Any]] = None, - budget_duration: Optional[str] = None, - allowed_models: Optional[List[str]] = None, + max_budget: float | None = None, + soft_budget: float | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + model_max_budget: dict[str, Any] | None = None, + budget_duration: str | None = None, + allowed_models: list[str] | None = None, ) -> LiteLLM_BudgetTable: """Create a new budget record.""" - data: Dict[str, Any] = { + data: dict[str, Any] = { "created_by": created_by, "updated_by": created_by, } @@ -62,17 +62,17 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): self, budget_id: str, updated_by: str, - max_budget: Optional[float] = None, - soft_budget: Optional[float] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - model_max_budget: Optional[Dict[str, Any]] = None, - budget_duration: Optional[str] = None, - allowed_models: Optional[List[str]] = None, - ) -> Optional[LiteLLM_BudgetTable]: + max_budget: float | None = None, + soft_budget: float | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + model_max_budget: dict[str, Any] | None = None, + budget_duration: str | None = None, + allowed_models: list[str] | None = None, + ) -> LiteLLM_BudgetTable | None: """Update an existing budget record.""" - data: Dict[str, Any] = {"updated_by": updated_by} + data: dict[str, Any] = {"updated_by": updated_by} if max_budget is not None: data["max_budget"] = max_budget if soft_budget is not None: @@ -92,6 +92,6 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): return await self.update(budget_id, data, id_field="budget_id") - async def delete_budget(self, budget_id: str) -> Optional[LiteLLM_BudgetTable]: + async def delete_budget(self, budget_id: str) -> LiteLLM_BudgetTable | None: """Delete a budget record.""" return await self.delete(budget_id, id_field="budget_id") diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index 5af78bf1a6c..bc3e7fbaf9f 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -10,7 +10,7 @@ import asyncio import copy import json import os -from typing import Any, Dict, List, Literal, Optional, cast +from typing import Any, Literal, cast from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper @@ -47,7 +47,7 @@ class ConfigRepository: def table(self) -> Any: return self.prisma_client.db.litellm_config - async def get_param(self, param_name: str) -> Optional[ConfigParam]: + async def get_param(self, param_name: str) -> ConfigParam | None: """Get a config parameter from the database.""" record = await self.table.find_unique(where={"param_name": param_name}) if record is None: @@ -77,7 +77,7 @@ class ConfigRepository: except Exception: return False - async def get_all_params(self) -> Dict[str, Any]: + async def get_all_params(self) -> dict[str, Any]: """Get all config parameters from the database.""" records = await self.table.find_many() result = {} @@ -107,9 +107,9 @@ class ConfigRepository: else: d[k] = v - def _decrypt_env_variables(self, env_vars: Dict[str, Any], return_original_value: bool = True) -> Dict[str, str]: + def _decrypt_env_variables(self, env_vars: dict[str, Any], return_original_value: bool = True) -> dict[str, str]: """Decrypt environment variables from database.""" - decrypted: Dict[str, str] = {} + decrypted: dict[str, str] = {} for key, value in env_vars.items(): if isinstance(value, str): decrypted_value = decrypt_value_helper( @@ -124,9 +124,9 @@ class ConfigRepository: decrypted[key] = str(value) return decrypted - def _normalize_env_variable_keys(self, env_vars: Dict[str, str]) -> Dict[str, str]: + def _normalize_env_variable_keys(self, env_vars: dict[str, str]) -> dict[str, str]: """Normalize env variable keys to include both original and uppercase versions.""" - normalized: Dict[str, str] = {} + normalized: dict[str, str] = {} for key, value in env_vars.items(): normalized[key] = value upper_key = key.upper() @@ -168,7 +168,7 @@ class ConfigRepository: async def reconcile_config( self, yaml_config: dict, - store_model_in_db: Optional[bool] = None, + store_model_in_db: bool | None = None, ) -> dict: """Reconcile config from YAML with database overrides. @@ -216,7 +216,7 @@ class ConfigRepository: return config - async def prefetch_params(self, param_names: List[str]) -> None: + async def prefetch_params(self, param_names: list[str]) -> None: """Prefetch config params to warm the cache. This can be called before reconcile_config to ensure all needed diff --git a/litellm/repositories/credentials_repository.py b/litellm/repositories/credentials_repository.py index b5a315d233c..ccb1f9b2467 100644 --- a/litellm/repositories/credentials_repository.py +++ b/litellm/repositories/credentials_repository.py @@ -6,7 +6,7 @@ credential values is the caller's responsibility (see ``CredentialHelperUtils``) so reads return the stored values verbatim. """ -from typing import Any, Dict, Optional +from typing import Any from litellm.models.credentials import CredentialItem @@ -28,7 +28,7 @@ class CredentialsRepository: return self.prisma_client.db.litellm_credentialstable @staticmethod - def _to_model(record: Any) -> Optional[CredentialItem]: + def _to_model(record: Any) -> CredentialItem | None: if record is None: return None data = record.dict() if hasattr(record, "dict") else dict(record) @@ -41,14 +41,14 @@ class CredentialsRepository: async def find_all(self) -> Any: return await self.table.find_many() - async def create(self, data: Dict[str, Any]) -> Any: + async def create(self, data: dict[str, Any]) -> Any: return await self.table.create(data=data) - async def find_by_name(self, credential_name: str) -> Optional[CredentialItem]: + async def find_by_name(self, credential_name: str) -> CredentialItem | None: record = await self.table.find_unique(where={"credential_name": credential_name}) return self._to_model(record) - async def update_by_name(self, credential_name: str, data: Dict[str, Any]) -> Any: + async def update_by_name(self, credential_name: str, data: dict[str, Any]) -> Any: return await self.table.update(where={"credential_name": credential_name}, data=data) async def delete_by_name(self, credential_name: str) -> Any: diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index 0da51519964..26b782be33d 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -3,20 +3,20 @@ Model repository for database operations on LiteLLM_ProxyModelTable. """ import json -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm.models.model import LiteLLM_ProxyModelTable -from litellm.repositories.base_repository import BaseRepository from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) +from litellm.repositories.base_repository import BaseRepository class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): """Repository for proxy model database operations with encryption support.""" - def __init__(self, prisma_client: Any, encryption_key: Optional[str] = None): + def __init__(self, prisma_client: Any, encryption_key: str | None = None): super().__init__(prisma_client) self._encryption_key = encryption_key @@ -25,10 +25,10 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): return self.prisma_client.db.litellm_proxymodeltable @property - def model_class(self) -> Type[LiteLLM_ProxyModelTable]: + def model_class(self) -> type[LiteLLM_ProxyModelTable]: return LiteLLM_ProxyModelTable - def _encrypt_litellm_params(self, litellm_params: Dict[str, Any]) -> Dict[str, Any]: + def _encrypt_litellm_params(self, litellm_params: dict[str, Any]) -> dict[str, Any]: """Encrypt sensitive values in litellm_params.""" encrypted = {} for key, value in litellm_params.items(): @@ -38,7 +38,7 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): encrypted[key] = value return encrypted - def _decrypt_litellm_params(self, litellm_params: Dict[str, Any]) -> Dict[str, Any]: + def _decrypt_litellm_params(self, litellm_params: dict[str, Any]) -> dict[str, Any]: """Decrypt sensitive values in litellm_params.""" decrypted = {} for key, value in litellm_params.items(): @@ -50,7 +50,7 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): decrypted[key] = value return decrypted - def _to_model(self, record: Any) -> Optional[LiteLLM_ProxyModelTable]: + def _to_model(self, record: Any) -> LiteLLM_ProxyModelTable | None: """Convert a database record to a Model with decryption.""" if record is None: return None @@ -67,25 +67,25 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): return LiteLLM_ProxyModelTable(**data) - async def find_by_id(self, model_id: str, id_field: str = "model_id") -> Optional[LiteLLM_ProxyModelTable]: + 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) - async def find_by_name(self, model_name: str) -> List[LiteLLM_ProxyModelTable]: + async def find_by_name(self, model_name: str) -> list[LiteLLM_ProxyModelTable]: """Find models by name.""" records = await self.table.find_many(where={"model_name": model_name}) return self._to_model_list(records) - async def find_all(self) -> List[LiteLLM_ProxyModelTable]: + async def find_all(self) -> list[LiteLLM_ProxyModelTable]: """Find all models.""" records = await self.table.find_many() return self._to_model_list(records) - async def find_unblocked(self) -> List[LiteLLM_ProxyModelTable]: + async def find_unblocked(self) -> list[LiteLLM_ProxyModelTable]: """Find all models that are not blocked.""" records = await self.table.find_many(where={"blocked": False}) return self._to_model_list(records) - async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProxyModelTable]: + async def find_by_team_id(self, team_id: str) -> list[LiteLLM_ProxyModelTable]: """Find models associated with a specific team. Note: This filters in-memory since team_id is stored within litellm_params @@ -98,16 +98,16 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): async def create_model( self, model_name: str, - litellm_params: Dict[str, Any], + litellm_params: dict[str, Any], created_by: str, - model_id: Optional[str] = None, - model_info: Optional[Dict[str, Any]] = None, + model_id: str | None = None, + model_info: dict[str, Any] | None = None, blocked: bool = False, ) -> LiteLLM_ProxyModelTable: """Create a new model with encryption.""" encrypted_params = self._encrypt_litellm_params(litellm_params) - data: Dict[str, Any] = { + data: dict[str, Any] = { "model_name": model_name, "litellm_params": json.dumps(encrypted_params), "created_by": created_by, @@ -128,13 +128,13 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): self, model_id: str, updated_by: str, - model_name: Optional[str] = None, - litellm_params: Optional[Dict[str, Any]] = None, - model_info: Optional[Dict[str, Any]] = None, - blocked: Optional[bool] = None, - ) -> Optional[LiteLLM_ProxyModelTable]: + model_name: str | None = None, + litellm_params: dict[str, Any] | None = None, + model_info: dict[str, Any] | None = None, + blocked: bool | None = None, + ) -> LiteLLM_ProxyModelTable | None: """Update a model with encryption.""" - data: Dict[str, Any] = {"updated_by": updated_by} + data: dict[str, Any] = {"updated_by": updated_by} if model_name is not None: data["model_name"] = model_name if litellm_params is not None: @@ -148,14 +148,14 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): record = await self.table.update(where={"model_id": model_id}, data=data) return self._to_model(record) - async def delete_model(self, model_id: str) -> Optional[LiteLLM_ProxyModelTable]: + async def delete_model(self, model_id: str) -> LiteLLM_ProxyModelTable | None: """Delete a model.""" return await self.delete(model_id, id_field="model_id") - async def block_model(self, model_id: str, updated_by: str) -> Optional[LiteLLM_ProxyModelTable]: + async def block_model(self, model_id: str, updated_by: str) -> LiteLLM_ProxyModelTable | None: """Block a model.""" return await self.update_model(model_id, updated_by, blocked=True) - async def unblock_model(self, model_id: str, updated_by: str) -> Optional[LiteLLM_ProxyModelTable]: + async def unblock_model(self, model_id: str, updated_by: str) -> LiteLLM_ProxyModelTable | None: """Unblock a model.""" return await self.update_model(model_id, updated_by, blocked=False) diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 063a17cfe92..d291185fa25 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 Any, Dict, List, Optional, Type +from typing import Any from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.repositories.base_repository import BaseRepository @@ -16,29 +16,29 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): return self.prisma_client.db.litellm_objectpermissiontable @property - def model_class(self) -> Type[LiteLLM_ObjectPermissionTable]: + def model_class(self) -> type[LiteLLM_ObjectPermissionTable]: return LiteLLM_ObjectPermissionTable async def find_by_id( self, object_permission_id: str, id_field: str = "object_permission_id" - ) -> Optional[LiteLLM_ObjectPermissionTable]: + ) -> LiteLLM_ObjectPermissionTable | None: return await super().find_by_id(object_permission_id, id_field) async def create_permission( self, - mcp_servers: Optional[List[str]] = None, - mcp_access_groups: Optional[List[str]] = None, - mcp_tool_permissions: Optional[Dict[str, List[str]]] = None, - vector_stores: Optional[List[str]] = None, - agents: Optional[List[str]] = None, - agent_access_groups: Optional[List[str]] = None, - models: Optional[List[str]] = None, - blocked_tools: Optional[List[str]] = None, - mcp_toolsets: Optional[List[str]] = None, - search_tools: Optional[List[str]] = None, + mcp_servers: list[str] | None = None, + mcp_access_groups: list[str] | None = None, + mcp_tool_permissions: dict[str, list[str]] | None = None, + vector_stores: list[str] | None = None, + agents: list[str] | None = None, + agent_access_groups: list[str] | None = None, + models: list[str] | None = None, + blocked_tools: list[str] | None = None, + mcp_toolsets: list[str] | None = None, + search_tools: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable: """Create a new object permission record.""" - data: Dict[str, Any] = {} + data: dict[str, Any] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: @@ -65,19 +65,19 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): async def update_permission( self, object_permission_id: str, - mcp_servers: Optional[List[str]] = None, - mcp_access_groups: Optional[List[str]] = None, - mcp_tool_permissions: Optional[Dict[str, List[str]]] = None, - vector_stores: Optional[List[str]] = None, - agents: Optional[List[str]] = None, - agent_access_groups: Optional[List[str]] = None, - models: Optional[List[str]] = None, - blocked_tools: Optional[List[str]] = None, - mcp_toolsets: Optional[List[str]] = None, - search_tools: Optional[List[str]] = None, - ) -> Optional[LiteLLM_ObjectPermissionTable]: + mcp_servers: list[str] | None = None, + mcp_access_groups: list[str] | None = None, + mcp_tool_permissions: dict[str, list[str]] | None = None, + vector_stores: list[str] | None = None, + agents: list[str] | None = None, + agent_access_groups: list[str] | None = None, + models: list[str] | None = None, + blocked_tools: list[str] | None = None, + mcp_toolsets: list[str] | None = None, + search_tools: list[str] | None = None, + ) -> LiteLLM_ObjectPermissionTable | None: """Update an object permission record.""" - data: Dict[str, Any] = {} + data: dict[str, Any] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: @@ -101,6 +101,6 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): return await self.update(object_permission_id, data, id_field="object_permission_id") - async def delete_permission(self, object_permission_id: str) -> Optional[LiteLLM_ObjectPermissionTable]: + async def delete_permission(self, object_permission_id: str) -> LiteLLM_ObjectPermissionTable | None: """Delete an object permission record.""" return await self.delete(object_permission_id, id_field="object_permission_id") diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index 99c4a881736..776126a888d 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -2,7 +2,7 @@ Organization repository for database operations on LiteLLM_OrganizationTable. """ -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm.models.organization import LiteLLM_OrganizationTable from litellm.repositories.base_repository import BaseRepository @@ -16,15 +16,15 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): return self.prisma_client.db.litellm_organizationtable @property - def model_class(self) -> Type[LiteLLM_OrganizationTable]: + def model_class(self) -> type[LiteLLM_OrganizationTable]: return LiteLLM_OrganizationTable async def find_by_id( self, organization_id: str, id_field: str = "organization_id" - ) -> Optional[LiteLLM_OrganizationTable]: + ) -> LiteLLM_OrganizationTable | None: return await super().find_by_id(organization_id, id_field) - async def find_by_alias(self, organization_alias: str) -> Optional[LiteLLM_OrganizationTable]: + async def find_by_alias(self, organization_alias: str) -> LiteLLM_OrganizationTable | None: """Find an organization by alias.""" organizations = await self.find_many(where={"organization_alias": organization_alias}) return organizations[0] if organizations else None @@ -34,13 +34,13 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): organization_alias: str, budget_id: str, created_by: str, - organization_id: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - models: Optional[List[str]] = None, - object_permission_id: Optional[str] = None, + organization_id: str | None = None, + metadata: dict[str, Any] | None = None, + models: list[str] | None = None, + object_permission_id: str | None = None, ) -> LiteLLM_OrganizationTable: """Create a new organization.""" - data: Dict[str, Any] = { + data: dict[str, Any] = { "organization_alias": organization_alias, "budget_id": budget_id, "created_by": created_by, @@ -61,14 +61,14 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): self, organization_id: str, updated_by: str, - organization_alias: Optional[str] = None, - budget_id: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - models: Optional[List[str]] = None, - object_permission_id: Optional[str] = None, - ) -> Optional[LiteLLM_OrganizationTable]: + organization_alias: str | None = None, + budget_id: str | None = None, + metadata: dict[str, Any] | None = None, + models: list[str] | None = None, + object_permission_id: str | None = None, + ) -> LiteLLM_OrganizationTable | None: """Update an organization.""" - data: Dict[str, Any] = {"updated_by": updated_by} + data: dict[str, Any] = {"updated_by": updated_by} if organization_alias is not None: data["organization_alias"] = organization_alias if budget_id is not None: @@ -82,10 +82,10 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): return await self.update(organization_id, data, id_field="organization_id") - async def delete_organization(self, organization_id: str) -> Optional[LiteLLM_OrganizationTable]: + async def delete_organization(self, organization_id: str) -> LiteLLM_OrganizationTable | None: """Delete an organization.""" return await self.delete(organization_id, id_field="organization_id") - async def update_spend(self, organization_id: str, spend: float) -> Optional[LiteLLM_OrganizationTable]: + async def update_spend(self, organization_id: str, spend: float) -> LiteLLM_OrganizationTable | None: """Update organization spend.""" return await self.update(organization_id, {"spend": spend}, id_field="organization_id") diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 27cb346e1b1..db0c54db56d 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -2,7 +2,7 @@ Project repository for database operations on LiteLLM_ProjectTable. """ -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm.models.project import LiteLLM_ProjectTable from litellm.repositories.base_repository import BaseRepository @@ -16,37 +16,37 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): return self.prisma_client.db.litellm_projecttable @property - def model_class(self) -> Type[LiteLLM_ProjectTable]: + def model_class(self) -> type[LiteLLM_ProjectTable]: return LiteLLM_ProjectTable - async def find_by_id(self, project_id: str, id_field: str = "project_id") -> Optional[LiteLLM_ProjectTable]: + async def find_by_id(self, project_id: str, id_field: str = "project_id") -> LiteLLM_ProjectTable | None: return await super().find_by_id(project_id, id_field) - async def find_by_alias(self, project_alias: str) -> Optional[LiteLLM_ProjectTable]: + async def find_by_alias(self, project_alias: str) -> LiteLLM_ProjectTable | None: """Find a project by alias.""" projects = await self.find_many(where={"project_alias": project_alias}) return projects[0] if projects else None - async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProjectTable]: + async def find_by_team_id(self, team_id: str) -> list[LiteLLM_ProjectTable]: """Find all projects belonging to a team.""" return await self.find_many(where={"team_id": team_id}) async def create_project( self, created_by: str, - project_id: Optional[str] = None, - project_alias: Optional[str] = None, - description: Optional[str] = None, - team_id: Optional[str] = None, - budget_id: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - models: Optional[List[str]] = None, - model_rpm_limit: Optional[Dict[str, int]] = None, - model_tpm_limit: Optional[Dict[str, int]] = None, - object_permission_id: Optional[str] = None, + project_id: str | None = None, + project_alias: str | None = None, + description: str | None = None, + team_id: str | None = None, + budget_id: str | None = None, + metadata: dict[str, Any] | None = None, + models: list[str] | None = None, + model_rpm_limit: dict[str, int] | None = None, + model_tpm_limit: dict[str, int] | None = None, + object_permission_id: str | None = None, ) -> LiteLLM_ProjectTable: """Create a new project.""" - data: Dict[str, Any] = { + data: dict[str, Any] = { "created_by": created_by, "updated_by": created_by, } @@ -77,19 +77,19 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): self, project_id: str, updated_by: str, - project_alias: Optional[str] = None, - description: Optional[str] = None, - team_id: Optional[str] = None, - budget_id: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - models: Optional[List[str]] = None, - model_rpm_limit: Optional[Dict[str, int]] = None, - model_tpm_limit: Optional[Dict[str, int]] = None, - blocked: Optional[bool] = None, - object_permission_id: Optional[str] = None, - ) -> Optional[LiteLLM_ProjectTable]: + project_alias: str | None = None, + description: str | None = None, + team_id: str | None = None, + budget_id: str | None = None, + metadata: dict[str, Any] | None = None, + models: list[str] | None = None, + model_rpm_limit: dict[str, int] | None = None, + model_tpm_limit: dict[str, int] | None = None, + blocked: bool | None = None, + object_permission_id: str | None = None, + ) -> LiteLLM_ProjectTable | None: """Update a project.""" - data: Dict[str, Any] = {"updated_by": updated_by} + data: dict[str, Any] = {"updated_by": updated_by} if project_alias is not None: data["project_alias"] = project_alias if description is not None: @@ -113,10 +113,10 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): return await self.update(project_id, data, id_field="project_id") - async def delete_project(self, project_id: str) -> Optional[LiteLLM_ProjectTable]: + async def delete_project(self, project_id: str) -> LiteLLM_ProjectTable | None: """Delete a project.""" return await self.delete(project_id, id_field="project_id") - async def update_spend(self, project_id: str, spend: float) -> Optional[LiteLLM_ProjectTable]: + async def update_spend(self, project_id: str, spend: float) -> LiteLLM_ProjectTable | None: """Update project spend.""" return await self.update(project_id, {"spend": spend}, id_field="project_id") diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 25437cfe49a..6248213f223 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -5,7 +5,7 @@ Team repository for database operations on LiteLLM_TeamTable. import json from collections.abc import Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type +from typing import TYPE_CHECKING, Any from pydantic import TypeAdapter @@ -42,10 +42,10 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return self.prisma_client.db.litellm_deletedteamtable @property - def model_class(self) -> Type[LiteLLM_TeamTable]: + def model_class(self) -> type[LiteLLM_TeamTable]: return LiteLLM_TeamTable - def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_TeamTable]: + def _to_model(self, record: DbRecord | None) -> LiteLLM_TeamTable | None: """Convert a database record to a Team model.""" if record is None: return None @@ -57,7 +57,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) - async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]: + async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member]: """Return the team's members_with_roles, locking the row FOR UPDATE. Must be called inside a transaction so the row lock is held until @@ -75,27 +75,27 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return [] return _MEMBERS_WITH_ROLES_ADAPTER.validate_python(parsed) - async def find_by_id(self, team_id: str, id_field: str = "team_id") -> Optional[LiteLLM_TeamTable]: + async def find_by_id(self, team_id: str, id_field: str = "team_id") -> LiteLLM_TeamTable | None: return await super().find_by_id(team_id, id_field) - async def find_by_alias(self, team_alias: str) -> Optional[LiteLLM_TeamTable]: + async def find_by_alias(self, team_alias: str) -> LiteLLM_TeamTable | None: """Find a team by alias.""" records = await self.table.find_many(where={"team_alias": team_alias}) if records: return self._to_model(records[0]) return None - async def find_by_organization_id(self, organization_id: str) -> List[LiteLLM_TeamTable]: + async def find_by_organization_id(self, organization_id: str) -> list[LiteLLM_TeamTable]: """Find all teams belonging to an organization.""" records = await self.table.find_many(where={"organization_id": organization_id}) return self._to_model_list(records) - async def find_by_member(self, user_id: str) -> List[LiteLLM_TeamTable]: + async def find_by_member(self, user_id: str) -> list[LiteLLM_TeamTable]: """Find all teams where user is a member.""" records = await self.table.find_many(where={"members": {"has": user_id}}) return self._to_model_list(records) - async def find_by_admin(self, user_id: str) -> List[LiteLLM_TeamTable]: + async def find_by_admin(self, user_id: str) -> list[LiteLLM_TeamTable]: """Find all teams where user is an admin.""" records = await self.table.find_many(where={"admins": {"has": user_id}}) return self._to_model_list(records) @@ -103,23 +103,23 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): async def create_team( self, team_id: str, - team_alias: Optional[str] = None, - organization_id: Optional[str] = None, - admins: Optional[List[str]] = None, - members: Optional[List[str]] = None, - members_with_roles: Optional[Mapping[str, object]] = None, - metadata: Optional[Mapping[str, object]] = None, - max_budget: Optional[float] = None, - soft_budget: Optional[float] = None, - models: Optional[List[str]] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - budget_duration: Optional[str] = None, - object_permission_id: Optional[str] = None, + team_alias: str | None = None, + organization_id: str | None = None, + admins: list[str] | None = None, + members: list[str] | None = None, + members_with_roles: Mapping[str, object] | None = None, + metadata: Mapping[str, object] | None = None, + max_budget: float | None = None, + soft_budget: float | None = None, + models: list[str] | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + budget_duration: str | None = None, + object_permission_id: str | None = None, ) -> LiteLLM_TeamTable: """Create a new team.""" - data: Dict[str, object] = {"team_id": team_id} + data: dict[str, object] = {"team_id": team_id} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -154,24 +154,24 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): async def update_team( self, team_id: str, - team_alias: Optional[str] = None, - organization_id: Optional[str] = None, - admins: Optional[List[str]] = None, - members: Optional[List[str]] = None, - members_with_roles: Optional[Mapping[str, object]] = None, - metadata: Optional[Mapping[str, object]] = None, - max_budget: Optional[float] = None, - soft_budget: Optional[float] = None, - models: Optional[List[str]] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - budget_duration: Optional[str] = None, - blocked: Optional[bool] = None, - object_permission_id: Optional[str] = None, - ) -> Optional[LiteLLM_TeamTable]: + team_alias: str | None = None, + organization_id: str | None = None, + admins: list[str] | None = None, + members: list[str] | None = None, + members_with_roles: Mapping[str, object] | None = None, + metadata: Mapping[str, object] | None = None, + max_budget: float | None = None, + soft_budget: float | None = None, + models: list[str] | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + budget_duration: str | None = None, + blocked: bool | None = None, + object_permission_id: str | None = None, + ) -> LiteLLM_TeamTable | None: """Update a team.""" - data: Dict[str, object] = {} + data: dict[str, object] = {} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -208,10 +208,10 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): async def delete_team( self, team_id: str, - deleted_by: Optional[str] = None, - deleted_by_api_key: Optional[str] = None, - litellm_changed_by: Optional[str] = None, - ) -> Optional[LiteLLM_TeamTable]: + deleted_by: str | None = None, + deleted_by_api_key: str | None = None, + litellm_changed_by: str | None = None, + ) -> LiteLLM_TeamTable | None: """Delete a team and archive it to the deleted teams table. Uses a transaction to ensure atomicity of the archive-then-delete operation. @@ -232,9 +232,9 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return team - def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, object]: + def _build_archive_data(self, team: LiteLLM_TeamTable) -> dict[str, object]: """Build archive data dict with only columns that exist in LiteLLM_DeletedTeamTable.""" - data: Dict[str, object] = {"team_id": team.team_id} + data: dict[str, object] = {"team_id": team.team_id} if team.team_alias is not None: data["team_alias"] = team.team_alias if team.organization_id is not None: @@ -278,11 +278,11 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): data["allow_team_guardrail_config"] = team.allow_team_guardrail_config return data - async def update_spend(self, team_id: str, spend: float) -> Optional[LiteLLM_TeamTable]: + async def update_spend(self, team_id: str, spend: float) -> LiteLLM_TeamTable | None: """Update team spend.""" return await self.update(team_id, {"spend": spend}, id_field="team_id") - async def add_member(self, team_id: str, user_id: str) -> Optional[LiteLLM_TeamTable]: + async def add_member(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None: """Add a member to a team using atomic array push operation.""" if not await self.exists(team_id, id_field="team_id"): return None @@ -293,7 +293,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): ) return self._to_model(record) - async def remove_member(self, team_id: str, user_id: str) -> Optional[LiteLLM_TeamTable]: + async def remove_member(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None: """Remove a member from a team. Note: Prisma doesn't support atomic array removal, so we use a @@ -307,7 +307,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): members = [m for m in team.members if m != user_id] return await self.update(team_id, {"members": members}, id_field="team_id") - async def add_admin(self, team_id: str, user_id: str) -> Optional[LiteLLM_TeamTable]: + async def add_admin(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None: """Add an admin to a team using atomic array push operation.""" if not await self.exists(team_id, id_field="team_id"): return None @@ -318,7 +318,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): ) return self._to_model(record) - async def remove_admin(self, team_id: str, user_id: str) -> Optional[LiteLLM_TeamTable]: + async def remove_admin(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None: """Remove an admin from a team. Note: Prisma doesn't support atomic array removal, so we use a @@ -332,7 +332,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): admins = [a for a in team.admins if a != user_id] return await self.update(team_id, {"admins": admins}, id_field="team_id") - async def add_models(self, team_id: str, models: List[str]) -> Optional[LiteLLM_TeamTable]: + async def add_models(self, team_id: str, models: list[str]) -> LiteLLM_TeamTable | None: """Add models to a team's allowed models list using atomic array push.""" if not await self.exists(team_id, id_field="team_id"): return None @@ -343,7 +343,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): ) return self._to_model(record) - async def remove_models(self, team_id: str, models: List[str]) -> Optional[LiteLLM_TeamTable]: + async def remove_models(self, team_id: str, models: list[str]) -> LiteLLM_TeamTable | None: """Remove models from a team's allowed models list. Note: Prisma doesn't support atomic array removal, so we use a diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 14b49f05304..5eb326bda18 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -4,7 +4,7 @@ User repository for database operations on LiteLLM_UserTable. import json from collections.abc import Mapping -from typing import Any, Dict, List, Optional, Type +from typing import Any from litellm.models.user import LiteLLM_UserTable from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict @@ -20,10 +20,10 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): return self.prisma_client.db.litellm_usertable @property - def model_class(self) -> Type[LiteLLM_UserTable]: + def model_class(self) -> type[LiteLLM_UserTable]: return LiteLLM_UserTable - def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_UserTable]: + def _to_model(self, record: DbRecord | None) -> LiteLLM_UserTable | None: """Convert a database record to a User model.""" if record is None: return None @@ -35,23 +35,23 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): } ) - async def find_by_id(self, id_value: str, id_field: str = "user_id") -> Optional[LiteLLM_UserTable]: + async def find_by_id(self, id_value: str, id_field: str = "user_id") -> LiteLLM_UserTable | None: return await super().find_by_id(id_value, id_field) - async def find_by_email(self, user_email: str) -> Optional[LiteLLM_UserTable]: + async def find_by_email(self, user_email: str) -> LiteLLM_UserTable | None: """Find a user by email.""" records = await self.find_many(where={"user_email": user_email}) return records[0] if records else None - async def find_by_sso_id(self, sso_user_id: str) -> Optional[LiteLLM_UserTable]: + async def find_by_sso_id(self, sso_user_id: str) -> LiteLLM_UserTable | None: """Find a user by SSO ID.""" return await self.find_by_id(sso_user_id, id_field="sso_user_id") - async def find_by_organization_id(self, organization_id: str) -> List[LiteLLM_UserTable]: + async def find_by_organization_id(self, organization_id: str) -> list[LiteLLM_UserTable]: """Find all users in an organization.""" return await self.find_many(where={"organization_id": organization_id}) - async def find_by_team_id(self, team_id: str) -> List[LiteLLM_UserTable]: + async def find_by_team_id(self, team_id: str) -> list[LiteLLM_UserTable]: """Find all users in a team.""" return await self.find_many(where={"teams": {"has": team_id}}) @@ -72,27 +72,27 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): async def create_user( self, user_id: str, - user_alias: Optional[str] = None, - team_id: Optional[str] = None, - sso_user_id: Optional[str] = None, - organization_id: Optional[str] = None, - password: Optional[str] = None, - teams: Optional[List[str]] = None, - user_role: Optional[str] = None, - max_budget: Optional[float] = None, - user_email: Optional[str] = None, - models: Optional[List[str]] = None, - metadata: Optional[Mapping[str, object]] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - budget_duration: Optional[str] = None, - allowed_cache_controls: Optional[List[str]] = None, - policies: Optional[List[str]] = None, - object_permission_id: Optional[str] = None, + user_alias: str | None = None, + team_id: str | None = None, + sso_user_id: str | None = None, + organization_id: str | None = None, + password: str | None = None, + teams: list[str] | None = None, + user_role: str | None = None, + max_budget: float | None = None, + user_email: str | None = None, + models: list[str] | None = None, + metadata: Mapping[str, object] | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + budget_duration: str | None = None, + allowed_cache_controls: list[str] | None = None, + policies: list[str] | None = None, + object_permission_id: str | None = None, ) -> LiteLLM_UserTable: """Create a new user.""" - data: Dict[str, object] = {"user_id": user_id} + data: dict[str, object] = {"user_id": user_id} if user_alias is not None: data["user_alias"] = user_alias if team_id is not None: @@ -135,27 +135,27 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): async def update_user( self, user_id: str, - user_alias: Optional[str] = None, - team_id: Optional[str] = None, - sso_user_id: Optional[str] = None, - organization_id: Optional[str] = None, - password: Optional[str] = None, - teams: Optional[List[str]] = None, - user_role: Optional[str] = None, - max_budget: Optional[float] = None, - user_email: Optional[str] = None, - models: Optional[List[str]] = None, - metadata: Optional[Mapping[str, object]] = None, - max_parallel_requests: Optional[int] = None, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - budget_duration: Optional[str] = None, - allowed_cache_controls: Optional[List[str]] = None, - policies: Optional[List[str]] = None, - object_permission_id: Optional[str] = None, - ) -> Optional[LiteLLM_UserTable]: + user_alias: str | None = None, + team_id: str | None = None, + sso_user_id: str | None = None, + organization_id: str | None = None, + password: str | None = None, + teams: list[str] | None = None, + user_role: str | None = None, + max_budget: float | None = None, + user_email: str | None = None, + models: list[str] | None = None, + metadata: Mapping[str, object] | None = None, + max_parallel_requests: int | None = None, + tpm_limit: int | None = None, + rpm_limit: int | None = None, + budget_duration: str | None = None, + allowed_cache_controls: list[str] | None = None, + policies: list[str] | None = None, + object_permission_id: str | None = None, + ) -> LiteLLM_UserTable | None: """Update a user.""" - data: Dict[str, object] = {} + data: dict[str, object] = {} if user_alias is not None: data["user_alias"] = user_alias if team_id is not None: @@ -195,22 +195,22 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): return await self.update(user_id, data, id_field="user_id") - async def delete_user(self, user_id: str) -> Optional[LiteLLM_UserTable]: + async def delete_user(self, user_id: str) -> LiteLLM_UserTable | None: """Delete a user.""" return await self.delete(user_id, id_field="user_id") - async def update_spend(self, user_id: str, spend: float) -> Optional[LiteLLM_UserTable]: + async def update_spend(self, user_id: str, spend: float) -> LiteLLM_UserTable | None: """Update user spend.""" return await self.update(user_id, {"spend": spend}, id_field="user_id") - async def add_to_team(self, user_id: str, team_id: str) -> Optional[LiteLLM_UserTable]: + async def add_to_team(self, user_id: str, team_id: str) -> LiteLLM_UserTable | None: """Add a user to a team using atomic array push operation.""" if not await self.exists(user_id, id_field="user_id"): return None return await self.update(user_id, {"teams": {"push": team_id}}, id_field="user_id") - async def remove_from_team(self, user_id: str, team_id: str) -> Optional[LiteLLM_UserTable]: + async def remove_from_team(self, user_id: str, team_id: str) -> LiteLLM_UserTable | None: """Remove a user from a team. Note: Prisma doesn't support atomic array removal, so we use a diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 75a35f8341e..03c13e504ac 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -2,7 +2,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Literal, Union +from typing import Any, Literal import litellm from litellm._logging import verbose_logger @@ -30,16 +30,16 @@ base_llm_http_handler = BaseLLMHTTPHandler() async def arerank( model: str, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: ( Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"] | None ) = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = None, max_chunks_per_doc: int | None = None, **kwargs, -) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: +) -> RerankResponse | Coroutine[Any, Any, RerankResponse]: """ Async: Reranks a list of documents based on their relevance to the query """ @@ -77,7 +77,7 @@ async def arerank( def rerank( model: str, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: ( Literal[ "cohere", @@ -94,12 +94,12 @@ def rerank( | None ) = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, **kwargs, -) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: +) -> RerankResponse | Coroutine[Any, Any, RerankResponse]: """ Reranks a list of documents based on their relevance to the query """ @@ -144,7 +144,7 @@ def rerank( present_version_params=present_version_params, ) - optional_rerank_params: Dict = get_optional_rerank_params( + optional_rerank_params: dict = get_optional_rerank_params( rerank_provider_config=rerank_provider_config, model=model, drop_params=kwargs.get("drop_params") or litellm.drop_params or False, @@ -534,5 +534,5 @@ def rerank( # Placeholder return return response except Exception as e: - verbose_logger.error(f"Error in rerank: {str(e)}") + verbose_logger.error(f"Error in rerank: {e!s}") raise exception_type(model=model, custom_llm_provider=custom_llm_provider, original_exception=e) diff --git a/litellm/rerank_api/rerank_utils.py b/litellm/rerank_api/rerank_utils.py index 856029e45ea..36c7835dca0 100644 --- a/litellm/rerank_api/rerank_utils.py +++ b/litellm/rerank_api/rerank_utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Union +from typing import Any from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig @@ -8,16 +8,16 @@ def get_optional_rerank_params( model: str, drop_params: bool, query: str, - documents: List[Union[str, Dict[str, Any]]], + documents: list[str | dict[str, Any]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: List[str] | None = None, + rank_fields: list[str] | None = None, return_documents: bool | None = True, max_chunks_per_doc: int | None = None, max_tokens_per_doc: int | None = None, instruction: str | None = None, non_default_params: dict | None = None, -) -> Dict: +) -> dict: all_non_default_params = non_default_params or {} if query is not None: all_non_default_params["query"] = query diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py index 0e01e7d977f..2b019c468ad 100644 --- a/litellm/responses/file_search/emulated_handler.py +++ b/litellm/responses/file_search/emulated_handler.py @@ -15,7 +15,7 @@ import json import time import uuid from collections.abc import Iterable -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, cast from litellm._internal_context import is_internal_call from litellm._logging import verbose_logger @@ -34,7 +34,7 @@ FILE_SEARCH_FUNCTION_NAME = "litellm_file_search" def should_use_emulated_file_search( - tools: Optional[Iterable[ToolParam]], + tools: Iterable[ToolParam] | None, provider_config: Any, # BaseResponsesAPIConfig ) -> bool: """Return True when there is a file_search tool and the provider can't handle it natively.""" @@ -51,7 +51,7 @@ def should_use_emulated_file_search( # --------------------------------------------------------------------------- -def _build_function_tool(vector_store_ids: List[str]) -> Dict[str, Any]: +def _build_function_tool(vector_store_ids: list[str]) -> dict[str, Any]: """ Create a Responses API function-tool definition that describes file search. The function accepts one or more natural-language queries (like OpenAI's native @@ -95,16 +95,16 @@ def _build_function_tool(vector_store_ids: List[str]) -> Dict[str, Any]: def _replace_file_search_tools( - tools: Optional[Iterable[ToolParam]], -) -> Tuple[List[Dict[str, Any]], List[str]]: + tools: Iterable[ToolParam] | None, +) -> tuple[list[dict[str, Any]], list[str]]: """ Replace all file_search tools with a single function tool. Returns: (new_tools_list, all_vector_store_ids) """ - non_file_search: List[Dict[str, Any]] = [] - vector_store_ids: List[str] = [] + non_file_search: list[dict[str, Any]] = [] + vector_store_ids: list[str] = [] for tool in tools or []: if isinstance(tool, dict) and tool.get("type") == "file_search": @@ -114,7 +114,7 @@ def _replace_file_search_tools( non_file_search.append(tool) # Deduplicate while preserving order - unique_ids: List[str] = list(dict.fromkeys(vector_store_ids)) + unique_ids: list[str] = list(dict.fromkeys(vector_store_ids)) if unique_ids: non_file_search.append(_build_function_tool(unique_ids)) @@ -127,9 +127,9 @@ def _replace_file_search_tools( async def _run_vector_searches( - queries: List[str], - vector_store_ids: List[str], -) -> Tuple[List[str], List[VectorStoreSearchResult]]: + queries: list[str], + vector_store_ids: list[str], +) -> tuple[list[str], list[VectorStoreSearchResult]]: """ Run `asearch` against all vector stores for all queries and collect results. @@ -142,7 +142,7 @@ async def _run_vector_searches( """ import litellm.vector_stores.main as vs_main - all_results: List[VectorStoreSearchResult] = [] + all_results: list[VectorStoreSearchResult] = [] ids_to_search = vector_store_ids # Execute each query against all vector stores @@ -180,13 +180,13 @@ def _get_field(result: Any, key: str, default: Any = None) -> Any: def _format_search_results_as_tool_output( - results: List[VectorStoreSearchResult], + results: list[VectorStoreSearchResult], ) -> str: """Serialize search results into a string to pass back as the tool's output.""" if not results: return "No results found in the vector store." - parts: List[str] = [] + parts: list[str] = [] for i, result in enumerate(results, 1): score = _get_field(result, "score") file_id = _get_field(result, "file_id") @@ -210,8 +210,8 @@ def _format_search_results_as_tool_output( def _build_search_results_for_include( - results: List[VectorStoreSearchResult], -) -> List[Dict[str, Any]]: + results: list[VectorStoreSearchResult], +) -> list[dict[str, Any]]: """ Convert VectorStoreSearchResult objects to the format expected in file_search_call.search_results (mirrors OpenAI's include= format). @@ -220,7 +220,7 @@ def _build_search_results_for_include( behaviour of OpenAI's native file_search which surfaces every relevant chunk even when multiple chunks originate from the same document. """ - formatted: List[Dict[str, Any]] = [] + formatted: list[dict[str, Any]] = [] for result in results: file_id = _get_field(result, "file_id") or "" content_items = _get_field(result, "content") or [] @@ -240,10 +240,10 @@ def _build_search_results_for_include( def _build_file_search_call_output( call_id: str, - queries: List[str], - results: Optional[List[VectorStoreSearchResult]] = None, + queries: list[str], + results: list[VectorStoreSearchResult] | None = None, include_search_results: bool = False, -) -> Dict[str, Any]: +) -> dict[str, Any]: """Build the file_search_call output item (mirrors OpenAI's format). Args: @@ -266,14 +266,14 @@ def _build_file_search_call_output( def _build_file_citation_annotations( - results: List[VectorStoreSearchResult], + results: list[VectorStoreSearchResult], text: str, -) -> List[Dict[str, Any]]: +) -> list[dict[str, Any]]: """ Build file_citation annotations for the text. Each result with a file_id gets a citation at the end of the text. """ - annotations: List[Dict[str, Any]] = [] + annotations: list[dict[str, Any]] = [] index = len(text) # cite at end of text block seen_file_ids: set = set() @@ -297,8 +297,8 @@ def _build_file_citation_annotations( def _build_message_output( response_text: str, - results: List[VectorStoreSearchResult], -) -> Dict[str, Any]: + results: list[VectorStoreSearchResult], +) -> dict[str, Any]: """Build the message output item with optional file_citation annotations.""" annotations = _build_file_citation_annotations(results, response_text) return { @@ -330,9 +330,9 @@ def _extract_text_from_responses_output(response: ResponsesAPIResponse) -> str: def _synthesize_responses_api_response( original_response: ResponsesAPIResponse, - file_search_call_output: Dict[str, Any], - message_output: Dict[str, Any], - first_response: Optional[ResponsesAPIResponse] = None, + file_search_call_output: dict[str, Any], + message_output: dict[str, Any], + first_response: ResponsesAPIResponse | None = None, ) -> ResponsesAPIResponse: """ Return a new ResponsesAPIResponse with: @@ -343,14 +343,14 @@ def _synthesize_responses_api_response( synthesized _hidden_params so that billing callbacks see the total cost of both provider calls that the emulated flow makes. """ - synthesized_output: List[Dict[str, Any]] = [file_search_call_output, message_output] + synthesized_output: list[dict[str, Any]] = [file_search_call_output, message_output] synthesized = ResponsesAPIResponse( id=getattr(original_response, "id", f"resp_{uuid.uuid4().hex}"), object="response", created_at=getattr(original_response, "created_at", int(time.time())), status="completed", model=getattr(original_response, "model", ""), - output=cast(List[Union[ResponseOutputItem, Dict[str, Any]]], synthesized_output), + output=cast(list[ResponseOutputItem | dict[str, Any]], synthesized_output), usage=getattr(original_response, "usage", None), error=None, ) @@ -382,9 +382,9 @@ async def _call_aresponses(input, model, tools, **kwargs): # pragma: no cover def _prepare_emulated_file_search_call( - kwargs: Dict[str, Any], -) -> Tuple[bool, Dict[str, Any]]: - include_items: List[str] = list(kwargs.get("include") or []) + kwargs: dict[str, Any], +) -> tuple[bool, dict[str, Any]]: + include_items: list[str] = list(kwargs.get("include") or []) include_search_results = "file_search_call.results" in include_items original_stream = kwargs.get("stream") @@ -398,7 +398,7 @@ def _prepare_emulated_file_search_call( return include_search_results, updated_kwargs -def _extract_tool_call_fields(tool_call: Any, fallback_call_id: str) -> Tuple[str, str]: +def _extract_tool_call_fields(tool_call: Any, fallback_call_id: str) -> tuple[str, str]: """Extract (call_id, raw_arguments_string) from a dict or Pydantic tool_call item.""" if isinstance(tool_call, dict): call_id = str(tool_call.get("call_id") or tool_call.get("id") or fallback_call_id) @@ -410,7 +410,7 @@ def _extract_tool_call_fields(tool_call: Any, fallback_call_id: str) -> Tuple[st return call_id, raw_args -def _resolve_queries_from_args(args: Dict[str, Any], input: Any) -> List[str]: +def _resolve_queries_from_args(args: dict[str, Any], input: Any) -> list[str]: """Pull the queries list out of parsed tool-call arguments, with backward-compat fallbacks.""" queries_from_call = args.get("queries") if not queries_from_call: @@ -423,15 +423,15 @@ def _resolve_queries_from_args(args: Dict[str, Any], input: Any) -> List[str]: async def _execute_file_search_tool_calls( - file_search_calls: List[Any], - all_vs_ids: List[str], + file_search_calls: list[Any], + all_vs_ids: list[str], input: Any, file_search_call_id: str, -) -> Tuple[List[Dict[str, Any]], List[str], List[VectorStoreSearchResult]]: +) -> tuple[list[dict[str, Any]], list[str], list[VectorStoreSearchResult]]: """Run the vector search for each file_search tool_call and collect results.""" - tool_results: List[Dict[str, Any]] = [] - all_queries: List[str] = [] - all_results: List[VectorStoreSearchResult] = [] + tool_results: list[dict[str, Any]] = [] + all_queries: list[str] = [] + all_results: list[VectorStoreSearchResult] = [] for tool_call in file_search_calls: call_id, raw_args = _extract_tool_call_fields(tool_call, fallback_call_id=file_search_call_id) @@ -467,8 +467,8 @@ async def _execute_file_search_tool_calls( def _build_follow_up_input( input: Any, first_response: ResponsesAPIResponse, - tool_results: List[Dict[str, Any]], -) -> List[Any]: + tool_results: list[dict[str, Any]], +) -> list[Any]: """Assemble the follow-up call input: original messages + first-response output + tool results. Including all output items (text blocks, reasoning, non-file-search calls) ensures providers @@ -478,7 +478,7 @@ def _build_follow_up_input( original_input_items = ( list(input) if isinstance(input, (list, tuple)) else [{"role": "user", "content": str(input)}] ) - first_response_output_items: List[Any] = [] + first_response_output_items: list[Any] = [] for _item in first_response.output: if isinstance(_item, dict): first_response_output_items.append(_item) @@ -493,7 +493,7 @@ def _build_follow_up_input( async def aresponses_with_emulated_file_search( input: Any, model: str, - tools: Optional[Iterable[ToolParam]] = None, + tools: Iterable[ToolParam] | None = None, # Pass-through params — forwarded as-is to the underlying aresponses call **kwargs: Any, ) -> ResponsesAPIResponse: diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index a11f38da150..9ccfcdf3bcb 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -3,7 +3,7 @@ Handler for transforming responses api requests to litellm.completion requests """ from collections.abc import Coroutine -from typing import Any, Dict, Optional, Union +from typing import Any import litellm from litellm.responses.litellm_completion_transformation.streaming_iterator import ( @@ -25,18 +25,18 @@ class LiteLLMCompletionTransformationHandler: def response_api_handler( self, model: str, - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, responses_api_request: ResponsesAPIOptionalRequestParams, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, _is_async: bool = False, - stream: Optional[bool] = None, - extra_headers: Optional[Dict[str, Any]] = None, + stream: bool | None = None, + extra_headers: dict[str, Any] | None = None, **kwargs, - ) -> Union[ - ResponsesAPIResponse, - BaseResponsesAPIStreamingIterator, - Coroutine[Any, Any, Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]], - ]: + ) -> ( + ResponsesAPIResponse + | BaseResponsesAPIStreamingIterator + | Coroutine[Any, Any, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator] + ): litellm_completion_request: dict = ( LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( model=model, @@ -62,7 +62,7 @@ class LiteLLMCompletionTransformationHandler: completion_args.update(litellm_completion_request) completion_args["_skip_responses_api_bridge"] = True - litellm_completion_response: Union[ModelResponse, litellm.CustomStreamWrapper] = litellm.completion( + litellm_completion_response: ModelResponse | litellm.CustomStreamWrapper = litellm.completion( **completion_args, ) @@ -91,11 +91,11 @@ class LiteLLMCompletionTransformationHandler: async def async_response_api_handler( self, litellm_completion_request: dict, - request_input: Union[str, ResponseInputParam], + request_input: str | ResponseInputParam, responses_api_request: ResponsesAPIOptionalRequestParams, **kwargs, - ) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: - previous_response_id: Optional[str] = responses_api_request.get("previous_response_id") + ) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: + previous_response_id: str | None = responses_api_request.get("previous_response_id") if previous_response_id: litellm_completion_request = await LiteLLMCompletionResponsesConfig.async_responses_api_session_handler( previous_response_id=previous_response_id, @@ -107,7 +107,7 @@ class LiteLLMCompletionTransformationHandler: acompletion_args.update(litellm_completion_request) acompletion_args["_skip_responses_api_bridge"] = True - litellm_completion_response: Union[ModelResponse, litellm.CustomStreamWrapper] = await litellm.acompletion( + litellm_completion_response: ModelResponse | litellm.CustomStreamWrapper = await litellm.acompletion( **acompletion_args, ) diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index 68637ea97b3..4faf75b3951 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -1,5 +1,5 @@ import json -from typing import TYPE_CHECKING, Any, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast import litellm from litellm._logging import verbose_proxy_logger @@ -41,23 +41,21 @@ class ResponsesSessionHandler: ) verbose_proxy_logger.debug("inside get_chat_completion_message_history_for_previous_response_id") - all_spend_logs: List[ + all_spend_logs: list[ SpendLogsPayload ] = await ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id(previous_response_id) verbose_proxy_logger.debug("found %s spend logs for this response id", len(all_spend_logs)) - litellm_session_id: Optional[str] = None + litellm_session_id: str | None = None if len(all_spend_logs) > 0: litellm_session_id = all_spend_logs[0].get("session_id") - chat_completion_message_history: List[ - Union[ - AllMessageValues, - GenericChatCompletionMessage, - ChatCompletionMessageToolCall, - ChatCompletionResponseMessage, - Message, - ] + chat_completion_message_history: list[ + AllMessageValues + | GenericChatCompletionMessage + | ChatCompletionMessageToolCall + | ChatCompletionResponseMessage + | Message ] = [] for spend_log in all_spend_logs: chat_completion_message_history = ( @@ -79,14 +77,12 @@ class ResponsesSessionHandler: @staticmethod async def extend_chat_completion_message_with_spend_log_payload( spend_log: SpendLogsPayload, - chat_completion_message_history: List[ - Union[ - AllMessageValues, - GenericChatCompletionMessage, - ChatCompletionMessageToolCall, - ChatCompletionResponseMessage, - Message, - ] + chat_completion_message_history: list[ + AllMessageValues + | GenericChatCompletionMessage + | ChatCompletionMessageToolCall + | ChatCompletionResponseMessage + | Message ], ): """ @@ -99,8 +95,8 @@ class ResponsesSessionHandler: proxy_server_request_dict = await ResponsesSessionHandler.get_proxy_server_request_from_spend_log( spend_log=spend_log, ) - response_input_param: Optional[Union[str, ResponseInputParam]] = None - _messages: Optional[Union[str, ResponseInputParam]] = None + response_input_param: str | ResponseInputParam | None = None + _messages: str | ResponseInputParam | None = None ############################################################ # Add Input messages for this Spend Log @@ -147,12 +143,12 @@ class ResponsesSessionHandler: @staticmethod async def get_proxy_server_request_from_spend_log( spend_log: SpendLogsPayload, - ) -> Optional[dict]: + ) -> dict | None: """ Get the parsed proxy server request from the spend log """ - proxy_server_request: Union[str, dict] = spend_log.get("proxy_server_request") or "{}" - proxy_server_request_dict: Optional[dict] = None + proxy_server_request: str | dict = spend_log.get("proxy_server_request") or "{}" + proxy_server_request_dict: dict | None = None if isinstance(proxy_server_request, dict): proxy_server_request_dict = proxy_server_request else: @@ -163,7 +159,7 @@ class ResponsesSessionHandler: ############################################################ if ResponsesSessionHandler._should_check_cold_storage_for_full_payload(proxy_server_request_dict): # Try to get cold storage object key from spend log metadata - _proxy_server_request_dict: Optional[dict] = None + _proxy_server_request_dict: dict | None = None cold_storage_object_key = ResponsesSessionHandler._get_cold_storage_object_key_from_spend_log(spend_log) if cold_storage_object_key: # Use the object key directly from metadata @@ -180,7 +176,7 @@ class ResponsesSessionHandler: @staticmethod def _get_cold_storage_object_key_from_spend_log( spend_log: SpendLogsPayload, - ) -> Optional[str]: + ) -> str | None: """ Extract the cold storage object key from spend log metadata. @@ -205,7 +201,7 @@ class ResponsesSessionHandler: @staticmethod async def get_proxy_server_request_from_cold_storage_with_object_key( object_key: str, - ) -> Optional[dict]: + ) -> dict | None: """ Get the proxy server request from cold storage using the object key directly. @@ -227,7 +223,7 @@ class ResponsesSessionHandler: @staticmethod def _should_check_cold_storage_for_full_payload( - proxy_server_request_dict: Optional[dict], + proxy_server_request_dict: dict | None, ) -> bool: """ Only check cold storage when both are true @@ -250,7 +246,7 @@ class ResponsesSessionHandler: @staticmethod async def get_all_spend_logs_for_previous_response_id( previous_response_id: str, - ) -> List[SpendLogsPayload]: + ) -> list[SpendLogsPayload]: """ Get all spend logs for a previous response id diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index cf69654d15d..c20b35b6bbc 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -313,7 +313,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def _default_response_created_event_data(self) -> dict: # Use cached response ID if available, otherwise generate a new one if self._cached_response_id is None: - self._cached_response_id = f"resp_{str(uuid.uuid4())}" + self._cached_response_id = f"resp_{uuid.uuid4()!s}" response_created_event_data = { "id": self._cached_response_id, @@ -386,7 +386,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def create_output_item_added_event(self) -> OutputItemAddedEvent: if self._cached_item_id is None: - self._cached_item_id = f"msg_{str(uuid.uuid4())}" + self._cached_item_id = f"msg_{uuid.uuid4()!s}" self._sequence_number += 1 event = OutputItemAddedEvent( @@ -407,7 +407,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def create_content_part_added_event(self) -> ContentPartAddedEvent: if self._cached_item_id is None: - self._cached_item_id = f"msg_{str(uuid.uuid4())}" + self._cached_item_id = f"msg_{uuid.uuid4()!s}" self._sequence_number += 1 event = ContentPartAddedEvent( @@ -528,7 +528,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def create_output_text_done_event(self, litellm_complete_object: ModelResponse) -> OutputTextDoneEvent: if self._cached_item_id is None: - self._cached_item_id = f"msg_{str(uuid.uuid4())}" + self._cached_item_id = f"msg_{uuid.uuid4()!s}" return OutputTextDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, @@ -541,7 +541,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def create_output_content_part_done_event(self, litellm_complete_object: ModelResponse) -> ContentPartDoneEvent: if self._cached_item_id is None: - self._cached_item_id = f"msg_{str(uuid.uuid4())}" + self._cached_item_id = f"msg_{uuid.uuid4()!s}" text = getattr(litellm_complete_object.choices[0].message, "content", "") or "" # type: ignore reasoning_content = getattr(litellm_complete_object.choices[0].message, "reasoning_content", "") or "" # type: ignore @@ -577,7 +577,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def create_output_item_done_event(self, litellm_complete_object: ModelResponse) -> OutputItemDoneEvent: if self._cached_item_id is None: - self._cached_item_id = f"msg_{str(uuid.uuid4())}" + self._cached_item_id = f"msg_{uuid.uuid4()!s}" text = self.litellm_model_response.choices[0].message.content or "" # type: ignore annotations = getattr(self.litellm_model_response.choices[0].message, "annotations", None) # type: ignore diff --git a/litellm/responses/main.py b/litellm/responses/main.py index baead17e58c..b5f2fca116d 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -5,12 +5,8 @@ from functools import partial from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Type, - Union, cast, ) @@ -76,7 +72,7 @@ litellm_completion_transformation_handler = LiteLLMCompletionTransformationHandl ################################################# -def _has_file_search_tool(tools: Optional[Any]) -> bool: +def _has_file_search_tool(tools: Any | None) -> bool: """Return True if any tool in the list has type 'file_search'.""" if not tools: return False @@ -136,36 +132,36 @@ def mock_responses_api_response( async def aresponses_api_with_mcp( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - include: Optional[List[ResponseIncludable]] = None, - instructions: Optional[str] = None, - max_output_tokens: Optional[int] = None, - prompt: Optional[PromptObject] = None, - metadata: Optional[Dict[str, Any]] = None, - parallel_tool_calls: Optional[bool] = None, - previous_response_id: Optional[str] = None, - reasoning: Optional[Reasoning] = None, - store: Optional[bool] = None, - background: Optional[bool] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, + include: list[ResponseIncludable] | None = None, + instructions: str | None = None, + max_output_tokens: int | None = None, + prompt: PromptObject | None = None, + metadata: dict[str, Any] | None = None, + parallel_tool_calls: bool | None = None, + previous_response_id: str | None = None, + reasoning: Reasoning | None = None, + store: bool | None = None, + background: bool | None = None, + stream: bool | None = None, + temperature: float | None = None, text: Optional["ResponseText"] = None, - tool_choice: Optional[ToolChoice] = None, - tools: Optional[Iterable[ToolParam]] = None, - top_p: Optional[float] = None, - truncation: Optional[Literal["auto", "disabled"]] = None, - user: Optional[str] = None, + tool_choice: ToolChoice | None = None, + tools: Iterable[ToolParam] | None = None, + top_p: float | None = None, + truncation: Literal["auto", "disabled"] | None = None, + user: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: +) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: """ Async version of responses API with MCP integration. @@ -190,8 +186,8 @@ async def aresponses_api_with_mcp( user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth") # Extract MCP auth headers from request (for dynamic auth when fetching tools) - mcp_auth_header: Optional[str] = None - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None + mcp_auth_header: str | None = None + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None secret_fields = kwargs.get("secret_fields") if secret_fields and isinstance(secret_fields, dict): ( @@ -401,39 +397,39 @@ async def aresponses_api_with_mcp( @client async def aresponses( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - include: Optional[List[ResponseIncludable]] = None, - instructions: Optional[str] = None, - max_output_tokens: Optional[int] = None, - prompt: Optional[PromptObject] = None, - metadata: Optional[Dict[str, Any]] = None, - parallel_tool_calls: Optional[bool] = None, - previous_response_id: Optional[str] = None, - reasoning: Optional[Reasoning] = None, - store: Optional[bool] = None, - background: Optional[bool] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, + include: list[ResponseIncludable] | None = None, + instructions: str | None = None, + max_output_tokens: int | None = None, + prompt: PromptObject | None = None, + metadata: dict[str, Any] | None = None, + parallel_tool_calls: bool | None = None, + previous_response_id: str | None = None, + reasoning: Reasoning | None = None, + store: bool | None = None, + background: bool | None = None, + stream: bool | None = None, + temperature: float | None = None, text: Optional["ResponseText"] = None, - text_format: Optional[Union[Type["BaseModel"], dict]] = None, - tool_choice: Optional[ToolChoice] = None, - tools: Optional[Iterable[ToolParam]] = None, - top_p: Optional[float] = None, - truncation: Optional[Literal["auto", "disabled"]] = None, - user: Optional[str] = None, - service_tier: Optional[str] = None, - safety_identifier: Optional[str] = None, + text_format: type["BaseModel"] | dict | None = None, + tool_choice: ToolChoice | None = None, + tools: Iterable[ToolParam] | None = None, + top_p: float | None = None, + truncation: Literal["auto", "disabled"] | None = None, + user: str | None = None, + service_tier: str | None = None, + safety_identifier: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: +) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: """ Async: Handles responses API requests by reusing the synchronous function """ @@ -465,15 +461,15 @@ async def aresponses( # can apply them to local_vars without re-invoking the hook. ######################################################### litellm_logging_obj = kwargs.get("litellm_logging_obj", None) - prompt_id = cast(Optional[str], kwargs.get("prompt_id", None)) - prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None)) + prompt_id = cast(str | None, kwargs.get("prompt_id", None)) + prompt_variables = cast(dict | None, kwargs.get("prompt_variables", None)) original_model = model if isinstance( litellm_logging_obj, LiteLLMLoggingObj ) and litellm_logging_obj.should_run_prompt_management_hooks(prompt_id=prompt_id, non_default_params=kwargs): if isinstance(input, str): - client_input: List[AllMessageValues] = [{"role": "user", "content": input}] + client_input: list[AllMessageValues] = [{"role": "user", "content": input}] else: client_input = [ item # type: ignore[misc] @@ -494,7 +490,7 @@ async def aresponses( prompt_version=kwargs.get("prompt_version", None), ) input = cast( - Union[str, ResponseInputParam], + str | ResponseInputParam, ResponsesAPIRequestUtils.merge_prompt_management_input( original_input=input, client_input=client_input, @@ -573,25 +569,25 @@ async def aresponses( def _apply_prompt_management_to_responses_call( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - custom_llm_provider: Optional[str], - litellm_logging_obj: Optional[LiteLLMLoggingObj], - kwargs: Dict[str, Any], - local_vars: Dict[str, Any], -) -> tuple[Union[str, ResponseInputParam], str, Optional[str]]: + custom_llm_provider: str | None, + litellm_logging_obj: LiteLLMLoggingObj | None, + kwargs: dict[str, Any], + local_vars: dict[str, Any], +) -> tuple[str | ResponseInputParam, str, str | None]: async_merged = kwargs.pop("_async_prompt_merged_params", None) if async_merged is not None: for key, value in async_merged.items(): local_vars[key] = value return input, model, custom_llm_provider - prompt_id = cast(Optional[str], kwargs.get("prompt_id", None)) - prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None)) + prompt_id = cast(str | None, kwargs.get("prompt_id", None)) + prompt_variables = cast(dict | None, kwargs.get("prompt_variables", None)) original_model = model if isinstance(input, str): - client_input: List[AllMessageValues] = [{"role": "user", "content": input}] + client_input: list[AllMessageValues] = [{"role": "user", "content": input}] else: client_input = [ item # type: ignore[misc] @@ -616,7 +612,7 @@ def _apply_prompt_management_to_responses_call( prompt_version=kwargs.get("prompt_version", None), ) input = cast( - Union[str, ResponseInputParam], + str | ResponseInputParam, ResponsesAPIRequestUtils.merge_prompt_management_input( original_input=input, client_input=client_input, @@ -651,7 +647,7 @@ def _normalize_openai_chat_completions_responses_model(model: str) -> tuple[str, return f"openai/{remainder}", True -def _pop_use_chat_completions_api_kw(kwargs: Dict[str, Any]) -> bool: +def _pop_use_chat_completions_api_kw(kwargs: dict[str, Any]) -> bool: """Pop use_chat_completions_api; True when the chat-completions bridge is requested.""" use_cc = kwargs.pop("use_chat_completions_api", None) return bool(use_cc) @@ -659,10 +655,10 @@ def _pop_use_chat_completions_api_kw(kwargs: Dict[str, Any]) -> bool: def _resolve_model_provider_for_responses( model: str, - custom_llm_provider: Optional[str], + custom_llm_provider: str | None, litellm_params: GenericLiteLLMParams, - local_vars: Dict[str, Any], -) -> tuple[str, Optional[str]]: + local_vars: dict[str, Any], +) -> tuple[str, str | None]: if custom_llm_provider is not None and not litellm_params.custom_llm_provider: litellm_params.custom_llm_provider = custom_llm_provider ( @@ -683,16 +679,16 @@ def _resolve_model_provider_for_responses( def _apply_managed_file_id_mapping( - input: Union[str, ResponseInputParam], - tools: Optional[Iterable[ToolParam]], - kwargs: Dict[str, Any], - local_vars: Dict[str, Any], -) -> tuple[Union[str, ResponseInputParam], Optional[Iterable[ToolParam]]]: + input: str | ResponseInputParam, + tools: Iterable[ToolParam] | None, + kwargs: dict[str, Any], + local_vars: dict[str, Any], +) -> tuple[str | ResponseInputParam, Iterable[ToolParam] | None]: model_file_id_mapping = kwargs.get("model_file_id_mapping") model_info_id = kwargs.get("model_info", {}).get("id") if isinstance(kwargs.get("model_info"), dict) else None input = cast( - Union[str, ResponseInputParam], + str | ResponseInputParam, update_responses_input_with_model_file_ids( input=input, model_id=model_info_id, @@ -703,9 +699,9 @@ def _apply_managed_file_id_mapping( if tools: tools = cast( - Optional[Iterable[ToolParam]], + Iterable[ToolParam] | None, update_responses_tools_with_model_file_ids( - tools=cast(Optional[List[Dict[str, Any]]], tools), + tools=cast(list[dict[str, Any]] | None, tools), model_id=model_info_id, model_file_id_mapping=model_file_id_mapping, ), @@ -717,34 +713,34 @@ def _apply_managed_file_id_mapping( def _responses_try_dispatch_mcp_gateway( *, - tools: Optional[Iterable[ToolParam]], - input: Union[str, ResponseInputParam], + tools: Iterable[ToolParam] | None, + input: str | ResponseInputParam, model: str, - include: Optional[List[ResponseIncludable]], - instructions: Optional[str], - max_output_tokens: Optional[int], - prompt: Optional[PromptObject], - metadata: Optional[Dict[str, Any]], - parallel_tool_calls: Optional[bool], - previous_response_id: Optional[str], - reasoning: Optional[Reasoning], - store: Optional[bool], - background: Optional[bool], - stream: Optional[bool], - temperature: Optional[float], + include: list[ResponseIncludable] | None, + instructions: str | None, + max_output_tokens: int | None, + prompt: PromptObject | None, + metadata: dict[str, Any] | None, + parallel_tool_calls: bool | None, + previous_response_id: str | None, + reasoning: Reasoning | None, + store: bool | None, + background: bool | None, + stream: bool | None, + temperature: float | None, text: Any, - tool_choice: Optional[ToolChoice], - top_p: Optional[float], - truncation: Optional[Literal["auto", "disabled"]], - user: Optional[str], - extra_headers: Optional[Dict[str, Any]], - extra_query: Optional[Dict[str, Any]], - extra_body: Optional[Dict[str, Any]], - timeout: Optional[Union[float, httpx.Timeout]], - custom_llm_provider: Optional[str], - kwargs: Dict[str, Any], + tool_choice: ToolChoice | None, + top_p: float | None, + truncation: Literal["auto", "disabled"] | None, + user: str | None, + extra_headers: dict[str, Any] | None, + extra_query: dict[str, Any] | None, + extra_body: dict[str, Any] | None, + timeout: float | httpx.Timeout | None, + custom_llm_provider: str | None, + kwargs: dict[str, Any], _is_async: bool, -) -> Optional[Any]: +) -> Any | None: """Return a response when MCP gateway handles the call; otherwise None.""" from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, @@ -787,40 +783,40 @@ def _responses_try_dispatch_mcp_gateway( def _responses_try_dispatch_emulated_file_search( *, - tools: Optional[Iterable[ToolParam]], - input: Union[str, ResponseInputParam], + tools: Iterable[ToolParam] | None, + input: str | ResponseInputParam, model: str, - responses_api_provider_config: Optional[BaseResponsesAPIConfig], + responses_api_provider_config: BaseResponsesAPIConfig | None, use_chat_completions_api: bool, - include: Optional[List[ResponseIncludable]], - instructions: Optional[str], - max_output_tokens: Optional[int], - prompt: Optional[PromptObject], - metadata: Optional[Dict[str, Any]], - parallel_tool_calls: Optional[bool], - previous_response_id: Optional[str], - reasoning: Optional[Reasoning], - store: Optional[bool], - background: Optional[bool], - stream: Optional[bool], - temperature: Optional[float], + include: list[ResponseIncludable] | None, + instructions: str | None, + max_output_tokens: int | None, + prompt: PromptObject | None, + metadata: dict[str, Any] | None, + parallel_tool_calls: bool | None, + previous_response_id: str | None, + reasoning: Reasoning | None, + store: bool | None, + background: bool | None, + stream: bool | None, + temperature: float | None, text: Any, - tool_choice: Optional[ToolChoice], - top_p: Optional[float], - truncation: Optional[Literal["auto", "disabled"]], - user: Optional[str], - service_tier: Optional[str], - safety_identifier: Optional[str], - text_format: Optional[Union[Type[BaseModel], dict]], - allowed_openai_params: Optional[List[str]], - extra_headers: Optional[Dict[str, Any]], - extra_query: Optional[Dict[str, Any]], - extra_body: Optional[Dict[str, Any]], - timeout: Optional[Union[float, httpx.Timeout]], - custom_llm_provider: Optional[str], - kwargs: Dict[str, Any], + tool_choice: ToolChoice | None, + top_p: float | None, + truncation: Literal["auto", "disabled"] | None, + user: str | None, + service_tier: str | None, + safety_identifier: str | None, + text_format: type[BaseModel] | dict | None, + allowed_openai_params: list[str] | None, + extra_headers: dict[str, Any] | None, + extra_query: dict[str, Any] | None, + extra_body: dict[str, Any] | None, + timeout: float | httpx.Timeout | None, + custom_llm_provider: str | None, + kwargs: dict[str, Any], _is_async: bool, -) -> Optional[Any]: +) -> Any | None: """Return a response when emulated file_search handles the call; otherwise None.""" if not _has_file_search_tool(tools) or not ( responses_api_provider_config is None @@ -860,12 +856,8 @@ def _responses_try_dispatch_emulated_file_search( "extra_body": extra_body, "timeout": timeout, "custom_llm_provider": custom_llm_provider, - **( - { - **({"use_chat_completions_api": True} if use_chat_completions_api else {}), - **{k: v for k, v in kwargs.items() if k not in _internal_skip}, - } - ), + **({"use_chat_completions_api": True} if use_chat_completions_api else {}), + **{k: v for k, v in kwargs.items() if k not in _internal_skip}, } if _is_async: return aresponses_with_emulated_file_search(input=input, model=model, tools=tools, **emulated_kwargs) @@ -880,38 +872,38 @@ def _responses_try_dispatch_emulated_file_search( @client def responses( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - include: Optional[List[ResponseIncludable]] = None, - instructions: Optional[str] = None, - max_output_tokens: Optional[int] = None, - prompt: Optional[PromptObject] = None, - metadata: Optional[Dict[str, Any]] = None, - parallel_tool_calls: Optional[bool] = None, - previous_response_id: Optional[str] = None, - reasoning: Optional[Reasoning] = None, - store: Optional[bool] = None, - background: Optional[bool] = None, - stream: Optional[bool] = None, - temperature: Optional[float] = None, + include: list[ResponseIncludable] | None = None, + instructions: str | None = None, + max_output_tokens: int | None = None, + prompt: PromptObject | None = None, + metadata: dict[str, Any] | None = None, + parallel_tool_calls: bool | None = None, + previous_response_id: str | None = None, + reasoning: Reasoning | None = None, + store: bool | None = None, + background: bool | None = None, + stream: bool | None = None, + temperature: float | None = None, text: Optional["ResponseText"] = None, - text_format: Optional[Union[Type["BaseModel"], dict]] = None, - tool_choice: Optional[ToolChoice] = None, - tools: Optional[Iterable[ToolParam]] = None, - top_p: Optional[float] = None, - truncation: Optional[Literal["auto", "disabled"]] = None, - user: Optional[str] = None, - service_tier: Optional[str] = None, - safety_identifier: Optional[str] = None, + text_format: type["BaseModel"] | dict | None = None, + tool_choice: ToolChoice | None = None, + tools: Iterable[ToolParam] | None = None, + top_p: float | None = None, + truncation: Literal["auto", "disabled"] | None = None, + user: str | None = None, + service_tier: str | None = None, + safety_identifier: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - allowed_openai_params: Optional[List[str]] = None, - custom_llm_provider: Optional[str] = None, + allowed_openai_params: list[str] | None = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -922,7 +914,7 @@ def responses( try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aresponses", False) is True use_chat_completions_api = _pop_use_chat_completions_api_kw(kwargs) @@ -1009,7 +1001,7 @@ def responses( return _mcp_dispatch # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] + responses_api_provider_config: BaseResponsesAPIConfig | None if custom_llm_provider is None: responses_api_provider_config = None else: @@ -1084,7 +1076,7 @@ def responses( # Get optional parameters for the responses API request_drop_params = kwargs.get("drop_params") - responses_api_request_params: Dict = ResponsesAPIRequestUtils.get_optional_params_responses_api( + responses_api_request_params: dict = ResponsesAPIRequestUtils.get_optional_params_responses_api( model=model, responses_api_provider_config=responses_api_provider_config, response_api_optional_params=response_api_optional_params, @@ -1162,12 +1154,12 @@ async def adelete_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> DeleteResponseResult: """ @@ -1223,14 +1215,14 @@ def delete_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[DeleteResponseResult, Coroutine[Any, Any, DeleteResponseResult]]: +) -> DeleteResponseResult | Coroutine[Any, Any, DeleteResponseResult]: """ Synchronous version of the DELETE Responses API @@ -1240,7 +1232,7 @@ def delete_responses( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete_responses", False) is True # get llm provider logic @@ -1257,7 +1249,7 @@ def delete_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( + responses_api_provider_config: BaseResponsesAPIConfig | None = ( ProviderConfigManager.get_provider_responses_api_config( model=None, provider=custom_llm_provider, @@ -1313,12 +1305,12 @@ async def aget_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ResponsesAPIResponse: """ @@ -1388,14 +1380,14 @@ def get_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: +) -> ResponsesAPIResponse | Coroutine[Any, Any, ResponsesAPIResponse]: """ Fetch a response by its ID. @@ -1411,7 +1403,7 @@ def get_responses( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aget_responses", False) is True # get llm provider logic @@ -1428,7 +1420,7 @@ def get_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( + responses_api_provider_config: BaseResponsesAPIConfig | None = ( ProviderConfigManager.get_provider_responses_api_config( model=None, provider=custom_llm_provider, @@ -1490,16 +1482,16 @@ def get_responses( @client async def alist_input_items( response_id: str, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Dict: +) -> dict: """Async: List input items for a response""" local_vars = locals() try: @@ -1546,21 +1538,21 @@ async def alist_input_items( @client def list_input_items( response_id: str, - after: Optional[str] = None, - before: Optional[str] = None, - include: Optional[List[str]] = None, + after: str | None = None, + before: str | None = None, + include: list[str] | None = None, limit: int = 20, order: Literal["asc", "desc"] = "desc", - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Dict, Coroutine[Any, Any, Dict]]: +) -> dict | Coroutine[Any, Any, dict]: """List input items for a response""" local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("alist_input_items", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -1572,7 +1564,7 @@ def list_input_items( if custom_llm_provider is None: raise ValueError("custom_llm_provider is required but passed as None") - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( + responses_api_provider_config: BaseResponsesAPIConfig | None = ( ProviderConfigManager.get_provider_responses_api_config( model=None, provider=custom_llm_provider, @@ -1626,12 +1618,12 @@ async def acancel_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ResponsesAPIResponse: """ @@ -1687,14 +1679,14 @@ def cancel_responses( response_id: str, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: +) -> ResponsesAPIResponse | Coroutine[Any, Any, ResponsesAPIResponse]: """ Synchronous version of the POST Responses API @@ -1704,7 +1696,7 @@ def cancel_responses( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acancel_responses", False) is True # get llm provider logic @@ -1721,7 +1713,7 @@ def cancel_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( + responses_api_provider_config: BaseResponsesAPIConfig | None = ( ProviderConfigManager.get_provider_responses_api_config( model=None, provider=custom_llm_provider, @@ -1774,18 +1766,18 @@ def cancel_responses( @client async def acompact_responses( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - instructions: Optional[str] = None, - previous_response_id: Optional[str] = None, + instructions: str | None = None, + previous_response_id: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ResponsesAPIResponse: """ @@ -1852,20 +1844,20 @@ async def acompact_responses( @client def compact_responses( - input: Union[str, ResponseInputParam], + input: str | ResponseInputParam, model: str, - instructions: Optional[str] = None, - previous_response_id: Optional[str] = None, + instructions: str | None = None, + previous_response_id: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: +) -> ResponsesAPIResponse | Coroutine[Any, Any, ResponsesAPIResponse]: """ Synchronous version of the POST Compact Responses API @@ -1876,7 +1868,7 @@ def compact_responses( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acompact_responses", False) is True # get llm provider logic @@ -1893,7 +1885,7 @@ def compact_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( + responses_api_provider_config: BaseResponsesAPIConfig | None = ( ProviderConfigManager.get_provider_responses_api_config( model=model, provider=custom_llm_provider, @@ -1912,7 +1904,7 @@ def compact_responses( # Get optional parameters for the responses API request_drop_params = kwargs.get("drop_params") - responses_api_request_params: Dict = ResponsesAPIRequestUtils.get_optional_params_responses_api( + responses_api_request_params: dict = ResponsesAPIRequestUtils.get_optional_params_responses_api( model=model, responses_api_provider_config=responses_api_provider_config, response_api_optional_params=response_api_optional_params, @@ -1990,9 +1982,9 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict: async def _aresponses_websocket( model: str, websocket: Any, - api_base: Optional[str] = None, - api_key: Optional[str] = None, - timeout: Optional[float] = None, + api_base: str | None = None, + api_key: str | None = None, + timeout: float | None = None, **kwargs, ): """ @@ -2034,7 +2026,7 @@ async def _aresponses_websocket( custom_llm_provider=_custom_llm_provider, ) - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = None + responses_api_provider_config: BaseResponsesAPIConfig | None = None if _custom_llm_provider is not None: responses_api_provider_config = ProviderConfigManager.get_provider_responses_api_config( model=model, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 5c3e0cf0902..8c0328fe5f0 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -3,9 +3,6 @@ import logging from typing import ( Any, - List, - Optional, - Union, cast, ) @@ -18,10 +15,10 @@ from litellm.utils import CustomStreamWrapper def _add_mcp_metadata_to_response( - response: Union[ModelResponse, CustomStreamWrapper], - openai_tools: Optional[List], - tool_calls: Optional[List] = None, - tool_results: Optional[List] = None, + response: ModelResponse | CustomStreamWrapper, + openai_tools: list | None, + tool_calls: list | None = None, + tool_results: list | None = None, ) -> None: """ Add MCP metadata to response's provider_specific_fields. @@ -80,10 +77,10 @@ def _add_mcp_metadata_to_response( async def acompletion_with_mcp( model: str, - messages: List, - tools: Optional[List] = None, + messages: list, + tools: list | None = None, **kwargs: Any, -) -> Union[ModelResponse, CustomStreamWrapper]: +) -> ModelResponse | CustomStreamWrapper: """ Async completion with MCP integration. @@ -229,10 +226,10 @@ async def acompletion_with_mcp( self.openai_tools = openai_tools self.base_call_args = base_call_args self.request_tags = request_tags - self.collected_chunks: List[ModelResponseStream] = [] - self.tool_calls: Optional[List] = None - self.tool_results: Optional[List] = None - self.complete_response: Optional[ModelResponse] = None + self.collected_chunks: list[ModelResponseStream] = [] + self.tool_calls: list | None = None + self.tool_results: list | None = None + self.complete_response: ModelResponse | None = None self.stream_exhausted = False self.tool_execution_done = False self.follow_up_stream = None diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 92cce711513..50744a7b93f 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -5,12 +5,8 @@ from datetime import datetime from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Union, ) from litellm._logging import verbose_logger @@ -69,7 +65,7 @@ class LiteLLM_Proxy_MCP_Handler: """ @staticmethod - def _get_parent_request_tags(kwargs: Optional[dict[str, Any]]) -> list[str]: + def _get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]: """Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent).""" if not kwargs: return [] @@ -83,7 +79,7 @@ class LiteLLM_Proxy_MCP_Handler: ) @staticmethod - def _should_use_litellm_mcp_gateway(tools: Optional[Iterable[ToolParam]]) -> bool: + def _should_use_litellm_mcp_gateway(tools: Iterable[ToolParam] | None) -> bool: """ Returns True if any MCP tool should be handled via the litellm proxy MCP gateway. This includes tools with server_url="litellm_proxy" as well as URLs ending in /mcp/. @@ -101,15 +97,15 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _parse_mcp_tools( tools: Iterable[Mapping[str, object]] | None, - ) -> Tuple[List[ToolParam], List[Any]]: + ) -> tuple[list[ToolParam], list[Any]]: """ Parse tools and separate MCP tools with litellm_proxy from other tools. Returns: Tuple of (mcp_tools_with_litellm_proxy, other_tools) """ - mcp_tools_with_litellm_proxy: List[ToolParam] = [] - other_tools: List[Any] = [] + mcp_tools_with_litellm_proxy: list[ToolParam] = [] + other_tools: list[Any] = [] if tools: for tool in tools: @@ -138,8 +134,8 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod async def _apply_toolset_permissions( - resolved_toolset_ids: List[str], - resolved_mcp_servers: List[str], + resolved_toolset_ids: list[str], + resolved_mcp_servers: list[str], user_api_key_auth: "UserAPIKeyAuth", ) -> "UserAPIKeyAuth": """Apply resolved toolset permissions to user_api_key_auth and return updated auth.""" @@ -182,11 +178,11 @@ class LiteLLM_Proxy_MCP_Handler: async def _get_mcp_tools_from_manager( user_api_key_auth: "UserAPIKeyAuth | None", mcp_tools_with_litellm_proxy: Iterable[Mapping[str, object]] | None, - litellm_trace_id: Optional[str] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, - request_tags: Optional[list[str]] = None, - ) -> tuple[List[MCPTool], List[str]]: + litellm_trace_id: str | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + request_tags: list[str] | None = None, + ) -> tuple[list[MCPTool], list[str]]: """ Get available tools from the MCP server manager. @@ -208,7 +204,7 @@ class LiteLLM_Proxy_MCP_Handler: _get_tools_from_mcp_servers, ) - mcp_servers: List[str] = [] + mcp_servers: list[str] = [] if mcp_tools_with_litellm_proxy: for _tool in mcp_tools_with_litellm_proxy: # if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github @@ -219,8 +215,8 @@ class LiteLLM_Proxy_MCP_Handler: # Resolve toolset names: collect all toolset IDs first, then apply their # combined permissions in a single pass so multiple toolsets are unioned # rather than the last one overwriting the others. - resolved_mcp_servers: List[str] = [] - resolved_toolset_ids: List[str] = [] + resolved_mcp_servers: list[str] = [] + resolved_toolset_ids: list[str] = [] for name in mcp_servers: if not global_mcp_server_manager.get_mcp_server_by_name(name): try: @@ -288,7 +284,7 @@ class LiteLLM_Proxy_MCP_Handler: allowed_mcp_servers=allowed_mcp_servers, ) - server_names: List[str] = [] + server_names: list[str] = [] for server in allowed_mcp_servers: if server is None: continue @@ -302,8 +298,8 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _deduplicate_mcp_tools( - mcp_tools: List[MCPTool], allowed_mcp_servers: List[str] - ) -> tuple[List[MCPTool], dict[str, str]]: + mcp_tools: list[MCPTool], allowed_mcp_servers: list[str] + ) -> tuple[list[MCPTool], dict[str, str]]: """ Deduplicate MCP tools by name, keeping the first occurrence of each tool. @@ -336,8 +332,8 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _filter_mcp_tools_by_allowed_tools( - mcp_tools: List[MCPTool], mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] - ) -> List[MCPTool]: + mcp_tools: list[MCPTool], mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] + ) -> list[MCPTool]: """Filter MCP tools based on allowed_tools parameter from the original tool configs.""" # Collect all allowed tool names from all MCP tool configs allowed_tool_names = set() @@ -376,9 +372,9 @@ class LiteLLM_Proxy_MCP_Handler: async def _process_mcp_tools_to_openai_format( user_api_key_auth: "UserAPIKeyAuth | None", mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], - litellm_trace_id: Optional[str] = None, - request_tags: Optional[list[str]] = None, - ) -> tuple[List[Any], dict[str, str]]: + litellm_trace_id: str | None = None, + request_tags: list[str] | None = None, + ) -> tuple[list[Any], dict[str, str]]: """ Centralized method to process MCP tools through the complete pipeline. @@ -409,11 +405,11 @@ class LiteLLM_Proxy_MCP_Handler: async def _process_mcp_tools_without_openai_transform( user_api_key_auth: Any, mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], - litellm_trace_id: Optional[str] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, - request_tags: Optional[list[str]] = None, - ) -> tuple[List[MCPTool], dict[str, str]]: + litellm_trace_id: str | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + request_tags: list[str] | None = None, + ) -> tuple[list[MCPTool], dict[str, str]]: """ Process MCP tools through filtering and deduplication pipeline without OpenAI transformation. This is useful for cases where we need the original MCP tool objects (e.g., for events). @@ -461,14 +457,14 @@ class LiteLLM_Proxy_MCP_Handler: def _transform_mcp_tools_to_openai( mcp_tools: Sequence[MCPTool], target_format: Literal["responses", "chat"] = "responses", - ) -> List[Any]: + ) -> list[Any]: """Transform MCP tools to OpenAI-compatible format.""" from litellm.experimental_mcp_client.tools import ( transform_mcp_tool_to_openai_responses_api_tool, transform_mcp_tool_to_openai_tool, ) - openai_tools: List[Any] = [] + openai_tools: list[Any] = [] for mcp_tool in mcp_tools: if target_format == "chat": openai_tool = transform_mcp_tool_to_openai_tool(mcp_tool) @@ -505,9 +501,9 @@ class LiteLLM_Proxy_MCP_Handler: return True @staticmethod - def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> List[Any]: + def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[Any]: """Extract tool calls from the response output.""" - tool_calls: List[Any] = [] + tool_calls: list[Any] = [] for output_item in response.output: # Check if this is a function call output item if isinstance(output_item, dict) and output_item.get("type") == "function_call": @@ -543,7 +539,7 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _extract_tool_call_details( tool_call, - ) -> Tuple[Optional[str], Optional[str], Optional[str]]: + ) -> tuple[str | None, str | None, str | None]: """Extract tool name, arguments, and call_id from a tool call.""" if isinstance(tool_call, dict): tool_call_id = tool_call.get("call_id") or tool_call.get("id") @@ -575,7 +571,7 @@ class LiteLLM_Proxy_MCP_Handler: return tool_name, tool_arguments, tool_call_id @staticmethod - def _parse_tool_arguments(tool_arguments: Any) -> Dict[str, Any]: + def _parse_tool_arguments(tool_arguments: Any) -> dict[str, Any]: """Parse tool arguments, handling both string and dict formats.""" import json @@ -633,14 +629,14 @@ class LiteLLM_Proxy_MCP_Handler: tool_server_map: dict[str, str], tool_calls: Sequence[object], user_api_key_auth: Any, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, - oauth2_headers: Optional[Dict[str, str]] = None, - raw_headers: Optional[Dict[str, str]] = None, - litellm_call_id: Optional[str] = None, - litellm_trace_id: Optional[str] = None, - request_tags: Optional[list[str]] = None, - ) -> List[Dict[str, Any]]: + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + litellm_call_id: str | None = None, + litellm_trace_id: str | None = None, + request_tags: list[str] | None = None, + ) -> list[dict[str, Any]]: """Execute tool calls and return results.""" from fastapi import HTTPException @@ -655,11 +651,11 @@ class LiteLLM_Proxy_MCP_Handler: from litellm.proxy.proxy_server import proxy_logging_obj tool_results = [] - tool_call_id: Optional[str] = None + tool_call_id: str | None = None rules_obj = Rules() for tool_call in tool_calls: - logging_request_data: Dict[str, Any] = {} - tool_name: Optional[str] = None + logging_request_data: dict[str, Any] = {} + tool_name: str | None = None try: ( tool_name, @@ -738,7 +734,7 @@ class LiteLLM_Proxy_MCP_Handler: if user_identifier: logging_request_data["user"] = user_identifier - litellm_logging_obj: Optional[LiteLLMLoggingObj] = None + litellm_logging_obj: LiteLLMLoggingObj | None = None try: litellm_logging_obj, _ = function_setup( original_function="call_mcp_tool", @@ -848,8 +844,8 @@ class LiteLLM_Proxy_MCP_Handler: request_data=logging_request_data, error=e, ) - verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") - error_message = f"Tool call blocked: PII entity '{getattr(e, 'entity_type', 'unknown')}' detected by guardrail '{getattr(e, 'guardrail_name', 'unknown')}'. {str(e)}" + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {e!s}") + error_message = f"Tool call blocked: PII entity '{getattr(e, 'entity_type', 'unknown')}' detected by guardrail '{getattr(e, 'guardrail_name', 'unknown')}'. {e!s}" tool_results.append( { "tool_call_id": tool_call_id, @@ -864,9 +860,9 @@ class LiteLLM_Proxy_MCP_Handler: request_data=logging_request_data, error=e, ) - verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {e!s}") error_message = ( - f"Tool call blocked: Guardrail '{getattr(e, 'guardrail_name', 'unknown')}' violation. {str(e)}" + f"Tool call blocked: Guardrail '{getattr(e, 'guardrail_name', 'unknown')}' violation. {e!s}" ) tool_results.append( { @@ -882,7 +878,7 @@ class LiteLLM_Proxy_MCP_Handler: request_data=logging_request_data, error=e, ) - verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + verbose_logger.error(f"HTTPException in MCP tool call: {e!s}") error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}" tool_results.append( { @@ -902,7 +898,7 @@ class LiteLLM_Proxy_MCP_Handler: tool_results.append( { "tool_call_id": tool_call_id, - "result": f"Error executing tool: {str(e)}", + "result": f"Error executing tool: {e!s}", "name": tool_name, } ) @@ -911,21 +907,21 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _create_follow_up_messages_for_chat( - original_messages: List[Any], + original_messages: list[Any], response: ModelResponse, tool_results: Sequence[Mapping[str, object]], - ) -> List[Any]: + ) -> list[Any]: """Create follow-up chat messages that include tool execution results.""" from copy import deepcopy from litellm.utils import convert_list_message_to_dict - follow_up_messages: List[Any] = convert_list_message_to_dict(deepcopy(original_messages)) + follow_up_messages: list[Any] = convert_list_message_to_dict(deepcopy(original_messages)) if not follow_up_messages: follow_up_messages = [] - message_to_append: Optional[dict] = None + message_to_append: dict | None = None try: first_choice = response.choices[0] if isinstance(first_choice, Choices) and getattr(first_choice, "message", None): @@ -959,9 +955,9 @@ class LiteLLM_Proxy_MCP_Handler: response: ResponsesAPIResponse, tool_results: Sequence[Mapping[str, object]], original_input: str | ResponseInputParam | None = None, - ) -> List[Any]: + ) -> list[Any]: """Create follow-up input with tool results in proper format.""" - follow_up_input: List[Any] = [] + follow_up_input: list[Any] = [] # Add original user input if available to maintain conversation context if original_input: @@ -973,8 +969,8 @@ class LiteLLM_Proxy_MCP_Handler: follow_up_input.append(original_input) # Add the assistant message with function calls - assistant_message_content: List[Any] = [] - function_calls: List[Dict[str, Any]] = [] + assistant_message_content: list[Any] = [] + function_calls: list[dict[str, Any]] = [] for output_item in response.output: if not isinstance(output_item, dict) and hasattr(output_item, "model_dump"): @@ -1034,12 +1030,12 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod async def _make_follow_up_call( - follow_up_input: List[Any], + follow_up_input: list[Any], model: str, - all_tools: Optional[List[Any]], + all_tools: list[Any] | None, response_id: str, **call_params: Any, - ) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]: + ) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator: """Make follow-up response API call with tool results.""" return await aresponses( input=follow_up_input, @@ -1081,8 +1077,8 @@ class LiteLLM_Proxy_MCP_Handler: all_tools: Sequence[object] | None, mcp_tools_with_litellm_proxy: list[Mapping[str, object]], mcp_discovery_events: list[ResponsesAPIStreamingResponse], - call_params: Dict[str, Any], - previous_response_id: Optional[str], + call_params: dict[str, Any], + previous_response_id: str | None, tool_server_map: dict[str, str], **kwargs, ) -> Any: @@ -1124,10 +1120,10 @@ class LiteLLM_Proxy_MCP_Handler: input: str | ResponseInputParam, model: str, all_tools: Sequence[object] | None, - call_params: Dict[str, Any], - previous_response_id: Optional[str], + call_params: dict[str, Any], + previous_response_id: str | None, **kwargs, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Build a clean request parameters dictionary for MCP streaming. @@ -1155,7 +1151,7 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod def _create_tool_execution_events( - tool_calls: Sequence[object], tool_results: List[Dict[str, Any]] + tool_calls: Sequence[object], tool_results: list[dict[str, Any]] ) -> list[ResponsesAPIStreamingResponse]: """ Create MCP tool execution events for streaming. @@ -1204,7 +1200,7 @@ class LiteLLM_Proxy_MCP_Handler: return tool_execution_events @staticmethod - def _prepare_initial_call_params(call_params: Dict[str, Any], should_auto_execute: bool) -> Dict[str, Any]: + def _prepare_initial_call_params(call_params: dict[str, Any], should_auto_execute: bool) -> dict[str, Any]: """ Prepare call parameters for the initial LLM call. @@ -1220,7 +1216,7 @@ class LiteLLM_Proxy_MCP_Handler: return initial_params @staticmethod - def _prepare_follow_up_call_params(call_params: Dict[str, Any], original_stream_setting: bool) -> Dict[str, Any]: + def _prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]: """ Prepare call parameters for the follow-up LLM call after tool execution. diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index d6b45855be5..c68628429da 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, cast from litellm._logging import verbose_logger from litellm._uuid import uuid @@ -30,14 +30,14 @@ MAX_MCP_TOOL_CALL_ROUNDS = 5 async def create_mcp_list_tools_events( - mcp_tools_with_litellm_proxy: List[ToolParam], + mcp_tools_with_litellm_proxy: list[ToolParam], user_api_key_auth: Any, base_item_id: str, - pre_processed_mcp_tools: List[Any], -) -> List[ResponsesAPIStreamingResponse]: + pre_processed_mcp_tools: list[Any], +) -> list[ResponsesAPIStreamingResponse]: """Create MCP discovery events using pre-processed tools from the parent""" - events: List[ResponsesAPIStreamingResponse] = [] + events: list[ResponsesAPIStreamingResponse] = [] try: # Extract MCP server names @@ -165,12 +165,12 @@ def create_mcp_call_events( tool_name: str, tool_call_id: str, arguments: str, - result: Optional[str] = None, - base_item_id: Optional[str] = None, + result: str | None = None, + base_item_id: str | None = None, sequence_start: int = 1, -) -> List[ResponsesAPIStreamingResponse]: +) -> list[ResponsesAPIStreamingResponse]: """Create MCP call events following OpenAI's specification""" - events: List[ResponsesAPIStreamingResponse] = [] + events: list[ResponsesAPIStreamingResponse] = [] item_id = base_item_id or f"mcp_{uuid.uuid4().hex[:8]}" # MCP call in progress event @@ -256,11 +256,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): def __init__( self, base_iterator: Any, # Can be None - will be created internally - mcp_events: List[ResponsesAPIStreamingResponse], + mcp_events: list[ResponsesAPIStreamingResponse], tool_server_map: dict[str, str], - mcp_tools_with_litellm_proxy: Optional[List[Any]] = None, + mcp_tools_with_litellm_proxy: list[Any] | None = None, user_api_key_auth: Any = None, - original_request_params: Optional[Dict[str, Any]] = None, + original_request_params: dict[str, Any] | None = None, ): # MCP setup self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or [] @@ -273,19 +273,19 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.finished = False # Event queues and generation flags - self.mcp_discovery_events: List[ResponsesAPIStreamingResponse] = ( + self.mcp_discovery_events: list[ResponsesAPIStreamingResponse] = ( mcp_events # Pre-generated MCP discovery events ) - self.tool_execution_events: List[ResponsesAPIStreamingResponse] = [] + self.tool_execution_events: list[ResponsesAPIStreamingResponse] = [] self.mcp_discovery_generated = True # Events are already generated self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility self.tool_server_map = tool_server_map # Iterator references - self.base_iterator: Optional[Union[Any, ResponsesAPIResponse]] = base_iterator # Will be created when needed + self.base_iterator: Any | ResponsesAPIResponse | None = base_iterator # Will be created when needed # Response collection for tool execution - self.collected_response: Optional[ResponsesAPIResponse] = None + self.collected_response: ResponsesAPIResponse | None = None # Counts completed rounds of tool execution, so a model that keeps # calling tools (e.g. retrying after an error) can't loop forever. @@ -298,7 +298,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # response had no tool calls (e.g. the model finally answered in # text) would incorrectly reuse tool_results left over from an # earlier round and keep looping instead of finishing. - self._tool_results_for_response: Optional[ResponsesAPIResponse] = None + self._tool_results_for_response: ResponsesAPIResponse | None = None # Set up model metadata (will be updated when we get the real iterator) self.model = self.original_request_params.get("model", "unknown") @@ -316,7 +316,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.initial_events_emitted = False # Cache the response ID to ensure consistency across all events - self._cached_response_id: Optional[str] = None + self._cached_response_id: str | None = None self._initial_creation_error: Exception | None = None self._stream_error: Exception | None = None @@ -325,23 +325,24 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): def _extract_mcp_headers_from_params(self) -> None: """Extract MCP headers from original request params to pass to tool calls""" - from typing import Dict, Optional + from starlette.datastructures import Headers + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) # Extract headers from secret_fields in original_request_params - raw_headers_from_request: Optional[Dict[str, str]] = None + raw_headers_from_request: dict[str, str] | None = None secret_fields = self.original_request_params.get("secret_fields") if secret_fields and isinstance(secret_fields, dict): raw_headers_from_request = secret_fields.get("raw_headers") # Extract MCP-specific headers - self.mcp_auth_header: Optional[str] = None - self.mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None - self.oauth2_headers: Optional[Dict[str, str]] = None - self.raw_headers: Optional[Dict[str, str]] = raw_headers_from_request + self.mcp_auth_header: str | None = None + self.mcp_server_auth_headers: dict[str, dict[str, str]] | None = None + self.oauth2_headers: dict[str, str] | None = None + self.raw_headers: dict[str, str] | None = raw_headers_from_request if raw_headers_from_request: headers_obj = Headers(raw_headers_from_request) @@ -483,7 +484,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): async def _handle_initial_response_phase( self, - ) -> Optional[ResponsesAPIStreamingResponse]: + ) -> ResponsesAPIStreamingResponse | None: """ Handle Phase 1: Initial Response Stream. diff --git a/litellm/responses/mcp/request_context.py b/litellm/responses/mcp/request_context.py index ce3e17dbd44..3b8e79ee5f2 100644 --- a/litellm/responses/mcp/request_context.py +++ b/litellm/responses/mcp/request_context.py @@ -10,7 +10,7 @@ still executes the tool, just with no credentials. from collections.abc import Iterable, Mapping, Sequence from dataclasses import dataclass -from typing import Any, Union +from typing import Any @dataclass(frozen=True, slots=True) @@ -18,19 +18,19 @@ class MCPRequestContext: """Everything a gateway handler must forward to MCP tool listing and execution.""" user_api_key_auth: Any # any-ok: UserAPIKeyAuth is proxy-only; importing it here would create a cycle - mcp_auth_header: Union[str, None] = None - mcp_server_auth_headers: Union[Mapping[str, Mapping[str, str]], None] = None - oauth2_headers: Union[Mapping[str, str], None] = None - raw_headers: Union[Mapping[str, str], None] = None - request_tags: Union[Sequence[str], None] = None - litellm_trace_id: Union[str, None] = None - litellm_call_id: Union[str, None] = None + mcp_auth_header: str | None = None + mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = None + oauth2_headers: Mapping[str, str] | None = None + raw_headers: Mapping[str, str] | None = None + request_tags: Sequence[str] | None = None + litellm_trace_id: str | None = None + litellm_call_id: str | None = None @classmethod def resolve( cls, kwargs: Mapping[str, Any], - tools: Union[Iterable[Any], None], + tools: Iterable[Any] | None, ) -> "MCPRequestContext": """ Build the context from a gateway handler's kwargs. diff --git a/litellm/responses/sse_output_recovery.py b/litellm/responses/sse_output_recovery.py index 1546b9414ea..d5985b24fe2 100644 --- a/litellm/responses/sse_output_recovery.py +++ b/litellm/responses/sse_output_recovery.py @@ -8,14 +8,14 @@ caller automatically applies to all of them. """ import json -from typing import Any, Dict, Optional +from typing import Any from litellm.constants import STREAM_SSE_DONE_STRING _MAX_CONTENT_INDEX = 1024 -def parse_sse_json_chunk(chunk: str) -> Optional[Dict[str, Any]]: +def parse_sse_json_chunk(chunk: str) -> dict[str, Any] | None: """Parse a single raw SSE line into a JSON object dict. Returns ``None`` for empty lines, ``event:`` lines, ``[DONE]`` markers, @@ -39,8 +39,8 @@ def parse_sse_json_chunk(chunk: str) -> Optional[Dict[str, Any]]: def record_output_item_chunk( - parsed_chunk: Dict[str, Any], - output_items: Dict[int, Dict[str, Any]], + parsed_chunk: dict[str, Any], + output_items: dict[int, dict[str, Any]], ) -> None: """Record an OUTPUT_ITEM_DONE chunk into ``output_items`` keyed by ``output_index`` (falling back to the next free slot when missing). @@ -59,9 +59,9 @@ def record_output_item_chunk( def record_output_text_chunk( - parsed_chunk: Dict[str, Any], - output_items: Dict[int, Dict[str, Any]], - text_only_items: Dict[int, Dict[str, Any]], + parsed_chunk: dict[str, Any], + output_items: dict[int, dict[str, Any]], + text_only_items: dict[int, dict[str, Any]], ) -> None: """Record an OUTPUT_TEXT_DONE chunk as a synthetic message item in ``text_only_items``. Real OUTPUT_ITEM_DONE events already captured in diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index d4356c80f19..3bcc19822a6 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -9,7 +9,7 @@ from collections.abc import Mapping from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import Any, Dict, List, Literal, Optional +from typing import Any, Literal import httpx from openai._streaming import SSEDecoder @@ -42,7 +42,7 @@ def _get_openai_response_types(): return openai_types -def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -> None: +def _log_background_task_failure(task: asyncio.Task[Any], *, task_name: str) -> None: if task.cancelled(): return exception = task.exception() @@ -79,7 +79,7 @@ _ERROR_CODE_HTTP_STATUS: Mapping[str, int] = MappingProxyType( ) -def _error_event_fields(error_obj: object) -> tuple[str, Optional[str], Optional[str]]: +def _error_event_fields(error_obj: object) -> tuple[str, str | None, str | None]: if isinstance(error_obj, dict): raw_message = error_obj.get("message") raw_type = error_obj.get("type") @@ -98,7 +98,7 @@ def _error_event_fields(error_obj: object) -> tuple[str, Optional[str], Optional return message, error_type, code -def _status_code_for_error_fields(error_type: Optional[str], error_code: Optional[str]) -> int: +def _status_code_for_error_fields(error_type: str | None, error_code: str | None) -> int: fields = tuple(field for field in (error_code, error_type) if field is not None) if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields): return 429 @@ -119,34 +119,34 @@ class BaseResponsesAPIStreamingIterator: self, response: httpx.Response, model: str, - responses_api_provider_config: Optional[BaseResponsesAPIConfig], + responses_api_provider_config: BaseResponsesAPIConfig | None, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, + request_data: dict[str, Any] | None = None, + call_type: str | None = None, ): self.response = response self.model = model self.logging_obj = logging_obj self.finished = False self.responses_api_provider_config = responses_api_provider_config - self.completed_response: Optional[Any] = None + self.completed_response: Any | None = None self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called self._yielded_first_chunk = False self._generated_content = "" self._completed_response_cached = False self._completed_response_logged = False - self._completed_response_cache_hit: Optional[bool] = None + self._completed_response_cache_hit: bool | None = None self._persist_completed_response_before_logging = True self._stream_created_time: float = time.time() # track request context for hooks self.litellm_metadata = litellm_metadata self.custom_llm_provider = custom_llm_provider - self.request_data: Dict[str, Any] = request_data or {} - self.call_type: Optional[str] = call_type + self.request_data: dict[str, Any] = request_data or {} + self.call_type: str | None = call_type # set hidden params for response headers (e.g., x-litellm-model-id) # This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py @@ -154,7 +154,7 @@ class BaseResponsesAPIStreamingIterator: model=model or "", optional_params=self.logging_obj.model_call_details.get("litellm_params", {}), ) - _model_info: Dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {} + _model_info: dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {} self._hidden_params = { "model_id": _model_info.get("id", None), "api_base": _api_base, @@ -176,7 +176,7 @@ class BaseResponsesAPIStreamingIterator: llm_provider=self.custom_llm_provider or "", ) - def _process_chunk(self, chunk) -> Optional[Any]: + def _process_chunk(self, chunk) -> Any | None: """Process a single chunk of data from the stream""" if not chunk: return None @@ -299,14 +299,12 @@ class BaseResponsesAPIStreamingIterator: self.completed_response = openai_responses_api_chunk # Add cost to usage object if include_cost_in_streaming_usage is True if litellm.include_cost_in_streaming_usage and self.logging_obj is not None: - response_obj: Optional[Any] = getattr(openai_responses_api_chunk, "response", None) + response_obj: Any | None = getattr(openai_responses_api_chunk, "response", None) if response_obj: - usage_obj: Optional[Any] = getattr(response_obj, "usage", None) + usage_obj: Any | None = getattr(response_obj, "usage", None) if usage_obj is not None: try: - cost: Optional[float] = self.logging_obj._response_cost_calculator( - result=response_obj - ) + cost: float | None = self.logging_obj._response_cost_calculator(result=response_obj) if cost is not None: setattr(usage_obj, "cost", cost) except Exception: @@ -381,7 +379,6 @@ class BaseResponsesAPIStreamingIterator: def _handle_logging_completed_response(self): """Base implementation - should be overridden by subclasses""" - pass def _handle_logging_failed_response(self): """ @@ -404,7 +401,7 @@ class BaseResponsesAPIStreamingIterator: ) self._handle_failure(exception) - def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None: + def _record_failed_response_usage(self, response_obj: Any | None) -> None: if response_obj is None or self.logging_obj is None: return usage_obj = getattr(response_obj, "usage", None) @@ -454,7 +451,7 @@ class BaseResponsesAPIStreamingIterator: is_pre_first_chunk=not self._yielded_first_chunk, ) - def _get_completed_response_object(self) -> Optional[Any]: + def _get_completed_response_object(self) -> Any | None: openai_types = _get_openai_response_types() completed_response = self.completed_response if isinstance(completed_response, openai_types.ResponsesAPIResponse): @@ -536,7 +533,7 @@ class BaseResponsesAPIStreamingIterator: """ try: # Align with chat pipeline: use logging_obj model_call_details + call_type - typed_call_type: Optional[CallTypes] = None + typed_call_type: CallTypes | None = None if self.call_type is not None: try: typed_call_type = CallTypes(self.call_type) @@ -580,7 +577,7 @@ class BaseResponsesAPIStreamingIterator: if self.completed_response is None: return - request_payload: Dict[str, Any] = {} + request_payload: dict[str, Any] = {} if isinstance(self.request_data, dict): request_payload.update(self.request_data) try: @@ -610,7 +607,7 @@ class BaseResponsesAPIStreamingIterator: pass try: - typed_call_type: Optional[CallTypes] = None + typed_call_type: CallTypes | None = None if self.call_type is not None: try: typed_call_type = CallTypes(self.call_type) @@ -690,10 +687,10 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): model: str, responses_api_provider_config: BaseResponsesAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, + request_data: dict[str, Any] | None = None, + call_type: str | None = None, ): super().__init__( response, @@ -772,10 +769,10 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): model: str, responses_api_provider_config: BaseResponsesAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, + request_data: dict[str, Any] | None = None, + call_type: str | None = None, ): super().__init__( response, @@ -859,10 +856,10 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): model: str, responses_api_provider_config: BaseResponsesAPIConfig, logging_obj: LiteLLMLoggingObj, - litellm_metadata: Optional[Dict[str, Any]] = None, - custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, + litellm_metadata: dict[str, Any] | None = None, + custom_llm_provider: str | None = None, + request_data: dict[str, Any] | None = None, + call_type: str | None = None, ): transformed = responses_api_provider_config.transform_response_api_response( model=model, @@ -928,8 +925,8 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): self, response: Any, logging_obj: LiteLLMLoggingObj, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, + request_data: dict[str, Any] | None = None, + call_type: str | None = None, ): BaseResponsesAPIStreamingIterator.__init__( self, @@ -944,7 +941,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): ) self._completed_response_cache_hit = True self._persist_completed_response_before_logging = False - self._events: List[Any] = [] + self._events: list[Any] = [] self._idx = 0 self._set_events_from_response(transformed=response, logging_obj=logging_obj) @@ -990,7 +987,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): return evt -def _dump_response_object(obj: Any) -> Dict[str, Any]: +def _dump_response_object(obj: Any) -> dict[str, Any]: if hasattr(obj, "model_dump"): return obj.model_dump() if isinstance(obj, dict): @@ -1020,8 +1017,8 @@ def _build_content_part_done_event( item_id: str, output_index: int, content_index: int, - part_payload: Dict[str, Any], -) -> Optional[Any]: + part_payload: dict[str, Any], +) -> Any | None: openai_types = _get_openai_response_types() part_type = part_payload.get("type") part: Any @@ -1060,11 +1057,11 @@ def _build_content_part_done_event( def _add_text_like_part_events( *, - events: List[Any], + events: list[Any], item_id: str, output_index: int, content_index: int, - part_payload: Dict[str, Any], + part_payload: dict[str, Any], chunk_size: int, ) -> None: openai_types = _get_openai_response_types() @@ -1129,19 +1126,19 @@ def _build_synthetic_response_events( transformed: Any, logging_obj: LiteLLMLoggingObj, chunk_size: int, -) -> List[Any]: +) -> list[Any]: openai_types = _get_openai_response_types() if litellm.include_cost_in_streaming_usage and logging_obj is not None: - usage_obj: Optional[Any] = getattr(transformed, "usage", None) + usage_obj: Any | None = getattr(transformed, "usage", None) if usage_obj is not None: try: - cost: Optional[float] = logging_obj._response_cost_calculator(result=transformed) + cost: float | None = logging_obj._response_cost_calculator(result=transformed) if cost is not None: setattr(usage_obj, "cost", cost) except Exception: pass - events: List[Any] = [ + events: list[Any] = [ _build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed), _build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed), ] @@ -1298,26 +1295,26 @@ class ResponsesWebSocketStreaming: websocket: Any, backend_ws: Any, logging_obj: LiteLLMLoggingObj, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[Dict] = None, - first_message: Optional[str] = None, - guardrail_callbacks: Optional[List[Any]] = None, - output_guardrail_callbacks: Optional[List[Any]] = None, - authorized_model: Optional[str] = None, + user_api_key_dict: Any | None = None, + request_data: dict | None = None, + first_message: str | None = None, + guardrail_callbacks: list[Any] | None = None, + output_guardrail_callbacks: list[Any] | None = None, + authorized_model: str | None = None, ): self.websocket = websocket self.backend_ws = backend_ws self.logging_obj = logging_obj self.user_api_key_dict = user_api_key_dict - self.request_data: Dict = request_data or {} - self.messages: list[Dict] = [] - self.input_messages: list[Dict[str, str]] = [] + self.request_data: dict = request_data or {} + self.messages: list[dict] = [] + self.input_messages: list[dict[str, str]] = [] self.first_message = first_message - self.guardrail_callbacks: List[Any] = guardrail_callbacks or [] - self.output_guardrail_callbacks: List[Any] = output_guardrail_callbacks or [] + self.guardrail_callbacks: list[Any] = guardrail_callbacks or [] + self.output_guardrail_callbacks: list[Any] = output_guardrail_callbacks or [] # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. - self.authorized_model: Optional[str] = authorized_model + self.authorized_model: str | None = authorized_model def _should_store_event(self, event_obj: dict) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -1593,7 +1590,7 @@ class ResponsesWebSocketStreaming: if not self.guardrail_callbacks: return response_str - pii_tokens: Dict[str, str] = (self.request_data.get("metadata") or {}).get("pii_tokens", {}) + pii_tokens: dict[str, str] = (self.request_data.get("metadata") or {}).get("pii_tokens", {}) if not pii_tokens: return response_str @@ -1798,22 +1795,22 @@ class ManagedResponsesWebSocketHandler: self, websocket: Any, model: str, - logging_obj: "LiteLLMLoggingObj", - user_api_key_dict: Optional[Any] = None, - litellm_metadata: Optional[Dict[str, Any]] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[float] = None, - custom_llm_provider: Optional[str] = None, - first_message: Optional[str] = None, + logging_obj: LiteLLMLoggingObj, + user_api_key_dict: Any | None = None, + litellm_metadata: dict[str, Any] | None = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | None = None, + custom_llm_provider: str | None = None, + first_message: str | None = None, **kwargs: Any, ) -> None: self.websocket = websocket self.model = model self.logging_obj = logging_obj self.user_api_key_dict = user_api_key_dict - self.litellm_metadata: Dict[str, Any] = litellm_metadata or {} - self.model_group: Optional[str] = self.litellm_metadata.get("model_group") or self.litellm_metadata.get( + self.litellm_metadata: dict[str, Any] = litellm_metadata or {} + self.model_group: str | None = self.litellm_metadata.get("model_group") or self.litellm_metadata.get( "deployment_model_name" ) self.api_key = api_key @@ -1823,19 +1820,19 @@ class ManagedResponsesWebSocketHandler: self._connection_provider = self._resolve_provider(model) or custom_llm_provider self.first_message = first_message # Carry through safe pass-through kwargs (e.g. extra_headers) - self.extra_kwargs: Dict[str, Any] = {k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS} + self.extra_kwargs: dict[str, Any] = {k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS} # In-memory session history: response_id → full accumulated message list. # Keyed by the DECODED (pre-encoding) response ID from response.completed. # This avoids the async DB-write race condition where spend logs haven't # been committed yet when the next response.create arrives. - self._session_history: Dict[str, List[Dict[str, Any]]] = {} + self._session_history: dict[str, list[dict[str, Any]]] = {} # ------------------------------------------------------------------ # Internal helpers # ------------------------------------------------------------------ @staticmethod - def _serialize_chunk(chunk: Any) -> Optional[str]: + def _serialize_chunk(chunk: Any) -> str | None: """Serialize a streaming chunk to a JSON string for WebSocket transmission.""" try: if hasattr(chunk, "model_dump_json"): @@ -1857,7 +1854,7 @@ class ManagedResponsesWebSocketHandler: except Exception: pass - def _get_history_messages(self, previous_response_id: str) -> List[Dict[str, Any]]: + def _get_history_messages(self, previous_response_id: str) -> list[dict[str, Any]]: """ Return accumulated message history for *previous_response_id*. @@ -1868,7 +1865,7 @@ class ManagedResponsesWebSocketHandler: raw_id = decoded.get("response_id", previous_response_id) return list(self._session_history.get(raw_id, [])) - def _store_history(self, response_id: str, messages: List[Dict[str, Any]]) -> None: + def _store_history(self, response_id: str, messages: list[dict[str, Any]]) -> None: """ Store the complete accumulated message history for *response_id*. @@ -1878,13 +1875,13 @@ class ManagedResponsesWebSocketHandler: self._session_history[response_id] = messages @staticmethod - def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]: + def _extract_response_id(completed_event: dict[str, Any]) -> str | None: """ Pull the raw (decoded) response ID out of a ``response.completed`` event. Returns *None* if the event doesn't contain a usable ID. """ resp_obj = completed_event.get("response", {}) - encoded_id: Optional[str] = resp_obj.get("id") if isinstance(resp_obj, dict) else None + encoded_id: str | None = resp_obj.get("id") if isinstance(resp_obj, dict) else None if not encoded_id: return None decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(encoded_id) @@ -1892,8 +1889,8 @@ class ManagedResponsesWebSocketHandler: @staticmethod def _extract_output_messages( - completed_event: Dict[str, Any], - ) -> List[Dict[str, Any]]: + completed_event: dict[str, Any], + ) -> list[dict[str, Any]]: """ Convert the output items in a ``response.completed`` event into Responses API message dicts suitable for the next turn's ``input``. @@ -1901,7 +1898,7 @@ class ManagedResponsesWebSocketHandler: resp_obj = completed_event.get("response", {}) if not isinstance(resp_obj, dict): return [] - messages: List[Dict[str, Any]] = [] + messages: list[dict[str, Any]] = [] for item in resp_obj.get("output", []) or []: if not isinstance(item, dict): continue @@ -1928,7 +1925,7 @@ class ManagedResponsesWebSocketHandler: return messages @staticmethod - def _input_to_messages(input_val: Any) -> List[Dict[str, Any]]: + def _input_to_messages(input_val: Any) -> list[dict[str, Any]]: """ Normalise the ``input`` field of a ``response.create`` event to a list of Responses API message dicts. @@ -1949,7 +1946,7 @@ class ManagedResponsesWebSocketHandler: # _process_response_create sub-methods # ------------------------------------------------------------------ - async def _parse_message(self, raw_message: str) -> Optional[Dict[str, Any]]: + async def _parse_message(self, raw_message: str) -> dict[str, Any] | None: """Parse raw WS text; return the message dict or None (JSON error / ignored type).""" try: msg_obj = json.loads(raw_message) @@ -1962,14 +1959,14 @@ class ManagedResponsesWebSocketHandler: return msg_obj @staticmethod - def _is_warmup_frame(msg_obj: Dict[str, Any]) -> bool: + def _is_warmup_frame(msg_obj: dict[str, Any]) -> bool: """Return True for a response.create whose generate flag is false.""" nested = msg_obj.get("response") source = nested if isinstance(nested, dict) and nested else msg_obj return source.get("generate") is False @staticmethod - def _is_warmup_response_id(response_id: Optional[str]) -> bool: + def _is_warmup_response_id(response_id: str | None) -> bool: """Return True for synthetic warmup IDs that only exist on this connection.""" if not response_id: return False @@ -1978,13 +1975,13 @@ class ManagedResponsesWebSocketHandler: return str(raw_id).startswith(_WARMUP_RESPONSE_ID_PREFIX) @staticmethod - def _warmup_source_params(msg_obj: Dict[str, Any]) -> Dict[str, Any]: + def _warmup_source_params(msg_obj: dict[str, Any]) -> dict[str, Any]: nested = msg_obj.get("response") if isinstance(nested, dict) and nested: return nested return {k: v for k, v in msg_obj.items() if k != "type"} - def _build_warmup_response(self, msg_obj: Dict[str, Any]) -> Dict[str, Any]: + def _build_warmup_response(self, msg_obj: dict[str, Any]) -> dict[str, Any]: """Build a minimal completed Responses API object for a warmup ack.""" source = self._warmup_source_params(msg_obj) wire_model = source.get("model") or self.model_group or self.model @@ -2002,7 +1999,7 @@ class ManagedResponsesWebSocketHandler: }, } - async def _send_warmup_ack(self, msg_obj: Dict[str, Any]) -> None: + async def _send_warmup_ack(self, msg_obj: dict[str, Any]) -> None: """ Acknowledge a generate=false prewarm without calling the provider. @@ -2025,14 +2022,14 @@ class ManagedResponsesWebSocketHandler: await self.websocket.send_text(serialized) @staticmethod - def _build_base_call_kwargs(msg_obj: Dict[str, Any]) -> Dict[str, Any]: + def _build_base_call_kwargs(msg_obj: dict[str, Any]) -> dict[str, Any]: """ Extract Responses API params from the event, handling both wire formats: Nested: {"type": "response.create", "response": {"input": [...], ...}} Flat: {"type": "response.create", "input": [...], "model": "...", ...} """ nested = msg_obj.get("response") - response_params: Dict[str, Any] = ( + response_params: dict[str, Any] = ( nested if isinstance(nested, dict) and nested else {k: v for k, v in msg_obj.items() if k != "type"} ) return { @@ -2043,10 +2040,10 @@ class ManagedResponsesWebSocketHandler: def _apply_history( self, - call_kwargs: Dict[str, Any], - previous_response_id: Optional[str], - current_messages: List[Dict[str, Any]], - prior_history: List[Dict[str, Any]], + call_kwargs: dict[str, Any], + previous_response_id: str | None, + current_messages: list[dict[str, Any]], + prior_history: list[dict[str, Any]], ) -> None: """Prepend in-memory turn history, or fall back to DB-based reconstruction.""" if not previous_response_id: @@ -2075,7 +2072,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs["previous_response_id"] = previous_response_id @staticmethod - def _resolve_provider(model: Optional[str]) -> Optional[str]: + def _resolve_provider(model: str | None) -> str | None: """Resolve the LLM provider for a model string, or None if unresolvable.""" if not model: return None @@ -2087,7 +2084,7 @@ class ManagedResponsesWebSocketHandler: except Exception: return None - def _same_provider(self, model: Optional[str]) -> bool: + def _same_provider(self, model: str | None) -> bool: """Return True if model uses the same LLM provider as the connection model.""" if model is None or model == self.model: return True @@ -2096,7 +2093,7 @@ class ManagedResponsesWebSocketHandler: return False return event_provider == self._connection_provider - def _inject_credentials(self, call_kwargs: Dict[str, Any], model: Optional[str] = None) -> None: + def _inject_credentials(self, call_kwargs: dict[str, Any], model: str | None = None) -> None: """Inject connection-level credentials and metadata into call_kwargs.""" if self.api_key is not None: call_kwargs["api_key"] = self.api_key @@ -2115,7 +2112,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs["litellm_metadata"] = dict(self.litellm_metadata) @staticmethod - def _update_proxy_request(call_kwargs: Dict[str, Any], model: str) -> None: + def _update_proxy_request(call_kwargs: dict[str, Any], model: str) -> None: """Update proxy_server_request body so spend logs record the full request.""" proxy_server_request = (call_kwargs.get("litellm_metadata") or {}).get("proxy_server_request") or {} if not isinstance(proxy_server_request, dict): @@ -2134,7 +2131,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs.setdefault("litellm_params", {}) call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request - async def _stream_and_forward(self, model: str, call_kwargs: Dict[str, Any]) -> Optional[Dict[str, Any]]: + async def _stream_and_forward(self, model: str, call_kwargs: dict[str, Any]) -> dict[str, Any] | None: """ Stream ``litellm.aresponses`` and forward every chunk over the WebSocket. @@ -2142,7 +2139,7 @@ class ManagedResponsesWebSocketHandler: directly (before serialization) to avoid a redundant JSON round-trip on every chunk. Returns the completed event dict, or ``None``. """ - completed_event: Optional[Dict[str, Any]] = None + completed_event: dict[str, Any] | None = None stream_response = await litellm.aresponses(model=model, **call_kwargs) async for chunk in stream_response: # type: ignore[union-attr] if chunk is None: @@ -2166,9 +2163,9 @@ class ManagedResponsesWebSocketHandler: def _save_turn_history( self, - completed_event: Optional[Dict[str, Any]], - prior_history: List[Dict[str, Any]], - current_messages: List[Dict[str, Any]], + completed_event: dict[str, Any] | None, + prior_history: list[dict[str, Any]], + current_messages: list[dict[str, Any]], ) -> None: """Store this turn in in-memory history for future previous_response_id lookups.""" if completed_event is None: @@ -2237,7 +2234,7 @@ class ManagedResponsesWebSocketHandler: else: model = requested_model - previous_response_id: Optional[str] = call_kwargs.pop("previous_response_id", None) + previous_response_id: str | None = call_kwargs.pop("previous_response_id", None) current_messages = self._input_to_messages(call_kwargs.get("input")) # Fetch history once; reused in both _apply_history and _save_turn_history diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index e46670b39e4..5d4cfcaf06f 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -3,10 +3,7 @@ import re from collections.abc import Iterable, Mapping from typing import ( Any, - Dict, - List, Optional, - Type, Union, cast, get_type_hints, @@ -103,16 +100,16 @@ class ResponsesAPIRequestUtils: @staticmethod def _check_valid_arg( - supported_params: Optional[List[str]], - non_default_params: Dict, - drop_params: Optional[bool], - custom_llm_provider: Optional[str], + supported_params: list[str] | None, + non_default_params: dict, + drop_params: bool | None, + custom_llm_provider: str | None, model: str, ): if supported_params is None: return unsupported_params = {} - for k in non_default_params.keys(): + for k in non_default_params: if k not in supported_params: unsupported_params[k] = non_default_params[k] if unsupported_params: @@ -129,9 +126,9 @@ class ResponsesAPIRequestUtils: model: str, responses_api_provider_config: BaseResponsesAPIConfig, response_api_optional_params: ResponsesAPIOptionalRequestParams, - allowed_openai_params: Optional[List[str]] = None, + allowed_openai_params: list[str] | None = None, drop_params: bool | None = None, - ) -> Dict: + ) -> dict: """ Get optional parameters for the responses API. @@ -151,7 +148,7 @@ class ResponsesAPIRequestUtils: should_drop_params = litellm.drop_params or drop_params is True - non_default_params = cast(Dict, response_api_optional_params) + non_default_params = cast(dict, response_api_optional_params) # Check for unsupported parameters ResponsesAPIRequestUtils._check_valid_arg( supported_params=supported_params + (allowed_openai_params or []), @@ -183,7 +180,7 @@ class ResponsesAPIRequestUtils: @staticmethod def get_requested_response_api_optional_param( - params: Dict[str, Any], + params: dict[str, Any], ) -> ResponsesAPIOptionalRequestParams: """ Filter parameters to only include those defined in ResponsesAPIOptionalRequestParams. @@ -235,35 +232,35 @@ class ResponsesAPIRequestUtils: @staticmethod def _update_responses_api_response_id_with_model_id( responses_api_response: ResponsesAPIResponse, - custom_llm_provider: Optional[str], - litellm_metadata: Optional[Dict[str, Any]] = None, + custom_llm_provider: str | None, + litellm_metadata: dict[str, Any] | None = None, ) -> ResponsesAPIResponse: ... @overload @staticmethod def _update_responses_api_response_id_with_model_id( - responses_api_response: Dict[str, Any], - custom_llm_provider: Optional[str], - litellm_metadata: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + responses_api_response: dict[str, Any], + custom_llm_provider: str | None, + litellm_metadata: dict[str, Any] | None = None, + ) -> dict[str, Any]: ... # fmt: on @staticmethod def _update_responses_api_response_id_with_model_id( - responses_api_response: Union[ResponsesAPIResponse, Dict[str, Any]], - custom_llm_provider: Optional[str], - litellm_metadata: Optional[Dict[str, Any]] = None, - ) -> Union[ResponsesAPIResponse, Dict[str, Any]]: + responses_api_response: ResponsesAPIResponse | dict[str, Any], + custom_llm_provider: str | None, + litellm_metadata: dict[str, Any] | None = None, + ) -> ResponsesAPIResponse | dict[str, Any]: """Update the responses_api_response_id with model_id and custom_llm_provider. Handles both ``ResponsesAPIResponse`` objects and plain dictionaries returned by some streaming providers. """ litellm_metadata = litellm_metadata or {} - model_info: Dict[str, Any] = litellm_metadata.get("model_info", {}) or {} + model_info: dict[str, Any] = litellm_metadata.get("model_info", {}) or {} model_id = model_info.get("id") # access the response id based on the object type @@ -316,7 +313,7 @@ class ResponsesAPIRequestUtils: return f"encitem_{encoded}" @staticmethod - def _decode_encrypted_item_id(encoded_id: str) -> Optional[Dict[str, str]]: + def _decode_encrypted_item_id(encoded_id: str) -> dict[str, str] | None: """Decode a litellm-encoded encrypted-content item ID. Returns a dict with ``model_id`` and ``item_id`` keys, or ``None`` if @@ -357,7 +354,7 @@ class ResponsesAPIRequestUtils: @staticmethod def _unwrap_encrypted_content_with_model_id( wrapped_content: str, - ) -> tuple[Optional[str], str]: + ) -> tuple[str | None, str]: """Unwrap encrypted_content to extract model_id and original content. Returns: @@ -389,9 +386,9 @@ class ResponsesAPIRequestUtils: @staticmethod def _update_encrypted_content_item_ids_in_response( - response: Union["ResponsesAPIResponse", Dict[str, Any]], - model_id: Optional[str], - ) -> Union["ResponsesAPIResponse", Dict[str, Any]]: + response: Union["ResponsesAPIResponse", dict[str, Any]], + model_id: str | None, + ) -> Union["ResponsesAPIResponse", dict[str, Any]]: """Rewrite item IDs for output items that contain ``encrypted_content``. Encodes ``model_id`` into the item ID so that follow-up requests can be @@ -403,7 +400,7 @@ class ResponsesAPIRequestUtils: if not model_id: return response - output: Optional[list] = None + output: list | None = None if isinstance(response, dict): output = response.get("output") else: @@ -481,8 +478,8 @@ class ResponsesAPIRequestUtils: @staticmethod def _build_responses_api_response_id( - custom_llm_provider: Optional[str], - model_id: Optional[str], + custom_llm_provider: str | None, + model_id: str | None, response_id: str, ) -> str: """Build the responses_api_response_id""" @@ -554,7 +551,7 @@ class ResponsesAPIRequestUtils: ) @staticmethod - def get_model_id_from_response_id(response_id: Optional[str]) -> Optional[str]: + def get_model_id_from_response_id(response_id: str | None) -> str | None: """Get the model_id from the response_id""" if response_id is None: return None @@ -583,8 +580,8 @@ class ResponsesAPIRequestUtils: @staticmethod def _build_container_id( - custom_llm_provider: Optional[str], - model_id: Optional[str], + custom_llm_provider: str | None, + model_id: str | None, container_id: str, ) -> str: """Build a managed container ID with provider and model info encoded. @@ -671,8 +668,8 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_ids_in_annotations( annotations: Any, - custom_llm_provider: Optional[str], - model_id: Optional[str], + custom_llm_provider: str | None, + model_id: str | None, ) -> None: """Encode ``container_id`` on each annotation (e.g. ``container_file_citation``).""" if not annotations or not isinstance(annotations, list): @@ -687,8 +684,8 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_ids_in_message_content( content: Any, - custom_llm_provider: Optional[str], - model_id: Optional[str], + custom_llm_provider: str | None, + model_id: str | None, ) -> None: """Walk message ``content`` parts and encode citation ``container_id`` values.""" if not content: @@ -711,8 +708,8 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_id_on_output_item( item: Any, - custom_llm_provider: Optional[str], - model_id: Optional[str], + custom_llm_provider: str | None, + model_id: str | None, ) -> None: """Mutate one output item (dict or object): wrap raw ``container_id`` as LiteLLM-managed. @@ -727,7 +724,7 @@ class ResponsesAPIRequestUtils: if item is None: return - def _maybe_encode(container_id: str) -> Optional[str]: + def _maybe_encode(container_id: str) -> str | None: decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) if decoded.get("custom_llm_provider") is not None: return None @@ -873,17 +870,17 @@ class ResponsesAPIRequestUtils: @staticmethod def _update_container_ids_in_response( - responses_api_response: Union[ResponsesAPIResponse, Dict[str, Any]], - custom_llm_provider: Optional[str], - litellm_metadata: Optional[Dict[str, Any]] = None, - ) -> Union[ResponsesAPIResponse, Dict[str, Any]]: + responses_api_response: ResponsesAPIResponse | dict[str, Any], + custom_llm_provider: str | None, + litellm_metadata: dict[str, Any] | None = None, + ) -> ResponsesAPIResponse | dict[str, Any]: """Encode container IDs in the response output with provider/model info. This walks through all output items and encodes any container_id fields so that follow-up container API calls can auto-route to the correct provider. """ litellm_metadata = litellm_metadata or {} - model_info: Dict[str, Any] = litellm_metadata.get("model_info", {}) or {} + model_info: dict[str, Any] = litellm_metadata.get("model_info", {}) or {} model_id = model_info.get("id") # Get the output list @@ -906,7 +903,7 @@ class ResponsesAPIRequestUtils: @staticmethod def convert_text_format_to_text_param( - text_format: Optional[Union[Type["BaseModel"], dict]], + text_format: type["BaseModel"] | dict | None, text: Optional["ResponseText"] = None, ) -> Optional["ResponseText"]: """ @@ -940,13 +937,13 @@ class ResponsesAPIRequestUtils: @staticmethod def extract_mcp_headers_from_request( - secret_fields: Optional[Dict[str, Any]], - tools: Optional[Iterable[Any]], + secret_fields: dict[str, Any] | None, + tools: Iterable[Any] | None, ) -> tuple[ - Optional[str], - Optional[Dict[str, Dict[str, str]]], - Optional[Dict[str, str]], - Optional[Dict[str, str]], + str | None, + dict[str, dict[str, str]] | None, + dict[str, str] | None, + dict[str, str] | None, ]: """ Extract MCP auth headers from the request to pass to MCP server. @@ -959,14 +956,14 @@ class ResponsesAPIRequestUtils: ) # Extract headers from secret_fields which contains the original request headers - raw_headers_from_request: Optional[Dict[str, str]] = None + raw_headers_from_request: dict[str, str] | None = None if secret_fields and isinstance(secret_fields, dict): raw_headers_from_request = secret_fields.get("raw_headers") # Extract MCP-specific headers using MCPRequestHandler methods - mcp_auth_header: Optional[str] = None - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None - oauth2_headers: Optional[Dict[str, str]] = None + mcp_auth_header: str | None = None + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None + oauth2_headers: dict[str, str] | None = None if raw_headers_from_request: headers_obj = Headers(raw_headers_from_request) @@ -1011,7 +1008,7 @@ class ResponsesAPIRequestUtils: class ResponseAPILoggingUtils: @staticmethod - def _is_response_api_usage(usage: Union[dict, ResponseAPIUsage]) -> bool: + def _is_response_api_usage(usage: dict | ResponseAPIUsage) -> bool: """returns True if usage is from OpenAI Response API""" if isinstance(usage, ResponseAPIUsage): return True @@ -1021,7 +1018,7 @@ class ResponseAPILoggingUtils: @staticmethod def _transform_response_api_usage_to_chat_usage( - usage_input: Optional[Union[dict, ResponseAPIUsage]], + usage_input: dict | ResponseAPIUsage | None, ) -> Usage: """ Transforms ResponseAPIUsage or ImageUsage to a Usage object. @@ -1055,7 +1052,7 @@ class ResponseAPILoggingUtils: response_api_usage = usage_input prompt_tokens: int = response_api_usage.input_tokens or 0 completion_tokens: int = response_api_usage.output_tokens or 0 - prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + prompt_tokens_details: PromptTokensDetailsWrapper | None = None if response_api_usage.input_tokens_details: if isinstance(response_api_usage.input_tokens_details, dict): prompt_tokens_details = PromptTokensDetailsWrapper(**response_api_usage.input_tokens_details) @@ -1067,7 +1064,7 @@ class ResponseAPILoggingUtils: image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None), cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None), ) - completion_tokens_details: Optional[CompletionTokensDetailsWrapper] = None + completion_tokens_details: CompletionTokensDetailsWrapper | None = None output_tokens_details = getattr(response_api_usage, "output_tokens_details", None) if output_tokens_details: completion_tokens_details = CompletionTokensDetailsWrapper( diff --git a/litellm/router.py b/litellm/router.py index 396f792a3f5..722dbf75509 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -24,13 +24,8 @@ from functools import lru_cache from typing import ( TYPE_CHECKING, Any, - Dict, - FrozenSet, - List, Literal, Optional, - Set, - Tuple, TypeVar, Union, cast, @@ -263,7 +258,7 @@ else: PreRoutingHookResponse = Any -def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]: +def _cost_value_as_float(value: str | float | None) -> float | None: if value is None: return None try: @@ -318,60 +313,60 @@ class RoutingArgs(enum.Enum): class Router: model_names: set = set() - cache_responses: Optional[bool] = False + cache_responses: bool | None = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour tenacity = None - leastbusy_logger: Optional[LeastBusyLoggingHandler] = None - lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None - optional_callbacks: Optional[List[Union[CustomLogger, Callable, str]]] = None + leastbusy_logger: LeastBusyLoggingHandler | None = None + lowesttpm_logger: LowestTPMLoggingHandler | None = None + optional_callbacks: list[CustomLogger | Callable | str] | None = None def __init__( self, - model_list: Optional[Union[List[DeploymentTypedDict], List[Dict[str, Any]]]] = None, + model_list: list[DeploymentTypedDict] | list[dict[str, Any]] | None = None, ## ASSISTANTS API ## - assistants_config: Optional[AssistantsTypedDict] = None, + assistants_config: AssistantsTypedDict | None = None, ## SEARCH API ## - search_tools: Optional[List[SearchToolTypedDict]] = None, + search_tools: list[SearchToolTypedDict] | None = None, ## GUARDRAIL API ## - guardrail_list: Optional[List[GuardrailTypedDict]] = None, + guardrail_list: list[GuardrailTypedDict] | None = None, ## CACHING ## - redis_url: Optional[str] = None, - redis_host: Optional[str] = None, - redis_port: Optional[int] = None, - redis_password: Optional[str] = None, - redis_db: Optional[int] = None, - cache_responses: Optional[bool] = False, + redis_url: str | None = None, + redis_host: str | None = None, + redis_port: int | None = None, + redis_password: str | None = None, + redis_db: int | None = None, + cache_responses: bool | None = False, cache_kwargs: dict = {}, # additional kwargs to pass to RedisCache (see caching.py) - caching_groups: Optional[List[tuple]] = None, # if you want to cache across model groups + caching_groups: list[tuple] | None = None, # if you want to cache across model groups client_ttl: int = 3600, # ttl for cached clients - will re-initialize after this time in seconds ## SCHEDULER ## - polling_interval: Optional[float] = None, - default_priority: Optional[int] = None, + polling_interval: float | None = None, + default_priority: int | None = None, ## RELIABILITY ## - num_retries: Optional[int] = None, - max_fallbacks: Optional[int] = None, # max fallbacks to try before exiting the call. Defaults to 5. - timeout: Optional[float] = None, - stream_timeout: Optional[float] = None, - default_litellm_params: Optional[dict] = None, # default params for Router.chat.completion.create - default_max_parallel_requests: Optional[int] = None, + num_retries: int | None = None, + max_fallbacks: int | None = None, # max fallbacks to try before exiting the call. Defaults to 5. + timeout: float | None = None, + stream_timeout: float | None = None, + default_litellm_params: dict | None = None, # default params for Router.chat.completion.create + default_max_parallel_requests: int | None = None, set_verbose: bool = False, debug_level: Literal["DEBUG", "INFO"] = "INFO", - default_fallbacks: Optional[List[str]] = None, # generic fallbacks, works across all deployments - fallbacks: List = [], - context_window_fallbacks: List = [], - content_policy_fallbacks: List = [], - model_group_alias: Optional[Dict[str, Union[str, RouterModelGroupAliasItem]]] = {}, + default_fallbacks: list[str] | None = None, # generic fallbacks, works across all deployments + fallbacks: list = [], + context_window_fallbacks: list = [], + content_policy_fallbacks: list = [], + model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = {}, enable_pre_call_checks: bool = False, enable_tag_filtering: bool = False, tag_filtering_match_any: bool = True, plugins: list[RoutingPlugin] | None = None, retry_after: int = 0, # min time to wait before retrying a failed request - retry_policy: Optional[Union[RetryPolicy, dict]] = None, # set custom retries for different exceptions - model_group_retry_policy: Dict[str, RetryPolicy] = {}, # set custom retry policies based on model group - allowed_fails: Optional[int] = None, # Number of times a deployment can failbefore being added to cooldown - allowed_fails_policy: Optional[AllowedFailsPolicy] = None, # set custom allowed fails policy - cooldown_time: Optional[float] = None, # (seconds) time to cooldown a deployment after failure - disable_cooldowns: Optional[bool] = None, + retry_policy: RetryPolicy | dict | None = None, # set custom retries for different exceptions + model_group_retry_policy: dict[str, RetryPolicy] = {}, # set custom retry policies based on model group + allowed_fails: int | None = None, # Number of times a deployment can failbefore being added to cooldown + allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy + cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure + disable_cooldowns: bool | None = None, routing_strategy: Literal[ "simple-shuffle", "least-busy", @@ -381,17 +376,17 @@ class Router: "usage-based-routing-v2", "lar1", ] = "simple-shuffle", - optional_pre_call_checks: Optional[OptionalPreCallChecks] = None, + optional_pre_call_checks: OptionalPreCallChecks | None = None, routing_strategy_args: dict = {}, # just for latency-based - routing_groups: Optional[List[Union[RoutingGroup, dict]]] = None, - provider_budget_config: Optional[GenericBudgetConfigType] = None, - alerting_config: Optional[AlertingConfig] = None, - router_general_settings: Optional[RouterGeneralSettings] = RouterGeneralSettings(), + routing_groups: list[RoutingGroup | dict] | None = None, + provider_budget_config: GenericBudgetConfigType | None = None, + alerting_config: AlertingConfig | None = None, + router_general_settings: RouterGeneralSettings | None = RouterGeneralSettings(), deployment_affinity_ttl_seconds: int = 3600, - model_group_affinity_config: Optional[Dict[str, List[str]]] = None, + model_group_affinity_config: dict[str, list[str]] | None = None, ignore_invalid_deployments: bool = False, enable_health_check_routing: bool = False, - health_check_staleness_threshold: Optional[int] = None, + health_check_staleness_threshold: int | None = None, health_check_ignore_transient_errors: bool = False, enable_weighted_failover: bool = False, ) -> None: @@ -487,12 +482,12 @@ class Router: self.assistants_config = assistants_config self.search_tools = search_tools or [] self.guardrail_list = guardrail_list or [] - self.deployment_names: List = [] # names of models under litellm_params. ex. azure/chatgpt-v-2 + self.deployment_names: list = [] # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = "local" # default to an in-memory cache redis_cache = None - cache_config: Dict[str, Any] = {} + cache_config: dict[str, Any] = {} self.client_ttl = client_ttl if redis_url is not None or (redis_host is not None and redis_port is not None): @@ -536,49 +531,49 @@ class Router: None # use this to track the users default deployment, when they want to use model = * ) self.default_max_parallel_requests = default_max_parallel_requests - self.provider_default_deployment_ids: List[str] = [] + self.provider_default_deployment_ids: list[str] = [] self.pattern_router = PatternMatchRouter() - self.team_pattern_routers: Dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter} - self.auto_routers: dict[str, list[TaggedPreRoutingStrategy["AutoRouter"]]] = {} - self.complexity_routers: dict[str, list[TaggedPreRoutingStrategy["ComplexityRouter"]]] = {} - self.adaptive_routers: dict[str, list[TaggedPreRoutingStrategy["AdaptiveRouter"]]] = {} - self.quality_routers: dict[str, list[TaggedPreRoutingStrategy["QualityRouter"]]] = {} + self.team_pattern_routers: dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter} + self.auto_routers: dict[str, list[TaggedPreRoutingStrategy[AutoRouter]]] = {} + self.complexity_routers: dict[str, list[TaggedPreRoutingStrategy[ComplexityRouter]]] = {} + self.adaptive_routers: dict[str, list[TaggedPreRoutingStrategy[AdaptiveRouter]]] = {} + self.quality_routers: dict[str, list[TaggedPreRoutingStrategy[QualityRouter]]] = {} self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else [] # Initialize model_group_alias early since it's used in set_model_list - self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = ( + self.model_group_alias: dict[str, str | RouterModelGroupAliasItem] = ( model_group_alias or {} ) # dict to store aliases for router, ex. {"gpt-4": "gpt-3.5-turbo"}, all requests with gpt-4 -> get routed to gpt-3.5-turbo group # Initialize model ID to deployment index mapping for O(1) lookups - self.model_id_to_deployment_index_map: Dict[str, int] = {} + self.model_id_to_deployment_index_map: dict[str, int] = {} # Initialize model name to deployment indices mapping for O(1) lookups # Maps model_name -> list of indices in model_list - self.model_name_to_deployment_indices: Dict[str, List[int]] = {} + self.model_name_to_deployment_indices: dict[str, list[int]] = {} # Maps (team_id, team_public_model_name) -> list of indices in model_list - self.team_model_to_deployment_indices: Dict[Tuple[str, str], List[int]] = {} - self.team_public_model_names: FrozenSet[str] = frozenset() + self.team_model_to_deployment_indices: dict[tuple[str, str], list[int]] = {} + self.team_public_model_names: frozenset[str] = frozenset() # Initialize cache attributes that ``_invalidate_model_group_info_cache`` # touches *before* the first ``set_model_list`` below (which calls # that invalidation as part of building the model index). - self._access_groups_cache: Optional[Dict[str, List[str]]] = None + self._access_groups_cache: dict[str, list[str]] | None = None # Per-router cache for the proxy auth-layer "is this model explicitly # zero-cost?" check. Lives on the router so it is invalidated alongside # ``_cached_get_model_group_info`` and dies with the router (no # ``id()``-reuse risk after GC). See # ``litellm.proxy.auth.auth_checks._is_model_cost_zero``. - self._zero_cost_cache: Dict[str, bool] = {} + self._zero_cost_cache: dict[str, bool] = {} if model_list is not None: # set_model_list will build indices automatically self.set_model_list(model_list) - self.healthy_deployments: List = self.model_list # type: ignore + self.healthy_deployments: list = self.model_list # type: ignore for m in model_list: if "model" in m["litellm_params"]: self.deployment_latency_map[m["litellm_params"]["model"]] = 0 else: - self.model_list: List = [] # initialize an empty list - to allow _add_deployment and delete_deployment to work + self.model_list: list = [] # initialize an empty list - to allow _add_deployment and delete_deployment to work if allowed_fails is not None: self.allowed_fails = allowed_fails @@ -620,7 +615,7 @@ class Router: self.retry_after = retry_after self.routing_strategy = self._normalize_strategy(routing_strategy) - self._routing_groups_input: Optional[List[Union[RoutingGroup, dict]]] = routing_groups + self._routing_groups_input: list[RoutingGroup | dict] | None = routing_groups ## SETTING FALLBACKS ## ### validate if it's set + in correct format @@ -645,7 +640,7 @@ class Router: self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model - self.previous_models: List = [] # list to store failed calls (passed in as metadata to next call) + self.previous_models: list = [] # list to store failed calls (passed in as metadata to next call) # make Router.chat.completions.create compatible for openai.chat.completions.create default_litellm_params = default_litellm_params or {} @@ -709,7 +704,7 @@ class Router: self.routing_strategy_args = routing_strategy_args self.provider_budget_config = provider_budget_config self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds - self.router_budget_logger: Optional[RouterBudgetLimiting] = None + self.router_budget_logger: RouterBudgetLimiting | None = None if RouterBudgetLimiting.should_init_router_budget_limiter( model_list=model_list, provider_budget_config=self.provider_budget_config ): @@ -717,7 +712,7 @@ class Router: optional_pre_call_checks.append("router_budget_limiting") else: optional_pre_call_checks = ["router_budget_limiting"] - self.retry_policy: Optional[RetryPolicy] = None + self.retry_policy: RetryPolicy | None = None if retry_policy is not None: if isinstance(retry_policy, dict): self.retry_policy = RetryPolicy(**retry_policy) @@ -725,15 +720,13 @@ class Router: self.retry_policy = retry_policy if self.retry_policy is not None: verbose_router_logger.info( - "\033[32mRouter Custom Retry Policy Set:\n{}\033[0m".format( - self.retry_policy.model_dump(exclude_none=True) - ) + f"\033[32mRouter Custom Retry Policy Set:\n{self.retry_policy.model_dump(exclude_none=True)}\033[0m" ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = model_group_retry_policy - self.model_group_affinity_config: Optional[Dict[str, List[str]]] = model_group_affinity_config + self.model_group_retry_policy: dict[str, RetryPolicy] | None = model_group_retry_policy + self.model_group_affinity_config: dict[str, list[str]] | None = model_group_affinity_config - self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None + self.allowed_fails_policy: AllowedFailsPolicy | None = None if allowed_fails_policy is not None: if isinstance(allowed_fails_policy, dict): self.allowed_fails_policy = AllowedFailsPolicy(**allowed_fails_policy) @@ -742,12 +735,10 @@ class Router: if self.allowed_fails_policy is not None: verbose_router_logger.info( - "\033[32mRouter Custom Allowed Fails Policy Set:\n{}\033[0m".format( - self.allowed_fails_policy.model_dump(exclude_none=True) - ) + f"\033[32mRouter Custom Allowed Fails Policy Set:\n{self.allowed_fails_policy.model_dump(exclude_none=True)}\033[0m" ) - self.alerting_config: Optional[AlertingConfig] = alerting_config + self.alerting_config: AlertingConfig | None = alerting_config if optional_pre_call_checks is not None: self.add_optional_pre_call_checks(optional_pre_call_checks) @@ -779,7 +770,7 @@ class Router: self.apply_default_settings() @staticmethod - def get_valid_args() -> List[str]: + def get_valid_args() -> list[str]: """ Returns a list of valid arguments for the Router.__init__ method. """ @@ -796,7 +787,6 @@ class Router: default_pre_call_checks: OptionalPreCallChecks = [] self.add_optional_pre_call_checks(default_pre_call_checks) - return None def discard(self): """ @@ -820,8 +810,8 @@ class Router: @staticmethod def _create_redis_cache( - cache_config: Dict[str, Any], - ) -> Union[RedisCache, RedisClusterCache]: + cache_config: dict[str, Any], + ) -> RedisCache | RedisClusterCache: """ Initializes either a RedisCache or RedisClusterCache based on the cache_config. """ @@ -853,7 +843,7 @@ class Router: # Maps a routing strategy string to the attribute on `self` that holds # the default group's strategy selector for that strategy. (The selectors # double as `CustomLogger` callbacks, hence the legacy `*_logger` attrs.) - _DEFAULT_SELECTOR_ATTR_BY_STRATEGY: Dict[str, str] = { + _DEFAULT_SELECTOR_ATTR_BY_STRATEGY: dict[str, str] = { "least-busy": "leastbusy_logger", "usage-based-routing": "lowesttpm_logger", "usage-based-routing-v2": "lowesttpm_logger_v2", @@ -863,15 +853,15 @@ class Router: @staticmethod def _normalize_strategy( - strategy: Union[RoutingStrategy, str, None], - ) -> Optional[str]: + strategy: RoutingStrategy | str | None, + ) -> str | None: if strategy is None: return None if isinstance(strategy, RoutingStrategy): return strategy.value return strategy - def _validate_routing_strategy(self, routing_strategy: Union[RoutingStrategy, str, None]) -> None: + def _validate_routing_strategy(self, routing_strategy: RoutingStrategy | str | None) -> None: # See: https://github.com/BerriAI/litellm/issues/11330 valid_strategy_strings = ["simple-shuffle", "lar1"] + [s.value for s in RoutingStrategy] if routing_strategy is None: @@ -888,16 +878,16 @@ class Router: def _build_strategy_selector( self, - strategy: Union[RoutingStrategy, str], + strategy: RoutingStrategy | str, routing_strategy_args: dict, register_callbacks: bool = True, - ) -> Optional[Any]: + ) -> Any | None: """ Constructs a strategy selector for a given strategy. Returns None for `simple-shuffle` (no selector needed) and unknown strategies. """ - selector: Optional[Any] = None + selector: Any | None = None match self._normalize_strategy(strategy): case RoutingStrategy.LEAST_BUSY.value: selector = LeastBusyLoggingHandler(router_cache=self.cache) @@ -934,7 +924,7 @@ class Router: return selector - def _unregister_router_selectors(self, selectors: List[Any]) -> None: + def _unregister_router_selectors(self, selectors: list[Any]) -> None: """ Drop router-owned strategy selectors from litellm's global callback lists by identity. Used before re-init (`routing_strategy_init` / @@ -949,7 +939,7 @@ class Router: if isinstance(litellm.input_callback, list): litellm.input_callback = [c for c in litellm.input_callback if id(c) not in selector_ids] - def routing_strategy_init(self, routing_strategy: Union[RoutingStrategy, str], routing_strategy_args: dict): + def routing_strategy_init(self, routing_strategy: RoutingStrategy | str, routing_strategy_args: dict): verbose_router_logger.info(f"Routing strategy: {routing_strategy}") self._validate_routing_strategy(routing_strategy) self._reset_custom_routing_strategy() @@ -960,11 +950,11 @@ class Router: ) self._override_selectors = {} - self.leastbusy_logger: Optional[LeastBusyLoggingHandler] = None - self.lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None - self.lowesttpm_logger_v2: Optional[LowestTPMLoggingHandler_v2] = None - self.lowestlatency_logger: Optional[LowestLatencyLoggingHandler] = None - self.lowestcost_logger: Optional[LowestCostLoggingHandler] = None + self.leastbusy_logger: LeastBusyLoggingHandler | None = None + self.lowesttpm_logger: LowestTPMLoggingHandler | None = None + self.lowesttpm_logger_v2: LowestTPMLoggingHandler_v2 | None = None + self.lowestlatency_logger: LowestLatencyLoggingHandler | None = None + self.lowestcost_logger: LowestCostLoggingHandler | None = None selector = self._build_strategy_selector( strategy=routing_strategy, @@ -980,7 +970,7 @@ class Router: def _init_routing_groups( self, - groups_input: Optional[List[Union[RoutingGroup, dict]]], + groups_input: list[RoutingGroup | dict] | None, ) -> None: """ Validates and indexes `routing_groups`. Each `model_name` may belong to @@ -995,9 +985,9 @@ class Router: [sel for selectors in getattr(self, "_group_selectors", {}).values() for sel in selectors.values()] ) - self._routing_groups: Dict[str, RoutingGroup] = {} - self._model_to_group: Dict[str, str] = {} - self._group_selectors: Dict[str, Dict[str, Any]] = {} + self._routing_groups: dict[str, RoutingGroup] = {} + self._model_to_group: dict[str, str] = {} + self._group_selectors: dict[str, dict[str, Any]] = {} if not groups_input: return @@ -1050,7 +1040,7 @@ class Router: _OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY}) - def _get_request_routing_strategy_override(self, request_kwargs: Optional[dict]) -> Optional[str]: + def _get_request_routing_strategy_override(self, request_kwargs: dict | None) -> str | None: """ Reads a per-request `routing_strategy` override (forwarded by the proxy from key/team `router_settings`) out of the request kwargs. @@ -1075,7 +1065,7 @@ class Router: return None return strategy - def _get_override_strategy_selector(self, strategy: str) -> Optional[Any]: + def _get_override_strategy_selector(self, strategy: str) -> Any | None: """ Returns the selector for a per-request strategy override. @@ -1096,9 +1086,7 @@ class Router: ) return self._override_selectors[strategy] - def _get_routing_context( - self, model: str, request_kwargs: Optional[dict] = None - ) -> tuple[Optional[str], Optional[Any]]: + def _get_routing_context(self, model: str, request_kwargs: dict | None = None) -> tuple[str | None, Any | None]: """ Resolves the routing strategy and selector to use for the given model. @@ -1138,14 +1126,14 @@ class Router: async def _select_deployment_async( self, *, - strategy: Optional[str], - selector: Optional[Any], + strategy: str | None, + selector: Any | None, model: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]], - input: Optional[Union[str, List]], - request_kwargs: Optional[Dict], - ) -> Optional[Any]: + messages: list[dict[str, str]] | None, + input: str | list | None, + request_kwargs: dict | None, + ) -> Any | None: """ Asks the strategy selector for a deployment. Caller handles `simple-shuffle` separately (it does not flow through a selector). @@ -1191,14 +1179,14 @@ class Router: def _select_deployment_sync( self, *, - strategy: Optional[str], - selector: Optional[Any], + strategy: str | None, + selector: Any | None, model: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]], - input: Optional[Union[str, List]], - request_kwargs: Optional[Dict], - ) -> Optional[Any]: + messages: list[dict[str, str]] | None, + input: str | list | None, + request_kwargs: dict | None, + ) -> Any | None: """ Sync sibling of `_select_deployment_async`. Caller handles `simple-shuffle` separately. @@ -1565,7 +1553,7 @@ class Router: self._initialize_core_endpoints() self._initialize_specialized_endpoints() - def validate_fallbacks(self, fallback_param: Optional[List]): + def validate_fallbacks(self, fallback_param: list | None): """ Validate the fallbacks parameter. """ @@ -1585,7 +1573,7 @@ class Router: ) def _move_before_deployment_affinity( - callback_list: List[Any], + callback_list: list[Any], callback_to_move: EncryptedContentAffinityCheck, ) -> None: if callback_to_move not in callback_list: @@ -1603,7 +1591,7 @@ class Router: if self.optional_callbacks is None: self.optional_callbacks = [] - existing_ec_callback: Optional[EncryptedContentAffinityCheck] = None + existing_ec_callback: EncryptedContentAffinityCheck | None = None for cb in self.optional_callbacks: if isinstance(cb, EncryptedContentAffinityCheck): existing_ec_callback = cb @@ -1628,7 +1616,7 @@ class Router: _move_before_deployment_affinity(self.optional_callbacks, ec_callback) _move_before_deployment_affinity(litellm.callbacks, ec_callback) - def add_optional_pre_call_checks(self, optional_pre_call_checks: Optional[OptionalPreCallChecks]): + def add_optional_pre_call_checks(self, optional_pre_call_checks: OptionalPreCallChecks | None): if optional_pre_call_checks is None: return @@ -1642,7 +1630,7 @@ class Router: if self.optional_callbacks is None: self.optional_callbacks = [] - existing_affinity_callback: Optional[DeploymentAffinityCheck] = None + existing_affinity_callback: DeploymentAffinityCheck | None = None for cb in self.optional_callbacks: if isinstance(cb, DeploymentAffinityCheck): existing_affinity_callback = cb @@ -1684,7 +1672,7 @@ class Router: # Remaining optional pre-call checks # --------------------------------------------------------------------- for pre_call_check in optional_pre_call_checks: - _callback: Optional[CustomLogger] = None + _callback: CustomLogger | None = None if pre_call_check in ( "deployment_affinity", "responses_api_deployment_check", @@ -1732,14 +1720,12 @@ class Router: return _deployment_copy except Exception as e: - verbose_router_logger.debug(f"Error occurred while printing deployment - {str(e)}") + verbose_router_logger.debug(f"Error occurred while printing deployment - {e!s}") raise e ### COMPLETION, EMBEDDING, IMG GENERATION FUNCTIONS - def completion( - self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ModelResponse, CustomStreamWrapper]: + def completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper: """ Example usage: response = router.completion(model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hey, how's it going?"}] @@ -1756,9 +1742,7 @@ class Router: except Exception as e: raise e - def _completion( - self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ModelResponse, CustomStreamWrapper]: + def _completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper: model_name = None deployment = None try: @@ -1844,7 +1828,7 @@ class Router: return response except Exception as e: - verbose_router_logger.info(f"litellm.completion(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.completion(model={model_name})\033[31m Exception {e!s}\033[0m") # Set per-deployment num_retries on exception for retry logic if deployment is not None: self._set_deployment_num_retries_on_exception(e, deployment) @@ -1896,7 +1880,7 @@ class Router: return silent_kwargs - def _silent_experiment_completion(self, silent_model: str, messages: List[Any], **kwargs): + def _silent_experiment_completion(self, silent_model: str, messages: list[Any], **kwargs): """ Run a silent experiment in the background (thread). """ @@ -1923,7 +1907,7 @@ class Router: async def _run_silent_completion(): await self.acompletion( model=silent_model, - messages=cast(List[AllMessageValues], messages), + messages=cast(list[AllMessageValues], messages), **silent_kwargs, ) # Drain any fire-and-forget tasks (e.g. alerting hooks) @@ -1939,26 +1923,26 @@ class Router: finally: loop.close() except Exception as e: - verbose_router_logger.error(f"Silent experiment failed for model {silent_model}: {str(e)}") + verbose_router_logger.error(f"Silent experiment failed for model {silent_model}: {e!s}") # fmt: off @overload async def acompletion( - self, model: str, messages: List[AllMessageValues], stream: Literal[True], **kwargs + self, model: str, messages: list[AllMessageValues], stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... @overload async def acompletion( - self, model: str, messages: List[AllMessageValues], stream: Literal[False] = False, **kwargs + self, model: str, messages: list[AllMessageValues], stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... @overload async def acompletion( - self, model: str, messages: List[AllMessageValues], stream: Union[Literal[True], Literal[False]] = False, **kwargs - ) -> Union[CustomStreamWrapper, ModelResponse]: + self, model: str, messages: list[AllMessageValues], stream: Literal[True, False] = False, **kwargs + ) -> CustomStreamWrapper | ModelResponse: ... # fmt: on @@ -1967,7 +1951,7 @@ class Router: async def acompletion( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], stream: bool = False, **kwargs, ): @@ -2020,12 +2004,12 @@ class Router: @staticmethod def _combine_fallback_usage( fallback_item: ModelResponseStream, - complete_response_object_usage: Optional[Usage], + complete_response_object_usage: Usage | None, ) -> None: """Merge partial-stream usage with fallback-stream usage on the chunk.""" from litellm.cost_calculator import BaseTokenUsageProcessor - usage = cast(Optional[Usage], getattr(fallback_item, "usage", None)) + usage = cast(Usage | None, getattr(fallback_item, "usage", None)) usage_objects = [usage] if usage is not None else [] if ( complete_response_object_usage is not None @@ -2069,7 +2053,7 @@ class Router: async def _acompletion_streaming_iterator( self, model_response: CustomStreamWrapper, - messages: List[Dict[str, str]], + messages: list[dict[str, str]], initial_kwargs: dict, ) -> CustomStreamWrapper: """ @@ -2112,17 +2096,17 @@ class Router: complete_response_object = stream_chunk_builder(chunks=model_response.chunks) complete_response_object_usage = cast( - Optional[Usage], + Usage | None, getattr(complete_response_object, "usage", None), ) try: # Use the router's fallback system model_group = cast(str, initial_kwargs.get("model")) - fallbacks: Optional[List] = initial_kwargs.get("fallbacks", self.fallbacks) - context_window_fallbacks: Optional[List] = initial_kwargs.get( + fallbacks: list | None = initial_kwargs.get("fallbacks", self.fallbacks) + context_window_fallbacks: list | None = initial_kwargs.get( "context_window_fallbacks", self.context_window_fallbacks ) - content_policy_fallbacks: Optional[List] = initial_kwargs.get( + content_policy_fallbacks: list | None = initial_kwargs.get( "content_policy_fallbacks", self.content_policy_fallbacks ) initial_kwargs["original_function"] = self._acompletion @@ -2328,7 +2312,7 @@ class Router: @staticmethod def _build_responses_continuation_input( - input_val: Optional[Union[str, "ResponseInputParam"]], + input_val: Union[str, "ResponseInputParam"] | None, generated_content: str, ) -> "ResponseInputParam": """ @@ -2350,7 +2334,7 @@ class Router: # ResponseOutputMessageParam, ...) — annotating as List[Dict[str, Any]] # rejects the list() spread of input_val. We cast the combined list to # ResponseInputParam at the return. - base: List[Any] + base: list[Any] if isinstance(input_val, str): base = [ { @@ -2363,7 +2347,7 @@ class Router: base = list(input_val) else: base = [] - continuation: List[Any] = [ + continuation: list[Any] = [ { "type": "message", "role": "developer", @@ -2390,7 +2374,7 @@ class Router: async def _aresponses_streaming_iterator( self, response: "BaseResponsesAPIStreamingIterator", - initial_kwargs: Dict[str, Any], + initial_kwargs: dict[str, Any], ) -> "BaseResponsesAPIStreamingIterator": """ Wrap a Responses-API streaming iterator so MidStreamFallbackError @@ -2539,11 +2523,11 @@ class Router: partial_usage = Router._extract_partial_responses_usage(source_iterator) try: model_group = cast(str, initial_kwargs.get("model")) - fallbacks: Optional[List] = initial_kwargs.get("fallbacks", self.fallbacks) - context_window_fallbacks: Optional[List] = initial_kwargs.get( + fallbacks: list | None = initial_kwargs.get("fallbacks", self.fallbacks) + context_window_fallbacks: list | None = initial_kwargs.get( "context_window_fallbacks", self.context_window_fallbacks ) - content_policy_fallbacks: Optional[List] = initial_kwargs.get( + content_policy_fallbacks: list | None = initial_kwargs.get( "content_policy_fallbacks", self.content_policy_fallbacks ) # Re-enter via the per-attempt helper so the fallback chain @@ -2625,7 +2609,7 @@ class Router: def _completion_streaming_iterator( self, model_response: CustomStreamWrapper, - messages: List[Dict[str, str]], + messages: list[dict[str, str]], initial_kwargs: dict, ) -> CustomStreamWrapper: """ @@ -2667,17 +2651,17 @@ class Router: complete_response_object = stream_chunk_builder(chunks=model_response.chunks) complete_response_object_usage = cast( - Optional[Usage], + Usage | None, getattr(complete_response_object, "usage", None), ) try: model_group = cast(str, initial_kwargs.get("model")) - fallbacks: Optional[List] = initial_kwargs.get("fallbacks", router_self.fallbacks) - context_window_fallbacks: Optional[List] = initial_kwargs.get( + fallbacks: list | None = initial_kwargs.get("fallbacks", router_self.fallbacks) + context_window_fallbacks: list | None = initial_kwargs.get( "context_window_fallbacks", router_self.context_window_fallbacks, ) - content_policy_fallbacks: Optional[List] = initial_kwargs.get( + content_policy_fallbacks: list | None = initial_kwargs.get( "content_policy_fallbacks", router_self.content_policy_fallbacks, ) @@ -2746,7 +2730,7 @@ class Router: return SyncFallbackStreamWrapper(stream_with_fallbacks()) - async def _silent_experiment_acompletion(self, silent_model: str, messages: List[Any], **kwargs): + async def _silent_experiment_acompletion(self, silent_model: str, messages: list[Any], **kwargs): """ Run a silent experiment in the background. """ @@ -2766,18 +2750,15 @@ class Router: # Trigger the silent request await self.acompletion( model=silent_model, - messages=cast(List[AllMessageValues], messages), + messages=cast(list[AllMessageValues], messages), **silent_kwargs, ) except Exception as e: - verbose_router_logger.error(f"Silent experiment failed for model {silent_model}: {str(e)}") + verbose_router_logger.error(f"Silent experiment failed for model {silent_model}: {e!s}") async def _acompletion( - self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ - ModelResponse, - CustomStreamWrapper, - ]: + self, model: str, messages: list[dict[str, str]], **kwargs + ) -> ModelResponse | CustomStreamWrapper: """ - Get an available deployment - call it with a semaphore over the call @@ -2858,7 +2839,7 @@ class Router: _response = litellm.acompletion(**input_kwargs) - logging_obj: Optional[LiteLLMLogging] = kwargs.get("litellm_logging_obj", None) + logging_obj: LiteLLMLogging | None = kwargs.get("litellm_logging_obj", None) rpm_semaphore = self._get_client( deployment=deployment, @@ -2926,7 +2907,7 @@ class Router: self._set_failed_deployment_id_on_exception(e, deployment) raise e except Exception as e: - verbose_router_logger.info(f"litellm.acompletion(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.acompletion(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 # Set per-deployment num_retries on exception for retry logic @@ -2939,7 +2920,7 @@ class Router: self, model: str, kwargs: dict, - metadata_variable_name: Optional[str] = "metadata", + metadata_variable_name: str | None = "metadata", ) -> None: """ Adds/updates to kwargs: @@ -2958,7 +2939,7 @@ class Router: else: kwargs["num_retries"] = self.num_retries if self.num_retries is not None else 0 kwargs.setdefault("litellm_trace_id", str(uuid.uuid4())) - model_group_alias: Optional[str] = None + model_group_alias: str | None = None if self._get_model_from_alias(model=model): model_group_alias = model kwargs.setdefault(metadata_variable_name, {}).update( @@ -3004,7 +2985,7 @@ class Router: pass def _update_kwargs_with_default_litellm_params( - self, kwargs: dict, metadata_variable_name: Optional[str] = "metadata" + self, kwargs: dict, metadata_variable_name: str | None = "metadata" ) -> None: """ Adds default litellm params to kwargs, if set. @@ -3025,7 +3006,7 @@ class Router: kwargs.setdefault(metadata_variable_name, {}).update(metadata_defaults) def _handle_clientside_credential( - self, deployment: dict, kwargs: dict, function_name: Optional[str] = None + self, deployment: dict, kwargs: dict, function_name: str | None = None ) -> Deployment: """ Handle clientside credential @@ -3074,7 +3055,7 @@ class Router: self, deployment: dict, kwargs: dict, - function_name: Optional[str] = None, + function_name: str | None = None, ) -> None: """ 3 jobs: @@ -3179,7 +3160,7 @@ class Router: return model_client - def _get_stream_timeout(self, kwargs: dict, data: dict) -> Optional[Union[float, int]]: + def _get_stream_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get stream timeout from kwargs or deployment params""" return ( kwargs.get("stream_timeout", None) # the params dynamically set by user @@ -3189,7 +3170,7 @@ class Router: or self.default_litellm_params.get("stream_timeout", None) ) - def _get_non_stream_timeout(self, kwargs: dict, data: dict) -> Optional[Union[float, int]]: + def _get_non_stream_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get non-stream timeout from kwargs or deployment params""" timeout = ( kwargs.get("timeout", None) # the params dynamically set by user @@ -3202,9 +3183,9 @@ class Router: ) return timeout - def _get_timeout(self, kwargs: dict, data: dict) -> Optional[Union[float, int]]: + def _get_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get timeout from kwargs or deployment params""" - timeout: Optional[Union[float, int]] = None + timeout: float | int | None = None if kwargs.get("stream", False): timeout = self._get_stream_timeout(kwargs=kwargs, data=data) if timeout is None: @@ -3215,8 +3196,8 @@ class Router: async def abatch_completion( self, - models: List[str], - messages: Union[List[Dict[str, str]], List[List[Dict[str, str]]]], + models: list[str], + messages: list[dict[str, str]] | list[list[dict[str, str]]], **kwargs, ): """ @@ -3249,7 +3230,7 @@ class Router: """ ############## Helpers for async completion ################## - async def _async_completion_no_exceptions(model: str, messages: List[AllMessageValues], **kwargs): + async def _async_completion_no_exceptions(model: str, messages: list[AllMessageValues], **kwargs): """ Wrapper around self.async_completion that catches exceptions and returns them as a result """ @@ -3260,7 +3241,7 @@ class Router: async def _async_completion_no_exceptions_return_idx( model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], idx: int, # index of message this response corresponds to **kwargs, ): @@ -3298,7 +3279,7 @@ class Router: ) ) responses = await asyncio.gather(*_tasks) - final_responses: List[List[Any]] = [[] for _ in range(len(messages))] + final_responses: list[list[Any]] = [[] for _ in range(len(messages))] for response in responses: if isinstance(response, tuple): final_responses[response[1]].append(response[0]) @@ -3307,7 +3288,7 @@ class Router: return final_responses async def abatch_completion_one_model_multiple_requests( - self, model: str, messages: List[List[AllMessageValues]], **kwargs + self, model: str, messages: list[list[AllMessageValues]], **kwargs ): """ Async Batch Completion - Batch Process multiple Messages to one model_group on litellm.Router @@ -3328,7 +3309,7 @@ class Router: ) """ - async def _async_completion_no_exceptions(model: str, messages: List[AllMessageValues], **kwargs): + async def _async_completion_no_exceptions(model: str, messages: list[AllMessageValues], **kwargs): """ Wrapper around self.async_completion that catches exceptions and returns them as a result """ @@ -3349,7 +3330,7 @@ class Router: @overload async def abatch_completion_fastest_response( - self, model: str, messages: List[Dict[str, str]], stream: Literal[True], **kwargs + self, model: str, messages: list[dict[str, str]], stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... @@ -3357,7 +3338,7 @@ class Router: @overload async def abatch_completion_fastest_response( - self, model: str, messages: List[Dict[str, str]], stream: Literal[False] = False, **kwargs + self, model: str, messages: list[dict[str, str]], stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... @@ -3366,7 +3347,7 @@ class Router: async def abatch_completion_fastest_response( self, model: str, - messages: List[Dict[str, str]], + messages: list[dict[str, str]], stream: bool = False, **kwargs, ): @@ -3378,8 +3359,8 @@ class Router: models = [m.strip() for m in model.split(",")] async def _async_completion_no_exceptions( - model: str, messages: List[Dict[str, str]], stream: bool, **kwargs: Any - ) -> Union[ModelResponse, CustomStreamWrapper, Exception]: + model: str, messages: list[dict[str, str]], stream: bool, **kwargs: Any + ) -> ModelResponse | CustomStreamWrapper | Exception: """ Wrapper around self.acompletion that catches exceptions and returns them as a result """ @@ -3387,7 +3368,7 @@ class Router: result = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) # type: ignore return result except asyncio.CancelledError: - verbose_router_logger.debug("Received 'task.cancel'. Cancelling call w/ model={}.".format(model)) + verbose_router_logger.debug(f"Received 'task.cancel'. Cancelling call w/ model={model}.") raise except Exception as e: return e @@ -3442,13 +3423,13 @@ class Router: @overload async def schedule_acompletion( - self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[False] = False, **kwargs + self, model: str, messages: list[AllMessageValues], priority: int, stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... @overload async def schedule_acompletion( - self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[True], **kwargs + self, model: str, messages: list[AllMessageValues], priority: int, stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... @@ -3457,7 +3438,7 @@ class Router: async def schedule_acompletion( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], priority: int, stream=False, **kwargs, @@ -3519,8 +3500,8 @@ class Router: model: str, priority: int, original_function: Callable, - args: Tuple[Any, ...], - kwargs: Dict[str, Any], + args: tuple[Any, ...], + kwargs: dict[str, Any], ): parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) ### FLOW ITEM ### @@ -3592,8 +3573,8 @@ class Router: async def _prompt_management_factory( self, model: str, - messages: List[AllMessageValues], - kwargs: Dict[str, Any], + messages: list[AllMessageValues], + kwargs: dict[str, Any], ): litellm_logging_object = kwargs.get("litellm_logging_obj", None) if litellm_logging_object is None: @@ -3722,9 +3703,7 @@ class Router: verbose_router_logger.info(f"litellm.image_generation(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info( - f"litellm.image_generation(model={model_name})\033[31m Exception {str(e)}\033[0m" - ) + verbose_router_logger.info(f"litellm.image_generation(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -3809,9 +3788,7 @@ class Router: verbose_router_logger.info(f"litellm.aimage_generation(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info( - f"litellm.aimage_generation(model={model_name})\033[31m Exception {str(e)}\033[0m" - ) + verbose_router_logger.info(f"litellm.aimage_generation(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -3915,7 +3892,7 @@ class Router: verbose_router_logger.info(f"litellm.atranscription(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.atranscription(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.atranscription(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -4029,7 +4006,7 @@ class Router: verbose_router_logger.info(f"litellm.aspeech(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.aspeech(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.aspeech(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -4087,7 +4064,7 @@ class Router: verbose_router_logger.info(f"litellm.arerank(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.arerank(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.arerank(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -4096,9 +4073,9 @@ class Router: self, model: str, prompt: str, - is_retry: Optional[bool] = False, - is_fallback: Optional[bool] = False, - is_async: Optional[bool] = False, + is_retry: bool | None = False, + is_fallback: bool | None = False, + is_async: bool | None = False, **kwargs, ): messages = [{"role": "user", "content": prompt}] @@ -4131,9 +4108,9 @@ class Router: self, model: str, prompt: str, - is_retry: Optional[bool] = False, - is_fallback: Optional[bool] = False, - is_async: Optional[bool] = False, + is_retry: bool | None = False, + is_fallback: bool | None = False, + is_async: bool | None = False, **kwargs, ): if kwargs.get("priority", None) is not None: @@ -4221,7 +4198,7 @@ class Router: verbose_router_logger.info(f"litellm.atext_completion(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.atext_completion(model={model})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.atext_completion(model={model})\033[31m Exception {e!s}\033[0m") if model is not None: self.fail_calls[model] += 1 raise e @@ -4230,9 +4207,9 @@ class Router: self, adapter_id: str, model: str, - is_retry: Optional[bool] = False, - is_fallback: Optional[bool] = False, - is_async: Optional[bool] = False, + is_retry: bool | None = False, + is_fallback: bool | None = False, + is_async: bool | None = False, **kwargs, ): try: @@ -4312,7 +4289,7 @@ class Router: verbose_router_logger.info(f"litellm.aadapter_completion(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.aadapter_completion(model={model})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.aadapter_completion(model={model})\033[31m Exception {e!s}\033[0m") if model is not None: self.fail_calls[model] += 1 raise e @@ -4461,8 +4438,8 @@ class Router: raise e def _add_deployment_model_to_endpoint_for_llm_passthrough_route( - self, kwargs: Dict[str, Any], model: str, model_name: str - ) -> Dict[str, Any]: + self, kwargs: dict[str, Any], model: str, model_name: str + ) -> dict[str, Any]: """ Add the deployment model to the endpoint for LLM passthrough route. @@ -4572,7 +4549,7 @@ class Router: return response except Exception as e: verbose_router_logger.info( - f"ageneric_api_call_with_fallbacks(model={model})\033[31m Exception {str(e)}\033[0m" + f"ageneric_api_call_with_fallbacks(model={model})\033[31m Exception {e!s}\033[0m" ) if model is not None: self.fail_calls[model] += 1 @@ -4610,7 +4587,7 @@ class Router: # fallback to the original reference for any non-picklable value. # The original_generic_function is preserved so the per-attempt # helper knows which underlying API to call on fallback. - fallback_kwargs: Dict[str, Any] = kwargs.copy() + fallback_kwargs: dict[str, Any] = kwargs.copy() if isinstance(fallback_kwargs.get("litellm_metadata"), dict): fallback_kwargs["litellm_metadata"] = safe_deep_copy(fallback_kwargs["litellm_metadata"]) if isinstance(fallback_kwargs.get("metadata"), dict): @@ -4693,7 +4670,7 @@ class Router: verbose_router_logger.info(f"{handler_name}(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"{handler_name}(model={model})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"{handler_name}(model={model})\033[31m Exception {e!s}\033[0m") if model is not None: self.fail_calls[model] += 1 raise e @@ -4701,8 +4678,8 @@ class Router: def embedding( self, model: str, - input: Union[str, List], - is_async: Optional[bool] = False, + input: str | list, + is_async: bool | None = False, **kwargs, ) -> EmbeddingResponse: try: @@ -4715,7 +4692,7 @@ class Router: except Exception as e: raise e - def _embedding(self, input: Union[str, List], model: str, **kwargs): + def _embedding(self, input: str | list, model: str, **kwargs): model_name = None try: verbose_router_logger.debug(f"Inside embedding()- model: {model}; kwargs: {kwargs}") @@ -4758,7 +4735,7 @@ class Router: verbose_router_logger.info(f"litellm.embedding(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.embedding(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.embedding(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -4766,8 +4743,8 @@ class Router: async def aembedding( self, model: str, - input: Union[str, List], - is_async: Optional[bool] = True, + input: str | list, + is_async: bool | None = True, **kwargs, ) -> EmbeddingResponse: try: @@ -4788,7 +4765,7 @@ class Router: ) raise e - async def _aembedding(self, input: Union[str, List], model: str, **kwargs): + async def _aembedding(self, input: str | list, model: str, **kwargs): model_name = None try: verbose_router_logger.debug(f"Inside _aembedding()- model: {model}; kwargs: {kwargs}") @@ -4845,7 +4822,7 @@ class Router: verbose_router_logger.info(f"litellm.aembedding(model={model_name})\033[32m 200 OK\033[0m") return response except Exception as e: - verbose_router_logger.info(f"litellm.aembedding(model={model_name})\033[31m Exception {str(e)}\033[0m") + verbose_router_logger.info(f"litellm.aembedding(model={model_name})\033[31m Exception {e!s}\033[0m") if model_name is not None: self.fail_calls[model_name] += 1 raise e @@ -4922,8 +4899,8 @@ class Router: custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider ## REPLACE MODEL IN FILE WITH SELECTED DEPLOYMENT ## - purpose = cast(Optional[OpenAIFilesPurpose], kwargs.get("purpose")) - file = cast(Optional[FileTypes], kwargs.get("file")) + purpose = cast(OpenAIFilesPurpose | None, kwargs.get("purpose")) + file = cast(FileTypes | None, kwargs.get("file")) if not file or not purpose: raise Exception("file and file_purpose are required for create_file") @@ -4999,7 +4976,7 @@ class Router: return returned_response except Exception as e: verbose_router_logger.exception( - f"litellm.acreate_file(model={model}, {kwargs})\033[31m Exception {str(e)}\033[0m" + f"litellm.acreate_file(model={model}, {kwargs})\033[31m Exception {e!s}\033[0m" ) if model is not None: self.fail_calls[model] += 1 @@ -5008,7 +4985,7 @@ class Router: #### VECTOR STORES API #### async def avector_store_create( self, - model: Union[str, None], + model: str | None, **kwargs, ): """ @@ -5095,7 +5072,7 @@ class Router: return response except Exception as e: verbose_router_logger.exception( - f"litellm.avector_store_create(model={model})\033[31m Exception {str(e)}\033[0m" + f"litellm.avector_store_create(model={model})\033[31m Exception {e!s}\033[0m" ) if model is not None: self.fail_calls[model] += 1 @@ -5110,7 +5087,7 @@ class Router: """ # Store references to the custom methods defined above # These methods handle proper routing through deployments - pass # The methods are already defined as instance methods above + # The methods are already defined as instance methods above async def acreate_batch( self, @@ -5212,7 +5189,7 @@ class Router: return response # type: ignore except Exception as e: verbose_router_logger.exception( - f"litellm._acreate_batch(model={model}, {kwargs})\033[31m Exception {str(e)}\033[0m" + f"litellm._acreate_batch(model={model}, {kwargs})\033[31m Exception {e!s}\033[0m" ) if model is not None: self.fail_calls[model] += 1 @@ -5220,7 +5197,7 @@ class Router: async def aretrieve_batch( self, - model: Optional[str] = None, + model: str | None = None, **kwargs, ) -> LiteLLMBatch: """ @@ -5231,9 +5208,9 @@ class Router: try: parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) if model is not None: - filtered_model_list: Optional[ - Union[List[DeploymentTypedDict], List[Dict], Dict] - ] = await self.async_get_healthy_deployments( + filtered_model_list: ( + list[DeploymentTypedDict] | list[dict] | dict | None + ) = await self.async_get_healthy_deployments( model=model, messages=[{"role": "user", "content": "retrieve-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), @@ -5313,7 +5290,7 @@ class Router: raise receieved_exceptions[0] # Raising the first exception encountered # If no exceptions were encountered, raise a generic exception - raise Exception("Unable to find batch in any model. Received errors - {}".format(receieved_exceptions)) + raise Exception(f"Unable to find batch in any model. Received errors - {receieved_exceptions}") except Exception as e: asyncio.create_task( send_llm_exception_alert( @@ -5435,7 +5412,7 @@ class Router: return response # type: ignore except Exception as e: verbose_router_logger.exception( - f"litellm._acancel_batch(model={model}, {kwargs})\033[31m Exception {str(e)}\033[0m" + f"litellm._acancel_batch(model={model}, {kwargs})\033[31m Exception {e!s}\033[0m" ) if model is not None: self.fail_calls[model] += 1 @@ -5464,7 +5441,7 @@ class Router: # Check all models in parallel results = await asyncio.gather(*[try_retrieve_batch(model) for model in filtered_model_list]) - final_results: Dict = { + final_results: dict = { "object": "list", "data": [], "first_id": None, @@ -5491,7 +5468,7 @@ class Router: async def _pass_through_moderation_endpoint_factory( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): # update kwargs with model_group @@ -5661,8 +5638,8 @@ class Router: ): def sync_wrapper( - custom_llm_provider: Optional[str] = None, - client: Optional[Any] = None, + custom_llm_provider: str | None = None, + client: Any | None = None, **kwargs, ): return self._generic_api_call_with_fallbacks(original_function=original_function, **kwargs) @@ -5677,8 +5654,8 @@ class Router: ): def vector_store_sync_wrapper( - custom_llm_provider: Optional[str] = None, - client: Optional[Any] = None, + custom_llm_provider: str | None = None, + client: Any | None = None, **kwargs, ): if custom_llm_provider and "custom_llm_provider" not in kwargs: @@ -5699,8 +5676,8 @@ class Router: ): def vector_store_file_sync_wrapper( - custom_llm_provider: Optional[str] = None, - client: Optional[Any] = None, + custom_llm_provider: str | None = None, + client: Any | None = None, **kwargs, ): return original_function( @@ -5720,8 +5697,8 @@ class Router: ): def managed_agents_sync_wrapper( - custom_llm_provider: Optional[str] = None, - client: Optional[Any] = None, + custom_llm_provider: str | None = None, + client: Any | None = None, **kwargs, ): if custom_llm_provider and "custom_llm_provider" not in kwargs: @@ -5734,8 +5711,8 @@ class Router: # Handle asynchronous call types async def async_wrapper( - custom_llm_provider: Optional[str] = None, - client: Optional[Any] = None, + custom_llm_provider: str | None = None, + client: Any | None = None, **kwargs, ): if call_type == "assistants": @@ -5897,7 +5874,7 @@ class Router: async def _init_vector_store_api_endpoints( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -5922,7 +5899,7 @@ class Router: async def _init_containers_api_endpoints( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -5989,7 +5966,7 @@ class Router: async def _init_interactions_api_endpoints( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -6018,7 +5995,7 @@ class Router: async def _init_managed_agents_api_endpoints( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -6036,8 +6013,8 @@ class Router: async def _pass_through_assistants_endpoint_factory( self, original_function: Callable, - custom_llm_provider: Optional[str] = None, - client: Optional[AsyncOpenAI] = None, + custom_llm_provider: str | None = None, + client: AsyncOpenAI | None = None, **kwargs, ): """Internal helper function to pass through the assistants endpoint""" @@ -6059,17 +6036,17 @@ class Router: self, exception: Exception, original_model_group: str, - all_deployments: List[DeploymentTypedDict], + all_deployments: list[DeploymentTypedDict], args: tuple, kwargs: dict, input_kwargs: dict, - ) -> Optional[Any]: + ) -> Any | None: """Same-model-group retry after a failed deployment; returns None if not applicable.""" strategy, _ = self._get_routing_context(original_model_group, kwargs) if strategy != "simple-shuffle": return None - failed_id: Optional[str] = getattr(exception, "failed_deployment_id", None) + failed_id: str | None = getattr(exception, "failed_deployment_id", None) if not failed_id: return None @@ -6133,11 +6110,11 @@ class Router: async def async_function_with_fallbacks_common_utils( self, e: Exception, - disable_fallbacks: Optional[bool], - fallbacks: Optional[List], - context_window_fallbacks: Optional[List], - content_policy_fallbacks: Optional[List], - model_group: Optional[str], + disable_fallbacks: bool | None, + fallbacks: list | None, + context_window_fallbacks: list | None, + content_policy_fallbacks: list | None, + model_group: str | None, args: tuple, kwargs: dict, include_fallback_errors: bool = False, @@ -6148,7 +6125,7 @@ class Router: verbose_router_logger.debug(f"Traceback{traceback.format_exc()}") original_exception = e fallback_model_group = None - original_model_group: Optional[str] = kwargs.get("model") # type: ignore + original_model_group: str | None = kwargs.get("model") # type: ignore fallback_failure_exception_str = "" if disable_fallbacks is True or original_model_group is None: @@ -6173,7 +6150,7 @@ class Router: e, (litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError), ) - _request_team_id: Optional[str] = (kwargs.get("metadata", {}) or {}).get("user_api_key_team_id") + _request_team_id: str | None = (kwargs.get("metadata", {}) or {}).get("user_api_key_team_id") # Use wildcard-aware lookup so order-based fallback also works for model # groups resolved via pattern routing (e.g. `openai/*` -> `openai/gpt-4.1-mini`). all_deployments = self.get_model_list(model_name=original_model_group, team_id=_request_team_id) or [] @@ -6188,11 +6165,11 @@ class Router: current_target = kwargs.get("_target_order") skip_up_to = current_target if current_target is not None else order_values[0] # Build order-based fallback entries (skip already-tried levels) - order_fallback_entries: List = [ + order_fallback_entries: list = [ {"model": original_model_group, "_target_order": o} for o in order_values if o > skip_up_to ] # Get external fallbacks — handle both standard and non-standard formats - external_fallback_group: Optional[List] = None + external_fallback_group: list | None = None if fallbacks is not None and model_group is not None: if _check_non_standard_fallback_format(fallbacks=fallbacks): # Non-standard formats (e.g. ["claude-3-haiku"] or @@ -6258,7 +6235,7 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - context_window_fallback_model_group: Optional[List[str]] = ( + context_window_fallback_model_group: list[str] | None = ( self._get_fallback_model_group_from_fallbacks( fallbacks=context_window_fallbacks, model_group=model_group, @@ -6281,21 +6258,17 @@ class Router: return response else: - error_message = "model={}. context_window_fallbacks={}. fallbacks={}.\n\nSet 'context_window_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( - model_group, - mask_sensitive_structure(context_window_fallbacks), - mask_sensitive_structure(fallbacks), - ) + error_message = f"model={model_group}. context_window_fallbacks={mask_sensitive_structure(context_window_fallbacks)}. fallbacks={mask_sensitive_structure(fallbacks)}.\n\nSet 'context_window_fallback' - https://docs.litellm.ai/docs/routing#fallbacks" verbose_router_logger.info( - msg="Got 'ContextWindowExceededError'. No context_window_fallback set. Defaulting \ - to fallbacks, if available.{}".format(error_message) + msg=f"Got 'ContextWindowExceededError'. No context_window_fallback set. Defaulting \ + to fallbacks, if available.{error_message}" ) if litellm.expose_router_debug_in_errors: - e.message += "\n{}".format(error_message) + e.message += f"\n{error_message}" elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - content_policy_fallback_model_group: Optional[List[str]] = ( + content_policy_fallback_model_group: list[str] | None = ( self._get_fallback_model_group_from_fallbacks( fallbacks=content_policy_fallbacks, model_group=model_group, @@ -6317,18 +6290,14 @@ class Router: ) return response else: - error_message = "model={}. content_policy_fallback={}. fallbacks={}.\n\nSet 'content_policy_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( - model_group, - mask_sensitive_structure(content_policy_fallbacks), - mask_sensitive_structure(fallbacks), - ) + error_message = f"model={model_group}. content_policy_fallback={mask_sensitive_structure(content_policy_fallbacks)}. fallbacks={mask_sensitive_structure(fallbacks)}.\n\nSet 'content_policy_fallback' - https://docs.litellm.ai/docs/routing#fallbacks" verbose_router_logger.info( - msg="Got 'ContentPolicyViolationError'. No content_policy_fallback set. Defaulting \ - to fallbacks, if available.{}".format(error_message) + msg=f"Got 'ContentPolicyViolationError'. No content_policy_fallback set. Defaulting \ + to fallbacks, if available.{error_message}" ) if litellm.expose_router_debug_in_errors: - e.message += "\n{}".format(error_message) + e.message += f"\n{error_message}" if fallbacks is not None and model_group is not None: verbose_router_logger.debug(f"inside model fallbacks: {mask_sensitive_structure(fallbacks)}") ( @@ -6386,7 +6355,7 @@ class Router: ) if len(fallback_failure_exception_str) > 0: original_exception.message += ( # type: ignore - "\nError doing the fallback: {}".format(fallback_failure_exception_str) + f"\nError doing the fallback: {fallback_failure_exception_str}" ) raise original_exception @@ -6397,12 +6366,12 @@ class Router: Try calling the function_with_retries If it fails after num_retries, fall back to another model group """ - model_group: Optional[str] = kwargs.get("model") + model_group: str | None = kwargs.get("model") include_fallback_errors = kwargs.get("include_fallback_errors", False) is True - disable_fallbacks: Optional[bool] = kwargs.pop("disable_fallbacks", False) - fallbacks: Optional[List] = kwargs.get("fallbacks", self.fallbacks) - context_window_fallbacks: Optional[List] = kwargs.get("context_window_fallbacks", self.context_window_fallbacks) - content_policy_fallbacks: Optional[List] = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) + disable_fallbacks: bool | None = kwargs.pop("disable_fallbacks", False) + fallbacks: list | None = kwargs.get("fallbacks", self.fallbacks) + context_window_fallbacks: list | None = kwargs.get("context_window_fallbacks", self.context_window_fallbacks) + content_policy_fallbacks: list | None = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) mock_timeout = kwargs.pop("mock_timeout", None) @@ -6442,10 +6411,10 @@ class Router: def _handle_mock_testing_fallbacks( self, kwargs: dict, - model_group: Optional[str] = None, - fallbacks: Optional[List] = None, - context_window_fallbacks: Optional[List] = None, - content_policy_fallbacks: Optional[List] = None, + model_group: str | None = None, + fallbacks: list | None = None, + context_window_fallbacks: list | None = None, + content_policy_fallbacks: list | None = None, ): """ Helper function to raise a litellm Error for mock testing purposes. @@ -6496,7 +6465,7 @@ class Router: content_policy_fallbacks = kwargs.pop("content_policy_fallbacks", self.content_policy_fallbacks) # Support per-request model_group_retry_policy override (from key/team settings) model_group_retry_policy = kwargs.pop("model_group_retry_policy", self.model_group_retry_policy) - model_group: Optional[str] = kwargs.get("model") + model_group: str | None = kwargs.get("model") num_retries = kwargs.pop("num_retries", None) if num_retries is None: # Fall back to the router setting (then 0) so the comparisons below never @@ -6616,7 +6585,7 @@ class Router: ## LOGGING kwargs = self.log_retry(kwargs=kwargs, e=e) remaining_retries = num_retries - current_attempt - 1 - _model: Optional[str] = kwargs.get("model") # type: ignore + _model: str | None = kwargs.get("model") # type: ignore if _model is not None: ( _healthy_deployments, @@ -6679,20 +6648,20 @@ class Router: return response - def _handle_mock_testing_rate_limit_error(self, kwargs: dict, model_group: Optional[str] = None): + def _handle_mock_testing_rate_limit_error(self, kwargs: dict, model_group: str | None = None): """ Helper function to raise a mock litellm.RateLimitError error for testing purposes. Raises: litellm.RateLimitError error when `mock_testing_rate_limit_error=True` passed in request params """ - mock_testing_rate_limit_error: Optional[bool] = kwargs.pop("mock_testing_rate_limit_error", None) + mock_testing_rate_limit_error: bool | None = kwargs.pop("mock_testing_rate_limit_error", None) available_models = self.get_model_list(model_name=model_group) - num_retries: Optional[int] = None + num_retries: int | None = None if available_models is not None and len(available_models) == 1: - num_retries = cast(Optional[int], available_models[0]["litellm_params"].get("num_retries")) + num_retries = cast(int | None, available_models[0]["litellm_params"].get("num_retries")) if mock_testing_rate_limit_error is not None and mock_testing_rate_limit_error is True: verbose_router_logger.info( @@ -6708,11 +6677,11 @@ class Router: def should_retry_this_error( self, error: Exception, - healthy_deployments: Optional[List] = None, - all_deployments: Optional[List] = None, - context_window_fallbacks: Optional[List] = None, - content_policy_fallbacks: Optional[List] = None, - regular_fallbacks: Optional[List] = None, + healthy_deployments: list | None = None, + all_deployments: list | None = None, + context_window_fallbacks: list | None = None, + content_policy_fallbacks: list | None = None, + regular_fallbacks: list | None = None, ): """ 1. raise an exception for ContextWindowExceededError if context_window_fallbacks is not None @@ -6779,9 +6748,9 @@ class Router: def _get_fallback_model_group_from_fallbacks( self, - fallbacks: List[Dict[str, List[str]]], - model_group: Optional[str] = None, - ) -> Optional[List[str]]: + fallbacks: list[dict[str, list[str]]], + model_group: str | None = None, + ) -> list[str] | None: """ Returns the list of fallback models to use for a given model group @@ -6795,14 +6764,14 @@ class Router: if model_group is None: return None - fallback_model_group: Optional[List[str]] = None + fallback_model_group: list[str] | None = None for item in fallbacks: # [{"gpt-3.5-turbo": ["gpt-4"]}] if list(item.keys())[0] == model_group: fallback_model_group = item[model_group] break return fallback_model_group - def _get_first_default_fallback(self) -> Optional[str]: + def _get_first_default_fallback(self) -> str | None: """ Returns the first model from the default_fallbacks list, if it exists. """ @@ -6820,9 +6789,9 @@ class Router: e: Exception, remaining_retries: int, num_retries: int, - healthy_deployments: Optional[List] = None, - all_deployments: Optional[List] = None, - ) -> Union[int, float]: + healthy_deployments: list | None = None, + all_deployments: list | None = None, + ) -> int | float: """ Calculate back-off, then retry @@ -6837,7 +6806,7 @@ class Router: elif healthy_deployments is not None and isinstance(healthy_deployments, list) and len(healthy_deployments) > 0: return 0 - response_headers: Optional[httpx.Headers] = None + response_headers: httpx.Headers | None = None if hasattr(e, "response") and hasattr(e.response, "headers"): # type: ignore response_headers = e.response.headers # type: ignore if hasattr(e, "litellm_response_headers"): @@ -6878,7 +6847,7 @@ class Router: # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): return - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: raise ValueError("standard_logging_object is None") if kwargs["litellm_params"].get("metadata") is None: @@ -6956,7 +6925,7 @@ class Router: # Update usage # ------------ # update cache - pipeline_operations: List[RedisPipelineIncrementOperation] = [] + pipeline_operations: list[RedisPipelineIncrementOperation] = [] ## TPM pipeline_operations.append( @@ -6986,9 +6955,8 @@ class Router: except Exception as e: verbose_router_logger.debug( - "litellm.router.Router::deployment_callback_on_success(): Exception occured - {}".format(str(e)) + f"litellm.router.Router::deployment_callback_on_success(): Exception occured - {e!s}" ) - pass def sync_deployment_callback_on_success( self, @@ -6996,7 +6964,7 @@ class Router: completion_response, # response from completion start_time, end_time, # start/end time - ) -> Optional[str]: + ) -> str | None: """ Tracks the number of successes for a deployment in the current minute (using in-memory cache) @@ -7084,7 +7052,7 @@ class Router: _time_to_cooldown = self.cooldown_time if isinstance(_model_info, dict): - deployment_id: Optional[str] = _model_info.get("id") + deployment_id: str | None = _model_info.get("id") if deployment_id is None: return False increment_deployment_failures_for_current_minute( @@ -7109,9 +7077,7 @@ class Router: except Exception as e: raise e - async def async_deployment_callback_on_failure( - self, kwargs, completion_response: Optional[Any], start_time, end_time - ): + async def async_deployment_callback_on_failure(self, kwargs, completion_response: Any | None, start_time, end_time): """ Update RPM usage for a deployment """ @@ -7186,7 +7152,7 @@ class Router: except Exception as e: raise e - def _update_usage(self, deployment_id: str, parent_otel_span: Optional[Span]) -> int: + def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int: """ Update deployment rpm for that minute @@ -7242,13 +7208,11 @@ class Router: return True verbose_router_logger.debug( - "Content Policy Error occurred. No available fallbacks. Returning original response. model={}, content_policy_fallbacks={}".format( - model, content_policy_fallbacks - ) + f"Content Policy Error occurred. No available fallbacks. Returning original response. model={model}, content_policy_fallbacks={content_policy_fallbacks}" ) return False - def _get_healthy_deployments(self, model: str, parent_otel_span: Optional[Span]): + def _get_healthy_deployments(self, model: str, parent_otel_span: Span | None): _all_deployments: list = [] try: _, _all_deployments = self._common_checks_available_deployment( # type: ignore @@ -7269,8 +7233,8 @@ class Router: return healthy_deployments, _all_deployments async def _async_get_healthy_deployments( - self, model: str, parent_otel_span: Optional[Span] - ) -> Tuple[List[Dict], List[Dict]]: + self, model: str, parent_otel_span: Span | None + ) -> tuple[list[dict], list[dict]]: """ Returns Tuple of: - Tuple[List[Dict], List[Dict]]: @@ -7317,8 +7281,8 @@ class Router: async def async_routing_strategy_pre_call_checks( self, deployment: dict, - parent_otel_span: Optional[Span], - logging_obj: Optional[LiteLLMLogging] = None, + parent_otel_span: Span | None, + logging_obj: LiteLLMLogging | None = None, ): """ For usage-based-routing-v2, enables running rpm checks before the call is made, inside the semaphore. @@ -7378,11 +7342,11 @@ class Router: async def async_callback_filter_deployments( self, model: str, - healthy_deployments: List[dict], - messages: Optional[List[AllMessageValues]], - parent_otel_span: Optional[Span], - request_kwargs: Optional[dict] = None, - logging_obj: Optional[LiteLLMLogging] = None, + healthy_deployments: list[dict], + messages: list[AllMessageValues] | None, + parent_otel_span: Span | None, + request_kwargs: dict | None = None, + logging_obj: LiteLLMLogging | None = None, ): """ For usage-based-routing-v2, enables running rpm checks before the call is made, inside the semaphore. @@ -7464,9 +7428,7 @@ class Router: return hash_object.hexdigest() @staticmethod - def _inherit_builtin_cache_pricing( - model_info: dict, backend_model: str, custom_llm_provider: Optional[str] - ) -> None: + def _inherit_builtin_cache_pricing(model_info: dict, backend_model: str, custom_llm_provider: str | None) -> None: """Fill missing cache pricing on a custom-priced deployment entry from the backend model's built-in cost map entry, so a deployment that only spells out ``input_cost_per_token``/``output_cost_per_token`` @@ -7500,7 +7462,7 @@ class Router: _model_name: str, _litellm_params: dict, _model_info: dict, - ) -> Optional[Deployment]: + ) -> Deployment | None: """ Create a deployment object and add it to the model list @@ -7550,7 +7512,7 @@ class Router: # name. Each deployment's full model_info is already stored under # its unique model_id above. _shared_model_info = shared_backend_model_info(_model_info) - _existing_shared_mode = (cast(Optional[dict], litellm.model_cost.get(_model_name, {})) or {}).get("mode") + _existing_shared_mode = (cast(dict | None, litellm.model_cost.get(_model_name, {})) or {}).get("mode") _deployment_mode = _shared_model_info.get("mode") # Keep the built-in bridge mode stable for shared backend keys. # Multiple aliases can point at the same provider/model backend, @@ -7638,20 +7600,20 @@ class Router: """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = deployment.litellm_params.auto_router_config_path - auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config + auto_router_config_path: str | None = deployment.litellm_params.auto_router_config_path + auto_router_config: str | None = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) - default_model: Optional[str] = deployment.litellm_params.auto_router_default_model + default_model: str | None = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) - embedding_model: Optional[str] = deployment.litellm_params.auto_router_embedding_model + embedding_model: str | None = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" @@ -7693,9 +7655,9 @@ class Router: ComplexityRouter, ) - complexity_router_config: Optional[dict] = deployment.litellm_params.complexity_router_config + complexity_router_config: dict | None = deployment.litellm_params.complexity_router_config - default_model: Optional[str] = deployment.litellm_params.complexity_router_default_model + default_model: str | None = deployment.litellm_params.complexity_router_default_model # If no default model specified, try to get from config tiers if default_model is None and complexity_router_config: @@ -7905,8 +7867,8 @@ class Router: config = AdaptiveRouterConfig(**raw_config) - model_to_prefs: Dict[str, AdaptiveRouterPreferences] = {} - model_to_cost: Dict[str, float] = {} + model_to_prefs: dict[str, AdaptiveRouterPreferences] = {} + model_to_cost: dict[str, float] = {} # O(k) via the name→indices map: only touch deployments whose name # is listed in `available_models`, instead of scanning model_list. for name in config.available_models: @@ -7915,14 +7877,14 @@ class Router: continue d = (self.model_list or [])[indices[0]] mi = d.get("model_info") if isinstance(d, dict) else d.model_info - mi_dict: Dict[str, Any] = mi if isinstance(mi, dict) else (mi.model_dump() if mi else {}) + mi_dict: dict[str, Any] = mi if isinstance(mi, dict) else (mi.model_dump() if mi else {}) prefs_raw = mi_dict.get("adaptive_router_preferences") if prefs_raw is not None: model_to_prefs[name] = AdaptiveRouterPreferences(**prefs_raw) # `input_cost_per_token` is a LiteLLM_Params field per types/router.py. lp = d.get("litellm_params") if isinstance(d, dict) else d.litellm_params - lp_dict: Dict[str, Any] = lp if isinstance(lp, dict) else (lp.model_dump() if lp else {}) + lp_dict: dict[str, Any] = lp if isinstance(lp, dict) else (lp.model_dump() if lp else {}) cost = lp_dict.get("input_cost_per_token") if cost is not None: model_to_cost[name] = float(cost) @@ -7970,9 +7932,9 @@ class Router: QualityRouter, ) - quality_router_config: Optional[dict] = deployment.litellm_params.quality_router_config + quality_router_config: dict | None = deployment.litellm_params.quality_router_config - default_model: Optional[str] = deployment.litellm_params.quality_router_default_model + default_model: str | None = deployment.litellm_params.quality_router_default_model if default_model is None and quality_router_config: default_model = quality_router_config.get("default_model") @@ -8243,10 +8205,8 @@ class Router: api_base=api_base, api_key=api_key, ) - pass - pass - def add_deployment(self, deployment: Deployment) -> Optional[Deployment]: + def add_deployment(self, deployment: Deployment) -> Deployment | None: """ Parameters: - deployment: Deployment - the deployment to be added to the Router @@ -8398,7 +8358,7 @@ class Router: if idx not in self.team_model_to_deployment_indices[key]: self.team_model_to_deployment_indices[key].append(idx) - def _add_model_to_list_and_index_map(self, model: dict, model_id: Optional[str] = None) -> None: + def _add_model_to_list_and_index_map(self, model: dict, model_id: str | None = None) -> None: """ Helper method to add a model to the model_list and update both indices. @@ -8427,7 +8387,7 @@ class Router: # Update team_model index for O(1) team-scoped lookup self._update_team_model_index(model, idx) - def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]: + def upsert_deployment(self, deployment: Deployment) -> Deployment | None: """ Add or update deployment Parameters: @@ -8440,7 +8400,7 @@ class Router: # check if deployment already exists _deployment_model_id = deployment.model_info.id or "" - _deployment_on_router: Optional[Deployment] = self.get_deployment(model_id=_deployment_model_id) + _deployment_on_router: Deployment | None = self.get_deployment(model_id=_deployment_model_id) if _deployment_on_router is not None: # deployment with this model_id exists on the router if ( @@ -8452,7 +8412,7 @@ class Router: # if there is a new litellm param -> then update the deployment # remove the previous deployment - removal_idx: Optional[int] = None + removal_idx: int | None = None deployment_id = deployment.model_info.id deployment_fast_mapping = self.model_id_to_deployment_index_map @@ -8493,7 +8453,7 @@ class Router: else: raise e - def delete_deployment(self, id: str) -> Optional[Deployment]: + def delete_deployment(self, id: str) -> Deployment | None: """ Parameters: - id: str - the id of the deployment to be deleted @@ -8534,7 +8494,7 @@ class Router: def _get_router_deployment_budget_limiter( self, - ) -> Optional[RouterBudgetLimiting]: + ) -> RouterBudgetLimiting | None: """ Return the router's deployment-budget callback. @@ -8577,7 +8537,7 @@ class Router: if _budget_limiter is not None: _budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True)) - def get_deployment(self, model_id: str) -> Optional[Deployment]: + def get_deployment(self, model_id: str) -> Deployment | None: """ Returns -> Deployment or None @@ -8592,11 +8552,11 @@ class Router: elif isinstance(model, Deployment): return model else: - raise Exception("Model invalid format - {}".format(type(model))) + raise Exception(f"Model invalid format - {type(model)}") return None - def get_deployment_credentials(self, model_id: str) -> Optional[dict]: + def get_deployment_credentials(self, model_id: str) -> dict | None: """ Returns -> dict of credentials for a given model id. @@ -8611,7 +8571,7 @@ class Router: deployment.litellm_params.model_dump(exclude_none=True) ).model_dump(exclude_none=True) - def get_deployment_by_model_group_name(self, model_group_name: str) -> Optional[Deployment]: + def get_deployment_by_model_group_name(self, model_group_name: str) -> Deployment | None: """ Returns -> Deployment or None @@ -8630,11 +8590,11 @@ class Router: elif isinstance(model, Deployment): return model else: - raise Exception("Model Name invalid - {}".format(type(model))) + raise Exception(f"Model Name invalid - {type(model)}") return None @staticmethod - def _deployment_usable_by_team(model: Union[Mapping, Deployment], team_id: str | None) -> bool: + def _deployment_usable_by_team(model: Mapping | Deployment, team_id: str | None) -> bool: """ A team-scoped deployment (``model_info.team_id`` set) is only usable by callers from that same team; deployments without a team owner are shared. @@ -8792,9 +8752,9 @@ class Router: def get_router_model_info( self, - deployment: Optional[Union[dict, "Deployment"]], + deployment: Union[dict, "Deployment"] | None, received_model_name: str, - id: Optional[str] = None, + id: str | None = None, ) -> ModelMapInfo: """ For a given model id, return the model info (max tokens, input cost, output cost, etc.). @@ -8869,8 +8829,8 @@ class Router: # Use the original model from litellm_params model = _model - if not model.startswith("{}/".format(custom_llm_provider)): - model_info_name = "{}/{}".format(custom_llm_provider, model) + if not model.startswith(f"{custom_llm_provider}/"): + model_info_name = f"{custom_llm_provider}/{model}" else: model_info_name = model @@ -8884,7 +8844,7 @@ class Router: return model_info - def get_model_info(self, id: str) -> Optional[dict]: + def get_model_info(self, id: str) -> dict | None: """ For a given model id, return the model info @@ -8900,7 +8860,7 @@ class Router: return self.model_list[idx] return None - def get_model_group(self, id: str) -> Optional[List]: + def get_model_group(self, id: str) -> list | None: """ Return list of all models in the same model group as that model id """ @@ -8912,7 +8872,7 @@ class Router: model_name = model_info["model_name"] return self.get_model_list(model_name=model_name) - def get_deployment_model_info(self, model_id: str, model_name: str) -> Optional[ModelInfo]: + def get_deployment_model_info(self, model_id: str, model_name: str) -> ModelInfo | None: """ For a given model id, return the model info @@ -8922,9 +8882,9 @@ class Router: """ from litellm.utils import _update_dictionary - model_info: Optional[ModelInfo] = None - custom_model_info: Optional[dict] = None - litellm_model_name_model_info: Optional[ModelInfo] = None + model_info: ModelInfo | None = None + custom_model_info: dict | None = None + litellm_model_name_model_info: ModelInfo | None = None try: custom_model_info = litellm.model_cost.get(model_id) @@ -8974,7 +8934,7 @@ class Router: return model_info - def _set_model_group_info(self, model_group: str, user_facing_model_group_name: str) -> Optional[ModelGroupInfo]: + def _set_model_group_info(self, model_group: str, user_facing_model_group_name: str) -> ModelGroupInfo | None: """ For a given model group name, return the combined model info @@ -8982,21 +8942,24 @@ class Router: - ModelGroupInfo if able to construct a model group - None if error constructing model group info """ - model_group_info: Optional[ModelGroupInfo] = None + model_group_info: ModelGroupInfo | None = None - total_tpm: Optional[int] = None - total_rpm: Optional[int] = None - total_itpm: Optional[int] = None - total_otpm: Optional[int] = None + total_tpm: int | None = None + total_rpm: int | None = None + total_itpm: int | None = None + total_otpm: int | None = None configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None model_list = self.get_model_list(model_name=model_group) if model_list is None: return None for model in model_list: is_match = False - if "model_name" in model and model["model_name"] == model_group: # exact match - is_match = True - elif "model_name" in model and self.pattern_router.route(model_group) is not None: # wildcard model + if ( + "model_name" in model + and model["model_name"] == model_group + or "model_name" in model + and self.pattern_router.route(model_group) is not None + ): # exact match is_match = True if not is_match: @@ -9011,7 +8974,7 @@ class Router: model_info_dict = model.get("model_info", {}) # get model tpm - _deployment_tpm: Optional[int] = None + _deployment_tpm: int | None = None if _deployment_tpm is None: _deployment_tpm = model.get("tpm", None) # type: ignore if _deployment_tpm is None: @@ -9020,7 +8983,7 @@ class Router: _deployment_tpm = model_info_dict.get("tpm", None) # type: ignore # get model rpm - _deployment_rpm: Optional[int] = None + _deployment_rpm: int | None = None if _deployment_rpm is None: _deployment_rpm = model.get("rpm", None) # type: ignore if _deployment_rpm is None: @@ -9028,13 +8991,13 @@ class Router: if _deployment_rpm is None: _deployment_rpm = model_info_dict.get("rpm", None) # type: ignore - _deployment_itpm: Optional[int] = model.get("itpm") # type: ignore + _deployment_itpm: int | None = model.get("itpm") # type: ignore if _deployment_itpm is None: _deployment_itpm = model_litellm_params.get("itpm", None) # type: ignore if _deployment_itpm is None: _deployment_itpm = model_info_dict.get("itpm", None) # type: ignore - _deployment_otpm: Optional[int] = model.get("otpm") # type: ignore + _deployment_otpm: int | None = model.get("otpm") # type: ignore if _deployment_otpm is None: _deployment_otpm = model_litellm_params.get("otpm", None) # type: ignore if _deployment_otpm is None: @@ -9058,7 +9021,7 @@ class Router: custom_llm_provider=litellm_params.custom_llm_provider, ) except litellm.exceptions.BadRequestError as e: - verbose_router_logger.error("litellm.router.py::get_model_group_info() - {}".format(str(e))) + verbose_router_logger.error(f"litellm.router.py::get_model_group_info() - {e!s}") if model_info is None: supported_openai_params = litellm.get_supported_openai_params( @@ -9212,7 +9175,7 @@ class Router: return model_group_info - def get_model_group_info(self, model_group: str) -> Optional[ModelGroupInfo]: + def get_model_group_info(self, model_group: str) -> ModelGroupInfo | None: """ For a given model group name, return the combined model info @@ -9241,7 +9204,7 @@ class Router: ## Check if actual model return self._set_model_group_info(model_group=model_group, user_facing_model_group_name=model_group) - async def get_model_group_usage(self, model_group: str) -> Tuple[Optional[int], Optional[int]]: + async def get_model_group_usage(self, model_group: str) -> tuple[int | None, int | None]: """ Returns current tpm/rpm usage for model group @@ -9253,16 +9216,16 @@ class Router: """ dt = get_utc_datetime() current_minute = dt.strftime("%H-%M") # use the same timezone regardless of system clock - tpm_keys: List[str] = [] - rpm_keys: List[str] = [] + tpm_keys: list[str] = [] + rpm_keys: list[str] = [] model_list = self.get_model_list(model_name=model_group) if model_list is None: # no matching deployments return None, None for model in model_list: - id: Optional[str] = model.get("model_info", {}).get("id") # type: ignore - litellm_model: Optional[str] = model["litellm_params"].get( + id: str | None = model.get("model_info", {}).get("id") # type: ignore + litellm_model: str | None = model["litellm_params"].get( "model" ) # USE THE MODEL SENT TO litellm.completion() - consistent with how global_router cache is written. if id is None or litellm_model is None: @@ -9287,11 +9250,11 @@ class Router: if combined_tpm_rpm_values is None: return None, None - tpm_usage_list: Optional[List] = combined_tpm_rpm_values[: len(tpm_keys)] - rpm_usage_list: Optional[List] = combined_tpm_rpm_values[len(tpm_keys) :] + tpm_usage_list: list | None = combined_tpm_rpm_values[: len(tpm_keys)] + rpm_usage_list: list | None = combined_tpm_rpm_values[len(tpm_keys) :] ## TPM - tpm_usage: Optional[int] = None + tpm_usage: int | None = None if tpm_usage_list is not None: for t in tpm_usage_list: if isinstance(t, int): @@ -9299,7 +9262,7 @@ class Router: tpm_usage = 0 tpm_usage += t ## RPM - rpm_usage: Optional[int] = None + rpm_usage: int | None = None if rpm_usage_list is not None: for t in rpm_usage_list: if isinstance(t, int): @@ -9308,7 +9271,7 @@ class Router: rpm_usage += t return tpm_usage, rpm_usage - async def get_model_group_io_token_usage(self, model_group: str) -> tuple[Optional[int], Optional[int]]: + async def get_model_group_io_token_usage(self, model_group: str) -> tuple[int | None, int | None]: """ Returns current ITPM/OTPM usage for a model group (sum across deployments). """ @@ -9322,8 +9285,8 @@ class Router: return None, None for model in model_list: - model_id: Optional[str] = model.get("model_info", {}).get("id") # type: ignore - litellm_model: Optional[str] = model["litellm_params"].get("model") + model_id: str | None = model.get("model_info", {}).get("id") # type: ignore + litellm_model: str | None = model["litellm_params"].get("model") if model_id is None or litellm_model is None: continue itpm_keys.append( @@ -9348,12 +9311,12 @@ class Router: itpm_values = combined_values[: len(itpm_keys)] otpm_values = combined_values[len(itpm_keys) :] - total_itpm: Optional[int] = None + total_itpm: int | None = None for value in itpm_values: if isinstance(value, int): total_itpm = (total_itpm or 0) + value - total_otpm: Optional[int] = None + total_otpm: int | None = None for value in otpm_values: if isinstance(value, int): total_otpm = (total_otpm or 0) + value @@ -9361,7 +9324,7 @@ class Router: return total_itpm, total_otpm @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) - def _cached_get_model_group_info(self, model_group: str) -> Optional[ModelGroupInfo]: + def _cached_get_model_group_info(self, model_group: str) -> ModelGroupInfo | None: """ Cached version of get_model_group_info, uses @lru_cache wrapper @@ -9408,8 +9371,8 @@ class Router: async def set_response_headers( self, response: Any, - model_group: Optional[str] = None, - request_kwargs: Optional[dict] = None, + model_group: str | None = None, + request_kwargs: dict | None = None, ) -> Any: """ Add the most accurate rate limit headers for a given model response. @@ -9486,7 +9449,7 @@ class Router: self._add_model_to_list_and_index_map(model=model, model_id=model_id) - def get_model_ids(self, model_name: Optional[str] = None, exclude_team_models: bool = False) -> List[str]: + def get_model_ids(self, model_name: str | None = None, exclude_team_models: bool = False) -> list[str]: """ if 'model_name' is none, returns all. @@ -9533,7 +9496,7 @@ class Router: """ return candidate_id in self.model_id_to_deployment_index_map - def resolve_model_name_from_model_id(self, model_id: Optional[str]) -> Optional[str]: + def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: """ Resolve model_name from model_id. @@ -9590,7 +9553,7 @@ class Router: # No match found return None - def map_team_model(self, team_model_name: Optional[str], team_id: str) -> Optional[str]: + def map_team_model(self, team_model_name: str | None, team_id: str) -> str | None: """ Check if team_model_name resolves to team-specific deployments. @@ -9623,7 +9586,7 @@ class Router: # handled downstream by the pattern_router in _common_checks_available_deployment. return None - def should_include_deployment(self, model_name: str, model: dict, team_id: Optional[str] = None) -> bool: + def should_include_deployment(self, model_name: str, model: dict, team_id: str | None = None) -> bool: """ Get the team-specific model name if team_id matches the deployment. """ @@ -9649,9 +9612,9 @@ class Router: def _get_all_deployments( self, model_name: str, - model_alias: Optional[str] = None, - team_id: Optional[str] = None, - ) -> List[DeploymentTypedDict]: + model_alias: str | None = None, + team_id: str | None = None, + ) -> list[DeploymentTypedDict]: """ Return all deployments of a model name @@ -9667,7 +9630,7 @@ class Router: name (for example, `model_name__`), this method falls back to the standard model-name index / scan path. """ - returned_models: List[DeploymentTypedDict] = [] + returned_models: list[DeploymentTypedDict] = [] # O(1) lookup in team_model index when team_id is provided if team_id is not None: @@ -9720,7 +9683,7 @@ class Router: return returned_models - def get_model_names(self, team_id: Optional[str] = None) -> List[str]: + def get_model_names(self, team_id: str | None = None) -> list[str]: """ Returns all possible model names for the router, including models defined via model_group_alias. @@ -9741,7 +9704,7 @@ class Router: return model_names - def get_fully_blocked_model_names(self) -> Set[str]: + def get_fully_blocked_model_names(self) -> set[str]: """ Returns the set of model_names where every backing deployment has `blocked=True`. @@ -9750,7 +9713,7 @@ class Router: one non-blocked deployment is still serviceable and remains visible. """ deployments = self.get_model_list() or [] - blocked_by_name: Dict[str, bool] = {} + blocked_by_name: dict[str, bool] = {} for deployment in deployments: name = deployment.get("model_name") or "" if not name: @@ -9764,7 +9727,7 @@ class Router: @staticmethod def _are_all_deployments_blocked( - deployments: List[DeploymentTypedDict], + deployments: list[DeploymentTypedDict], ) -> bool: return len(deployments) > 0 and all( (deployment.get("model_info") or {}).get("blocked") is True for deployment in deployments @@ -9774,7 +9737,7 @@ class Router: deployments = self.get_model_list(model_name=model) or [] return self._are_all_deployments_blocked(deployments=deployments) - async def async_get_fully_unhealthy_model_names(self) -> Set[str]: + async def async_get_fully_unhealthy_model_names(self) -> set[str]: """ Returns the set of model names where every backing deployment is currently marked unhealthy by background health checks (and the health state is not stale). @@ -9808,7 +9771,7 @@ class Router: if not unhealthy_ids: return set() deployments = self.get_model_list() or [] - unhealthy_by_name: Dict[str, bool] = {} + unhealthy_by_name: dict[str, bool] = {} for deployment in deployments: model_info = deployment.get("model_info") or {} names = [deployment.get("model_name") or ""] @@ -9825,7 +9788,7 @@ class Router: unhealthy_by_name[name] = is_unhealthy return {name for name, fully_unhealthy in unhealthy_by_name.items() if fully_unhealthy} - def _get_team_specific_model(self, deployment: DeploymentTypedDict, team_id: Optional[str] = None) -> Optional[str]: + def _get_team_specific_model(self, deployment: DeploymentTypedDict, team_id: str | None = None) -> str | None: """ Get the team-specific model name if team_id matches the deployment. @@ -9837,14 +9800,14 @@ class Router: str: The `team_public_model_name` if team_id matches None: If team_id doesn't match or no team info exists """ - model_info: Optional[Dict] = deployment.get("model_info") or {} + model_info: dict | None = deployment.get("model_info") or {} if model_info is None: return None if team_id == model_info.get("team_id"): return model_info.get("team_public_model_name") return None - def _is_team_specific_model(self, model_info: Optional[Dict]) -> bool: + def _is_team_specific_model(self, model_info: dict | None) -> bool: """ Check if model info contains team-specific configuration. @@ -9856,13 +9819,13 @@ class Router: """ return bool(model_info and model_info.get("team_id")) - def get_model_list_from_model_alias(self, model_name: Optional[str] = None) -> List[DeploymentTypedDict]: + def get_model_list_from_model_alias(self, model_name: str | None = None) -> list[DeploymentTypedDict]: """ Helper function to get model list from model alias. Used by `.get_model_list` to get model list from model alias. """ - returned_models: List[DeploymentTypedDict] = [] + returned_models: list[DeploymentTypedDict] = [] if model_name is not None: # Fast path: direct dict lookup avoids scanning all aliases for non-alias model names. @@ -9889,8 +9852,8 @@ class Router: return returned_models def get_model_list( - self, model_name: Optional[str] = None, team_id: Optional[str] = None - ) -> Optional[List[DeploymentTypedDict]]: + self, model_name: str | None = None, team_id: str | None = None + ) -> list[DeploymentTypedDict] | None: """ Includes router model_group_alias'es as well @@ -9898,7 +9861,7 @@ class Router: """ # Note: model_list and model_group_alias are always initialized in __init__ # so hasattr checks are unnecessary - returned_models: List[DeploymentTypedDict] = [] + returned_models: list[DeploymentTypedDict] = [] if model_name is not None: returned_models.extend(self._get_all_deployments(model_name=model_name, team_id=team_id)) @@ -9945,10 +9908,10 @@ class Router: def get_model_access_groups( self, - model_name: Optional[str] = None, - model_access_group: Optional[str] = None, - team_id: Optional[str] = None, - ) -> Dict[str, List[str]]: + model_name: str | None = None, + model_access_group: str | None = None, + team_id: str | None = None, + ) -> dict[str, list[str]]: """ If model_name is provided, only return access groups for that model. @@ -10120,7 +10083,7 @@ class Router: relink_lar1_from_args = True setattr(self, var, value) else: - verbose_router_logger.debug("Setting {} is not allowed".format(var)) + verbose_router_logger.debug(f"Setting {var} is not allowed") if relink_lar1_from_args and self._normalize_strategy(self.routing_strategy) == "lar1": from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy @@ -10144,9 +10107,9 @@ class Router: The appropriate client based on the given client_type and kwargs. """ model_id = deployment["model_info"]["id"] - parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(kwargs) if client_type == "max_parallel_requests": - cache_key = "{}_max_parallel_requests_client".format(model_id) + cache_key = f"{model_id}_max_parallel_requests_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) if client is None: InitalizeCachedClient.set_max_parallel_requests_client(litellm_router_instance=self, model=deployment) @@ -10206,10 +10169,10 @@ class Router: def _pre_call_checks( self, model: str, - healthy_deployments: List, + healthy_deployments: list, messages: list[dict[str, str]] | None = None, input: str | list | None = None, - request_kwargs: Optional[dict] = None, + request_kwargs: dict | None = None, ): """ Filter out model in model group, if: @@ -10231,7 +10194,7 @@ class Router: # Token counting (tiktoken) is the dominant on-loop cost for large prompts. # Only count when a deployment actually declares max_input_tokens, and count # at most once; for model groups with no context-window limit it is skipped. - input_tokens: Optional[int] = None + input_tokens: int | None = None _context_window_error = False _potential_error_str = "" @@ -10272,22 +10235,18 @@ class Router: ) except Exception as e: verbose_router_logger.error( - "litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - {}".format( - str(e) - ) + f"litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - {e!s}" ) return _returned_deployments if input_tokens > max_input_tokens: invalid_model_indices.add(idx) _context_window_error = True - _potential_error_str += "Model={}, Max Input Tokens={}, Got={}".format( - _deployment_model, - max_input_tokens, - input_tokens, + _potential_error_str += ( + f"Model={_deployment_model}, Max Input Tokens={max_input_tokens}, Got={input_tokens}" ) continue except Exception as e: - verbose_router_logger.exception("An error occurs - {}".format(str(e))) + verbose_router_logger.exception(f"An error occurs - {e!s}") model_id = _model_info.get("id", "") ## RPM CHECK ## @@ -10365,9 +10324,7 @@ class Router: elif _context_window_error is True: raise litellm.ContextWindowExceededError( - message="litellm._pre_call_checks: Context Window exceeded for given call. No models have context window large enough for this call.\n{}".format( - _potential_error_str - ), + message=f"litellm._pre_call_checks: Context Window exceeded for given call. No models have context window large enough for this call.\n{_potential_error_str}", model=model, llm_provider="", ) @@ -10377,7 +10334,7 @@ class Router: return _returned_deployments - def _get_model_from_alias(self, model: str) -> Optional[str]: + def _get_model_from_alias(self, model: str) -> str | None: """ Get the model from the alias. @@ -10396,7 +10353,7 @@ class Router: return model - def _get_deployment_by_litellm_model(self, model: str) -> List: + def _get_deployment_by_litellm_model(self, model: str) -> list: """ Get the deployment by litellm model. """ @@ -10405,9 +10362,9 @@ class Router: def _try_early_resolve_deployments_for_model_not_in_names( self, model: str, - request_team_id: Optional[str], + request_team_id: str | None, include_team_models: bool = False, - ) -> Optional[Tuple[str, Union[List, Dict]]]: + ) -> tuple[str, list | dict] | None: """ When ``model`` is not in ``self.model_names``, try team routes, pattern routes, team pattern routers, then default deployment. Returns None if none apply. @@ -10472,11 +10429,11 @@ class Router: def _common_checks_available_deployment( self, model: str, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - request_kwargs: Optional[Dict] = None, - ) -> Tuple[str, Union[List, Dict]]: + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + request_kwargs: dict | None = None, + ) -> tuple[str, list | dict]: """ Common checks for 'get_available_deployment' across sync + async call. @@ -10488,7 +10445,7 @@ class Router: - Dict, if specific model chosen """ - request_team_id: Optional[str] = None + request_team_id: str | None = None if request_kwargs is not None: metadata = request_kwargs.get("metadata") or {} litellm_metadata = request_kwargs.get("litellm_metadata") or {} @@ -10595,10 +10552,10 @@ class Router: def _filter_deployments_by_model_access_groups( self, model: str, - healthy_deployments: List, - request_kwargs: Optional[Dict], - request_team_id: Optional[str], - ) -> List: + healthy_deployments: list, + request_kwargs: dict | None, + request_team_id: str | None, + ) -> list: """ Restrict candidate deployments to caller-authorized model access groups. @@ -10649,12 +10606,12 @@ class Router: async def async_get_healthy_deployments( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - parent_otel_span: Optional[Span] = None, - ) -> Union[List[Dict], Dict]: + request_kwargs: dict, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + parent_otel_span: Span | None = None, + ) -> list[dict] | dict: """ Get the healthy deployments for a model. @@ -10729,7 +10686,7 @@ class Router: healthy_deployments = await self.async_callback_filter_deployments( model=model, healthy_deployments=healthy_deployments, - messages=(cast(List[AllMessageValues], messages) if messages is not None else None), + messages=(cast(list[AllMessageValues], messages) if messages is not None else None), request_kwargs=request_kwargs, parent_otel_span=parent_otel_span, ) @@ -10737,7 +10694,7 @@ class Router: if self.enable_pre_call_checks and (messages is not None or input is not None): healthy_deployments = self._pre_call_checks( model=model, - healthy_deployments=cast(List[Dict], healthy_deployments), + healthy_deployments=cast(list[dict], healthy_deployments), messages=messages, input=input, request_kwargs=request_kwargs, @@ -10760,7 +10717,7 @@ class Router: ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) _target_order = (request_kwargs or {}).pop("_target_order", None) healthy_deployments = litellm.utils._get_order_filtered_deployments( - cast(List[Dict], healthy_deployments), target_order=_target_order + cast(list[dict], healthy_deployments), target_order=_target_order ) ## WEIGHTED FAILOVER EXCLUSION ## -> drop deployments already tried in @@ -10768,7 +10725,7 @@ class Router: ## router-level flag, so a stale exclusion key on kwargs cannot escape. _excluded_deployment_ids = (request_kwargs or {}).pop("_excluded_deployment_ids", None) healthy_deployments = litellm.utils._get_excluded_filtered_deployments( - cast(List[Dict], healthy_deployments), + cast(list[dict], healthy_deployments), excluded_deployment_ids=_excluded_deployment_ids, ) @@ -10785,10 +10742,10 @@ class Router: async def async_get_available_deployment( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, + request_kwargs: dict, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, ): """ Async implementation of 'get_available_deployments'. @@ -10910,10 +10867,10 @@ class Router: async def async_get_available_deployment_for_pass_through( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, + request_kwargs: dict, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, ): """ Async version of get_available_deployment_for_pass_through @@ -11074,9 +11031,9 @@ class Router: def _filter_by_routing_plugin_candidates( self, - healthy_deployments: Union[list[dict], dict], + healthy_deployments: list[dict] | dict, request_kwargs: dict, - ) -> Union[list[dict], dict]: + ) -> list[dict] | dict: """ Narrow `healthy_deployments` to whatever `self.routing_plugins` left in `context.candidate_models`. Raises rather than silently falling back to @@ -11102,7 +11059,7 @@ class Router: return filtered - def _select_pre_routing_strategy(self, model: str, request_kwargs: Dict) -> "PreRoutingStrategy | None": + def _select_pre_routing_strategy(self, model: str, request_kwargs: dict) -> "PreRoutingStrategy | None": """ Resolve the pre-routing strategy for `model`, disambiguating deployments that share a `model_name` by matching the request's tags against each @@ -11134,11 +11091,11 @@ class Router: async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - ) -> Optional[PreRoutingHookResponse]: + request_kwargs: dict, + messages: list[dict[str, Any]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + ) -> PreRoutingHookResponse | None: """ This hook is called before the routing decision is made. @@ -11252,10 +11209,10 @@ class Router: def get_available_deployment( self, model: str, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - request_kwargs: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + request_kwargs: dict | None = None, ): """ Returns the deployment based on routing strategy @@ -11292,7 +11249,7 @@ class Router: ) return healthy_deployments - parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(request_kwargs) # Health-check-based filtering (before cooldown) healthy_deployments = self._filter_health_check_unhealthy_deployments( @@ -11394,10 +11351,10 @@ class Router: def get_available_deployment_for_pass_through( self, model: str, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - request_kwargs: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + request_kwargs: dict | None = None, ): """ Returns deployments available for pass-through endpoints (based on load balancing strategy) @@ -11457,7 +11414,7 @@ class Router: ) # 4. Apply health-check and cooldown filtering - parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(request_kwargs) pass_through_deployments = self._filter_health_check_unhealthy_deployments( healthy_deployments=pass_through_deployments, parent_otel_span=parent_otel_span, @@ -11534,8 +11491,8 @@ class Router: return deployment def _filter_cooldown_deployments( - self, healthy_deployments: List[Dict], cooldown_deployments: List[str] - ) -> List[Dict]: + self, healthy_deployments: list[dict], cooldown_deployments: list[str] + ) -> list[dict]: """ Filters out the deployments currently cooling down from the list of healthy deployments @@ -11552,7 +11509,7 @@ class Router: cooldown_set = set(cooldown_deployments) return [deployment for deployment in healthy_deployments if deployment["model_info"]["id"] not in cooldown_set] - def _filter_blocked_deployments(self, healthy_deployments: List[Dict]) -> List[Dict]: + def _filter_blocked_deployments(self, healthy_deployments: list[dict]) -> list[dict]: """ Filters out deployments that an admin has paused via `LiteLLM_ProxyModelTable.blocked`. @@ -11582,9 +11539,9 @@ class Router: async def _async_filter_health_check_unhealthy_deployments( self, - healthy_deployments: List[Dict], - parent_otel_span: Optional[Span] = None, - ) -> List[Dict]: + healthy_deployments: list[dict], + parent_otel_span: Span | None = None, + ) -> list[dict]: """ Filter out deployments marked unhealthy by background health checks. No-op when enable_health_check_routing is False. @@ -11616,9 +11573,9 @@ class Router: def _filter_health_check_unhealthy_deployments( self, - healthy_deployments: List[Dict], - parent_otel_span: Optional[Span] = None, - ) -> List[Dict]: + healthy_deployments: list[dict], + parent_otel_span: Span | None = None, + ) -> list[dict]: """Sync version of _async_filter_health_check_unhealthy_deployments.""" if not self.enable_health_check_routing: return healthy_deployments @@ -11638,7 +11595,7 @@ class Router: return filtered - def _filter_pass_through_deployments(self, healthy_deployments: List[Dict]) -> List[Dict]: + def _filter_pass_through_deployments(self, healthy_deployments: list[dict]) -> list[dict]: """ Filter out deployments configured with use_in_pass_through=True @@ -11662,7 +11619,7 @@ class Router: return pass_through_deployments - def _track_deployment_metrics(self, deployment, parent_otel_span: Optional[Span], response=None): + def _track_deployment_metrics(self, deployment, parent_otel_span: Span | None, response=None): """ Tracks successful requests rpm usage. """ @@ -11673,9 +11630,9 @@ class Router: if model_id is not None: self._update_usage(model_id, parent_otel_span) # update in-memory cache for tracking except Exception as e: - verbose_router_logger.error(f"Error in _track_deployment_metrics: {str(e)}") + verbose_router_logger.error(f"Error in _track_deployment_metrics: {e!s}") - def get_num_retries_from_retry_policy(self, exception: Exception, model_group: Optional[str] = None): + def get_num_retries_from_retry_policy(self, exception: Exception, model_group: str | None = None): return _get_num_retries_from_retry_policy( exception=exception, model_group=model_group, @@ -11692,7 +11649,7 @@ class Router: ContentPolicyViolationErrorRetries: Optional[int] = None """ # if we can find the exception then in the retry policy -> return the number of retries - allowed_fails_policy: Optional[AllowedFailsPolicy] = self.allowed_fails_policy + allowed_fails_policy: AllowedFailsPolicy | None = self.allowed_fails_policy if allowed_fails_policy is None: return None diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index e8fcec2667d..086c7eb73c6 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -17,13 +17,12 @@ import asyncio import time from collections import OrderedDict from dataclasses import asdict, dataclass -from typing import Any, Union, cast +from typing import Any, cast from litellm._logging import verbose_router_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_last_user_message, ) -from litellm.types.utils import StandardLoggingRoutingDecision from litellm.router_strategy.adaptive_router.bandit import ( BanditCell, apply_delta, @@ -51,6 +50,7 @@ from litellm.router_strategy.adaptive_router.signals import ( from litellm.router_strategy.adaptive_router.update_queue import ( AdaptiveRouterUpdateQueue, ) +from litellm.types.utils import StandardLoggingRoutingDecision # Sweep session-state cache when it exceeds this many live entries. Expired # entries are dropped in bulk; amortizes to O(1) per insert. @@ -160,7 +160,7 @@ class AdaptiveRouter: model: str, request_kwargs: dict[str, Any], messages: list[dict[str, Any]] | None = None, - input: Union[str, list] | None = None, + input: str | list | None = None, specific_deployment: bool | None = False, ) -> PreRoutingHookResponse | None: """ diff --git a/litellm/router_strategy/adaptive_router/bandit.py b/litellm/router_strategy/adaptive_router/bandit.py index 4c914e00826..14a2d55c70f 100644 --- a/litellm/router_strategy/adaptive_router/bandit.py +++ b/litellm/router_strategy/adaptive_router/bandit.py @@ -12,7 +12,6 @@ Hot path: thompson_sample() — pure function, no I/O. import random from dataclasses import dataclass -from typing import Dict, List, Optional from litellm.router_strategy.adaptive_router.config import ( BASE_TIER_WEIGHT, @@ -75,13 +74,13 @@ def apply_delta(cell: BanditCell, delta_alpha: float, delta_beta: float) -> Band return BanditCell(alpha=new_alpha, beta=new_beta) -def thompson_sample(cell: BanditCell, rng: Optional[random.Random] = None) -> float: +def thompson_sample(cell: BanditCell, rng: random.Random | None = None) -> float: """Draw a sample from Beta(alpha, beta). Returns a quality estimate in [0, 1].""" r = rng if rng is not None else random return r.betavariate(cell.alpha, cell.beta) -def normalized_cost(model_cost: float, all_costs: List[float]) -> float: +def normalized_cost(model_cost: float, all_costs: list[float]) -> float: """ Map a raw $/1k-token cost into [0, 1] where 0 = most expensive, 1 = cheapest. Returns 0.5 when there's no spread. @@ -97,7 +96,7 @@ def normalized_cost(model_cost: float, all_costs: List[float]) -> float: def score( quality_sample: float, model_cost: float, - all_costs: List[float], + all_costs: list[float], quality_weight: float = DEFAULT_QUALITY_WEIGHT, cost_weight: float = DEFAULT_COST_WEIGHT, ) -> float: @@ -110,11 +109,11 @@ def score( def pick_best( - cells: Dict[str, BanditCell], - model_costs: Dict[str, float], + cells: dict[str, BanditCell], + model_costs: dict[str, float], quality_weight: float = DEFAULT_QUALITY_WEIGHT, cost_weight: float = DEFAULT_COST_WEIGHT, - rng: Optional[random.Random] = None, + rng: random.Random | None = None, ) -> str: """ Sample once per model, score each, return the model with highest score. @@ -125,7 +124,7 @@ def pick_best( if not cells: raise ValueError("pick_best called with no models") all_costs = list(model_costs.values()) - best_model: Optional[str] = None + best_model: str | None = None best_score = float("-inf") for model, cell in cells.items(): q = thompson_sample(cell, rng=rng) diff --git a/litellm/router_strategy/adaptive_router/classifier.py b/litellm/router_strategy/adaptive_router/classifier.py index 7164d8a794e..f3da35e438e 100644 --- a/litellm/router_strategy/adaptive_router/classifier.py +++ b/litellm/router_strategy/adaptive_router/classifier.py @@ -9,11 +9,10 @@ Order matters: we check more specific types first, falling back to GENERAL. import re from re import Pattern -from typing import List, Tuple from litellm.types.router import RequestType -_RULES: List[Tuple[Pattern[str], RequestType]] = [ +_RULES: list[tuple[Pattern[str], RequestType]] = [ ( re.compile( r"\b(write|create|generate|implement|build)\s+(?:a |an |the |me )?(?:python|javascript|typescript|java|rust|go|c\+\+|sql|bash|shell)\b", diff --git a/litellm/router_strategy/adaptive_router/config.py b/litellm/router_strategy/adaptive_router/config.py index e72826cc056..8405023c812 100644 --- a/litellm/router_strategy/adaptive_router/config.py +++ b/litellm/router_strategy/adaptive_router/config.py @@ -5,8 +5,6 @@ All magic numbers are first-pass guesses (D3-D6 in the handoff plan). Expect to retune after first 1000 sessions of real traffic. """ -from typing import Dict - from litellm.types.router import RequestType # re-export for convenience # noqa: F401 # D3 — Score weights (default; user-overridable via AdaptiveRouterConfig.weights) @@ -15,7 +13,7 @@ DEFAULT_COST_WEIGHT: float = 0.3 # UNVALIDATED — calibrated against [0] sessi # D4 — Cold-start prior: (alpha + beta) total mass = COLD_START_MASS # Mean of Beta = base_tier_weight + (strength_bonus if declared) -BASE_TIER_WEIGHT: Dict[int, float] = {1: 0.3, 2: 0.5, 3: 0.7} # UNVALIDATED +BASE_TIER_WEIGHT: dict[int, float] = {1: 0.3, 2: 0.5, 3: 0.7} # UNVALIDATED STRENGTH_BONUS: float = 0.3 # UNVALIDATED COLD_START_MASS: float = 10.0 diff --git a/litellm/router_strategy/adaptive_router/hooks.py b/litellm/router_strategy/adaptive_router/hooks.py index 89ae28be227..3805436e95c 100644 --- a/litellm/router_strategy/adaptive_router/hooks.py +++ b/litellm/router_strategy/adaptive_router/hooks.py @@ -13,7 +13,7 @@ from __future__ import annotations import hashlib import json -from typing import Any, Dict, List, Optional +from typing import Any from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger @@ -37,7 +37,7 @@ _IDENTITY_FIELDS = ( ) -def _resolve_session_key(kwargs: Dict[str, Any]) -> Optional[str]: +def _resolve_session_key(kwargs: dict[str, Any]) -> str | None: """Pick a stable per-conversation key for owner-cache attribution. Order: @@ -82,7 +82,7 @@ def _resolve_session_key(kwargs: Dict[str, Any]) -> Optional[str]: return hashlib.sha256(payload.encode("utf-8")).hexdigest() -def _last_user_content(messages: Optional[List[Dict[str, Any]]]) -> Optional[str]: +def _last_user_content(messages: list[dict[str, Any]] | None) -> str | None: if not messages: return None for msg in reversed(messages): @@ -100,8 +100,8 @@ def _last_user_content(messages: Optional[List[Dict[str, Any]]]) -> Optional[str def _recent_tool_results( - messages: Optional[List[Dict[str, Any]]], -) -> List[Dict[str, Any]]: + messages: list[dict[str, Any]] | None, +) -> list[dict[str, Any]]: """Extract the current turn's tool result payloads from the request messages. Tool results are `role == "tool"` messages that sit at the tail of the @@ -115,7 +115,7 @@ def _recent_tool_results( """ if not messages: return [] - results: List[Dict[str, Any]] = [] + results: list[dict[str, Any]] = [] for msg in reversed(messages): if not isinstance(msg, dict): break @@ -154,7 +154,7 @@ def _assistant_content_and_tool_calls(response_obj: Any) -> tuple: raw_tool_calls = getattr(msg, "tool_calls", None) if raw_tool_calls is None and isinstance(msg, dict): raw_tool_calls = msg.get("tool_calls") - tool_calls: List[Dict[str, Any]] = [] + tool_calls: list[dict[str, Any]] = [] for tc in raw_tool_calls or []: if isinstance(tc, dict): tool_calls.append(tc) @@ -174,12 +174,12 @@ class AdaptiveRouterPostCallHook(CustomLogger): async def async_post_call_response_headers_hook( self, - data: Dict[str, Any], + data: dict[str, Any], user_api_key_dict: Any, response: Any, - request_headers: Optional[Dict[str, str]] = None, - litellm_call_info: Optional[Dict[str, Any]] = None, - ) -> Optional[Dict[str, str]]: + request_headers: dict[str, str] | None = None, + litellm_call_info: dict[str, Any] | None = None, + ) -> dict[str, str] | None: """ Surface the chosen logical model as the `x-litellm-adaptive-router-model` response header for both streaming and non-streaming responses. @@ -208,7 +208,7 @@ class AdaptiveRouterPostCallHook(CustomLogger): async def _record( self, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], response_obj: Any, response_status: int, ) -> None: diff --git a/litellm/router_strategy/adaptive_router/update_queue.py b/litellm/router_strategy/adaptive_router/update_queue.py index 505d243202c..4e373237de7 100644 --- a/litellm/router_strategy/adaptive_router/update_queue.py +++ b/litellm/router_strategy/adaptive_router/update_queue.py @@ -19,7 +19,7 @@ to the in-memory aggregator). Flush is async and batched. from __future__ import annotations import asyncio -from typing import Any, Dict, Tuple +from typing import Any from litellm._logging import verbose_router_logger from litellm.repositories.table_repositories import ( @@ -27,8 +27,8 @@ from litellm.repositories.table_repositories import ( AdaptiveRouterStateRepository, ) -StateKey = Tuple[str, str, str] # (router_name, request_type, model_name) -SessionKey = Tuple[str, str, str] # (session_id, router_name, model_name) +StateKey = tuple[str, str, str] # (router_name, request_type, model_name) +SessionKey = tuple[str, str, str] # (session_id, router_name, model_name) class AdaptiveRouterUpdateQueue: @@ -38,8 +38,8 @@ class AdaptiveRouterUpdateQueue: """ def __init__(self) -> None: - self._state_agg: Dict[StateKey, Dict[str, float]] = {} - self._session_agg: Dict[SessionKey, Dict[str, Any]] = {} + self._state_agg: dict[StateKey, dict[str, float]] = {} + self._session_agg: dict[SessionKey, dict[str, Any]] = {} self._lock = asyncio.Lock() self._max_state_size_seen = 0 self._max_session_size_seen = 0 @@ -68,8 +68,7 @@ class AdaptiveRouterUpdateQueue: current["delta_alpha"] += delta_alpha current["delta_beta"] += delta_beta current["samples_added"] += 1 - if len(self._state_agg) > self._max_state_size_seen: - self._max_state_size_seen = len(self._state_agg) + self._max_state_size_seen = max(self._max_state_size_seen, len(self._state_agg)) # ---- Hot-path: session snapshot -------------------------------------- @@ -78,7 +77,7 @@ class AdaptiveRouterUpdateQueue: session_id: str, router_name: str, model_name: str, - state_dict: Dict[str, Any], + state_dict: dict[str, Any], ) -> None: """ Last-write-wins per session row. The state_dict is a snapshot of the @@ -88,8 +87,7 @@ class AdaptiveRouterUpdateQueue: key: SessionKey = (session_id, router_name, model_name) async with self._lock: self._session_agg[key] = state_dict - if len(self._session_agg) > self._max_session_size_seen: - self._max_session_size_seen = len(self._session_agg) + self._max_session_size_seen = max(self._max_session_size_seen, len(self._session_agg)) # ---- Flushers (called by background task) ---------------------------- @@ -203,7 +201,7 @@ class AdaptiveRouterUpdateQueue: # ---- Observability --------------------------------------------------- - async def queue_size(self) -> Dict[str, int]: + async def queue_size(self) -> dict[str, int]: async with self._lock: return { "state_pending": len(self._state_agg), diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index c01e2f10d2c..d959eb3ef73 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -2,7 +2,7 @@ Auto-Routing Strategy that works with a Semantic Router Config """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger @@ -27,8 +27,8 @@ class AutoRouter(CustomLogger): default_model: str, embedding_model: str, litellm_router_instance: "Router", - auto_router_config_path: Optional[str] = None, - auto_router_config: Optional[str] = None, + auto_router_config_path: str | None = None, + auto_router_config: str | None = None, ): """ Auto-Router class that uses a semantic router to route requests to the appropriate model. @@ -43,16 +43,16 @@ class AutoRouter(CustomLogger): """ from semantic_router.routers import SemanticRouter - self.auto_router_config_path: Optional[str] = auto_router_config_path - self.auto_router_config: Optional[str] = auto_router_config + self.auto_router_config_path: str | None = auto_router_config_path + self.auto_router_config: str | None = auto_router_config self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE - self.loaded_routes: List[Route] = self._load_semantic_routing_routes() - self.routelayer: Optional[SemanticRouter] = None + self.loaded_routes: list[Route] = self._load_semantic_routing_routes() + self.routelayer: SemanticRouter | None = None self.default_model = default_model self.embedding_model: str = embedding_model - self.litellm_router_instance: "Router" = litellm_router_instance + self.litellm_router_instance: Router = litellm_router_instance - def _load_semantic_routing_routes(self) -> List[Route]: + def _load_semantic_routing_routes(self) -> list[Route]: from semantic_router.routers import SemanticRouter if self.auto_router_config_path: @@ -62,14 +62,14 @@ class AutoRouter(CustomLogger): else: raise ValueError("No router config provided") - def _load_auto_router_routes_from_config_json(self) -> List[Route]: + def _load_auto_router_routes_from_config_json(self) -> list[Route]: import json from semantic_router.routers.base import Route if self.auto_router_config is None: raise ValueError("No auto router config provided") - auto_router_routes: List[Route] = [] + auto_router_routes: list[Route] = [] loaded_config = json.loads(self.auto_router_config) for route in loaded_config.get("routes", []): auto_router_routes.append( @@ -83,7 +83,7 @@ class AutoRouter(CustomLogger): return auto_router_routes @staticmethod - def _extract_text_from_messages(messages: List[Dict[str, Any]]) -> str: + def _extract_text_from_messages(messages: list[dict[str, Any]]) -> str: """ Extract text content from the last user message for routing. @@ -108,10 +108,10 @@ class AutoRouter(CustomLogger): async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, + request_kwargs: dict, + messages: list[dict[str, Any]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, ) -> Optional["PreRoutingHookResponse"]: """ This hook is called before the routing decision is made. @@ -146,7 +146,7 @@ class AutoRouter(CustomLogger): self.routelayer = routelayer message_content = self._extract_text_from_messages(messages) - route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = routelayer(text=message_content) + route_choice: RouteChoice | list[RouteChoice] | None = routelayer(text=message_content) verbose_router_logger.debug(f"route_choice: {route_choice}") if isinstance(route_choice, RouteChoice): model = route_choice.name or self.default_model diff --git a/litellm/router_strategy/auto_router/litellm_encoder.py b/litellm/router_strategy/auto_router/litellm_encoder.py index 7e163ba16a6..c3523af2eb1 100644 --- a/litellm/router_strategy/auto_router/litellm_encoder.py +++ b/litellm/router_strategy/auto_router/litellm_encoder.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Optional from pydantic import ConfigDict from semantic_router.encoders import DenseEncoder @@ -49,7 +49,7 @@ class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): self, litellm_router_instance: "Router", model_name: str, - score_threshold: Union[float, None] = None, + score_threshold: float | None = None, ): """Initialize the LiteLLMEncoder. diff --git a/litellm/router_strategy/base_routing_strategy.py b/litellm/router_strategy/base_routing_strategy.py index 7cc83b0feb1..ff395828b2a 100644 --- a/litellm/router_strategy/base_routing_strategy.py +++ b/litellm/router_strategy/base_routing_strategy.py @@ -4,7 +4,6 @@ Base class across routing strategies to abstract commmon functions like batch in import asyncio from abc import ABC -from typing import Dict, List, Optional, Set, Tuple, Union from litellm._logging import verbose_router_logger from litellm.caching.caching import DualCache @@ -17,17 +16,17 @@ class BaseRoutingStrategy(ABC): self, dual_cache: DualCache, should_batch_redis_writes: bool, - default_sync_interval: Optional[Union[int, float]], + default_sync_interval: float | None, ): self.dual_cache = dual_cache - self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = [] - self._sync_task: Optional[asyncio.Task[None]] = None + self.redis_increment_operation_queue: list[RedisPipelineIncrementOperation] = [] + self._sync_task: asyncio.Task[None] | None = None if should_batch_redis_writes: self.setup_sync_task(default_sync_interval) self.in_memory_keys_to_update: set[str] = set() # Set with max size of 1000 keys - def setup_sync_task(self, default_sync_interval: Optional[Union[int, float]]): + def setup_sync_task(self, default_sync_interval: float | None): """Setup the sync task in a way that's compatible with FastAPI""" try: loop = asyncio.get_running_loop() @@ -49,8 +48,8 @@ class BaseRoutingStrategy(ABC): pass async def _increment_value_list_in_current_window( - self, increment_list: List[Tuple[str, int]], ttl: int - ) -> List[float]: + self, increment_list: list[tuple[str, int]], ttl: int + ) -> list[float]: """ Increment a list of values in the current window """ @@ -60,7 +59,7 @@ class BaseRoutingStrategy(ABC): results.append(result) return results - async def _increment_value_in_current_window(self, key: str, value: Union[int, float], ttl: int): + async def _increment_value_in_current_window(self, key: str, value: float, ttl: int): """ Increment spend within existing budget window @@ -84,7 +83,7 @@ class BaseRoutingStrategy(ABC): self.add_to_in_memory_keys_to_update(key=key) return result - async def periodic_sync_in_memory_spend_with_redis(self, default_sync_interval: Optional[Union[int, float]]): + async def periodic_sync_in_memory_spend_with_redis(self, default_sync_interval: float | None): """ Handler that triggers sync_in_memory_spend_with_redis every DEFAULT_REDIS_SYNC_INTERVAL seconds @@ -98,7 +97,7 @@ class BaseRoutingStrategy(ABC): default_sync_interval ) # Wait for DEFAULT_REDIS_SYNC_INTERVAL seconds before next sync except Exception as e: - verbose_router_logger.error(f"Error in periodic sync task: {str(e)}") + verbose_router_logger.error(f"Error in periodic sync task: {e!s}") await asyncio.sleep( default_sync_interval ) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying @@ -118,7 +117,7 @@ class BaseRoutingStrategy(ABC): if len(self.redis_increment_operation_queue) > 0: # Compress operations for the same key - compressed_ops: Dict[str, RedisPipelineIncrementOperation] = {} + compressed_ops: dict[str, RedisPipelineIncrementOperation] = {} ops_to_remove = [] for idx, op in enumerate(self.redis_increment_operation_queue): if op["key"] in compressed_ops: @@ -147,22 +146,22 @@ class BaseRoutingStrategy(ABC): return return_result except Exception as e: - verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {str(e)}") + verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {e!s}") self.redis_increment_operation_queue = [] def add_to_in_memory_keys_to_update(self, key: str): self.in_memory_keys_to_update.add(key) - def get_key_pattern_to_sync(self) -> Optional[str]: + def get_key_pattern_to_sync(self) -> str | None: """ Get the key pattern to sync """ return None - def get_in_memory_keys_to_update(self) -> Set[str]: + def get_in_memory_keys_to_update(self) -> set[str]: return self.in_memory_keys_to_update - def get_and_reset_in_memory_keys_to_update(self) -> Set[str]: + def get_and_reset_in_memory_keys_to_update(self) -> set[str]: """Atomic get and reset in-memory keys to update""" keys = self.in_memory_keys_to_update self.in_memory_keys_to_update = set() @@ -227,4 +226,4 @@ class BaseRoutingStrategy(ABC): await self.dual_cache.in_memory_cache.async_set_cache(key=key, value=merged) except Exception as e: - verbose_router_logger.exception(f"Error syncing in-memory cache with Redis: {str(e)}") + verbose_router_logger.exception(f"Error syncing in-memory cache with Redis: {e!s}") diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 067f38ab11c..619f1fc4629 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -19,27 +19,27 @@ anthropic: """ import asyncio +import builtins from datetime import datetime, timedelta, timezone -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any import litellm from litellm._logging import verbose_router_logger from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisPipelineIncrementOperation from litellm.integrations.custom_logger import CustomLogger, Span -from litellm.litellm_core_utils.duration_parser import duration_in_seconds -from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, ) +from litellm.litellm_core_utils.duration_parser import duration_in_seconds +from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs from litellm.router_utils.cooldown_callbacks import ( _get_prometheus_logger_from_callbacks, ) from litellm.types.llms.openai import AllMessageValues from litellm.types.router import DeploymentTypedDict, LiteLLM_Params, RouterErrors -from litellm.types.utils import BudgetConfig +from litellm.types.utils import BudgetConfig, GenericBudgetConfigType, StandardLoggingPayload from litellm.types.utils import BudgetConfig as GenericBudgetInfo -from litellm.types.utils import GenericBudgetConfigType, StandardLoggingPayload DEFAULT_REDIS_SYNC_INTERVAL = 1 @@ -54,7 +54,7 @@ class _LiteLLMParamsDictView: __slots__ = ("_params",) - def __init__(self, params: Dict[str, Any]): + def __init__(self, params: dict[str, Any]): self._params = params def __getattr__(self, key: str) -> Any: @@ -84,10 +84,10 @@ class _LiteLLMParamsDictView: def __len__(self) -> int: return len(self._params) - def dict(self) -> Dict[str, Any]: + def dict(self) -> dict[str, Any]: return dict(self._params) - def model_dump(self) -> Dict[str, Any]: + def model_dump(self) -> builtins.dict[str, Any]: return dict(self._params) @@ -95,15 +95,15 @@ class RouterBudgetLimiting(CustomLogger): def __init__( self, dual_cache: DualCache, - provider_budget_config: Optional[dict], - model_list: Optional[List[Union[DeploymentTypedDict, Dict[str, Any]]]] = None, + provider_budget_config: dict | None, + model_list: list[DeploymentTypedDict | dict[str, Any]] | None = None, ): self.dual_cache = dual_cache - self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = [] + self.redis_increment_operation_queue: list[RedisPipelineIncrementOperation] = [] asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis()) - self.provider_budget_config: Optional[GenericBudgetConfigType] = provider_budget_config - self.deployment_budget_config: Optional[GenericBudgetConfigType] = None - self.tag_budget_config: Optional[GenericBudgetConfigType] = None + self.provider_budget_config: GenericBudgetConfigType | None = provider_budget_config + self.deployment_budget_config: GenericBudgetConfigType | None = None + self.tag_budget_config: GenericBudgetConfigType | None = None self._init_provider_budgets() self._init_deployment_budgets(model_list=model_list) self._init_tag_budgets() @@ -115,11 +115,11 @@ class RouterBudgetLimiting(CustomLogger): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, # type: ignore - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, # type: ignore + ) -> list[dict]: """ Filter out deployments that have exceeded their provider budget limit. @@ -138,7 +138,7 @@ class RouterBudgetLimiting(CustomLogger): if len(healthy_deployments) == 0: return healthy_deployments - potential_deployments: List[Dict] = [] + potential_deployments: list[dict] = [] ( cache_keys, @@ -156,10 +156,10 @@ class RouterBudgetLimiting(CustomLogger): keys=cache_keys, parent_otel_span=parent_otel_span, ) - current_spends: List = _current_spends or [0.0] * len(cache_keys) + current_spends: list = _current_spends or [0.0] * len(cache_keys) # Map spends to their respective keys - spend_map: Dict[str, float] = {} + spend_map: dict[str, float] = {} for idx, key in enumerate(cache_keys): spend_map[key] = float(current_spends[idx] or 0.0) @@ -190,14 +190,14 @@ class RouterBudgetLimiting(CustomLogger): def _filter_out_deployments_above_budget( self, - potential_deployments: List[Dict[str, Any]], - healthy_deployments: List[Dict[str, Any]], - provider_configs: Dict[str, GenericBudgetInfo], - deployment_configs: Dict[str, GenericBudgetInfo], - deployment_providers: List[Optional[str]], - spend_map: Dict[str, float], - request_tags: List[str], - ) -> Tuple[List[Dict[str, Any]], str]: + potential_deployments: list[dict[str, Any]], + healthy_deployments: list[dict[str, Any]], + provider_configs: dict[str, GenericBudgetInfo], + deployment_configs: dict[str, GenericBudgetInfo], + deployment_providers: list[str | None], + spend_map: dict[str, float], + request_tags: list[str], + ) -> tuple[list[dict[str, Any]], str]: """ Filter out deployments that have exceeded their budget limit. Follow budget checks are run here: @@ -274,13 +274,13 @@ class RouterBudgetLimiting(CustomLogger): async def _async_get_cache_keys_for_router_budget_limiting( self, - healthy_deployments: List[Dict[str, Any]], - request_kwargs: Optional[Dict] = None, - ) -> Tuple[ - List[str], - Dict[str, GenericBudgetInfo], - Dict[str, GenericBudgetInfo], - List[Optional[str]], + healthy_deployments: list[dict[str, Any]], + request_kwargs: dict | None = None, + ) -> tuple[ + list[str], + dict[str, GenericBudgetInfo], + dict[str, GenericBudgetInfo], + list[str | None], ]: """ Returns list of cache keys to fetch from router cache for budget limiting and provider and deployment configs @@ -292,13 +292,13 @@ class RouterBudgetLimiting(CustomLogger): - Dict of deployment budget configs `deployment_configs` - List of resolved providers aligned by deployment index `deployment_providers` """ - cache_keys: List[str] = [] - provider_configs: Dict[str, GenericBudgetInfo] = {} - deployment_configs: Dict[str, GenericBudgetInfo] = {} - deployment_providers: List[Optional[str]] = [] + cache_keys: list[str] = [] + provider_configs: dict[str, GenericBudgetInfo] = {} + deployment_configs: dict[str, GenericBudgetInfo] = {} + deployment_providers: list[str | None] = [] # Resolve tags once before the loop (loop-invariant) - _request_tags: List[str] = [] + _request_tags: list[str] = [] if self.tag_budget_config: _request_tags = _get_tags_from_request_kwargs( request_kwargs=request_kwargs, @@ -401,7 +401,7 @@ class RouterBudgetLimiting(CustomLogger): # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): return - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: raise ValueError("standard_logging_payload is required") @@ -514,7 +514,7 @@ class RouterBudgetLimiting(CustomLogger): DEFAULT_REDIS_SYNC_INTERVAL ) # Wait for DEFAULT_REDIS_SYNC_INTERVAL seconds before next sync except Exception as e: - verbose_router_logger.error(f"Error in periodic sync task: {str(e)}") + verbose_router_logger.error(f"Error in periodic sync task: {e!s}") await asyncio.sleep( DEFAULT_REDIS_SYNC_INTERVAL ) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying @@ -545,7 +545,7 @@ class RouterBudgetLimiting(CustomLogger): self.redis_increment_operation_queue = [] except Exception as e: - verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {str(e)}") + verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {e!s}") async def _sync_in_memory_spend_with_redis(self): """ @@ -600,27 +600,27 @@ class RouterBudgetLimiting(CustomLogger): verbose_router_logger.debug(f"Updated in-memory cache for {key}: {value}") except Exception as e: - verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {str(e)}") + verbose_router_logger.error(f"Error syncing in-memory cache with Redis: {e!s}") def _get_budget_config_for_deployment( self, model_id: str, - ) -> Optional[GenericBudgetInfo]: + ) -> GenericBudgetInfo | None: if self.deployment_budget_config is None: return None return self.deployment_budget_config.get(model_id, None) - def _get_budget_config_for_provider(self, provider: str) -> Optional[GenericBudgetInfo]: + def _get_budget_config_for_provider(self, provider: str) -> GenericBudgetInfo | None: if self.provider_budget_config is None: return None return self.provider_budget_config.get(provider, None) - def _get_budget_config_for_tag(self, tag: str) -> Optional[GenericBudgetInfo]: + def _get_budget_config_for_tag(self, tag: str) -> GenericBudgetInfo | None: if self.tag_budget_config is None: return None return self.tag_budget_config.get(tag, None) - def _get_llm_provider_for_deployment(self, deployment: Dict) -> Optional[str]: + def _get_llm_provider_for_deployment(self, deployment: dict) -> str | None: try: deployment_litellm_params = deployment.get("litellm_params") or {} @@ -658,7 +658,7 @@ class RouterBudgetLimiting(CustomLogger): budget_limit=budget_limit, ) - async def _get_current_provider_spend(self, provider: str) -> Optional[float]: + async def _get_current_provider_spend(self, provider: str) -> float | None: """ GET the current spend for a provider from cache @@ -684,7 +684,7 @@ class RouterBudgetLimiting(CustomLogger): current_spend = await self.dual_cache.async_get_cache(spend_key) return float(current_spend) if current_spend is not None else 0.0 - async def _get_current_provider_budget_reset_at(self, provider: str) -> Optional[str]: + async def _get_current_provider_budget_reset_at(self, provider: str) -> str | None: budget_config = self._get_budget_config_for_provider(provider) if budget_config is None: return None @@ -710,7 +710,7 @@ class RouterBudgetLimiting(CustomLogger): spend_key = f"provider_spend:{provider}:{budget_config.budget_duration}" start_time_key = f"provider_budget_start_time:{provider}" - ttl_seconds: Optional[int] = None + ttl_seconds: int | None = None if budget_config.budget_duration is not None: ttl_seconds = duration_in_seconds(budget_config.budget_duration) @@ -725,8 +725,8 @@ class RouterBudgetLimiting(CustomLogger): @staticmethod def should_init_router_budget_limiter( - provider_budget_config: Optional[dict], - model_list: Optional[Union[List[DeploymentTypedDict], List[Dict[str, Any]]]] = None, + provider_budget_config: dict | None, + model_list: list[DeploymentTypedDict] | list[dict[str, Any]] | None = None, ): """ Returns `True` if the router budget routing settings are set and RouterBudgetLimiting should be initialized @@ -776,13 +776,13 @@ class RouterBudgetLimiting(CustomLogger): def _init_deployment_budgets( self, - model_list: Optional[List[Union[DeploymentTypedDict, Dict[str, Any]]]] = None, + model_list: list[DeploymentTypedDict | dict[str, Any]] | None = None, ): if model_list is None: return for _model in model_list: _litellm_params = _model.get("litellm_params", {}) - _model_info: Dict = _model.get("model_info") or {} + _model_info: dict = _model.get("model_info") or {} _model_id = _model_info.get("id") _max_budget = _litellm_params.get("max_budget") _budget_duration = _litellm_params.get("budget_duration") @@ -803,7 +803,7 @@ class RouterBudgetLimiting(CustomLogger): def register_deployment_budget( self, - deployment: Union[Dict[str, Any], DeploymentTypedDict], + deployment: dict[str, Any] | DeploymentTypedDict, ) -> None: """ Register or refresh deployment-level budget config for a runtime-added deployment. diff --git a/litellm/router_strategy/complexity_router/__init__.py b/litellm/router_strategy/complexity_router/__init__.py index b6f84d7cfa6..98f6ce399a8 100644 --- a/litellm/router_strategy/complexity_router/__init__.py +++ b/litellm/router_strategy/complexity_router/__init__.py @@ -9,14 +9,14 @@ No external API calls - all scoring is local and <1ms. from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter from litellm.router_strategy.complexity_router.config import ( - ComplexityTier, DEFAULT_COMPLEXITY_CONFIG, ComplexityRouterConfig, + ComplexityTier, ) __all__ = [ - "ComplexityRouter", - "ComplexityTier", "DEFAULT_COMPLEXITY_CONFIG", + "ComplexityRouter", "ComplexityRouterConfig", + "ComplexityTier", ] diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 1c6a83c6cd9..465698408db 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -20,7 +20,7 @@ import random import re from collections.abc import Iterator, Mapping, Sequence from itertools import islice -from typing import TYPE_CHECKING, Any, Literal, NamedTuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, NamedTuple, cast from pydantic import BaseModel @@ -1303,7 +1303,7 @@ class ComplexityRouter(CustomLogger): model: str, request_kwargs: dict, messages: list[dict[str, Any]] | None = None, - input: Union[str, list] | None = None, + input: str | list | None = None, specific_deployment: bool | None = False, ) -> PreRoutingHookResponse | None: """ @@ -1395,7 +1395,7 @@ class ComplexityRouter(CustomLogger): model: str, request_kwargs: dict, messages: list[dict[str, Any]] | None = None, - input: Union[str, list] | None = None, + input: str | list | None = None, specific_deployment: bool | None = False, ) -> PreRoutingHookResponse | None: """ diff --git a/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py b/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py index b46071ccee8..37e281f66dc 100644 --- a/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py +++ b/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py @@ -12,7 +12,6 @@ import sys # ruff: noqa: T201 from dataclasses import dataclass -from typing import List, Optional, Tuple from unittest.mock import MagicMock sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) @@ -28,14 +27,14 @@ class EvalCase: prompt: str expected_tier: ComplexityTier description: str - system_prompt: Optional[str] = None + system_prompt: str | None = None # Allow some flexibility - if actual tier is in acceptable_tiers, still passes - acceptable_tiers: Optional[List[ComplexityTier]] = None + acceptable_tiers: list[ComplexityTier] | None = None # ─── Evaluation Dataset ─── -EVAL_CASES: List[EvalCase] = [ +EVAL_CASES: list[EvalCase] = [ # === SIMPLE tier cases === EvalCase( prompt="Hello!", @@ -230,7 +229,7 @@ EVAL_CASES: List[EvalCase] = [ ] -def run_eval() -> Tuple[int, int, List[dict]]: +def run_eval() -> tuple[int, int, list[dict]]: """ Run the evaluation suite. diff --git a/litellm/router_strategy/lar1_routing.py b/litellm/router_strategy/lar1_routing.py index b53165f0969..acd7ac63225 100644 --- a/litellm/router_strategy/lar1_routing.py +++ b/litellm/router_strategy/lar1_routing.py @@ -10,7 +10,7 @@ LAR-1 metadata passed via request_kwargs["metadata"]["lar1"] from __future__ import annotations from collections.abc import Mapping -from typing import TYPE_CHECKING, Optional, Union +from typing import TYPE_CHECKING from pydantic import ValidationError @@ -31,7 +31,7 @@ def _coerce_threshold(value: object, default: float) -> float: def lar1_thresholds_from_args( - routing_strategy_args: Optional[Mapping[str, object]] = None, + routing_strategy_args: Mapping[str, object] | None = None, ) -> dict[str, float]: args = routing_strategy_args or {} return { @@ -43,7 +43,7 @@ def lar1_thresholds_from_args( def apply_lar1_routing_strategy( router: Router, - routing_strategy_args: Optional[Mapping[str, object]] = None, + routing_strategy_args: Mapping[str, object] | None = None, ) -> None: strategy = LAR1RoutingStrategy( router_instance=router, @@ -53,7 +53,7 @@ def apply_lar1_routing_strategy( router.set_custom_routing_strategy(strategy) -def _normalize_thresholds(thresholds: Optional[dict[str, float]]) -> dict[str, float]: +def _normalize_thresholds(thresholds: dict[str, float] | None) -> dict[str, float]: merged = {**DEFAULT_THRESHOLDS, **(thresholds or {})} low = merged["low"] medium = merged["medium"] @@ -80,8 +80,8 @@ def _parse_lar1_metadata(request_kwargs: dict) -> LAR1Metadata: class LAR1RoutingStrategy(CustomRoutingStrategyBase): def __init__( self, - router_instance: Optional[Router] = None, - thresholds: Optional[dict[str, float]] = None, + router_instance: Router | None = None, + thresholds: dict[str, float] | None = None, ): self._router = router_instance self.thresholds = _normalize_thresholds(thresholds) @@ -89,10 +89,10 @@ class LAR1RoutingStrategy(CustomRoutingStrategyBase): async def async_get_available_deployment( self, model: str, - messages: Optional[list[dict[str, str]]] = None, - input: Optional[Union[str, list]] = None, - specific_deployment: Optional[bool] = False, - request_kwargs: Optional[dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, + request_kwargs: dict | None = None, ): if request_kwargs is None: request_kwargs = {} @@ -156,7 +156,7 @@ class LAR1RoutingStrategy(CustomRoutingStrategyBase): self, target_type: str, deployments: list[dict], - ) -> tuple[Optional[dict], bool]: + ) -> tuple[dict | None, bool]: if not deployments: return None, False diff --git a/litellm/router_strategy/least_busy.py b/litellm/router_strategy/least_busy.py index 819fedde991..5fca0b568c6 100644 --- a/litellm/router_strategy/least_busy.py +++ b/litellm/router_strategy/least_busy.py @@ -7,7 +7,6 @@ # - in get_available_deployment, for a given model group name -> pick based on traffic import random -from typing import Optional from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger @@ -63,7 +62,7 @@ class LeastBusyLoggingHandler(CustomLogger): request_count_api_key = f"{model_group}_request_count" # decrement count in cache request_count_dict = self.router_cache.get_cache(key=request_count_api_key) or {} - request_count_value: Optional[int] = request_count_dict.get(id, 0) + request_count_value: int | None = request_count_dict.get(id, 0) if request_count_value is None: return request_count_dict[id] = request_count_value - 1 @@ -90,7 +89,7 @@ class LeastBusyLoggingHandler(CustomLogger): request_count_api_key = f"{model_group}_request_count" # decrement count in cache request_count_dict = self.router_cache.get_cache(key=request_count_api_key) or {} - request_count_value: Optional[int] = request_count_dict.get(id, 0) + request_count_value: int | None = request_count_dict.get(id, 0) if request_count_value is None: return request_count_dict[id] = request_count_value - 1 @@ -118,7 +117,7 @@ class LeastBusyLoggingHandler(CustomLogger): request_count_api_key = f"{model_group}_request_count" # decrement count in cache request_count_dict = await self.router_cache.async_get_cache(key=request_count_api_key) or {} - request_count_value: Optional[int] = request_count_dict.get(id, 0) + request_count_value: int | None = request_count_dict.get(id, 0) if request_count_value is None: return request_count_dict[id] = request_count_value - 1 @@ -145,7 +144,7 @@ class LeastBusyLoggingHandler(CustomLogger): request_count_api_key = f"{model_group}_request_count" # decrement count in cache request_count_dict = await self.router_cache.async_get_cache(key=request_count_api_key) or {} - request_count_value: Optional[int] = request_count_dict.get(id, 0) + request_count_value: int | None = request_count_dict.get(id, 0) if request_count_value is None: return request_count_dict[id] = request_count_value - 1 diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 67bdbfbe0ac..12820ae1237 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -1,7 +1,6 @@ #### What this does #### # picks based on response time (for streaming, this is time to first token) from datetime import datetime, timedelta -from typing import Dict, List, Optional, Union import litellm from litellm import ModelResponse, token_counter, verbose_logger @@ -92,9 +91,8 @@ class LowestCostLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.router_strategy.lowest_cost.py::log_success_event(): Exception occured - {}".format(str(e)) + f"litellm.router_strategy.lowest_cost.py::log_success_event(): Exception occured - {e!s}" ) - pass async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: @@ -172,19 +170,16 @@ class LowestCostLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {e!s}" ) - pass async def async_get_available_deployments( self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - request_kwargs: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + request_kwargs: dict | None = None, ): """ Returns a deployment with the lowest cost diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index ffe8245b012..2f73450b8d2 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -2,14 +2,13 @@ # picks based on response time (for streaming, this is time to first token) import random from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union import litellm from litellm import ModelResponse, token_counter, verbose_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import safe_divide_seconds -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds from litellm.types.utils import LiteLLMPydanticObjectBase if TYPE_CHECKING: @@ -87,7 +86,7 @@ class LowestLatencyLoggingHandler(CustomLogger): time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time final_value: float = response_ms - time_to_first_token: Optional[float] = None + time_to_first_token: float | None = None total_tokens = 0 if isinstance(response_obj, ModelResponse): @@ -161,11 +160,8 @@ class LowestLatencyLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {e!s}" ) - pass async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """ @@ -221,11 +217,8 @@ class LowestLatencyLoggingHandler(CustomLogger): return except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - {e!s}" ) - pass async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: @@ -280,7 +273,7 @@ class LowestLatencyLoggingHandler(CustomLogger): final_value: float = response_ms total_tokens = 0 - time_to_first_token: Optional[float] = None + time_to_first_token: float | None = None if isinstance(response_obj, ModelResponse): _usage = getattr(response_obj, "usage", None) @@ -357,20 +350,17 @@ class LowestLatencyLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.router_strategy.lowest_latency.py::async_log_success_event(): Exception occured - {}".format( - str(e) - ) + f"litellm.router_strategy.lowest_latency.py::async_log_success_event(): Exception occured - {e!s}" ) - pass def _get_available_deployments( self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - request_kwargs: Optional[Dict] = None, - request_count_dict: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + request_kwargs: dict | None = None, + request_count_dict: dict | None = None, ): """Common logic for both sync and async get_available_deployments""" @@ -504,14 +494,14 @@ class LowestLatencyLoggingHandler(CustomLogger): self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - request_kwargs: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + request_kwargs: dict | None = None, ): # get list of potential deployments latency_key = f"{model_group}_map" - parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(request_kwargs) request_count_dict = ( await self.router_cache.async_get_cache(key=latency_key, parent_otel_span=parent_otel_span) or {} ) @@ -529,9 +519,9 @@ class LowestLatencyLoggingHandler(CustomLogger): self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - request_kwargs: Optional[Dict] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + request_kwargs: dict | None = None, ): """ Returns a deployment with the lowest latency @@ -539,7 +529,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # get list of potential deployments latency_key = f"{model_group}_map" - parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Span | None = _get_parent_otel_span_from_kwargs(request_kwargs) request_count_dict = self.router_cache.get_cache(key=latency_key, parent_otel_span=parent_otel_span) or {} return self._get_available_deployments( diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index 89ad7526f20..4a4352fe19d 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -2,7 +2,6 @@ # identifies lowest tpm deployment import traceback from datetime import datetime -from typing import Dict, List, Optional, Union from litellm import token_counter from litellm._logging import verbose_router_logger @@ -74,12 +73,9 @@ class LowestTPMLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_router_logger.error( - "litellm.router_strategy.lowest_tpm_rpm.py::async_log_success_event(): Exception occured - {}".format( - str(e) - ) + f"litellm.router_strategy.lowest_tpm_rpm.py::async_log_success_event(): Exception occured - {e!s}" ) verbose_router_logger.debug(traceback.format_exc()) - pass async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: @@ -139,19 +135,16 @@ class LowestTPMLoggingHandler(CustomLogger): self.logged_success += 1 except Exception as e: verbose_router_logger.exception( - "litellm.router_strategy.lowest_tpm_rpm.py::async_log_success_event(): Exception occured - {}".format( - str(e) - ) + f"litellm.router_strategy.lowest_tpm_rpm.py::async_log_success_event(): Exception occured - {e!s}" ) verbose_router_logger.debug(traceback.format_exc()) - pass def get_available_deployments( self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, ): """ Returns a deployment with the lowest TPM/RPM usage. @@ -222,9 +215,11 @@ class LowestTPMLoggingHandler(CustomLogger): if _deployment_rpm is None: _deployment_rpm = float("inf") - if item_tpm + input_tokens > _deployment_tpm: - continue - elif (rpm_dict is not None and item in rpm_dict) and (rpm_dict[item] + 1 >= _deployment_rpm): + if ( + item_tpm + input_tokens > _deployment_tpm + or (rpm_dict is not None and item in rpm_dict) + and (rpm_dict[item] + 1 >= _deployment_rpm) + ): continue elif item_tpm < lowest_tpm: lowest_tpm = item_tpm diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index d63621a7d77..03793c5577c 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -1,7 +1,7 @@ #### What this does #### # identifies lowest tpm deployment import random -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -57,7 +57,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): default_sync_interval=0.1, ) - def pre_call_check(self, deployment: Dict) -> Optional[Dict]: + def pre_call_check(self, deployment: dict) -> dict | None: """ Pre-call check + update model rpm @@ -90,9 +90,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): if local_result is not None and local_result >= deployment_rpm: raise litellm.RateLimitError( - message="Deployment over defined rpm limit={}. current usage={}".format( - deployment_rpm, local_result - ), + message=f"Deployment over defined rpm limit={deployment_rpm}. current usage={local_result}", llm_provider="", model=deployment.get("litellm_params", {}).get("model"), response=httpx.Response( @@ -116,16 +114,12 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): result = self.router_cache.increment_cache(key=rpm_key, value=1, ttl=self.routing_args.ttl) if result is not None and result > deployment_rpm: raise litellm.RateLimitError( - message="Deployment over defined rpm limit={}. current usage={}".format(deployment_rpm, result), + message=f"Deployment over defined rpm limit={deployment_rpm}. current usage={result}", llm_provider="", model=deployment.get("litellm_params", {}).get("model"), response=httpx.Response( status_code=429, - content="{} rpm limit={}. current usage={}".format( - RouterErrors.user_defined_ratelimit_error.value, - deployment_rpm, - result, - ), + content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={deployment_rpm}. current usage={result}", request=httpx.Request( method="tpm_rpm_limits", url="https://github.com/BerriAI/litellm", @@ -138,7 +132,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): raise e return deployment # don't fail calls if eg. redis fails to connect - async def async_pre_call_check(self, deployment: Dict, parent_otel_span: Optional[Span]) -> Optional[Dict]: + async def async_pre_call_check(self, deployment: dict, parent_otel_span: Span | None) -> dict | None: """ Pre-call check + update model rpm - Used inside semaphore @@ -175,18 +169,12 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): deployment_rpm = float("inf") if local_result is not None and local_result >= deployment_rpm: raise litellm.RateLimitError( - message="Deployment over defined rpm limit={}. current usage={}".format( - deployment_rpm, local_result - ), + message=f"Deployment over defined rpm limit={deployment_rpm}. current usage={local_result}", llm_provider="", model=deployment.get("litellm_params", {}).get("model"), response=httpx.Response( status_code=429, - content="{} rpm limit={}. current usage={}".format( - RouterErrors.user_defined_ratelimit_error.value, - deployment_rpm, - local_result, - ), + content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={deployment_rpm}. current usage={local_result}", headers={"retry-after": str(60)}, # type: ignore request=httpx.Request( method="tpm_rpm_limits", @@ -200,16 +188,12 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): result = await self._increment_value_in_current_window(key=rpm_key, value=1, ttl=self.routing_args.ttl) if result is not None and result > deployment_rpm: raise litellm.RateLimitError( - message="Deployment over defined rpm limit={}. current usage={}".format(deployment_rpm, result), + message=f"Deployment over defined rpm limit={deployment_rpm}. current usage={result}", llm_provider="", model=deployment.get("litellm_params", {}).get("model"), response=httpx.Response( status_code=429, - content="{} rpm limit={}. current usage={}".format( - RouterErrors.user_defined_ratelimit_error.value, - deployment_rpm, - result, - ), + content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={deployment_rpm}. current usage={result}", headers={"retry-after": str(60)}, # type: ignore request=httpx.Request( method="tpm_rpm_limits", @@ -229,7 +213,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): """ Update TPM/RPM usage on success """ - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_object is None: raise ValueError("standard_logging_object not passed in.") model_group = standard_logging_object.get("model_group") @@ -261,16 +245,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.lowest_tpm_rpm_v2.py::log_success_event(): Exception occured - {}".format(str(e)) + f"litellm.proxy.hooks.lowest_tpm_rpm_v2.py::log_success_event(): Exception occured - {e!s}" ) - pass async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: """ Update TPM usage on success """ - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_object is None: raise ValueError("standard_logging_object not passed in.") model_group = standard_logging_object.get("model_group") @@ -306,18 +289,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): self.logged_success += 1 except Exception as e: verbose_logger.exception( - "litellm.proxy.hooks.lowest_tpm_rpm_v2.py::async_log_success_event(): Exception occured - {}".format( - str(e) - ) + f"litellm.proxy.hooks.lowest_tpm_rpm_v2.py::async_log_success_event(): Exception occured - {e!s}" ) - pass def _return_potential_deployments( self, - healthy_deployments: List[Dict], - all_deployments: Dict, + healthy_deployments: list[dict], + all_deployments: dict, input_tokens: int, - rpm_dict: Dict, + rpm_dict: dict, ): lowest_tpm = float("inf") potential_deployments = [] # if multiple deployments have the same low value @@ -352,9 +332,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): _deployment_rpm = _deployment.get("model_info", {}).get("rpm") if _deployment_rpm is None: _deployment_rpm = float("inf") - if item_tpm + input_tokens > _deployment_tpm: - continue - elif ( + if item_tpm + input_tokens > _deployment_tpm or ( (rpm_dict is not None and item in rpm_dict) and rpm_dict[item] is not None and (rpm_dict[item] + 1 >= _deployment_rpm) @@ -372,12 +350,12 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): model_group: str, healthy_deployments: list, tpm_keys: list, - tpm_values: Optional[list], + tpm_values: list | None, rpm_keys: list, - rpm_values: Optional[list], - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - ) -> Optional[dict]: + rpm_values: list | None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + ) -> dict | None: """ Common checks for get available deployment, across sync + async implementations """ @@ -432,8 +410,8 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, ): """ Async implementation of get deployments. @@ -456,8 +434,8 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): "id" ) # a deployment should always have an 'id'. this is set in router.py deployment_name = m.get("litellm_params", {}).get("model") - tpm_key = "{}:{}:tpm:{}".format(id, deployment_name, current_minute) - rpm_key = "{}:{}:rpm:{}".format(id, deployment_name, current_minute) + tpm_key = f"{id}:{deployment_name}:tpm:{current_minute}" + rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}" tpm_keys.append(tpm_key) rpm_keys.append(rpm_key) @@ -548,9 +526,9 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): self, model_group: str, healthy_deployments: list, - messages: Optional[List[Dict[str, str]]] = None, - input: Optional[Union[str, List]] = None, - parent_otel_span: Optional[Span] = None, + messages: list[dict[str, str]] | None = None, + input: str | list | None = None, + parent_otel_span: Span | None = None, ): """ Returns a deployment with the lowest TPM/RPM usage. @@ -570,8 +548,8 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): "id" ) # a deployment should always have an 'id'. this is set in router.py deployment_name = m.get("litellm_params", {}).get("model") - tpm_key = "{}:{}:tpm:{}".format(id, deployment_name, current_minute) - rpm_key = "{}:{}:rpm:{}".format(id, deployment_name, current_minute) + tpm_key = f"{id}:{deployment_name}:tpm:{current_minute}" + rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}" tpm_keys.append(tpm_key) rpm_keys.append(rpm_key) diff --git a/litellm/router_strategy/quality_router/__init__.py b/litellm/router_strategy/quality_router/__init__.py index 5728943448a..86d8f88ac5a 100644 --- a/litellm/router_strategy/quality_router/__init__.py +++ b/litellm/router_strategy/quality_router/__init__.py @@ -14,8 +14,8 @@ from .config import ( from .quality_router import QualityRouter __all__ = [ + "DEFAULT_COMPLEXITY_TO_QUALITY", "QualityRouter", "QualityRouterConfig", "RoutingPreferences", - "DEFAULT_COMPLEXITY_TO_QUALITY", ] diff --git a/litellm/router_strategy/quality_router/config.py b/litellm/router_strategy/quality_router/config.py index 125ecd5bb9b..89a36415ad2 100644 --- a/litellm/router_strategy/quality_router/config.py +++ b/litellm/router_strategy/quality_router/config.py @@ -2,13 +2,11 @@ Configuration models for the QualityRouter. """ -from typing import Dict, List, Optional - from pydantic import BaseModel, ConfigDict, Field # Default mapping from ComplexityTier name (string) to quality tier (int). # Higher tier = higher capability requirement. -DEFAULT_COMPLEXITY_TO_QUALITY: Dict[str, int] = { +DEFAULT_COMPLEXITY_TO_QUALITY: dict[str, int] = { "SIMPLE": 1, "MEDIUM": 2, "COMPLEX": 3, @@ -19,7 +17,7 @@ DEFAULT_COMPLEXITY_TO_QUALITY: Dict[str, int] = { class QualityRouterConfig(BaseModel): """Configuration for the QualityRouter.""" - available_models: List[str] = Field( + available_models: list[str] = Field( default_factory=list, description=( "List of candidate model names this router may route to. Each model " @@ -27,12 +25,12 @@ class QualityRouterConfig(BaseModel): ), ) - default_model: Optional[str] = Field( + default_model: str | None = Field( default=None, description="Fallback model when no quality tier resolves.", ) - complexity_to_quality: Dict[str, int] = Field( + complexity_to_quality: dict[str, int] = Field( default_factory=lambda: DEFAULT_COMPLEXITY_TO_QUALITY.copy(), description="Mapping from ComplexityTier name to quality tier (int).", ) @@ -48,7 +46,7 @@ class RoutingPreferences(BaseModel): description="The quality tier this deployment satisfies.", ) - keywords: List[str] = Field( + keywords: list[str] = Field( default_factory=list, description=( "Substring keywords (case-insensitive) that, when present in the " @@ -58,7 +56,7 @@ class RoutingPreferences(BaseModel): ), ) - order: Optional[int] = Field( + order: int | None = Field( default=None, description=( "Explicit priority used to break ties between deployments at the " diff --git a/litellm/router_strategy/quality_router/quality_router.py b/litellm/router_strategy/quality_router/quality_router.py index fd4a91d76bd..da6825a5741 100644 --- a/litellm/router_strategy/quality_router/quality_router.py +++ b/litellm/router_strategy/quality_router/quality_router.py @@ -16,7 +16,7 @@ then cheapest `model_info.input_cost_per_token`). """ import math -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger @@ -45,8 +45,8 @@ class QualityRouter(CustomLogger): self, model_name: str, litellm_router_instance: "Router", - default_model: Optional[str] = None, - quality_router_config: Optional[Dict[str, Any]] = None, + default_model: str | None = None, + quality_router_config: dict[str, Any] | None = None, ): self.model_name = model_name self.litellm_router_instance = litellm_router_instance @@ -71,10 +71,10 @@ class QualityRouter(CustomLogger): # lowercased user message in O(total-keyword-count). `_model_quality`, # `_model_cost`, and `_model_order` drive tiebreaking — `_model_order` # is the explicit priority (lower wins, unset = +inf). - self._model_keywords: Dict[str, List[str]] = {} - self._model_quality: Dict[str, int] = {} - self._model_cost: Dict[str, Optional[float]] = {} - self._model_order: Dict[str, Optional[int]] = {} + self._model_keywords: dict[str, list[str]] = {} + self._model_quality: dict[str, int] = {} + self._model_cost: dict[str, float | None] = {} + self._model_order: dict[str, int | None] = {} # Tier → models index. Built lazily on first access so the QualityRouter # deployment does NOT need to appear after all its referenced models in @@ -82,7 +82,7 @@ class QualityRouter(CustomLogger): # router instance's `model_list` is still being assembled incrementally # by `_create_deployment`, and any `available_models` defined AFTER the # router entry in config.yaml would silently be reported as missing. - self._tier_to_models_cache: Optional[Dict[int, List[str]]] = None + self._tier_to_models_cache: dict[int, list[str]] | None = None verbose_router_logger.debug( f"QualityRouter initialized for {model_name} with " @@ -91,13 +91,13 @@ class QualityRouter(CustomLogger): ) @property - def _tier_to_models(self) -> Dict[int, List[str]]: + def _tier_to_models(self) -> dict[int, list[str]]: """Lazy tier→models index; built on first access.""" if self._tier_to_models_cache is None: self._tier_to_models_cache = self._build_tier_index() return self._tier_to_models_cache - def _get_routing_preferences(self, deployment: Any) -> Optional[Dict[str, Any]]: + def _get_routing_preferences(self, deployment: Any) -> dict[str, Any] | None: """ Extract litellm_routing_preferences from a deployment, handling both dict-shaped and Pydantic-object-shaped deployments. @@ -118,7 +118,7 @@ class QualityRouter(CustomLogger): return model_info.get("litellm_routing_preferences") return getattr(model_info, "litellm_routing_preferences", None) - def _get_deployment_input_cost(self, deployment: Any) -> Optional[float]: + def _get_deployment_input_cost(self, deployment: Any) -> float | None: """ Extract `input_cost_per_token` from a deployment's model_info. @@ -143,13 +143,13 @@ class QualityRouter(CustomLogger): except (TypeError, ValueError): return None - def _get_deployment_model_name(self, deployment: Any) -> Optional[str]: + def _get_deployment_model_name(self, deployment: Any) -> str | None: """Extract `model_name` from a dict- or object-shaped deployment.""" if isinstance(deployment, dict): return deployment.get("model_name") return getattr(deployment, "model_name", None) - def _build_tier_index(self) -> Dict[int, List[str]]: + def _build_tier_index(self) -> dict[int, list[str]]: """ Build {quality_tier: [model_name, ...]} for every model in `available_models`, plus side indices `_model_keywords`, @@ -160,8 +160,8 @@ class QualityRouter(CustomLogger): available = set(self.config.available_models) # Track which available models we've matched so we can error on missing. - seen: Dict[str, bool] = {name: False for name in available} - tier_to_models: Dict[int, List[str]] = {} + seen: dict[str, bool] = {name: False for name in available} + tier_to_models: dict[int, list[str]] = {} for deployment in model_list: name = self._get_deployment_model_name(deployment) @@ -226,7 +226,7 @@ class QualityRouter(CustomLogger): cost = self._model_cost.get(model_name) return float(cost) if cost is not None else math.inf - def _keyword_override(self, user_message: str) -> Optional[Tuple[str, str]]: + def _keyword_override(self, user_message: str) -> tuple[str, str] | None: """ Find a deployment whose declared keywords appear in `user_message`. @@ -244,7 +244,7 @@ class QualityRouter(CustomLogger): text = user_message.lower() - matches: List[Tuple[str, str]] = [] # (model_name, matched_keyword) + matches: list[tuple[str, str]] = [] # (model_name, matched_keyword) for model_name, keywords in self._model_keywords.items(): for kw in keywords: if kw and kw in text: @@ -254,7 +254,7 @@ class QualityRouter(CustomLogger): if not matches: return None - def sort_key(match: Tuple[str, str]) -> Tuple[int, float, float, str]: + def sort_key(match: tuple[str, str]) -> tuple[int, float, float, str]: name = match[0] quality = self._model_quality.get(name, 0) order_val = self._order_key(name) @@ -303,8 +303,8 @@ class QualityRouter(CustomLogger): def _stash_decision( self, - request_kwargs: Optional[Dict[str, Any]], - decision: Dict[str, Any], + request_kwargs: dict[str, Any] | None, + decision: dict[str, Any], ) -> None: """ Stash the routing decision in request_kwargs.metadata so the Router can @@ -320,10 +320,10 @@ class QualityRouter(CustomLogger): async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, + request_kwargs: dict, + messages: list[dict[str, Any]] | None = None, + input: str | list | None = None, + specific_deployment: bool | None = False, ) -> Optional["PreRoutingHookResponse"]: """Try keyword override first; fall back to complexity-tier routing.""" from litellm.types.router import PreRoutingHookResponse @@ -334,8 +334,8 @@ class QualityRouter(CustomLogger): # Extract last user message and last system prompt — same rules as # ComplexityRouter.async_pre_routing_hook. - user_message: Optional[str] = None - system_prompt: Optional[str] = None + user_message: str | None = None + system_prompt: str | None = None for msg in reversed(messages): role = msg.get("role", "") diff --git a/litellm/router_strategy/simple_shuffle.py b/litellm/router_strategy/simple_shuffle.py index d3349ee29ce..ab2fab09a0d 100644 --- a/litellm/router_strategy/simple_shuffle.py +++ b/litellm/router_strategy/simple_shuffle.py @@ -6,7 +6,7 @@ If weights are provided, it will return a deployment based on the weights. """ import random -from typing import TYPE_CHECKING, Any, Dict, List, Union +from typing import TYPE_CHECKING, Any from litellm._logging import verbose_router_logger @@ -20,9 +20,9 @@ else: def simple_shuffle( llm_router_instance: LitellmRouter, - healthy_deployments: Union[List[Any], Dict[Any, Any]], + healthy_deployments: list[Any] | dict[Any, Any], model: str, -) -> Dict: +) -> dict: """ Returns a random deployment from the list of healthy deployments. diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 710c2199107..5a97df14332 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -7,7 +7,7 @@ Use this to route requests between Teams """ import re -from typing import TYPE_CHECKING, Any, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Literal from litellm._logging import verbose_logger from litellm.types.router import RouterErrors @@ -23,7 +23,7 @@ else: def _is_valid_deployment_tag_regex( tag_regexes: list[str], header_strings: list[str], -) -> Optional[str]: +) -> str | None: """ Test compiled regex patterns against "Header-Name: value" strings. @@ -71,10 +71,10 @@ def is_valid_deployment_tag(deployment_tags: list[str], request_tags: list[str], def _match_deployment( deployment: Any, - request_tags: Optional[list[str]], + request_tags: list[str] | None, header_strings: list[str], match_any: bool, -) -> Optional[dict[str, str]]: +) -> dict[str, str] | None: """ Determine whether *deployment* matches the current request. @@ -87,8 +87,8 @@ def _match_deployment( ran and failed, so the regex cannot override strict-tag policy. """ litellm_params = deployment.get("litellm_params", {}) - deployment_tags: Optional[list[str]] = litellm_params.get("tags") - deployment_tag_regex: Optional[list[str]] = litellm_params.get("tag_regex") + deployment_tags: list[str] | None = litellm_params.get("tags") + deployment_tag_regex: list[str] | None = litellm_params.get("tag_regex") # 1. Exact tag match (existing behaviour). if deployment_tags and request_tags: @@ -121,7 +121,7 @@ def _split_tags(tags: list[str]) -> tuple[list[str], list[str]]: def _exclude_deployments( - deployments: Union[list[Any], dict[Any, Any]], + deployments: list[Any] | dict[Any, Any], excluded_set: frozenset[str], ) -> list[Any]: if not excluded_set: @@ -142,7 +142,7 @@ def _require_candidates( def _ban_only_base_pool( - deployments: Union[list[Any], dict[Any, Any]], + deployments: list[Any] | dict[Any, Any], ) -> list[Any]: # Mirrors untagged-request semantics so callers can't use !tags to escape the default pool. defaults = [d for d in deployments if "default" in (d.get("litellm_params", {}).get("tags") or [])] @@ -152,8 +152,8 @@ def _ban_only_base_pool( async def get_deployments_for_tag( llm_router_instance: LitellmRouter, model: str, # used to raise the correct error - healthy_deployments: Union[list[Any], dict[Any, Any]], - request_kwargs: Optional[dict[Any, Any]] = None, + healthy_deployments: list[Any] | dict[Any, Any], + request_kwargs: dict[Any, Any] | None = None, metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata", ): """ @@ -270,7 +270,7 @@ async def get_deployments_for_tag( def _get_tags_from_request_kwargs( - request_kwargs: Optional[dict[Any, Any]] = None, + request_kwargs: dict[Any, Any] | None = None, metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata", ) -> list[str]: """ diff --git a/litellm/router_utils/batch_utils.py b/litellm/router_utils/batch_utils.py index b09c5eada79..eefacb7c9e6 100644 --- a/litellm/router_utils/batch_utils.py +++ b/litellm/router_utils/batch_utils.py @@ -1,7 +1,6 @@ import io import json from os import PathLike -from typing import List, Optional from litellm._logging import verbose_logger from litellm.types.llms.openai import FileTypes, OpenAIFilesPurpose @@ -14,7 +13,7 @@ class InMemoryFile(io.BytesIO): self.content_type = content_type -def parse_jsonl_with_embedded_newlines(content: str) -> List[dict]: +def parse_jsonl_with_embedded_newlines(content: str) -> list[dict]: """ Parse JSONL content that may contain JSON objects with embedded newlines in string values. @@ -149,7 +148,7 @@ def replace_model_in_jsonl(file_content: FileTypes, new_model_name: str) -> File return file_content -def _get_router_metadata_variable_name(function_name: Optional[str]) -> str: +def _get_router_metadata_variable_name(function_name: str | None) -> str: """ Helper to return what the "metadata" field should be called in the request data diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index 8234d89e248..d635892e21b 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -11,12 +11,10 @@ If given, generate a unique model_id for the deployment. Ensures cooldowns are applied correctly. """ -from typing import List - clientside_credential_keys = ["api_key", "api_base", "base_url"] -def _admin_config_fields_to_clear_on_base_override() -> List[str]: +def _admin_config_fields_to_clear_on_base_override() -> list[str]: """ Provider-specific credential / endpoint-targeting fields that must NOT flow through to a client-redirected upstream. diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 5cfea5e3bf2..189296a8955 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -1,17 +1,17 @@ import hashlib import json from collections.abc import Mapping -from typing import TYPE_CHECKING, Dict, List, Optional, Union +from typing import TYPE_CHECKING if TYPE_CHECKING: from litellm.types.llms.openai import OpenAIFileObject +from litellm._logging import verbose_logger from litellm.exceptions import BadRequestError from litellm.types.router import CredentialLiteLLMParams -from litellm._logging import verbose_logger -def _is_proxy_admin_request(request_kwargs: Optional[Mapping[str, object]]) -> bool: +def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool: if request_kwargs is None: return False metadata_value = request_kwargs.get("metadata") @@ -30,9 +30,7 @@ def get_litellm_params_sensitive_credential_hash(litellm_params: dict) -> str: return hashlib.sha256(json.dumps(sensitive_params.model_dump()).encode()).hexdigest() -def add_model_file_id_mappings( - healthy_deployments: Union[List[Dict], Dict], responses: List["OpenAIFileObject"] -) -> dict: +def add_model_file_id_mappings(healthy_deployments: list[dict] | dict, responses: list["OpenAIFileObject"]) -> dict: """ Create a mapping of model id to file id { @@ -46,8 +44,8 @@ def add_model_file_id_mappings( `model_info.id`). Both shapes must be handled by extracting `model_info.id` from each deployment. """ - model_file_id_mapping: Dict[str, str] = {} - deployments_list: List[Dict] = ( + model_file_id_mapping: dict[str, str] = {} + deployments_list: list[dict] = ( healthy_deployments if isinstance(healthy_deployments, list) else [healthy_deployments] ) for deployment, response in zip(deployments_list, responses): @@ -58,9 +56,9 @@ def add_model_file_id_mappings( def filter_team_based_models( - healthy_deployments: Union[List[Dict], Dict], - request_kwargs: Optional[Dict] = None, -) -> Union[List[Dict], Dict]: + healthy_deployments: list[dict] | dict, + request_kwargs: dict | None = None, +) -> list[dict] | dict: """ If a model has a team_id @@ -124,7 +122,7 @@ def filter_team_based_models( ] -def _deployment_supports_web_search(deployment: Dict) -> bool: +def _deployment_supports_web_search(deployment: dict) -> bool: """ Check if a deployment supports web search. @@ -145,9 +143,9 @@ def _deployment_supports_web_search(deployment: Dict) -> bool: def filter_web_search_deployments( - healthy_deployments: Union[List[Dict], Dict], - request_kwargs: Optional[Dict] = None, -) -> Union[List[Dict], Dict]: + healthy_deployments: list[dict] | dict, + request_kwargs: dict | None = None, +) -> list[dict] | dict: """ If the request is websearch, filter out deployments that don't support web search """ diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index 3f3284b315c..ef62a5d8c6c 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic import functools import time -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Union from typing_extensions import TypedDict @@ -43,7 +43,7 @@ class CooldownCache: def _common_add_cooldown_logic( self, model_id: str, original_exception, exception_status, cooldown_time: float - ) -> Tuple[str, CooldownCacheValue]: + ) -> tuple[str, CooldownCacheValue]: try: current_time = time.time() cooldown_key = CooldownCache.get_cooldown_cache_key(model_id) @@ -58,7 +58,7 @@ class CooldownCache: return cooldown_key, cooldown_data except Exception as e: - verbose_logger.error("CooldownCache::_common_add_cooldown_logic - Exception occurred - {}".format(str(e))) + verbose_logger.error(f"CooldownCache::_common_add_cooldown_logic - Exception occurred - {e!s}") raise e def add_deployment_to_cooldown( @@ -66,7 +66,7 @@ class CooldownCache: model_id: str, original_exception: Exception, exception_status: int, - cooldown_time: Optional[float], + cooldown_time: float | None, ): try: ######################################################### @@ -92,7 +92,7 @@ class CooldownCache: ttl=_cooldown_time, ) except Exception as e: - verbose_logger.error("CooldownCache::add_deployment_to_cooldown - Exception occurred - {}".format(str(e))) + verbose_logger.error(f"CooldownCache::add_deployment_to_cooldown - Exception occurred - {e!s}") raise e @staticmethod @@ -101,8 +101,8 @@ class CooldownCache: return "deployment:" + model_id + ":cooldown" async def async_get_active_cooldowns( - self, model_ids: List[str], parent_otel_span: Optional[Span] - ) -> List[Tuple[str, CooldownCacheValue]]: + self, model_ids: list[str], parent_otel_span: Span | None + ) -> list[tuple[str, CooldownCacheValue]]: # Generate the keys for the deployments keys = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] @@ -112,7 +112,7 @@ class CooldownCache: ## check in memory cache first results = await self.cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) - active_cooldowns: List[Tuple[str, CooldownCacheValue]] = [] + active_cooldowns: list[tuple[str, CooldownCacheValue]] = [] if results is None or all(v is None for v in results): return active_cooldowns @@ -126,8 +126,8 @@ class CooldownCache: return active_cooldowns def get_active_cooldowns( - self, model_ids: List[str], parent_otel_span: Optional[Span] - ) -> List[Tuple[str, CooldownCacheValue]]: + self, model_ids: list[str], parent_otel_span: Span | None + ) -> list[tuple[str, CooldownCacheValue]]: # Generate the keys for the deployments keys = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] # Retrieve the values for the keys using mget @@ -142,7 +142,7 @@ class CooldownCache: return active_cooldowns - def get_min_cooldown(self, model_ids: List[str], parent_otel_span: Optional[Span]) -> float: + def get_min_cooldown(self, model_ids: list[str], parent_otel_span: Span | None) -> float: """Return min cooldown time required for a group of model id's.""" # Generate the keys for the deployments @@ -151,14 +151,12 @@ class CooldownCache: # Retrieve the values for the keys using mget results = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or [] - min_cooldown_time: Optional[float] = None + min_cooldown_time: float | None = None # Process the results for model_id, result in zip(model_ids, results): if result and isinstance(result, dict): cooldown_cache_value = CooldownCacheValue(**result) # type: ignore - if min_cooldown_time is None: - min_cooldown_time = cooldown_cache_value["cooldown_time"] - elif cooldown_cache_value["cooldown_time"] < min_cooldown_time: + if min_cooldown_time is None or cooldown_cache_value["cooldown_time"] < min_cooldown_time: min_cooldown_time = cooldown_cache_value["cooldown_time"] return min_cooldown_time or self.default_cooldown_time diff --git a/litellm/router_utils/cooldown_callbacks.py b/litellm/router_utils/cooldown_callbacks.py index 313037a6364..acd1c5b47ad 100644 --- a/litellm/router_utils/cooldown_callbacks.py +++ b/litellm/router_utils/cooldown_callbacks.py @@ -3,7 +3,7 @@ Callbacks triggered on cooling down deployments """ import copy -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_logger @@ -21,8 +21,8 @@ else: async def router_cooldown_event_callback( litellm_router_instance: LitellmRouter, deployment_id: str, - exception_status: Union[str, int], - cooldown_time: Optional[float], + exception_status: str | int, + cooldown_time: float | None, ): """ Callback triggered when a deployment is put into cooldown by litellm @@ -56,7 +56,7 @@ async def router_cooldown_event_callback( pass # get the prometheus logger from in memory loggers - prometheusLogger: Optional[PrometheusLogger] = _get_prometheus_logger_from_callbacks() + prometheusLogger: PrometheusLogger | None = _get_prometheus_logger_from_callbacks() if prometheusLogger is not None: prometheusLogger.set_deployment_complete_outage( @@ -77,7 +77,7 @@ async def router_cooldown_event_callback( return -def _get_prometheus_logger_from_callbacks() -> Optional[PrometheusLogger]: +def _get_prometheus_logger_from_callbacks() -> PrometheusLogger | None: """ Checks if prometheus is a initalized callback, if yes returns it """ diff --git a/litellm/router_utils/cooldown_handlers.py b/litellm/router_utils/cooldown_handlers.py index c1fc939880a..380689e653b 100644 --- a/litellm/router_utils/cooldown_handlers.py +++ b/litellm/router_utils/cooldown_handlers.py @@ -8,7 +8,7 @@ Router cooldown handlers import asyncio import math -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, Union import litellm from litellm._logging import verbose_router_logger @@ -61,8 +61,8 @@ def is_advisor_orchestration_failure(exception: BaseException | None) -> bool: def _is_cooldown_required( litellm_router_instance: LitellmRouter, model_id: str, - exception_status: Union[str, int], - exception_str: Optional[str] = None, + exception_status: str | int, + exception_str: str | None = None, ) -> bool: """ A function to determine if a cooldown is required based on the exception status. @@ -95,10 +95,7 @@ def _is_cooldown_required( # Cool down 401 Auth Errors return True - elif exception_status == 408: - return True - - elif exception_status == 404: + elif exception_status == 408 or exception_status == 404: return True else: @@ -116,10 +113,10 @@ def _is_cooldown_required( def _should_run_cooldown_logic( litellm_router_instance: LitellmRouter, - deployment: Optional[str], - exception_status: Union[str, int], + deployment: str | None, + exception_status: str | int, original_exception: Any, - time_to_cooldown: Optional[float] = None, + time_to_cooldown: float | None = None, ) -> bool: """ Helper that decides if cooldown logic should be run @@ -172,7 +169,7 @@ def _should_run_cooldown_logic( def _should_cooldown_deployment( litellm_router_instance: LitellmRouter, deployment: str, - exception_status: Union[str, int], + exception_status: str | int, original_exception: Any, ) -> bool: """ @@ -252,9 +249,9 @@ def _should_cooldown_deployment( def _set_cooldown_deployments( litellm_router_instance: LitellmRouter, original_exception: Any, - exception_status: Union[str, int], - deployment: Optional[str] = None, - time_to_cooldown: Optional[float] = None, + exception_status: str | int, + deployment: str | None = None, + time_to_cooldown: float | None = None, ) -> bool: """ Add a model to the list of models being cooled down for that minute, if it exceeds the allowed fails / minute @@ -314,8 +311,8 @@ def _set_cooldown_deployments( async def _async_get_cooldown_deployments( litellm_router_instance: LitellmRouter, - parent_otel_span: Optional[Span], -) -> List[str]: + parent_otel_span: Span | None, +) -> list[str]: """ Async implementation of '_get_cooldown_deployments' """ @@ -340,8 +337,8 @@ async def _async_get_cooldown_deployments( async def _async_get_cooldown_deployments_with_debug_info( litellm_router_instance: LitellmRouter, - parent_otel_span: Optional[Span], -) -> List[tuple]: + parent_otel_span: Span | None, +) -> list[tuple]: """ Async implementation of '_get_cooldown_deployments' """ @@ -354,7 +351,7 @@ async def _async_get_cooldown_deployments_with_debug_info( return cooldown_models -def _get_cooldown_deployments(litellm_router_instance: LitellmRouter, parent_otel_span: Optional[Span]) -> List[str]: +def _get_cooldown_deployments(litellm_router_instance: LitellmRouter, parent_otel_span: Span | None) -> list[str]: """ Get the list of models being cooled down for this minute """ @@ -429,7 +426,7 @@ def _is_allowed_fails_set_on_router( return False -def cast_exception_status_to_int(exception_status: Union[str, int]) -> int: +def cast_exception_status_to_int(exception_status: str | int) -> int: if isinstance(exception_status, str): try: exception_status = int(exception_status) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index eb92e9b1c29..0c92a6fa2ab 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any import litellm from litellm._logging import verbose_router_logger @@ -44,7 +44,7 @@ def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool: return False -def get_fallback_model_group(fallbacks: List[Any], model_group: str) -> Tuple[Optional[List[str]], Optional[int]]: +def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[list[str] | None, int | None]: """ Returns: - fallback_model_group: List[str] of fallback model groups. example: ["gpt-4", "gpt-3.5-turbo"] @@ -55,9 +55,9 @@ def get_fallback_model_group(fallbacks: List[Any], model_group: str) -> Tuple[Op - stripped model group match - generic fallback """ - generic_fallback_idx: Optional[int] = None - stripped_model_fallback: Optional[List[str]] = None - fallback_model_group: Optional[List[str]] = None + generic_fallback_idx: int | None = None + stripped_model_fallback: list[str] | None = None + fallback_model_group: list[str] | None = None ## check for specific model group-specific fallbacks for idx, item in enumerate(fallbacks): if isinstance(item, dict): @@ -83,9 +83,9 @@ def get_fallback_model_group(fallbacks: List[Any], model_group: str) -> Tuple[Op async def run_async_fallback( - *args: Tuple[Any], + *args: tuple[Any], litellm_router: LitellmRouter, - fallback_model_group: List[str], + fallback_model_group: list[str], original_model_group: str, original_exception: Exception, max_fallbacks: int, @@ -190,7 +190,7 @@ async def log_success_fallback_event(original_model_group: str, kwargs: dict, or original_exception=original_exception, ) except Exception as e: - verbose_router_logger.error(f"Error in log_success_fallback_event: {str(e)}") + verbose_router_logger.error(f"Error in log_success_fallback_event: {e!s}") async def log_failure_fallback_event(original_model_group: str, kwargs: dict, original_exception: Exception): @@ -218,10 +218,10 @@ async def log_failure_fallback_event(original_model_group: str, kwargs: dict, or original_exception=original_exception, ) except Exception as e: - verbose_router_logger.error(f"Error in log_failure_fallback_event: {str(e)}") + verbose_router_logger.error(f"Error in log_failure_fallback_event: {e!s}") -def _check_non_standard_fallback_format(fallbacks: Optional[List[Any]]) -> bool: +def _check_non_standard_fallback_format(fallbacks: list[Any] | None) -> bool: """ Checks if the fallbacks list is a list of strings or a list of dictionaries. @@ -247,5 +247,5 @@ def _check_non_standard_fallback_format(fallbacks: Optional[List[Any]]) -> bool: return False -def run_non_standard_fallback_format(fallbacks: Union[List[str], List[Dict[str, Any]]], model_group: str): +def run_non_standard_fallback_format(fallbacks: list[str] | list[dict[str, Any]], model_group: str): pass diff --git a/litellm/router_utils/get_retry_from_policy.py b/litellm/router_utils/get_retry_from_policy.py index 314917d3f56..1645e6776fc 100644 --- a/litellm/router_utils/get_retry_from_policy.py +++ b/litellm/router_utils/get_retry_from_policy.py @@ -4,8 +4,6 @@ Get num retries for an exception. - Account for retry policy by exception type. """ -from typing import Dict, Optional, Union - from litellm.exceptions import ( AuthenticationError, BadRequestError, @@ -18,9 +16,9 @@ from litellm.types.router import RetryPolicy def get_num_retries_from_retry_policy( exception: Exception, - retry_policy: Optional[Union[RetryPolicy, dict]] = None, - model_group: Optional[str] = None, - model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = None, + retry_policy: RetryPolicy | dict | None = None, + model_group: str | None = None, + model_group_retry_policy: dict[str, RetryPolicy] | None = None, ): """ BadRequestErrorRetries: Optional[int] = None diff --git a/litellm/router_utils/handle_error.py b/litellm/router_utils/handle_error.py index 05bde3d50d1..b38d6605ed2 100644 --- a/litellm/router_utils/handle_error.py +++ b/litellm/router_utils/handle_error.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union from litellm._logging import redact_secrets, verbose_router_logger from litellm.constants import MAX_EXCEPTION_MESSAGE_LENGTH @@ -69,7 +69,7 @@ async def send_llm_exception_alert( async def async_raise_no_deployment_exception( - litellm_router_instance: LitellmRouter, model: str, parent_otel_span: Optional[Span] + litellm_router_instance: LitellmRouter, model: str, parent_otel_span: Span | None ): """ Raises a RouterRateLimitError if no deployment is found for the given model. diff --git a/litellm/router_utils/health_state_cache.py b/litellm/router_utils/health_state_cache.py index aeb2ba945f5..d9ea7cdbaaf 100644 --- a/litellm/router_utils/health_state_cache.py +++ b/litellm/router_utils/health_state_cache.py @@ -6,7 +6,7 @@ and exposes it for router candidate filtering. """ import time -from typing import TYPE_CHECKING, Any, Dict, Optional, Set, Union +from typing import TYPE_CHECKING, Any, Union from typing_extensions import TypedDict @@ -42,7 +42,7 @@ class DeploymentHealthCache: self.cache = cache self.staleness_threshold = staleness_threshold - def set_deployment_health_states(self, states: Dict[str, DeploymentHealthStateValue]) -> None: + def set_deployment_health_states(self, states: dict[str, DeploymentHealthStateValue]) -> None: """Bulk-write all deployment health states as a single cache entry.""" try: self.cache.set_cache( @@ -56,7 +56,7 @@ class DeploymentHealthCache: str(e), ) - def _extract_unhealthy_ids(self, raw: Any) -> Set[str]: + def _extract_unhealthy_ids(self, raw: Any) -> set[str]: """Given raw cache value, return set of non-stale unhealthy deployment IDs.""" if not raw or not isinstance(raw, dict): return set() @@ -69,7 +69,7 @@ class DeploymentHealthCache: and (now - state.get("timestamp", 0)) < self.staleness_threshold } - async def async_get_unhealthy_deployment_ids(self, parent_otel_span: Optional[Span] = None) -> Set[str]: + async def async_get_unhealthy_deployment_ids(self, parent_otel_span: Span | None = None) -> set[str]: """Return set of deployment IDs currently marked unhealthy and not stale.""" try: raw = await self.cache.async_get_cache(key=self.CACHE_KEY) @@ -81,7 +81,7 @@ class DeploymentHealthCache: ) return set() - def get_unhealthy_deployment_ids(self, parent_otel_span: Optional[Span] = None) -> Set[str]: + def get_unhealthy_deployment_ids(self, parent_otel_span: Span | None = None) -> set[str]: """Sync version: return set of deployment IDs currently marked unhealthy and not stale.""" try: raw = self.cache.get_cache(key=self.CACHE_KEY) diff --git a/litellm/router_utils/pattern_match_deployments.py b/litellm/router_utils/pattern_match_deployments.py index 7e1ed739ef8..004f7b53869 100644 --- a/litellm/router_utils/pattern_match_deployments.py +++ b/litellm/router_utils/pattern_match_deployments.py @@ -5,15 +5,14 @@ Class to handle llm wildcard routing and regex pattern matching import copy import re from re import Match -from typing import Dict, List, Optional, Tuple -from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm._logging import verbose_router_logger +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider class PatternUtils: @staticmethod - def calculate_pattern_specificity(pattern: str) -> Tuple[int, int]: + def calculate_pattern_specificity(pattern: str) -> tuple[int, int]: """ Calculate pattern specificity based on length and complexity. @@ -32,8 +31,8 @@ class PatternUtils: @staticmethod def sorted_patterns( - patterns: Dict[str, List[Dict]], - ) -> List[Tuple[str, List[Dict]]]: + patterns: dict[str, list[dict]], + ) -> list[tuple[str, list[dict]]]: """ Cached property for patterns sorted by specificity. @@ -57,9 +56,9 @@ class PatternMatchRouter: """ def __init__(self): - self.patterns: Dict[str, List] = {} + self.patterns: dict[str, list] = {} - def add_pattern(self, pattern: str, llm_deployment: Dict): + def add_pattern(self, pattern: str, llm_deployment: dict): """ Add a regex pattern and the corresponding llm deployments to the patterns @@ -108,7 +107,7 @@ class PatternMatchRouter: # return f"^{regex}$" return re.escape(pattern).replace(r"\*", "(.*)") - def _return_pattern_matched_deployments(self, matched_pattern: Match, deployments: List[Dict]) -> List[Dict]: + def _return_pattern_matched_deployments(self, matched_pattern: Match, deployments: list[dict]) -> list[dict]: new_deployments = [] for deployment in deployments: new_deployment = copy.deepcopy(deployment) @@ -120,7 +119,7 @@ class PatternMatchRouter: return new_deployments - def route(self, request: Optional[str], filtered_model_names: Optional[List[str]] = None) -> Optional[List[Dict]]: + def route(self, request: str | None, filtered_model_names: list[str] | None = None) -> list[dict] | None: """ Route a requested model to the corresponding llm deployments based on the regex pattern @@ -151,7 +150,7 @@ class PatternMatchRouter: matched_pattern=pattern_match, deployments=llm_deployments ) except Exception as e: - verbose_router_logger.debug(f"Error in PatternMatchRouter.route: {str(e)}") + verbose_router_logger.debug(f"Error in PatternMatchRouter.route: {e!s}") return None # No matching pattern found @@ -204,7 +203,7 @@ class PatternMatchRouter: return litellm_deployment_litellm_model - def get_pattern(self, model: str, custom_llm_provider: Optional[str] = None) -> Optional[List[Dict]]: + def get_pattern(self, model: str, custom_llm_provider: str | None = None) -> list[dict] | None: """ Check if a pattern exists for the given model and custom llm provider @@ -228,7 +227,7 @@ class PatternMatchRouter: pass return self.route(model) or self.route(f"{custom_llm_provider}/{model}") - def get_deployments_by_pattern(self, model: str, custom_llm_provider: Optional[str] = None) -> List[Dict]: + def get_deployments_by_pattern(self, model: str, custom_llm_provider: str | None = None) -> list[dict]: """ Get the deployments by pattern diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index c3f58935ef0..a84cacc7e7d 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -13,7 +13,7 @@ where routing to a consistent deployment is still beneficial. """ import hashlib -from typing import Any, Dict, List, Optional, Tuple, cast +from typing import Any, cast from typing_extensions import TypedDict @@ -54,7 +54,7 @@ class DeploymentAffinityCheck(CustomLogger): enable_user_key_affinity: bool, enable_responses_api_affinity: bool, enable_session_id_affinity: bool = False, - model_group_affinity_config: Optional[Dict[str, List[str]]] = None, + model_group_affinity_config: dict[str, list[str]] | None = None, ): super().__init__() self.cache = cache @@ -62,7 +62,7 @@ class DeploymentAffinityCheck(CustomLogger): self.enable_user_key_affinity = enable_user_key_affinity self.enable_responses_api_affinity = enable_responses_api_affinity self.enable_session_id_affinity = enable_session_id_affinity - self.model_group_affinity_config: Dict[str, List[str]] = model_group_affinity_config or {} + self.model_group_affinity_config: dict[str, list[str]] = model_group_affinity_config or {} for group, flags in self.model_group_affinity_config.items(): unknown = set(flags) - self.VALID_FLAGS if unknown: @@ -73,7 +73,7 @@ class DeploymentAffinityCheck(CustomLogger): self.VALID_FLAGS, ) - def _get_effective_flags(self, model_group: str) -> Tuple[bool, bool, bool]: + def _get_effective_flags(self, model_group: str) -> tuple[bool, bool, bool]: """ Return (enable_user_key_affinity, enable_responses_api_affinity, enable_session_id_affinity) for the given model group. @@ -122,7 +122,7 @@ class DeploymentAffinityCheck(CustomLogger): @staticmethod def _get_model_map_key_from_litellm_model_name( litellm_model_name: str, - ) -> Optional[str]: + ) -> str | None: """ Best-effort derivation of a stable "model map key" for affinity scoping. @@ -148,7 +148,7 @@ class DeploymentAffinityCheck(CustomLogger): return remainder @staticmethod - def _get_model_map_key_from_deployment(deployment: dict) -> Optional[str]: + def _get_model_map_key_from_deployment(deployment: dict) -> str | None: """ Derive a stable model-map key from a router deployment dict. @@ -183,8 +183,8 @@ class DeploymentAffinityCheck(CustomLogger): @staticmethod def _get_stable_model_map_key_from_deployments( - healthy_deployments: List[dict], - ) -> Optional[str]: + healthy_deployments: list[dict], + ) -> str | None: """ Only use model-map key scoping when it is stable across the deployment set. @@ -194,7 +194,7 @@ class DeploymentAffinityCheck(CustomLogger): if not healthy_deployments: return None - keys: List[str] = [] + keys: list[str] = [] for deployment in healthy_deployments: key = DeploymentAffinityCheck._get_model_map_key_from_deployment(deployment) if key is None: @@ -222,7 +222,7 @@ class DeploymentAffinityCheck(CustomLogger): return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{session_id}" @staticmethod - def _get_user_key_from_metadata_dict(metadata: dict) -> Optional[str]: + def _get_user_key_from_metadata_dict(metadata: dict) -> str | None: # NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the # OpenAI `user` parameter, which is an end-user identifier). user_key = metadata.get("user_api_key_hash") @@ -231,21 +231,21 @@ class DeploymentAffinityCheck(CustomLogger): return str(user_key) @staticmethod - def _get_session_id_from_metadata_dict(metadata: dict) -> Optional[str]: + def _get_session_id_from_metadata_dict(metadata: dict) -> str | None: session_id = metadata.get("session_id") if session_id is None: return None return str(session_id) @staticmethod - def _iter_metadata_dicts(request_kwargs: dict) -> List[dict]: + def _iter_metadata_dicts(request_kwargs: dict) -> list[dict]: """ Return all metadata dicts available on the request. Depending on the endpoint, Router may populate `metadata` or `litellm_metadata`. Users may also send one or both, so we check both (rather than using `or`). """ - metadata_dicts: List[dict] = [] + metadata_dicts: list[dict] = [] for key in ("litellm_metadata", "metadata"): md = request_kwargs.get(key) if isinstance(md, dict): @@ -253,7 +253,7 @@ class DeploymentAffinityCheck(CustomLogger): return metadata_dicts @staticmethod - def _get_user_key_from_request_kwargs(request_kwargs: dict) -> Optional[str]: + def _get_user_key_from_request_kwargs(request_kwargs: dict) -> str | None: """ Extract a stable affinity key from request kwargs. @@ -271,7 +271,7 @@ class DeploymentAffinityCheck(CustomLogger): return None @staticmethod - def _get_session_id_from_request_kwargs(request_kwargs: dict) -> Optional[str]: + def _get_session_id_from_request_kwargs(request_kwargs: dict) -> str | None: for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs): session_id = DeploymentAffinityCheck._get_session_id_from_metadata_dict(metadata=metadata) if session_id is not None: @@ -279,7 +279,7 @@ class DeploymentAffinityCheck(CustomLogger): return None @staticmethod - def _find_deployment_by_model_id(healthy_deployments: List[dict], model_id: str) -> Optional[dict]: + def _find_deployment_by_model_id(healthy_deployments: list[dict], model_id: str) -> dict | None: for deployment in healthy_deployments: model_info = deployment.get("model_info") if not isinstance(model_info, dict): @@ -292,18 +292,18 @@ class DeploymentAffinityCheck(CustomLogger): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, + ) -> list[dict]: """ Optionally filter healthy deployments based on: 1. `previous_response_id` (Responses API continuity) [highest priority] 2. cached API-key deployment affinity """ request_kwargs = request_kwargs or {} - typed_healthy_deployments = cast(List[dict], healthy_deployments) + typed_healthy_deployments = cast(list[dict], healthy_deployments) ( enable_user_key, @@ -343,9 +343,9 @@ class DeploymentAffinityCheck(CustomLogger): ) session_cache_result = await self.cache.async_get_cache(key=session_cache_key) - session_model_id: Optional[str] = None + session_model_id: str | None = None if isinstance(session_cache_result, dict): - session_model_id = cast(Optional[str], session_cache_result.get("model_id")) + session_model_id = cast(str | None, session_cache_result.get("model_id")) elif isinstance(session_cache_result, str): session_model_id = session_cache_result @@ -378,9 +378,9 @@ class DeploymentAffinityCheck(CustomLogger): cache_key = self.get_affinity_cache_key(model_group=stable_model_map_key, user_key=user_key) cache_result = await self.cache.async_get_cache(key=cache_key) - model_id: Optional[str] = None + model_id: str | None = None if isinstance(cache_result, dict): - model_id = cast(Optional[str], cache_result.get("model_id")) + model_id = cast(str | None, cache_result.get("model_id")) elif isinstance(cache_result, str): # Backwards / safety: allow raw string values. model_id = cache_result @@ -406,9 +406,7 @@ class DeploymentAffinityCheck(CustomLogger): ) return [deployment] - async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: """ Persist/update the API-key -> deployment mapping for this request. @@ -420,7 +418,7 @@ class DeploymentAffinityCheck(CustomLogger): # Extract deployment_model_name first — needed for both per-group flag resolution # and cache key scoping. - deployment_model_name: Optional[str] = None + deployment_model_name: str | None = None for metadata in metadata_dicts: maybe_deployment_model_name = metadata.get("deployment_model_name") if isinstance(maybe_deployment_model_name, str) and maybe_deployment_model_name: diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py index 9788f5f2299..4a5a0cb0ae1 100644 --- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py @@ -37,7 +37,7 @@ Safe to enable globally: """ import time -from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast +from typing import TYPE_CHECKING, Any, Optional, cast import httpx @@ -72,12 +72,12 @@ class EncryptedContentAffinityCheck(CustomLogger): self, router: Optional["Router"] = None, enable_global_affinity: bool = True, - model_group_affinity_config: Optional[Dict[str, List[str]]] = None, + model_group_affinity_config: dict[str, list[str]] | None = None, ) -> None: super().__init__() self.router = router self.enable_global_affinity = enable_global_affinity - self.model_group_affinity_config: Dict[str, List[str]] = model_group_affinity_config or {} + self.model_group_affinity_config: dict[str, list[str]] = model_group_affinity_config or {} # ------------------------------------------------------------------ # Helpers @@ -85,7 +85,7 @@ class EncryptedContentAffinityCheck(CustomLogger): @staticmethod def has_model_group_affinity_enabled( - model_group_affinity_config: Optional[Dict[str, List[str]]], + model_group_affinity_config: dict[str, list[str]] | None, ) -> bool: if not model_group_affinity_config: return False @@ -99,7 +99,7 @@ class EncryptedContentAffinityCheck(CustomLogger): ) @staticmethod - def _extract_model_id_from_input(request_input: Any) -> Optional[str]: + def _extract_model_id_from_input(request_input: Any) -> str | None: """ Scan ``input`` items for litellm-encoded encrypted-content markers and return the ``model_id`` embedded in the first one found. @@ -139,7 +139,7 @@ class EncryptedContentAffinityCheck(CustomLogger): return None @staticmethod - def _find_deployment_by_model_id(healthy_deployments: List[dict], model_id: str) -> Optional[dict]: + def _find_deployment_by_model_id(healthy_deployments: list[dict], model_id: str) -> dict | None: for deployment in healthy_deployments: model_info = deployment.get("model_info") if not isinstance(model_info, dict): @@ -152,7 +152,7 @@ class EncryptedContentAffinityCheck(CustomLogger): @staticmethod def _encryption_boundary_key( litellm_params: Any, - ) -> Optional[tuple]: + ) -> tuple | None: """ ``(api_base, api_key)`` pair identifying an Azure resource. Two deployments sharing both are interchangeable for ``encrypted_content`` @@ -177,9 +177,9 @@ class EncryptedContentAffinityCheck(CustomLogger): def _find_deployments_on_same_encryption_boundary( self, - healthy_deployments: List[dict], + healthy_deployments: list[dict], model_id: str, - ) -> tuple[List[dict], Any]: + ) -> tuple[list[dict], Any]: """ Deployments in ``healthy_deployments`` sharing the originating deployment's ``(api_base, api_key)``, alongside the originating @@ -208,11 +208,11 @@ class EncryptedContentAffinityCheck(CustomLogger): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, + ) -> list[dict]: """ If the request ``input`` contains litellm-encoded item IDs, decode the embedded ``model_id`` and pin the request to that deployment. Raises @@ -225,7 +225,7 @@ class EncryptedContentAffinityCheck(CustomLogger): retry after the deployment is eligible again. """ request_kwargs = request_kwargs or {} - typed_healthy_deployments = cast(List[dict], healthy_deployments) + typed_healthy_deployments = cast(list[dict], healthy_deployments) if not self._is_enabled_for_model_group(model): return typed_healthy_deployments @@ -290,7 +290,7 @@ class EncryptedContentAffinityCheck(CustomLogger): model: str, model_id: str, originating: Any, - parent_otel_span: Optional[Span], + parent_otel_span: Span | None, ) -> Exception: # Public error messages intentionally omit the originating ``model_id`` so # an authenticated caller forging encrypted-content markers cannot use the @@ -343,8 +343,8 @@ class EncryptedContentAffinityCheck(CustomLogger): async def _get_origin_cooldown( self, model_id: str, - parent_otel_span: Optional[Span], - ) -> Optional[CooldownCacheValue]: + parent_otel_span: Span | None, + ) -> CooldownCacheValue | None: if self.router is None: return None cooldown_cache = getattr(self.router, "cooldown_cache", None) diff --git a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py index 803fdc4b353..9561eafa900 100644 --- a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py @@ -12,7 +12,7 @@ from __future__ import annotations import contextlib import contextvars -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import httpx @@ -32,7 +32,7 @@ else: RoutingArgsTTL = 60 -_io_token_rate_limit_request_kwargs: contextvars.ContextVar[Optional[dict[str, Any]]] = contextvars.ContextVar( +_io_token_rate_limit_request_kwargs: contextvars.ContextVar[dict[str, Any] | None] = contextvars.ContextVar( "io_token_rate_limit_request_kwargs", default=None, ) @@ -43,7 +43,7 @@ ITPM_CACHE_KEY = "_litellm_itpm_cache_key" OTPM_CACHE_KEY = "_litellm_otpm_cache_key" -def set_io_token_rate_limit_request_kwargs(kwargs: Optional[dict[str, Any]], store_in_context: bool = True) -> None: +def set_io_token_rate_limit_request_kwargs(kwargs: dict[str, Any] | None, store_in_context: bool = True) -> None: # The reservation sentinels are server-only, but `metadata` is caller # controlled on proxy requests. Strip any client-supplied copies here (this # runs before the router stashes its own reservation) so a forged @@ -60,7 +60,7 @@ def set_io_token_rate_limit_request_kwargs(kwargs: Optional[dict[str, Any]], sto _io_token_rate_limit_request_kwargs.set(kwargs if store_in_context else None) -def get_io_token_rate_limit_request_kwargs() -> Optional[dict[str, Any]]: +def get_io_token_rate_limit_request_kwargs() -> dict[str, Any] | None: return _io_token_rate_limit_request_kwargs.get() @@ -71,7 +71,7 @@ def seconds_until_minute_reset() -> int: def get_deployment_io_token_limits( deployment: dict, -) -> tuple[Optional[int], Optional[int]]: +) -> tuple[int | None, int | None]: itpm = deployment.get("itpm") otpm = deployment.get("otpm") litellm_params = deployment.get("litellm_params") or {} @@ -92,7 +92,7 @@ def deployment_has_io_token_limits(deployment: dict) -> bool: return itpm is not None or otpm is not None -def _get_cache_keys(deployment: dict, current_minute: str) -> Optional[tuple[str, str]]: +def _get_cache_keys(deployment: dict, current_minute: str) -> tuple[str, str] | None: model_id = deployment.get("model_info", {}).get("id") deployment_name = deployment.get("litellm_params", {}).get("model") # Without both a deployment id and model name the key would collapse to a @@ -104,7 +104,7 @@ def _get_cache_keys(deployment: dict, current_minute: str) -> Optional[tuple[str return itpm_key, otpm_key -def _estimate_input_tokens(request_kwargs: Optional[dict[str, Any]], model: str = "") -> int: +def _estimate_input_tokens(request_kwargs: dict[str, Any] | None, model: str = "") -> int: if not request_kwargs: return 0 messages = request_kwargs.get("messages") @@ -120,7 +120,7 @@ def _estimate_input_tokens(request_kwargs: Optional[dict[str, Any]], model: str return 0 -def _model_max_output_tokens(model_name: str) -> Optional[int]: +def _model_max_output_tokens(model_name: str) -> int | None: # litellm.get_model_info raises a bare Exception for an unrecognized model; # this lookup is a fallback default and must never fail the request. with contextlib.suppress(Exception): @@ -131,7 +131,7 @@ def _model_max_output_tokens(model_name: str) -> Optional[int]: return None -def _resolve_max_tokens(request_kwargs: Optional[dict[str, Any]], deployment: dict) -> int: +def _resolve_max_tokens(request_kwargs: dict[str, Any] | None, deployment: dict) -> int: if request_kwargs: # An explicit max_tokens=0 must be honored, not treated as absent and # replaced by the model default. @@ -233,12 +233,12 @@ def _resolve_reconcile_usage_tokens( def _stash_reservation_in_metadata( - request_kwargs: Optional[dict[str, Any]], + request_kwargs: dict[str, Any] | None, *, itpm_reserved: int, otpm_reserved: int, - itpm_cache_key: Optional[str], - otpm_cache_key: Optional[str], + itpm_cache_key: str | None, + otpm_cache_key: str | None, ) -> None: if not request_kwargs: return @@ -256,7 +256,7 @@ def _stash_reservation_in_metadata( request_kwargs[channel] = dict(reservation) -def _extract_reservation(reservation: dict[str, Any]) -> tuple[int, int, Optional[str], Optional[str]]: +def _extract_reservation(reservation: dict[str, Any]) -> tuple[int, int, str | None, str | None]: itpm_cache_key = reservation.get(ITPM_CACHE_KEY) otpm_cache_key = reservation.get(OTPM_CACHE_KEY) return ( @@ -285,7 +285,7 @@ def _reservation_channels(kwargs: Any) -> tuple[Any, ...]: return tuple(channels) -def _read_reservation_from_kwargs(kwargs: Any) -> tuple[int, int, Optional[str], Optional[str]]: +def _read_reservation_from_kwargs(kwargs: Any) -> tuple[int, int, str | None, str | None]: for channel_dict in _reservation_channels(kwargs): if isinstance(channel_dict, dict) and ITPM_RESERVED_KEY in channel_dict: return _extract_reservation(channel_dict) @@ -303,7 +303,7 @@ def _clear_reservation_from_kwargs(kwargs: Any) -> None: channel_dict.pop(key, None) -def _reservation_value(value: int, limit: Optional[int]) -> int: +def _reservation_value(value: int, limit: int | None) -> int: if limit is None: return 0 if value > 0: @@ -341,7 +341,7 @@ def _sync_increment_with_rollback( dual_cache: DualCache, key: str, value: int, - limit: Optional[int], + limit: int | None, *, limit_label: str, ) -> None: @@ -365,9 +365,9 @@ async def _increment_with_rollback( dual_cache: DualCache, key: str, value: int, - limit: Optional[int], + limit: int | None, *, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, limit_label: str, ) -> None: if value <= 0 or limit is None: @@ -391,7 +391,7 @@ async def _increment_with_rollback( def io_token_pre_call_check( dual_cache: DualCache, deployment: dict, -) -> Optional[dict]: +) -> dict | None: itpm_limit, otpm_limit = get_deployment_io_token_limits(deployment) if itpm_limit is None and otpm_limit is None: return deployment @@ -453,8 +453,8 @@ def io_token_pre_call_check( async def async_io_token_pre_call_check( dual_cache: DualCache, deployment: dict, - parent_otel_span: Optional[Span] = None, -) -> Optional[dict]: + parent_otel_span: Span | None = None, +) -> dict | None: itpm_limit, otpm_limit = get_deployment_io_token_limits(deployment) if itpm_limit is None and otpm_limit is None: return deployment @@ -569,7 +569,7 @@ async def async_io_token_reconcile_success( kwargs: Any, response_obj: Any, *, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> None: itpm_reserved, otpm_reserved, itpm_key, otpm_key = _read_reservation_from_kwargs(kwargs) if itpm_key is None and otpm_key is None: @@ -642,7 +642,7 @@ def io_token_refund_failure( verbose_router_logger.debug(f"[IO TOKEN LIMIT] refunded ITPM={itpm_reserved} OTPM={otpm_reserved}") -def refund_stale_reservation_before_retry(dual_cache: DualCache, kwargs: Optional[dict[str, Any]]) -> None: +def refund_stale_reservation_before_retry(dual_cache: DualCache, kwargs: dict[str, Any] | None) -> None: """ Synchronously refund and clear any reservation a previous deployment attempt stashed in ``kwargs``, before it's overwritten for the next @@ -673,7 +673,7 @@ async def async_io_token_refund_failure( dual_cache: DualCache, kwargs: Any, *, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> None: itpm_reserved, otpm_reserved, itpm_key, otpm_key = _read_reservation_from_kwargs(kwargs) if itpm_key is None and otpm_key is None: @@ -698,10 +698,10 @@ async def async_io_token_refund_failure( def build_io_token_rate_limit_headers( *, - itpm_limit: Optional[int], - otpm_limit: Optional[int], - current_itpm: Optional[int], - current_otpm: Optional[int], + itpm_limit: int | None, + otpm_limit: int | None, + current_itpm: int | None, + current_otpm: int | None, ) -> dict[str, int]: headers: dict[str, int] = {} reset = seconds_until_minute_reset() diff --git a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py index 373563ca442..d67f2a2bf47 100644 --- a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py @@ -11,7 +11,7 @@ is logged the first time such a deployment is seen. """ import contextlib -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Union import httpx @@ -20,6 +20,7 @@ from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import ( + ITPM_RESERVED_KEY, async_io_token_pre_call_check, async_io_token_reconcile_success, async_io_token_refund_failure, @@ -28,7 +29,6 @@ from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import ( io_token_pre_call_check, io_token_reconcile_success, io_token_refund_failure, - ITPM_RESERVED_KEY, ) from litellm.types.router import RouterErrors from litellm.types.utils import StandardLoggingPayload @@ -86,7 +86,7 @@ class ModelRateLimitingCheck(CustomLogger): async def _async_refund_io_token_reservation_if_any( self, - parent_otel_span: Optional[Span] = None, + parent_otel_span: Span | None = None, ) -> None: request_kwargs = get_io_token_rate_limit_request_kwargs() if request_kwargs is not None: @@ -96,7 +96,7 @@ class ModelRateLimitingCheck(CustomLogger): parent_otel_span=parent_otel_span, ) - def _get_deployment_limits(self, deployment: Dict) -> tuple[Optional[int], Optional[int]]: + def _get_deployment_limits(self, deployment: dict) -> tuple[int | None, int | None]: """ Extract TPM and RPM limits from a deployment configuration. @@ -126,7 +126,7 @@ class ModelRateLimitingCheck(CustomLogger): return tpm, rpm - def _get_cache_keys(self, deployment: Dict, current_minute: str) -> tuple[str, str]: + def _get_cache_keys(self, deployment: dict, current_minute: str) -> tuple[str, str]: """Get the cache keys for TPM and RPM tracking.""" model_id = deployment.get("model_info", {}).get("id") deployment_name = deployment.get("litellm_params", {}).get("model") @@ -136,7 +136,7 @@ class ModelRateLimitingCheck(CustomLogger): return tpm_key, rpm_key - def pre_call_check(self, deployment: Dict) -> Optional[Dict]: + def pre_call_check(self, deployment: dict) -> dict | None: """ Synchronous pre-call check for model rate limits. @@ -212,11 +212,11 @@ class ModelRateLimitingCheck(CustomLogger): self._refund_io_token_reservation_if_any() raise except Exception as e: - verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.pre_call_check: {str(e)}") + verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.pre_call_check: {e!s}") # Don't fail the request if rate limit check fails return deployment - async def async_pre_call_check(self, deployment: Dict, parent_otel_span: Optional[Span] = None) -> Optional[Dict]: + async def async_pre_call_check(self, deployment: dict, parent_otel_span: Span | None = None) -> dict | None: """ Async pre-call check for model rate limits. @@ -300,7 +300,7 @@ class ModelRateLimitingCheck(CustomLogger): await self._async_refund_io_token_reservation_if_any(parent_otel_span=parent_otel_span) raise except Exception as e: - verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.async_pre_call_check: {str(e)}") + verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.async_pre_call_check: {e!s}") # Don't fail the request if rate limit check fails return deployment @@ -310,7 +310,7 @@ class ModelRateLimitingCheck(CustomLogger): ) try: - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") # IO token reconciliation works purely from the cache keys stashed in # kwargs/metadata, so it must run before the model_id guard below @@ -360,7 +360,7 @@ class ModelRateLimitingCheck(CustomLogger): ) except Exception as e: - verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.async_log_success_event: {str(e)}") + verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.async_log_success_event: {e!s}") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): from litellm.litellm_core_utils.core_helpers import ( @@ -381,7 +381,7 @@ class ModelRateLimitingCheck(CustomLogger): Always tracks tokens - the pre-call check handles enforcement. """ try: - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") slo_metadata = (standard_logging_object.get("metadata") or {}) if standard_logging_object else {} kwargs_metadata = kwargs.get("metadata") or {} if ITPM_RESERVED_KEY in slo_metadata or ITPM_RESERVED_KEY in kwargs_metadata: @@ -418,7 +418,7 @@ class ModelRateLimitingCheck(CustomLogger): ) except Exception as e: - verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.log_success_event: {str(e)}") + verbose_router_logger.debug(f"Error in ModelRateLimitingCheck.log_success_event: {e!s}") def log_failure_event(self, kwargs, response_obj, start_time, end_time): with contextlib.suppress(Exception): diff --git a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py index d6412c95da0..fa4aa119966 100644 --- a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py @@ -4,7 +4,7 @@ Check if prompt caching is valid for a given deployment Route to previously cached model id, if valid """ -from typing import List, Optional, cast +from typing import cast from litellm import verbose_logger from litellm.caching.dual_cache import DualCache @@ -49,11 +49,11 @@ class PromptCachingDeploymentCheck(CustomLogger): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, + ) -> list[dict]: if messages is not None and is_prompt_caching_valid_prompt( messages=messages, model=model, @@ -64,7 +64,7 @@ class PromptCachingDeploymentCheck(CustomLogger): ) model_id_dict = await prompt_cache.async_get_model_id( - messages=cast(List[AllMessageValues], messages), + messages=cast(list[AllMessageValues], messages), tools=None, ) if model_id_dict is not None: @@ -76,7 +76,7 @@ class PromptCachingDeploymentCheck(CustomLogger): return healthy_deployments async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None) + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -111,7 +111,7 @@ class PromptCachingDeploymentCheck(CustomLogger): ## PROMPT CACHING - cache model id, if prompt caching valid prompt + provider if is_prompt_caching_valid_prompt( model=model, - messages=cast(List[AllMessageValues], messages), + messages=cast(list[AllMessageValues], messages), ): cache = PromptCachingCache( cache=self.cache, diff --git a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py index 5ae3c20baf3..3c0052b4c91 100644 --- a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py @@ -11,7 +11,6 @@ If previous_response_id is provided, route to the deployment that returned the p """ import warnings -from typing import List, Optional from litellm.integrations.custom_logger import CustomLogger, Span from litellm.responses.utils import ResponsesAPIRequestUtils @@ -33,11 +32,11 @@ class ResponsesApiDeploymentCheck(CustomLogger): async def async_filter_deployments( self, model: str, - healthy_deployments: List, - messages: Optional[List[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, - ) -> List[dict]: + healthy_deployments: list, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, + ) -> list[dict]: request_kwargs = request_kwargs or {} previous_response_id = request_kwargs.get("previous_response_id", None) if previous_response_id is None: diff --git a/litellm/router_utils/prompt_caching_cache.py b/litellm/router_utils/prompt_caching_cache.py index faf632f3b40..fe62f1af63c 100644 --- a/litellm/router_utils/prompt_caching_cache.py +++ b/litellm/router_utils/prompt_caching_cache.py @@ -4,7 +4,7 @@ Wrapper around router cache. Meant to store model id when prompt caching support import hashlib import json -from typing import TYPE_CHECKING, Any, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Union, cast from typing_extensions import TypedDict @@ -52,8 +52,8 @@ class PromptCachingCache: @staticmethod def extract_cacheable_prefix( - messages: List[AllMessageValues], - ) -> List[AllMessageValues]: + messages: list[AllMessageValues], + ) -> list[AllMessageValues]: """ Extract the cacheable prefix from messages. @@ -141,9 +141,9 @@ class PromptCachingCache: @staticmethod def get_prompt_caching_cache_key( - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]], - ) -> Optional[str]: + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None, + ) -> str | None: if messages is None and tools is None: return None @@ -178,46 +178,46 @@ class PromptCachingCache: def add_model_id( self, model_id: str, - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]], + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None, ) -> None: if messages is None and tools is None: - return None + return cache_key = PromptCachingCache.get_prompt_caching_cache_key(messages, tools) # If no cacheable prefix found, don't cache (can't generate cache key) if cache_key is None: - return None + return self.cache.set_cache(cache_key, PromptCachingCacheValue(model_id=model_id), ttl=300) - return None + return async def async_add_model_id( self, model_id: str, - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]], + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None, ) -> None: if messages is None and tools is None: - return None + return cache_key = PromptCachingCache.get_prompt_caching_cache_key(messages, tools) # If no cacheable prefix found, don't cache (can't generate cache key) if cache_key is None: - return None + return await self.cache.async_set_cache( cache_key, PromptCachingCacheValue(model_id=model_id), ttl=300, # store for 5 minutes ) - return None + return async def async_get_model_id( self, - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]], - ) -> Optional[PromptCachingCacheValue]: + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None, + ) -> PromptCachingCacheValue | None: """ Get model ID from cache using the cacheable prefix. @@ -239,9 +239,9 @@ class PromptCachingCache: def get_model_id( self, - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]], - ) -> Optional[PromptCachingCacheValue]: + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None, + ) -> PromptCachingCacheValue | None: if messages is None and tools is None: return None diff --git a/litellm/router_utils/search_api_router.py b/litellm/router_utils/search_api_router.py index f0bad959da4..531b2b577b1 100644 --- a/litellm/router_utils/search_api_router.py +++ b/litellm/router_utils/search_api_router.py @@ -9,7 +9,7 @@ import random import traceback from collections.abc import Callable from functools import partial -from typing import Any, Dict, Optional, Tuple +from typing import Any from litellm._logging import verbose_router_logger @@ -24,8 +24,8 @@ class SearchAPIRouter: @staticmethod def _resolve_search_provider_credentials( *, - tool_litellm_params: Dict[str, Any], - ) -> Tuple[Optional[str], Optional[str]]: + tool_litellm_params: dict[str, Any], + ) -> tuple[str | None, str | None]: """ Resolve search provider credentials from tool configuration ONLY. @@ -38,8 +38,8 @@ class SearchAPIRouter: Returns: Tuple of (api_key, api_base) from tool configuration """ - resolved_api_key: Optional[str] = tool_litellm_params.get("api_key") - resolved_api_base: Optional[str] = tool_litellm_params.get("api_base") + resolved_api_key: str | None = tool_litellm_params.get("api_key") + resolved_api_base: str | None = tool_litellm_params.get("api_base") return resolved_api_key, resolved_api_base @@ -77,7 +77,7 @@ class SearchAPIRouter: verbose_router_logger.info(f"Successfully updated router with {len(router_search_tools)} search tool(s)") except Exception as e: - verbose_router_logger.exception(f"Error updating router with search tools: {str(e)}") + verbose_router_logger.exception(f"Error updating router with search tools: {e!s}") raise e @staticmethod @@ -226,6 +226,6 @@ class SearchAPIRouter: except Exception as e: verbose_router_logger.error( - f"Error in SearchAPIRouter.async_search_with_fallbacks_helper for {search_tool_name}: {str(e)}" + f"Error in SearchAPIRouter.async_search_with_fallbacks_helper for {search_tool_name}: {e!s}" ) raise e diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 728a16edad3..660cc2633fb 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -4,7 +4,7 @@ from __future__ import annotations from collections.abc import Awaitable from dataclasses import dataclass -from typing import Final, Protocol, Union, cast +from typing import Final, Protocol, cast import httpx @@ -96,7 +96,7 @@ def messages( api_base: str | None, custom_llm_provider: str | None, extra_headers: dict[str, object] | None, - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_messages = load_rust_messages() if rust_messages is None: @@ -120,7 +120,7 @@ async def amessages( api_base: str | None, custom_llm_provider: str | None, extra_headers: dict[str, object] | None, - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_amessages = load_rust_amessages() if rust_amessages is None: diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 388da00f7f0..6e387156fc4 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -4,7 +4,7 @@ from __future__ import annotations import os from collections.abc import Awaitable -from typing import TYPE_CHECKING, Any, Final, Protocol, Union, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, cast import httpx @@ -147,7 +147,7 @@ def ocr( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_ocr = load_rust_ocr() if rust_ocr is None: @@ -173,7 +173,7 @@ async def aocr( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_aocr = load_rust_aocr() if rust_aocr is None: diff --git a/litellm/rust_bridge/timeouts.py b/litellm/rust_bridge/timeouts.py index 4407986c3da..5fde4397b55 100644 --- a/litellm/rust_bridge/timeouts.py +++ b/litellm/rust_bridge/timeouts.py @@ -2,12 +2,10 @@ from __future__ import annotations -from typing import Union - import httpx -def timeout_to_seconds(timeout: Union[float, httpx.Timeout] | None) -> float | None: +def timeout_to_seconds(timeout: float | httpx.Timeout | None) -> float | None: if timeout is None: return None if isinstance(timeout, httpx.Timeout): diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 0ed50e5d1dd..cc16d22643b 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Awaitable from dataclasses import dataclass -from typing import Final, Protocol, Union, cast +from typing import Final, Protocol, cast import httpx @@ -106,7 +106,7 @@ def transcription( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_transcription = load_rust_transcription() if rust_transcription is None: @@ -132,7 +132,7 @@ async def atranscription( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], - timeout: Union[float, httpx.Timeout] | None, + timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_atranscription = load_rust_atranscription() if rust_atranscription is None: diff --git a/litellm/sandbox/main.py b/litellm/sandbox/main.py index c59552fbd9c..eae7e5c097a 100644 --- a/litellm/sandbox/main.py +++ b/litellm/sandbox/main.py @@ -11,8 +11,6 @@ Each entrypoint is `@client`-decorated, so every operation is logged the same way `litellm.asearch` is. """ -from typing import Union - import litellm from litellm.llms.base_llm.sandbox.transformation import ( BaseSandboxConfig, @@ -23,10 +21,10 @@ from litellm.types.utils import SandboxProviders from litellm.utils import ProviderConfigManager, client __all__ = [ - "acreate_sandbox", - "arun_code", - "adelete_sandbox", "acode_interpreter_tool", + "acreate_sandbox", + "adelete_sandbox", + "arun_code", ] _LITELLM_INTERNAL_KWARGS = { @@ -85,7 +83,7 @@ async def acreate_sandbox( @client async def arun_code( provider: str, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, code: str, api_key: str | None = None, **kwargs, @@ -102,7 +100,7 @@ async def arun_code( @client async def adelete_sandbox( provider: str, - container: Union[ContainerHandle, str], + container: ContainerHandle | str, api_key: str | None = None, api_base: str | None = None, **kwargs, diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 814fd74333a..7de9f556d97 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -1,6 +1,5 @@ import enum import heapq -from typing import Optional from pydantic import BaseModel @@ -25,14 +24,14 @@ class Scheduler: def __init__( self, - polling_interval: Optional[float] = None, - redis_cache: Optional[RedisCache] = None, + polling_interval: float | None = None, + redis_cache: RedisCache | None = None, ): """ polling_interval: float or null - frequency of polling queue. Default is 3ms. """ self.queue: list = [] - default_in_memory_ttl: Optional[float] = None + default_in_memory_ttl: float | None = None if redis_cache is not None: # if redis-cache available frequently poll that instead of using in-memory. default_in_memory_ttl = SchedulerCacheKeys.default_in_memory_ttl.value @@ -63,7 +62,7 @@ class Scheduler: """ queue = await self.get_queue(model_name=model_name) if not queue: - raise Exception("Incorrectly setup. Queue is invalid. Queue={}".format(queue)) + raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") # ------------ # Setup values @@ -99,7 +98,7 @@ class Scheduler: """Return if the id is at the top of the queue. Don't pop the value from heap.""" queue = await self.get_queue(model_name=model_name) if not queue: - raise Exception("Incorrectly setup. Queue is invalid. Queue={}".format(queue)) + raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") # ------------ # Setup values @@ -120,7 +119,7 @@ class Scheduler: Return a queue for that specific model group """ if self.cache is not None: - _cache_key = "{}:{}".format(SchedulerCacheKeys.queue.value, model_name) + _cache_key = f"{SchedulerCacheKeys.queue.value}:{model_name}" response = await self.cache.async_get_cache(key=_cache_key) if response is None or not isinstance(response, list): return [] @@ -133,6 +132,5 @@ class Scheduler: Save the updated queue of the model group """ if self.cache is not None: - _cache_key = "{}:{}".format(SchedulerCacheKeys.queue.value, model_name) + _cache_key = f"{SchedulerCacheKeys.queue.value}:{model_name}" await self.cache.async_set_cache(key=_cache_key, value=queue) - return None diff --git a/litellm/search/__init__.py b/litellm/search/__init__.py index 51f311618e3..b71a1221adf 100644 --- a/litellm/search/__init__.py +++ b/litellm/search/__init__.py @@ -5,4 +5,4 @@ LiteLLM Search API module. from litellm.search.cost_calculator import search_provider_cost_per_query from litellm.search.main import asearch, search -__all__ = ["search", "asearch", "search_provider_cost_per_query"] +__all__ = ["asearch", "search", "search_provider_cost_per_query"] diff --git a/litellm/search/cost_calculator.py b/litellm/search/cost_calculator.py index 9680446064a..752bf1a7d0e 100644 --- a/litellm/search/cost_calculator.py +++ b/litellm/search/cost_calculator.py @@ -2,17 +2,15 @@ Cost calculation for search providers. """ -from typing import Optional, Tuple - from litellm.utils import get_model_info def search_provider_cost_per_query( model: str, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, number_of_queries: int = 1, - optional_params: Optional[dict] = None, -) -> Tuple[float, float]: + optional_params: dict | None = None, +) -> tuple[float, float]: """ Calculate cost for search-only providers. diff --git a/litellm/search/main.py b/litellm/search/main.py index 8954b626151..932a73c0955 100644 --- a/litellm/search/main.py +++ b/litellm/search/main.py @@ -6,7 +6,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -25,11 +25,11 @@ base_llm_http_handler = BaseLLMHTTPHandler() def _build_search_optional_params( - max_results: Optional[int] = None, - search_domain_filter: Optional[List[str]] = None, - max_tokens_per_page: Optional[int] = None, - country: Optional[str] = None, -) -> Dict[str, Any]: + max_results: int | None = None, + search_domain_filter: list[str] | None = None, + max_tokens_per_page: int | None = None, + country: str | None = None, +) -> dict[str, Any]: """ Helper function to build optional_params dict from Perplexity Search API parameters. @@ -42,7 +42,7 @@ def _build_search_optional_params( Returns: Dict with non-None optional parameters """ - optional_params: Dict[str, Any] = {} + optional_params: dict[str, Any] = {} if max_results is not None: optional_params["max_results"] = max_results @@ -58,16 +58,16 @@ def _build_search_optional_params( @client async def asearch( - query: Union[str, List[str]], + query: str | list[str], search_provider: str, - max_results: Optional[int] = None, - search_domain_filter: Optional[List[str]] = None, - max_tokens_per_page: Optional[int] = None, - country: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - extra_headers: Optional[Dict[str, Any]] = None, + max_results: int | None = None, + search_domain_filter: list[str] | None = None, + max_tokens_per_page: int | None = None, + country: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + extra_headers: dict[str, Any] | None = None, **kwargs, ) -> SearchResponse: """ @@ -161,18 +161,18 @@ async def asearch( @client def search( - query: Union[str, List[str]], + query: str | list[str], search_provider: str, - max_results: Optional[int] = None, - search_domain_filter: Optional[List[str]] = None, - max_tokens_per_page: Optional[int] = None, - country: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - extra_headers: Optional[Dict[str, Any]] = None, + max_results: int | None = None, + search_domain_filter: list[str] | None = None, + max_tokens_per_page: int | None = None, + country: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + extra_headers: dict[str, Any] | None = None, **kwargs, -) -> Union[SearchResponse, Coroutine[Any, Any, SearchResponse]]: +) -> SearchResponse | Coroutine[Any, Any, SearchResponse]: """ Synchronous Search function. @@ -229,7 +229,7 @@ def search( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("asearch", False) is True # Validate query parameter @@ -240,7 +240,7 @@ def search( raise ValueError("All items in query list must be strings") # Get provider config - search_provider_config: Optional[BaseSearchConfig] = ProviderConfigManager.get_provider_search_config( + search_provider_config: BaseSearchConfig | None = ProviderConfigManager.get_provider_search_config( provider=SearchProviders(search_provider), ) diff --git a/litellm/secret_managers/aws_secret_manager.py b/litellm/secret_managers/aws_secret_manager.py index 090e3ebcfa2..70d75fcdf98 100644 --- a/litellm/secret_managers/aws_secret_manager.py +++ b/litellm/secret_managers/aws_secret_manager.py @@ -12,7 +12,7 @@ import ast import base64 import os import re -from typing import Any, Dict, Optional +from typing import Any import litellm from litellm.proxy._types import KeyManagementSystem @@ -23,7 +23,7 @@ def validate_environment(): raise ValueError("Missing required environment variable - AWS_REGION_NAME") -def load_aws_kms(use_aws_kms: Optional[bool]): +def load_aws_kms(use_aws_kms: bool | None): if use_aws_kms is None or use_aws_kms is False: return try: @@ -59,16 +59,17 @@ class AWSKeyManagementService_V2: ## CHECK IF LICENSE IN ENV ## - premium feature is_litellm_license_in_env: bool = False - if os.getenv("LITELLM_LICENSE", None) is not None: - is_litellm_license_in_env = True - elif os.getenv("LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE", None) is not None: + if ( + os.getenv("LITELLM_LICENSE", None) is not None + or os.getenv("LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE", None) is not None + ): is_litellm_license_in_env = True if is_litellm_license_in_env is False: raise ValueError( "AWSKeyManagementService V2 is an Enterprise Feature. Please add a valid LITELLM_LICENSE to your envionment." ) - def load_aws_kms(self, use_aws_kms: Optional[bool]): + def load_aws_kms(self, use_aws_kms: bool | None): if use_aws_kms is None or use_aws_kms is False: return try: @@ -88,7 +89,7 @@ class AWSKeyManagementService_V2: raise ValueError("kms_client is None") encrypted_value = os.getenv(secret_name, None) if encrypted_value is None: - raise Exception("AWS KMS - Encrypted Value of Key={} is None".format(secret_name)) + raise Exception(f"AWS KMS - Encrypted Value of Key={secret_name} is None") if isinstance(encrypted_value, str) and encrypted_value.startswith("aws_kms/"): encrypted_value = encrypted_value.replace("aws_kms/", "") @@ -122,7 +123,7 @@ class AWSKeyManagementService_V2: """ -def decrypt_env_var() -> Dict[str, Any]: +def decrypt_env_var() -> dict[str, Any]: # setup client class aws_kms = AWSKeyManagementService_V2() # iterate through env - for `aws_kms/` diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 2b24ea1ce61..1bfeafa51c9 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -16,7 +16,7 @@ Requires: import json import os -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -37,13 +37,13 @@ from .base_secret_manager import BaseSecretManager class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): def __init__( self, - aws_region_name: Optional[str] = None, - aws_role_name: Optional[str] = None, - aws_session_name: Optional[str] = None, - aws_external_id: Optional[str] = None, - aws_profile_name: Optional[str] = None, - aws_web_identity_token: Optional[str] = None, - aws_sts_endpoint: Optional[str] = None, + aws_region_name: str | None = None, + aws_role_name: str | None = None, + aws_session_name: str | None = None, + aws_external_id: str | None = None, + aws_profile_name: str | None = None, + aws_web_identity_token: str | None = None, + aws_sts_endpoint: str | None = None, replica_regions: list[str] | None = None, **kwargs, ): @@ -77,7 +77,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): @classmethod def load_aws_secret_manager( cls, - use_aws_secret_manager: Optional[bool], + use_aws_secret_manager: bool | None, key_management_settings: KeyManagementSettings | None = None, ): """ @@ -113,10 +113,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): async def async_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - primary_secret_name: Optional[str] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + primary_secret_name: str | None = None, + ) -> str | None: """ Async function to read a secret from AWS Secrets Manager @@ -158,10 +158,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): def sync_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - primary_secret_name: Optional[str] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + primary_secret_name: str | None = None, + ) -> str | None: """ Sync function to read a secret from AWS Secrets Manager @@ -212,7 +212,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): ) return None - def _parse_primary_secret(self, primary_secret_json_str: Optional[str]) -> dict: + def _parse_primary_secret(self, primary_secret_json_str: str | None) -> dict: """ Parse the primary secret JSON string into a dictionary @@ -224,7 +224,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): """ return json.loads(primary_secret_json_str or "{}") - def sync_read_secret_from_primary_secret(self, secret_name: str, primary_secret_name: str) -> Optional[str]: + def sync_read_secret_from_primary_secret(self, secret_name: str, primary_secret_name: str) -> str | None: """ Read a secret from the primary secret """ @@ -232,7 +232,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): primary_secret_kv_pairs = self._parse_primary_secret(primary_secret_json_str) return primary_secret_kv_pairs.get(secret_name) - async def async_read_secret_from_primary_secret(self, secret_name: str, primary_secret_name: str) -> Optional[str]: + async def async_read_secret_from_primary_secret(self, secret_name: str, primary_secret_name: str) -> str | None: """ Read a secret from the primary secret """ @@ -244,10 +244,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self, secret_name: str, secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None, + description: str | None = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + tags: dict | list | None = None, ) -> dict: """ Async function to write a secret to AWS Secrets Manager @@ -264,7 +264,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): """ from litellm._uuid import uuid - data: Dict[str, Any] = { + data: dict[str, Any] = { "Name": secret_name, "SecretString": secret_value, "ClientRequestToken": str(uuid.uuid4()), @@ -393,8 +393,8 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self, secret_name: str, secret_value: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Async function to update an existing secret's value in AWS Secrets Manager. @@ -413,7 +413,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): """ from litellm._uuid import uuid - data: Dict[str, Any] = { + data: dict[str, Any] = { "SecretId": secret_name, "SecretString": secret_value, "ClientRequestToken": str(uuid.uuid4()), @@ -446,8 +446,8 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): current_secret_name: str, new_secret_name: str, new_secret_value: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Rotate a secret. When current_secret_name == new_secret_name (in-place @@ -478,9 +478,9 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): async def async_delete_secret( self, secret_name: str, - recovery_window_in_days: Optional[int] = 7, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + recovery_window_in_days: int | None = 7, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Async function to delete a secret from AWS Secrets Manager @@ -525,9 +525,9 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self, action: str, # "GetSecretValue" or "PutSecretValue" secret_name: str, - secret_value: Optional[str] = None, - optional_params: Optional[dict] = None, - request_data: Optional[dict] = None, + secret_value: str | None = None, + optional_params: dict | None = None, + request_data: dict | None = None, ) -> tuple[str, Any, bytes]: """Prepare the AWS Secrets Manager request""" try: diff --git a/litellm/secret_managers/base_secret_manager.py b/litellm/secret_managers/base_secret_manager.py index 2bb8dc73138..d8fb083374c 100644 --- a/litellm/secret_managers/base_secret_manager.py +++ b/litellm/secret_managers/base_secret_manager.py @@ -1,6 +1,6 @@ import re from abc import ABC, abstractmethod -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -30,9 +30,9 @@ class BaseSecretManager(ABC): async def async_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Asynchronously read a secret from the secret manager. @@ -44,15 +44,14 @@ class BaseSecretManager(ABC): Returns: Optional[str]: The secret value if found, None otherwise """ - pass @abstractmethod def sync_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Synchronously read a secret from the secret manager. @@ -64,18 +63,17 @@ class BaseSecretManager(ABC): Returns: Optional[str]: The secret value if found, None otherwise """ - pass @abstractmethod async def async_write_secret( self, secret_name: str, secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None, - ) -> Dict[str, Any]: + description: str | None = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + tags: dict | list | None = None, + ) -> dict[str, Any]: """ Asynchronously write a secret to the secret manager. @@ -91,15 +89,14 @@ class BaseSecretManager(ABC): Returns: Dict[str, Any]: Response from the secret manager containing write operation details """ - pass @abstractmethod async def async_delete_secret( self, secret_name: str, - recovery_window_in_days: Optional[int] = 7, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + recovery_window_in_days: int | None = 7, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Async function to delete a secret from the secret manager @@ -113,15 +110,14 @@ class BaseSecretManager(ABC): Returns: dict: Response from the secret manager containing deletion details """ - pass async def async_rotate_secret( self, current_secret_name: str, new_secret_name: str, new_secret_value: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Async function to rotate a secret by creating a new one and deleting the old one. diff --git a/litellm/secret_managers/custom_secret_manager_loader.py b/litellm/secret_managers/custom_secret_manager_loader.py index ab8153c9436..b296690b025 100644 --- a/litellm/secret_managers/custom_secret_manager_loader.py +++ b/litellm/secret_managers/custom_secret_manager_loader.py @@ -6,7 +6,6 @@ Handles dynamic loading of user-defined secret manager classes from Python files import importlib.util import os -from typing import Optional import litellm from litellm._logging import verbose_proxy_logger @@ -14,7 +13,7 @@ from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.types.secret_managers.main import KeyManagementSystem -def load_custom_secret_manager(config_file_path: Optional[str] = None) -> None: +def load_custom_secret_manager(config_file_path: str | None = None) -> None: """ Load and initialize a custom secret manager from a python file. diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index 2b888cb85f6..6e7eb742088 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -1,6 +1,6 @@ import base64 import os -from typing import Any, Dict, Optional, Union +from typing import Any from urllib.parse import quote import httpx @@ -172,9 +172,9 @@ class CyberArkSecretManager(BaseSecretManager): async def async_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Reads a secret from CyberArk Conjur using an async HTTPX client. @@ -218,9 +218,9 @@ class CyberArkSecretManager(BaseSecretManager): def sync_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Reads a secret from CyberArk Conjur using a sync HTTPX client. @@ -262,11 +262,11 @@ class CyberArkSecretManager(BaseSecretManager): self, secret_name: str, secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None, - ) -> Dict[str, Any]: + description: str | None = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + tags: dict | list | None = None, + ) -> dict[str, Any]: """ Writes a secret to CyberArk Conjur using an async HTTPX client. @@ -309,9 +309,9 @@ class CyberArkSecretManager(BaseSecretManager): async def async_delete_secret( self, secret_name: str, - recovery_window_in_days: Optional[int] = 7, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + recovery_window_in_days: int | None = 7, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ CyberArk Conjur does not support direct secret deletion via API. diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index 4c54166a877..ed348865859 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -1,6 +1,6 @@ import os from collections.abc import Callable -from typing import Any, Optional, Union +from typing import Any from litellm._logging import verbose_logger from litellm.types.secret_managers.get_azure_ad_token_provider import ( @@ -18,23 +18,23 @@ def infer_credential_type_from_environment() -> AzureCredentialType: elif os.environ.get("AZURE_CLIENT_ID"): return AzureCredentialType.ManagedIdentityCredential elif ( - os.environ.get("AZURE_CLIENT_ID") - and os.environ.get("AZURE_TENANT_ID") - and os.environ.get("AZURE_CERTIFICATE_PATH") - and os.environ.get("AZURE_CERTIFICATE_PASSWORD") + ( + os.environ.get("AZURE_CLIENT_ID") + and os.environ.get("AZURE_TENANT_ID") + and os.environ.get("AZURE_CERTIFICATE_PATH") + and os.environ.get("AZURE_CERTIFICATE_PASSWORD") + ) + or os.environ.get("AZURE_CERTIFICATE_PASSWORD") + or os.environ.get("AZURE_CERTIFICATE_PATH") ): return AzureCredentialType.CertificateCredential - elif os.environ.get("AZURE_CERTIFICATE_PASSWORD"): - return AzureCredentialType.CertificateCredential - elif os.environ.get("AZURE_CERTIFICATE_PATH"): - return AzureCredentialType.CertificateCredential else: return AzureCredentialType.DefaultAzureCredential def get_azure_ad_token_provider( - azure_scope: Optional[str] = None, - azure_credential: Optional[AzureCredentialType] = None, + azure_scope: str | None = None, + azure_credential: AzureCredentialType | None = None, ) -> Callable[[], str]: """ Get Azure AD token provider based on Service Principal with Secret workflow. @@ -52,7 +52,7 @@ def get_azure_ad_token_provider( Returns: Callable that returns a temporary authentication token. """ - import azure.identity as identity + from azure import identity from azure.identity import ( CertificateCredential, ClientSecretCredential, @@ -70,15 +70,9 @@ def get_azure_ad_token_provider( else None or os.environ.get("AZURE_CREDENTIAL") or infer_credential_type_from_environment() ) verbose_logger.info(f"For Azure AD Token Provider, choosing credential type: {cred}") - credential: Optional[ - Union[ - ClientSecretCredential, - ManagedIdentityCredential, - CertificateCredential, - DefaultAzureCredential, - Any, - ] - ] = None + credential: ( + ClientSecretCredential | ManagedIdentityCredential | CertificateCredential | DefaultAzureCredential | Any | None + ) = None if cred == AzureCredentialType.ClientSecretCredential: credential = ClientSecretCredential( client_id=os.environ["AZURE_CLIENT_ID"], diff --git a/litellm/secret_managers/google_kms.py b/litellm/secret_managers/google_kms.py index d22cd0b38b3..4a39241d7d8 100644 --- a/litellm/secret_managers/google_kms.py +++ b/litellm/secret_managers/google_kms.py @@ -9,7 +9,6 @@ Requires: """ import os -from typing import Optional import litellm from litellm.proxy._types import KeyManagementSystem @@ -22,7 +21,7 @@ def validate_environment(): raise ValueError("Missing required environment variable - GOOGLE_KMS_RESOURCE_NAME") -def load_google_kms(use_google_kms: Optional[bool]): +def load_google_kms(use_google_kms: bool | None): if use_google_kms is None or use_google_kms is False: return try: diff --git a/litellm/secret_managers/google_secret_manager.py b/litellm/secret_managers/google_secret_manager.py index 91284d5eb30..268689cca01 100644 --- a/litellm/secret_managers/google_secret_manager.py +++ b/litellm/secret_managers/google_secret_manager.py @@ -1,6 +1,5 @@ import base64 import os -from typing import Optional import litellm from litellm._logging import verbose_logger @@ -14,8 +13,8 @@ from litellm.proxy._types import CommonProxyErrors, KeyManagementSystem class GoogleSecretManager(GCSBucketBase): def __init__( self, - refresh_interval: Optional[int] = SECRET_MANAGER_REFRESH_INTERVAL, - always_read_secret_manager: Optional[bool] = False, + refresh_interval: int | None = SECRET_MANAGER_REFRESH_INTERVAL, + always_read_secret_manager: bool | None = False, ) -> None: """ Args: @@ -50,7 +49,7 @@ class GoogleSecretManager(GCSBucketBase): # by default this should be False, we want to use in memory caching for this. It's a bad idea to fetch from secret manager for all requests self.always_read_secret_manager = always_read_secret_manager or False - def get_secret_from_google_secret_manager(self, secret_name: str) -> Optional[str]: + def get_secret_from_google_secret_manager(self, secret_name: str) -> str | None: """ Retrieve a secret from Google Secret Manager or cache. diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index 039aecb9e58..3f15a4fe5f5 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -1,5 +1,5 @@ import os -from typing import Any, Dict, Optional, Union +from typing import Any import httpx @@ -205,9 +205,9 @@ class HashicorpSecretManager(BaseSecretManager): def get_url( self, secret_name: str, - namespace: Optional[str] = None, - mount_name: Optional[str] = None, - path_prefix: Optional[str] = None, + namespace: str | None = None, + mount_name: str | None = None, + path_prefix: str | None = None, ) -> str: """ Constructs the Vault URL for KV v2 secrets. @@ -238,7 +238,7 @@ class HashicorpSecretManager(BaseSecretManager): _url += secret_name return _url - def _sanitize_plain_value(self, value: Optional[Union[str, int]]) -> Optional[str]: + def _sanitize_plain_value(self, value: str | int | None) -> str | None: if value is None: return None value_str = str(value).strip() @@ -246,14 +246,14 @@ class HashicorpSecretManager(BaseSecretManager): return None return value_str - def _sanitize_path_component(self, value: Optional[Union[str, int]]) -> Optional[str]: + def _sanitize_path_component(self, value: str | int | None) -> str | None: sanitized_value = self._sanitize_plain_value(value) if sanitized_value is None: return None sanitized_value = sanitized_value.strip("/") return sanitized_value or None - def _extract_secret_manager_settings(self, optional_params: Optional[dict]) -> Dict[str, Any]: + def _extract_secret_manager_settings(self, optional_params: dict | None) -> dict[str, Any]: if not isinstance(optional_params, dict): return {} @@ -262,7 +262,7 @@ class HashicorpSecretManager(BaseSecretManager): allowed_keys = {"namespace", "mount", "path_prefix", "data"} return {k: source[k] for k in allowed_keys if k in source} - def _build_secret_target(self, secret_name: str, optional_params: Optional[dict]) -> Dict[str, Any]: + def _build_secret_target(self, secret_name: str, optional_params: dict | None) -> dict[str, Any]: settings = self._extract_secret_manager_settings(optional_params) namespace = settings.get("namespace", self.vault_namespace) @@ -308,9 +308,9 @@ class HashicorpSecretManager(BaseSecretManager): async def async_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Reads a secret from Vault KV v2 using an async HTTPX client. secret_name is just the path inside the KV mount (e.g., 'myapp/config'). @@ -343,9 +343,9 @@ class HashicorpSecretManager(BaseSecretManager): def sync_read_secret( self, secret_name: str, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> Optional[str]: + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: """ Reads a secret from Vault KV v2 using a sync HTTPX client. secret_name is just the path inside the KV mount (e.g., 'myapp/config'). @@ -375,11 +375,11 @@ class HashicorpSecretManager(BaseSecretManager): self, secret_name: str, secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None, - ) -> Dict[str, Any]: + description: str | None = None, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + tags: dict | list | None = None, + ) -> dict[str, Any]: """ Writes a secret to Vault KV v2 using an async HTTPX client. @@ -423,9 +423,9 @@ class HashicorpSecretManager(BaseSecretManager): current_secret_name: str, new_secret_name: str, new_secret_value: str, - optional_params: Dict | None = None, + optional_params: dict | None = None, timeout: float | httpx.Timeout | None = None, - ) -> Dict: + ) -> dict: """ Rotates a secret by creating a new one and deleting the old one. Uses _build_secret_target to handle optional_params for namespace, mount, path_prefix customization. @@ -567,9 +567,9 @@ class HashicorpSecretManager(BaseSecretManager): async def async_delete_secret( self, secret_name: str, - recovery_window_in_days: Optional[int] = 7, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + recovery_window_in_days: int | None = 7, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, ) -> dict: """ Async function to delete a secret from Hashicorp Vault. @@ -607,7 +607,7 @@ class HashicorpSecretManager(BaseSecretManager): verbose_logger.exception(f"Error deleting secret from Hashicorp Vault: {e}") return {"status": "error", "message": str(e)} - def _get_secret_value_from_json_response(self, json_resp: Optional[dict]) -> Optional[str]: + def _get_secret_value_from_json_response(self, json_resp: dict | None) -> str | None: """ Get the secret value from the JSON response diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index 88e3ad16cc3..a05ea367b19 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -3,7 +3,6 @@ import base64 import os import time import traceback -from typing import Optional, Union import httpx from pydantic import BaseModel, ValidationError @@ -93,7 +92,7 @@ def _resolve_oidc_file_path(requested_path: str) -> str: ) -def _get_oidc_http_handler(timeout: Optional[httpx.Timeout] = None) -> HTTPHandler: +def _get_oidc_http_handler(timeout: httpx.Timeout | None = None) -> HTTPHandler: """ Factory function to create HTTPHandler for OIDC requests. This function can be mocked in tests. @@ -112,7 +111,7 @@ def _get_oidc_http_handler(timeout: Optional[httpx.Timeout] = None) -> HTTPHandl ######### Secret Manager ############################ # checks if user has passed in a secret manager client # if passed in then checks the secret there -def str_to_bool(value: Optional[str]) -> Optional[bool]: +def str_to_bool(value: str | None) -> bool | None: """ Converts a string to a boolean if it's a recognized boolean string. Returns None if the string is not a recognized boolean value. @@ -138,8 +137,8 @@ def str_to_bool(value: Optional[str]) -> Optional[bool]: def get_secret_str( secret_name: str, - default_value: Optional[Union[str, bool]] = None, -) -> Optional[str]: + default_value: str | bool | None = None, +) -> str | None: """ Guarantees response from 'get_secret' is either string or none. Used for fixing linting errors. """ @@ -150,7 +149,7 @@ def get_secret_str( return value -def normalize_nonempty_secret_str(val: Optional[str]) -> Optional[str]: +def normalize_nonempty_secret_str(val: str | None) -> str | None: """ Strip whitespace and treat None, '', and whitespace-only strings as unset. @@ -165,8 +164,8 @@ def normalize_nonempty_secret_str(val: Optional[str]) -> Optional[str]: def get_secret_bool( secret_name: str, - default_value: Optional[bool] = None, -) -> Optional[bool]: + default_value: bool | None = None, +) -> bool | None: """ Guarantees response from 'get_secret' is either boolean or none. Used for fixing linting errors. @@ -188,7 +187,7 @@ def get_secret_bool( def get_secret( secret_name: str, - default_value: Optional[Union[str, bool]] = None, + default_value: str | bool | None = None, ): key_management_system = litellm._key_management_system key_management_settings = litellm._key_management_settings @@ -283,7 +282,7 @@ def get_secret( raise ValueError("Azure OIDC provider returned None token") return oidc_token except Exception as e: - error_msg = f"Azure OIDC provider failed: {str(e)}" + error_msg = f"Azure OIDC provider failed: {e!s}" verbose_logger.error(error_msg) raise ValueError(error_msg) with open(azure_federated_token_file, "r") as f: @@ -336,7 +335,7 @@ def get_secret( ) except Exception as e: # check if it's in os.environ verbose_logger.error( - f"Defaulting to os.environ value for key={secret_name}. An exception occurred - {str(e)}.\n\n{traceback.format_exc()}" + f"Defaulting to os.environ value for key={secret_name}. An exception occurred - {e!s}.\n\n{traceback.format_exc()}" ) secret = os.getenv(secret_name) try: diff --git a/litellm/secret_managers/secret_manager_handler.py b/litellm/secret_managers/secret_manager_handler.py index e2fb0b900b8..2acb154dd59 100644 --- a/litellm/secret_managers/secret_manager_handler.py +++ b/litellm/secret_managers/secret_manager_handler.py @@ -6,7 +6,7 @@ Handles retrieving secrets from different secret management systems. import base64 import os -from typing import Any, Optional +from typing import Any import litellm from litellm._logging import print_verbose @@ -27,8 +27,8 @@ def get_secret_from_manager( client: Any, key_manager: str, secret_name: str, - key_management_settings: Optional[Any] = None, -) -> Optional[str]: + key_management_settings: Any | None = None, +) -> str | None: """ Get a secret from the configured secret manager. @@ -81,7 +81,7 @@ def get_secret_from_manager( """ encrypted_value = os.getenv(secret_name, None) if encrypted_value is None: - raise Exception("AWS KMS - Encrypted Value of Key={} is None".format(secret_name)) + raise Exception(f"AWS KMS - Encrypted Value of Key={secret_name} is None") # Decode the base64 encoded ciphertext ciphertext_blob = base64.b64decode(encrypted_value) @@ -119,7 +119,7 @@ def get_secret_from_manager( if secret is None: raise ValueError(f"No secret found in Google Secret Manager for {secret_name}") except Exception as e: - print_verbose(f"An error occurred - {str(e)}") + print_verbose(f"An error occurred - {e!s}") raise e elif key_manager == KeyManagementSystem.HASHICORP_VAULT.value: @@ -128,7 +128,7 @@ def get_secret_from_manager( if secret is None: raise ValueError(f"No secret found in Hashicorp Secret Manager for {secret_name}") except Exception as e: - print_verbose(f"An error occurred - {str(e)}") + print_verbose(f"An error occurred - {e!s}") raise e elif key_manager == KeyManagementSystem.CYBERARK.value: @@ -137,7 +137,7 @@ def get_secret_from_manager( if secret is None: raise ValueError(f"No secret found in CyberArk Secret Manager for {secret_name}") except Exception as e: - print_verbose(f"An error occurred - {str(e)}") + print_verbose(f"An error occurred - {e!s}") raise e elif key_manager == KeyManagementSystem.CUSTOM.value: diff --git a/litellm/setup_wizard.py b/litellm/setup_wizard.py index c6d0c1717a9..41de3eb69dd 100644 --- a/litellm/setup_wizard.py +++ b/litellm/setup_wizard.py @@ -14,7 +14,6 @@ import secrets import sys import sysconfig from pathlib import Path -from typing import Dict, List, Optional, Set # termios / tty are Unix-only; fall back gracefully on Windows try: @@ -39,7 +38,7 @@ from litellm.utils import check_valid_key # `models` — default models written into the generated config # --------------------------------------------------------------------------- -PROVIDERS: List[Dict] = [ +PROVIDERS: list[dict] = [ { "id": "openai", "name": "OpenAI", @@ -267,7 +266,7 @@ class SetupWizard: # ── provider selector ─────────────────────────────────────────────────── @staticmethod - def _select_providers() -> List[Dict]: + def _select_providers() -> list[dict]: """Arrow-key multi-select. Falls back to number input if /dev/tty unavailable.""" if not _HAS_RAW_TERMINAL: return SetupWizard._select_fallback() @@ -297,7 +296,7 @@ class SetupWizard: termios.tcsetattr(fd, termios.TCSADRAIN, old) @staticmethod - def _render_selector(cursor: int, selected: Set[int], first_render: bool) -> int: + def _render_selector(cursor: int, selected: set[int], first_render: bool) -> int: """Draw or redraw the provider list. Returns the number of lines printed.""" lines = [ f"\n {bold('Add your first model')}\n", @@ -319,7 +318,7 @@ class SetupWizard: return content.count("\n") @staticmethod - def _select_interactive() -> List[Dict]: + def _select_interactive() -> list[dict]: cursor = 0 selected: set[int] = set() @@ -356,7 +355,7 @@ class SetupWizard: return [PROVIDERS[i] for i in sorted(selected)] @staticmethod - def _select_fallback() -> List[Dict]: + def _select_fallback() -> list[dict]: """Number-based fallback when raw terminal input is unavailable.""" print() print(f" {bold('Add your first model')}") @@ -384,8 +383,8 @@ class SetupWizard: # ── key collection ─────────────────────────────────────────────────────── @staticmethod - def _collect_keys(providers: List[Dict]) -> Dict[str, str]: - env_vars: Dict[str, str] = {} + def _collect_keys(providers: list[dict]) -> dict[str, str]: + env_vars: dict[str, str] = {} print() print(_divider()) print() @@ -422,7 +421,7 @@ class SetupWizard: return env_vars @staticmethod - def _prompt_key(provider: Dict) -> str: + def _prompt_key(provider: dict) -> str: """Prompt for a provider's API key, with skip option. Returns the key or ''.""" hint = grey(provider.get("key_hint", "")) while True: @@ -434,12 +433,12 @@ class SetupWizard: return "" @staticmethod - def _validate_and_report(provider: Dict, api_key: str) -> str: + def _validate_and_report(provider: dict, api_key: str) -> str: """ Validate credentials using litellm.utils.check_valid_key and print result. Offers a re-entry loop on failure. Returns the final (possibly re-entered) key. """ - test_model: Optional[str] = provider.get("test_model") + test_model: str | None = provider.get("test_model") if not test_model: return api_key # Azure / Bedrock / Ollama — skip validation @@ -489,8 +488,8 @@ class SetupWizard: @staticmethod def _build_config( - providers: List[Dict], - env_vars: Dict[str, str], + providers: list[dict], + env_vars: dict[str, str], master_key: str, ) -> str: env_copy = dict(env_vars) # work on a copy — do not mutate caller's dict diff --git a/litellm/skills/__init__.py b/litellm/skills/__init__.py index 96147d5a10f..c6f5005b7f0 100644 --- a/litellm/skills/__init__.py +++ b/litellm/skills/__init__.py @@ -12,12 +12,12 @@ from .main import ( ) __all__ = [ - "create_skill", "acreate_skill", - "list_skills", - "alist_skills", - "get_skill", - "aget_skill", - "delete_skill", "adelete_skill", + "aget_skill", + "alist_skills", + "create_skill", + "delete_skill", + "get_skill", + "list_skills", ] diff --git a/litellm/skills/main.py b/litellm/skills/main.py index 382aceca3af..8450a54e21e 100644 --- a/litellm/skills/main.py +++ b/litellm/skills/main.py @@ -7,7 +7,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -35,7 +35,7 @@ DEFAULT_ANTHROPIC_API_BASE = "https://api.anthropic.com/v1" _litellm_skills_handler = None -def _get_user_api_key_auth_from_kwargs(kwargs: Dict[str, Any]) -> Optional[Any]: +def _get_user_api_key_auth_from_kwargs(kwargs: dict[str, Any]) -> Any | None: for metadata_key in ("metadata", "litellm_metadata"): metadata = kwargs.get(metadata_key) if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None: @@ -44,9 +44,9 @@ def _get_user_api_key_auth_from_kwargs(kwargs: Dict[str, Any]) -> Optional[Any]: def _get_skill_request_metadata( - kwargs: Dict[str, Any], - extra_body: Optional[Dict[str, Any]], -) -> Optional[Dict[str, Any]]: + kwargs: dict[str, Any], + extra_body: dict[str, Any] | None, +) -> dict[str, Any] | None: if extra_body and isinstance(extra_body.get("metadata"), dict): return extra_body["metadata"] @@ -70,13 +70,13 @@ def _get_litellm_skills_handler(): @client async def acreate_skill( - files: Optional[List[Any]] = None, - display_title: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + files: list[Any] | None = None, + display_title: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Skill: """ @@ -133,15 +133,15 @@ async def acreate_skill( @client def create_skill( - files: Optional[List[Any]] = None, - display_title: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + files: list[Any] | None = None, + display_title: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Skill, Coroutine[Any, Any, Skill]]: +) -> Skill | Coroutine[Any, Any, Skill]: """ Create a new skill @@ -161,7 +161,7 @@ def create_skill( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acreate_skill", False) is True # Get LiteLLM parameters @@ -196,10 +196,8 @@ def create_skill( ) # Get provider config for external providers (Anthropic, etc.) - skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( - ProviderConfigManager.get_provider_skills_api_config( - provider=litellm.LlmProviders(custom_llm_provider), - ) + skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( + provider=litellm.LlmProviders(custom_llm_provider), ) if skills_api_provider_config is None: @@ -261,13 +259,13 @@ def create_skill( @client async def alist_skills( - limit: Optional[int] = None, - page: Optional[str] = None, - source: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + page: str | None = None, + source: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> ListSkillsResponse: """ @@ -324,15 +322,15 @@ async def alist_skills( @client def list_skills( - limit: Optional[int] = None, - page: Optional[str] = None, - source: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + limit: int | None = None, + page: str | None = None, + source: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[ListSkillsResponse, Coroutine[Any, Any, ListSkillsResponse]]: +) -> ListSkillsResponse | Coroutine[Any, Any, ListSkillsResponse]: """ List all skills @@ -352,7 +350,7 @@ def list_skills( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("alist_skills", False) is True # Get LiteLLM parameters @@ -374,10 +372,8 @@ def list_skills( ) # Get provider config for external providers (Anthropic, etc.) - skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( - ProviderConfigManager.get_provider_skills_api_config( - provider=litellm.LlmProviders(custom_llm_provider), - ) + skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( + provider=litellm.LlmProviders(custom_llm_provider), ) if skills_api_provider_config is None: @@ -447,10 +443,10 @@ def list_skills( @client async def aget_skill( skill_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> Skill: """ @@ -504,12 +500,12 @@ async def aget_skill( @client def get_skill( skill_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[Skill, Coroutine[Any, Any, Skill]]: +) -> Skill | Coroutine[Any, Any, Skill]: """ Get a skill by ID @@ -527,7 +523,7 @@ def get_skill( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aget_skill", False) is True # Get LiteLLM parameters @@ -548,10 +544,8 @@ def get_skill( ) # Get provider config for external providers (Anthropic, etc.) - skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( - ProviderConfigManager.get_provider_skills_api_config( - provider=litellm.LlmProviders(custom_llm_provider), - ) + skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( + provider=litellm.LlmProviders(custom_llm_provider), ) if skills_api_provider_config is None: @@ -613,10 +607,10 @@ def get_skill( @client async def adelete_skill( skill_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> DeleteSkillResponse: """ @@ -670,12 +664,12 @@ async def adelete_skill( @client def delete_skill( skill_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[DeleteSkillResponse, Coroutine[Any, Any, DeleteSkillResponse]]: +) -> DeleteSkillResponse | Coroutine[Any, Any, DeleteSkillResponse]: """ Delete a skill by ID @@ -693,7 +687,7 @@ def delete_skill( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete_skill", False) is True # Get LiteLLM parameters @@ -714,10 +708,8 @@ def delete_skill( ) # Get provider config for external providers (Anthropic, etc.) - skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( - ProviderConfigManager.get_provider_skills_api_config( - provider=litellm.LlmProviders(custom_llm_provider), - ) + skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( + provider=litellm.LlmProviders(custom_llm_provider), ) if skills_api_provider_config is None: diff --git a/litellm/utils.py b/litellm/utils.py index 9bcf0739361..6ef3871a3c1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -238,12 +238,8 @@ from collections.abc import Callable, Iterable, Mapping from typing import ( TYPE_CHECKING, Any, - Dict, - List, Literal, Optional, - Tuple, - Type, Union, cast, get_args, @@ -446,10 +442,10 @@ greenscaleLogger = None lunaryLogger = None aispendLogger = None supabaseClient = None -callback_list: Optional[List[str]] = [] +callback_list: list[str] | None = [] user_logger_fn = None -additional_details: Optional[Dict[str, str]] = {} -local_cache: Optional[Dict[str, str]] = {} +additional_details: dict[str, str] | None = {} +local_cache: dict[str, str] | None = {} last_fetched_at = None last_fetched_at_keys = None ######## Model Response ######################### @@ -545,7 +541,7 @@ def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: def _custom_logger_class_exists_in_success_callbacks( - callback_class: "CustomLogger", + callback_class: CustomLogger, ) -> bool: """ Returns True if an instance of the custom logger exists in litellm.success_callback or litellm._async_success_callback @@ -560,7 +556,7 @@ def _custom_logger_class_exists_in_success_callbacks( def _custom_logger_class_exists_in_failure_callbacks( - callback_class: "CustomLogger", + callback_class: CustomLogger, ) -> bool: """ Returns True if an instance of the custom logger exists in litellm.failure_callback or litellm._async_failure_callback @@ -574,7 +570,7 @@ def _custom_logger_class_exists_in_failure_callbacks( return any(type(cb) is type(callback_class) for cb in litellm.failure_callback + litellm._async_failure_callback) -def get_request_guardrails(kwargs: Dict[str, Any]) -> List[str]: +def get_request_guardrails(kwargs: dict[str, Any]) -> list[str]: """ Get the request guardrails from the kwargs """ @@ -584,7 +580,7 @@ def get_request_guardrails(kwargs: Dict[str, Any]) -> List[str]: return applied_guardrails -def get_applied_guardrails(kwargs: Dict[str, Any]) -> List[str]: +def get_applied_guardrails(kwargs: dict[str, Any]) -> list[str]: """ - Add 'default_on' guardrails to the list - Add request guardrails to the list @@ -596,9 +592,7 @@ def get_applied_guardrails(kwargs: Dict[str, Any]) -> List[str]: for callback in litellm.callbacks: if callback is not None and isinstance(callback, CustomGuardrail): if callback.guardrail_name is not None: - if callback.default_on is True: - applied_guardrails.append(callback.guardrail_name) - elif callback.guardrail_name in request_guardrails: + if callback.default_on is True or callback.guardrail_name in request_guardrails: applied_guardrails.append(callback.guardrail_name) return applied_guardrails @@ -620,15 +614,15 @@ def load_credentials_from_list(kwargs: dict): def get_dynamic_callbacks( - dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]], -) -> List: + dynamic_callbacks: list[str | Callable | CustomLogger] | None, +) -> list: returned_callbacks = litellm.callbacks.copy() if dynamic_callbacks: returned_callbacks.extend(dynamic_callbacks) # type: ignore return returned_callbacks -def _is_gemini_model(model: Optional[str], custom_llm_provider: Optional[str]) -> bool: +def _is_gemini_model(model: str | None, custom_llm_provider: str | None) -> bool: """ Check if the target model is a Gemini or Vertex AI Gemini model. """ @@ -695,7 +689,7 @@ def _process_tool_message_id(msg_copy: dict, thought_signature_separator: str) - return msg_copy -def _remove_thought_signatures_from_messages(messages: List, thought_signature_separator: str) -> List: +def _remove_thought_signatures_from_messages(messages: list, thought_signature_separator: str) -> list: """ Remove thought signatures from tool call IDs in all messages. """ @@ -741,14 +735,14 @@ def function_setup( applied_guardrails = get_applied_guardrails(kwargs) ## LOGGING SETUP - function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None + function_id: str | None = kwargs["id"] if "id" in kwargs else None ## LAZY LOAD COROUTINE CHECKER ## get_coroutine_checker_fn = getattr(sys.modules[__name__], "get_coroutine_checker") coroutine_checker = get_coroutine_checker_fn() ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = kwargs.pop("callbacks", None) + dynamic_callbacks: list[str | Callable | CustomLogger] | None = kwargs.pop("callbacks", None) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -832,10 +826,10 @@ def function_setup( for index in reversed(removed_async_items): litellm.failure_callback.pop(index) ### DYNAMIC CALLBACKS ### - dynamic_success_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = None - dynamic_async_success_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = None - dynamic_failure_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = None - dynamic_async_failure_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = None + dynamic_success_callbacks: list[str | Callable | CustomLogger] | None = None + dynamic_async_success_callbacks: list[str | Callable | CustomLogger] | None = None + dynamic_failure_callbacks: list[str | Callable | CustomLogger] | None = None + dynamic_async_failure_callbacks: list[str | Callable | CustomLogger] | None = None if kwargs.get("success_callback", None) is not None and isinstance(kwargs["success_callback"], list): removed_async_items = [] for index, callback in enumerate(kwargs["success_callback"]): @@ -953,7 +947,7 @@ def function_setup( except Exception as e: # Log the error but don't fail the request - verbose_logger.warning(f"Error removing thought signatures from tool call IDs: {str(e)}") + verbose_logger.warning(f"Error removing thought signatures from tool call IDs: {e!s}") elif call_type == CallTypes.embedding.value or call_type == CallTypes.aembedding.value: messages = args[1] if len(args) > 1 else kwargs.get("input", None) elif call_type == CallTypes.image_generation.value or call_type == CallTypes.aimage_generation.value: @@ -1010,7 +1004,7 @@ def function_setup( else: messages = "default-message-value" except Exception as e: - verbose_logger.debug(f"Error extracting messages from Google contents: {str(e)}") + verbose_logger.debug(f"Error extracting messages from Google contents: {e!s}") messages = "default-message-value" else: messages = "default-message-value" @@ -1039,7 +1033,7 @@ def function_setup( ) ## check if metadata is passed in - litellm_params: Dict[str, Any] = {"api_base": ""} + litellm_params: dict[str, Any] = {"api_base": ""} if "metadata" in kwargs: litellm_params["metadata"] = kwargs["metadata"] if "litellm_metadata" in kwargs and isinstance(kwargs["litellm_metadata"], dict): @@ -1096,7 +1090,7 @@ async def _client_async_logging_helper( ) -def _get_wrapper_num_retries(kwargs: Dict[str, Any], exception: Exception) -> Tuple[Optional[int], Dict[str, Any]]: +def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]: """ Get the number of retries from the kwargs and the retry policy. Used for the wrapper functions. @@ -1119,13 +1113,13 @@ def _get_wrapper_num_retries(kwargs: Dict[str, Any], exception: Exception) -> Tu return num_retries, kwargs -def _get_wrapper_timeout(kwargs: Dict[str, Any], exception: Exception) -> Optional[Union[float, int, httpx.Timeout]]: +def _get_wrapper_timeout(kwargs: dict[str, Any], exception: Exception) -> float | int | httpx.Timeout | None: """ Get the timeout from the kwargs Used for the wrapper functions. """ - timeout = cast(Optional[Union[float, int, httpx.Timeout]], kwargs.get("timeout", None)) + timeout = cast(float | int | httpx.Timeout | None, kwargs.get("timeout", None)) return timeout @@ -1135,7 +1129,7 @@ def check_coroutine(value) -> bool: return get_coroutine_checker().is_async_callable(value) -async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str): +async def async_pre_call_deployment_hook(kwargs: dict[str, Any], call_type: str): """ Allow modifying the request just before it's sent to the deployment. @@ -1159,8 +1153,8 @@ async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str) async def async_post_call_success_deployment_hook( - request_data: dict, response: Any, call_type: Optional[CallTypes] -) -> Optional[Any]: + request_data: dict, response: Any, call_type: CallTypes | None +) -> Any | None: """ Allow modifying / reviewing the response just after it's received from the deployment. """ @@ -1184,7 +1178,7 @@ async def async_post_call_success_deployment_hook( def post_call_processing( original_response, model, - optional_params: Optional[dict], + optional_params: dict | None, original_function, rules_obj, ): @@ -1199,7 +1193,7 @@ def post_call_processing( pass else: if isinstance(original_response, ModelResponse) and len(original_response.choices) > 0: - model_response: Optional[str] = original_response.choices[0].message.content # type: ignore + model_response: str | None = original_response.choices[0].message.content # type: ignore if model_response is not None: ### POST-CALL RULES ### rules_obj.post_call_rules(input=model_response, model=model) @@ -1222,7 +1216,7 @@ def post_call_processing( and "response_format" in optional_params and optional_params["response_format"] is not None ): - json_response_format: Optional[dict] = None + json_response_format: dict | None = None if ( isinstance( optional_params["response_format"], @@ -1306,13 +1300,13 @@ def client(original_function): print_args_passed_to_litellm(original_function, args, kwargs) start_time = datetime.datetime.now() result = None - logging_obj: Optional[LiteLLMLoggingObject] = kwargs.get("litellm_logging_obj", None) + logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) # only set litellm_call_id if its not in kwargs if "litellm_call_id" not in kwargs: kwargs["litellm_call_id"] = str(uuid.uuid4()) - model: Optional[str] = args[0] if len(args) > 0 else kwargs.get("model", None) + model: str | None = args[0] if len(args) > 0 else kwargs.get("model", None) try: if logging_obj is None: @@ -1325,7 +1319,7 @@ def client(original_function): load_credentials_from_list(kwargs) kwargs["litellm_logging_obj"] = logging_obj LLMCachingHandler = _get_cached_llm_caching_handler() - _llm_caching_handler: "LLMCachingHandler" = LLMCachingHandler( + _llm_caching_handler: LLMCachingHandler = LLMCachingHandler( original_function=original_function, request_kwargs=kwargs, start_time=start_time, @@ -1371,7 +1365,7 @@ def client(original_function): ): # allow users to control returning cached responses from the completion function # checking cache verbose_logger.debug("INSIDE CHECKING SYNC CACHE") - caching_handler_response: "CachingHandlerResponse" = _llm_caching_handler._sync_get_cache( + caching_handler_response: CachingHandlerResponse = _llm_caching_handler._sync_get_cache( model=model or "", original_function=original_function, logging_obj=logging_obj, @@ -1416,7 +1410,7 @@ def client(original_function): ) kwargs["max_tokens"] = modified_max_tokens except Exception as e: - print_verbose(f"Error while checking max token limit: {str(e)}") + print_verbose(f"Error while checking max token limit: {e!s}") # MODEL CALL result = original_function(*args, **kwargs) end_time = datetime.datetime.now() @@ -1441,17 +1435,19 @@ def client(original_function): end_time=end_time, ) return result - elif "acompletion" in kwargs and kwargs["acompletion"] is True: - return result - elif "aembedding" in kwargs and kwargs["aembedding"] is True: - return result - elif "aimg_generation" in kwargs and kwargs["aimg_generation"] is True: - return result - elif "atranscription" in kwargs and kwargs["atranscription"] is True: - return result - elif "aspeech" in kwargs and kwargs["aspeech"] is True: - return result - elif asyncio.iscoroutine(result): # bubble up to relevant async function + elif ( + "acompletion" in kwargs + and kwargs["acompletion"] is True + or "aembedding" in kwargs + and kwargs["aembedding"] is True + or "aimg_generation" in kwargs + and kwargs["aimg_generation"] is True + or "atranscription" in kwargs + and kwargs["atranscription"] is True + or "aspeech" in kwargs + and kwargs["aspeech"] is True + or asyncio.iscoroutine(result) + ): return result ### POST-CALL RULES ### @@ -1578,9 +1574,9 @@ def client(original_function): start_time = datetime.datetime.now() result = None _update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") - logging_obj: Optional[LiteLLMLoggingObject] = kwargs.get("litellm_logging_obj", None) + logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) LLMCachingHandler = _get_cached_llm_caching_handler() - _llm_caching_handler: "LLMCachingHandler" = LLMCachingHandler( + _llm_caching_handler: LLMCachingHandler = LLMCachingHandler( original_function=original_function, request_kwargs=kwargs, start_time=start_time, @@ -1590,7 +1586,7 @@ def client(original_function): if "litellm_call_id" not in kwargs: kwargs["litellm_call_id"] = str(uuid.uuid4()) - model: Optional[str] = args[0] if len(args) > 0 else kwargs.get("model", None) + model: str | None = args[0] if len(args) > 0 else kwargs.get("model", None) is_completion_with_fallbacks = kwargs.get("fallbacks") is not None kwargs.pop("_is_litellm_internal_call", None) # discard if injected _is_litellm_internal_call = is_internal_call.get() @@ -1628,7 +1624,7 @@ def client(original_function): print_verbose( f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}" ) - _caching_handler_response: "Optional[CachingHandlerResponse]" = await _llm_caching_handler._async_get_cache( + _caching_handler_response: CachingHandlerResponse | None = await _llm_caching_handler._async_get_cache( model=model or "", original_function=original_function, logging_obj=logging_obj, @@ -1679,7 +1675,7 @@ def client(original_function): ) kwargs["max_tokens"] = modified_max_tokens except Exception as e: - print_verbose(f"Error while checking max token limit: {str(e)}") + print_verbose(f"Error while checking max token limit: {e!s}") # MODEL CALL result = await original_function(*args, **kwargs) @@ -1878,7 +1874,7 @@ def client(original_function): def _is_async_request( - kwargs: Optional[dict], + kwargs: dict | None, is_pass_through: bool = False, ) -> bool: """ @@ -1920,8 +1916,8 @@ _STREAMING_CALL_TYPES = frozenset( def _is_streaming_request( - kwargs: Dict[str, Any], - call_type: Union[CallTypes, str], + kwargs: dict[str, Any], + call_type: CallTypes | str, ) -> bool: """ Returns True if the call type is a streaming request. @@ -1934,7 +1930,7 @@ def _is_streaming_request( return call_type in _STREAMING_CALL_TYPES -def _select_tokenizer(model: str, custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None): +def _select_tokenizer(model: str, custom_tokenizer: CustomHuggingfaceTokenizer | None = None): if custom_tokenizer is not None: _tokenizer = create_pretrained_tokenizer( identifier=custom_tokenizer["identifier"], @@ -1965,7 +1961,7 @@ def _return_openai_tokenizer(model: str) -> SelectTokenizerResponse: return {"type": "openai_tokenizer", "tokenizer": _get_default_encoding()} -def _return_huggingface_tokenizer(model: str) -> Optional[SelectTokenizerResponse]: +def _return_huggingface_tokenizer(model: str) -> SelectTokenizerResponse | None: if model in litellm.cohere_models and "command-r" in model: # cohere cohere_tokenizer = Tokenizer.from_pretrained("Xenova/c4ai-command-r-v01-tokenizer") @@ -1986,7 +1982,7 @@ def _return_huggingface_tokenizer(model: str) -> Optional[SelectTokenizerRespons return None -def encode(model="", text="", custom_tokenizer: Optional[dict] = None): +def encode(model="", text="", custom_tokenizer: dict | None = None): """ Encodes the given text using the specified model. @@ -2012,8 +2008,8 @@ def encode(model="", text="", custom_tokenizer: Optional[dict] = None): def decode( model="", - tokens: List[int] = [], - custom_tokenizer: Optional[dict] = None, + tokens: list[int] = [], + custom_tokenizer: dict | None = None, skip_special_tokens: bool = True, ): """ @@ -2034,7 +2030,7 @@ def decode( return dec -def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: List[int]) -> List[int]: +def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: list[int]) -> list[int]: try: added_tokens_decoder = tokenizer.get_added_tokens_decoder() except Exception: @@ -2048,7 +2044,7 @@ def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: List[int] return [token for token in tokens if token not in special_token_ids] -def create_pretrained_tokenizer(identifier: str, revision="main", auth_token: Optional[str] = None): +def create_pretrained_tokenizer(identifier: str, revision="main", auth_token: str | None = None): """ Creates a tokenizer from an existing file on a HuggingFace repository to be used with `token_counter`. @@ -2090,14 +2086,14 @@ def create_tokenizer(json: str): def token_counter( model="", - custom_tokenizer: Optional[Union[dict, SelectTokenizerResponse]] = None, - text: Optional[Union[str, List[str]]] = None, - messages: Optional[List] = None, - count_response_tokens: Optional[bool] = False, - tools: Optional[List[ChatCompletionToolParam]] = None, - tool_choice: Optional[ChatCompletionNamedToolChoiceParam] = None, - use_default_image_token_count: Optional[bool] = False, - default_token_count: Optional[int] = None, + custom_tokenizer: dict | SelectTokenizerResponse | None = None, + text: str | list[str] | None = None, + messages: list | None = None, + count_response_tokens: bool | None = False, + tools: list[ChatCompletionToolParam] | None = None, + tool_choice: ChatCompletionNamedToolChoiceParam | None = None, + use_default_image_token_count: bool | None = False, + default_token_count: int | None = None, ) -> int: """ The same as `litellm.litellm_core_utils.token_counter`. @@ -2139,7 +2135,7 @@ def supports_httpx_timeout(custom_llm_provider: str) -> bool: return False -def supports_system_messages(model: str, custom_llm_provider: Optional[str]) -> bool: +def supports_system_messages(model: str, custom_llm_provider: str | None) -> bool: """ Check if the given model supports system messages and return a boolean value. @@ -2160,7 +2156,7 @@ def supports_system_messages(model: str, custom_llm_provider: Optional[str]) -> ) -def supports_web_search(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_web_search(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports web search and return a boolean value. @@ -2181,7 +2177,7 @@ def supports_web_search(model: str, custom_llm_provider: Optional[str] = None) - ) -def supports_url_context(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_url_context(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports URL context and return a boolean value. @@ -2202,7 +2198,7 @@ def supports_url_context(model: str, custom_llm_provider: Optional[str] = None) ) -def supports_native_streaming(model: str, custom_llm_provider: Optional[str]) -> bool: +def supports_native_streaming(model: str, custom_llm_provider: str | None) -> bool: """ Check if the given model supports native streaming and return a boolean value. @@ -2228,12 +2224,12 @@ def supports_native_streaming(model: str, custom_llm_provider: Optional[str]) -> return supports_native_streaming except Exception as e: verbose_logger.debug( - f"Model not found or error in checking supports_native_streaming support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {str(e)}" + f"Model not found or error in checking supports_native_streaming support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {e!s}" ) return False -def supports_response_schema(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_response_schema(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model + provider supports 'response_schema' as a param. @@ -2252,7 +2248,7 @@ def supports_response_schema(model: str, custom_llm_provider: Optional[str] = No model, custom_llm_provider, _, _ = get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) except Exception as e: verbose_logger.debug( - f"Model not found or error in checking response schema support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {str(e)}" + f"Model not found or error in checking response schema support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {e!s}" ) return False @@ -2274,7 +2270,7 @@ def supports_response_schema(model: str, custom_llm_provider: Optional[str] = No ) -def supports_parallel_function_calling(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_parallel_function_calling(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports parallel tool calls and return a boolean value. """ @@ -2285,7 +2281,7 @@ def supports_parallel_function_calling(model: str, custom_llm_provider: Optional ) -def supports_function_calling(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_function_calling(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports function calling and return a boolean value. @@ -2306,16 +2302,14 @@ def supports_function_calling(model: str, custom_llm_provider: Optional[str] = N ) -def supports_tool_choice(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_tool_choice(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports `tool_choice` and return a boolean value. """ return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_tool_choice") -def _supports_provider_info_factory( - model: str, custom_llm_provider: Optional[str], key: str -) -> Optional[Literal[True]]: +def _supports_provider_info_factory(model: str, custom_llm_provider: str | None, key: str) -> Literal[True] | None: """ Check if the given model supports a provider specific model info and return a boolean value. """ @@ -2327,7 +2321,7 @@ def _supports_provider_info_factory( return None -def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str) -> bool: +def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: """ Check if the given model supports function calling and return a boolean value. @@ -2368,7 +2362,7 @@ def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str) return False except Exception as e: verbose_logger.debug( - f"Model not found or error in checking {key} support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {str(e)}" + f"Model not found or error in checking {key} support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {e!s}" ) supported_by_provider = _supports_provider_info_factory(model, custom_llm_provider, key) @@ -2378,7 +2372,7 @@ def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str) return False -def _is_explicitly_disabled_factory(model: str, custom_llm_provider: Optional[str], key: str) -> bool: +def _is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: """Return True only when the model map explicitly sets *key* to ``False``. This is the opt-out mirror of :func:`_supports_factory`. Where @@ -2410,27 +2404,27 @@ def _is_explicitly_disabled_factory(model: str, custom_llm_provider: Optional[st verbose_logger.debug( f"Model not found or error in checking {key} disabled state. " f"You passed model={model}, custom_llm_provider={custom_llm_provider}. " - f"Error: {str(e)}" + f"Error: {e!s}" ) return False -def supports_audio_input(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_audio_input(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports audio input in a chat completion call""" return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") -def supports_pdf_input(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_pdf_input(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports pdf input in a chat completion call""" return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_pdf_input") -def supports_audio_output(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_audio_output(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports audio output in a chat completion call""" return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") -def supports_prompt_caching(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports prompt caching and return a boolean value. @@ -2451,7 +2445,7 @@ def supports_prompt_caching(model: str, custom_llm_provider: Optional[str] = Non ) -def supports_computer_use(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_computer_use(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports computer use and return a boolean value. @@ -2472,7 +2466,7 @@ def supports_computer_use(model: str, custom_llm_provider: Optional[str] = None) ) -def supports_vision(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_vision(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports vision and return a boolean value. @@ -2490,14 +2484,14 @@ def supports_vision(model: str, custom_llm_provider: Optional[str] = None) -> bo ) -def supports_reasoning(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_reasoning(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports reasoning and return a boolean value. """ return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_reasoning") -def supports_native_structured_output(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_native_structured_output(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports native structured outputs and return a boolean value. """ @@ -2508,7 +2502,7 @@ def supports_native_structured_output(model: str, custom_llm_provider: Optional[ ) -def get_supported_regions(model: str, custom_llm_provider: Optional[str] = None) -> Optional[List[str]]: +def get_supported_regions(model: str, custom_llm_provider: str | None = None) -> list[str] | None: """ Get a list of supported regions for a given model and provider. @@ -2543,12 +2537,12 @@ def get_supported_regions(model: str, custom_llm_provider: Optional[str] = None) return None except Exception as e: verbose_logger.debug( - f"Model not found or error in checking supported_regions support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {str(e)}" + f"Model not found or error in checking supported_regions support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {e!s}" ) return None -def supports_embedding_image_input(model: str, custom_llm_provider: Optional[str] = None) -> bool: +def supports_embedding_image_input(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model supports embedding image input and return a boolean value. """ @@ -2560,7 +2554,7 @@ def supports_embedding_image_input(model: str, custom_llm_provider: Optional[str ####### HELPER FUNCTIONS ################ -def _update_dictionary(existing_dict: Dict, new_dict: dict) -> dict: +def _update_dictionary(existing_dict: dict, new_dict: dict) -> dict: for k, v in new_dict.items(): if v is not None: # Convert stringified numbers to appropriate numeric types @@ -2615,7 +2609,7 @@ _CACHE_PRICING_FIELDS = ( ) -def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[str, Any]]: +def _resolve_builtin_model_cost_entry(key: str, provider: str) -> dict[str, Any] | None: """Best-effort lookup of a built-in ``model_cost`` entry for a custom key whose shape ``get_model_info`` cannot resolve (repeated provider prefixes like ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region @@ -2625,7 +2619,7 @@ def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[ (most importantly cache pricing) without mutating the shared built-in. Returns ``None`` when no safe match exists. """ - candidates: List[str] = [] + candidates: list[str] = [] segments = key.split("/") idx = 0 while idx < len(segments) - 1 and segments[idx] in LlmProvidersSet: @@ -2649,7 +2643,7 @@ def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[ return None -def _get_builtin_model_info_for_registration(model: str) -> Optional[ModelInfo]: +def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None: """Resolve ``model`` to its built-in cost-map entry for registration merging. Returns ``None`` when the lookup raises or when it resolved via a @@ -2669,7 +2663,7 @@ def _get_builtin_model_info_for_registration(model: str) -> Optional[ModelInfo]: return None -def register_model(model_cost: Union[str, dict]): +def register_model(model_cost: str | dict): """ Register new / Override existing models (and their pricing) to specific providers. Provide EITHER a model cost dictionary or a url to a hosted json blob @@ -2818,7 +2812,7 @@ def _should_drop_param(k, additional_drop_params) -> bool: return False -def _get_non_default_params(passed_params: dict, default_params: dict, additional_drop_params: Optional[list]) -> dict: +def _get_non_default_params(passed_params: dict, default_params: dict, additional_drop_params: list | None) -> dict: non_default_params = {} for k, v in passed_params.items(): if ( @@ -2834,12 +2828,12 @@ def _get_non_default_params(passed_params: dict, default_params: dict, additiona def get_optional_params_transcription( model: str, custom_llm_provider: str, - language: Optional[str] = None, - prompt: Optional[str] = None, - response_format: Optional[str] = None, - temperature: Optional[int] = None, - timestamp_granularities: Optional[List[Literal["word", "segment"]]] = None, - drop_params: Optional[bool] = None, + language: str | None = None, + prompt: str | None = None, + response_format: str | None = None, + temperature: int | None = None, + timestamp_granularities: list[Literal["word", "segment"]] | None = None, + drop_params: bool | None = None, **kwargs, ): from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS @@ -2881,7 +2875,7 @@ def get_optional_params_transcription( ) return non_default_params - provider_config: Optional[BaseAudioTranscriptionConfig] = None + provider_config: BaseAudioTranscriptionConfig | None = None if custom_llm_provider is not None: provider_config = ProviderConfigManager.get_provider_audio_transcription_config( model=model, @@ -2920,7 +2914,7 @@ def get_optional_params_transcription( return optional_params -def _map_openai_size_to_vertex_ai_aspect_ratio(size: Optional[str]) -> str: +def _map_openai_size_to_vertex_ai_aspect_ratio(size: str | None) -> str: """Map OpenAI size parameter to Vertex AI aspectRatio.""" if size is None: return "1:1" @@ -2938,18 +2932,18 @@ def _map_openai_size_to_vertex_ai_aspect_ratio(size: Optional[str]) -> str: def get_optional_params_image_gen( - model: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[str] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, - imageConfig: Optional[dict] = None, - custom_llm_provider: Optional[str] = None, - additional_drop_params: Optional[list] = None, - provider_config: Optional[BaseImageGenerationConfig] = None, - drop_params: Optional[bool] = None, + model: str | None = None, + n: int | None = None, + quality: str | None = None, + response_format: str | None = None, + size: str | None = None, + style: str | None = None, + user: str | None = None, + imageConfig: dict | None = None, + custom_llm_provider: str | None = None, + additional_drop_params: list | None = None, + provider_config: BaseImageGenerationConfig | None = None, + drop_params: bool | None = None, **kwargs, ): # retrieve all parameters passed to the function @@ -2961,12 +2955,13 @@ def get_optional_params_image_gen( additional_drop_params = passed_params.pop("additional_drop_params", None) special_params = passed_params.pop("kwargs") for k, v in special_params.items(): - if k.startswith("aws_") and ( - custom_llm_provider != "bedrock" and custom_llm_provider != "sagemaker" + if ( + k.startswith("aws_") + and (custom_llm_provider != "bedrock" and custom_llm_provider != "sagemaker") + or k == "hf_model_name" + and custom_llm_provider != "sagemaker" ): # allow dynamically setting boto3 init logic continue - elif k == "hf_model_name" and custom_llm_provider != "sagemaker": - continue elif ( k.startswith("vertex_") and custom_llm_provider != "vertex_ai" and custom_llm_provider != "vertex_ai_beta" ): # allow dynamically setting vertex ai init logic @@ -2990,7 +2985,7 @@ def get_optional_params_image_gen( default_params=default_params, additional_drop_params=additional_drop_params, ) - optional_params: Dict[str, Any] = {} + optional_params: dict[str, Any] = {} ## raise exception if non-default value passed for non-openai/azure embedding calls def _check_valid_arg(supported_params): @@ -3064,13 +3059,13 @@ def get_optional_params_image_gen( def get_optional_params_embeddings( # 2 optional params model: str, - user: Optional[str] = None, - encoding_format: Optional[str] = None, - dimensions: Optional[int] = None, + user: str | None = None, + encoding_format: str | None = None, + dimensions: int | None = None, custom_llm_provider="", - drop_params: Optional[bool] = None, - additional_drop_params: Optional[List[str]] = None, - allowed_openai_params: Optional[List[str]] = None, + drop_params: bool | None = None, + additional_drop_params: list[str] | None = None, + allowed_openai_params: list[str] | None = None, **kwargs, ): # Lazy load get_supported_openai_params @@ -3087,7 +3082,7 @@ def get_optional_params_embeddings( # Remove function objects from passed_params to avoid JSON serialization errors passed_params.pop("get_supported_openai_params", None) - def _check_valid_arg(supported_params: Optional[list]): + def _check_valid_arg(supported_params: list | None): if supported_params is None: return unsupported_params = {} @@ -3111,7 +3106,7 @@ def get_optional_params_embeddings( model=model, ) - provider_config: Optional[BaseEmbeddingConfig] = None + provider_config: BaseEmbeddingConfig | None = None optional_params = {} if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values(): @@ -3121,7 +3116,7 @@ def get_optional_params_embeddings( ) if provider_config is not None: - supported_params: Optional[list] = provider_config.get_supported_openai_params(model=model) + supported_params: list | None = provider_config.get_supported_openai_params(model=model) _check_valid_arg(supported_params=supported_params) optional_params = provider_config.map_openai_params( non_default_params=non_default_params, @@ -3486,14 +3481,14 @@ def _remove_json_schema_refs(schema, max_depth=10): return schema -def _remove_unsupported_params(non_default_params: dict, supported_openai_params: Optional[List[str]]) -> dict: +def _remove_unsupported_params(non_default_params: dict, supported_openai_params: list[str] | None) -> dict: """ Remove unsupported params from non_default_params """ remove_keys = [] if supported_openai_params is None: return {} # no supported params, so no optional openai params to send - for param in non_default_params.keys(): + for param in non_default_params: if param not in supported_openai_params: remove_keys.append(param) for key in remove_keys: @@ -3540,23 +3535,24 @@ class PreProcessNonDefaultParams: passed_params: dict, special_params: dict, custom_llm_provider: str, - additional_drop_params: Optional[List[str]], + additional_drop_params: list[str] | None, default_param_values: dict, - additional_endpoint_specific_params: List[str], + additional_endpoint_specific_params: list[str], ) -> dict: for k, v in special_params.items(): if k == "aws_bedrock_project_id": # sent as a request header (read from litellm_params by the # bedrock-mantle configs), never as a request body field continue - if k.startswith("aws_") and ( - custom_llm_provider != "bedrock" and not custom_llm_provider.startswith("sagemaker") + if ( + k.startswith("aws_") + and (custom_llm_provider != "bedrock" and not custom_llm_provider.startswith("sagemaker")) + or k == "hf_model_name" + and custom_llm_provider != "sagemaker" + or k.startswith("vertex_") + and not _provider_supports_vertex_params(custom_llm_provider) ): # allow dynamically setting boto3 init logic continue - elif k == "hf_model_name" and custom_llm_provider != "sagemaker": - continue - elif k.startswith("vertex_") and not _provider_supports_vertex_params(custom_llm_provider): - continue passed_params[k] = v # filter out those parameters that were passed with non-default values @@ -3584,7 +3580,7 @@ class PreProcessNonDefaultParams: passed_params: dict, special_params: dict, custom_llm_provider: str, - additional_drop_params: Optional[List[str]], + additional_drop_params: list[str] | None, model: str, remove_sensitive_keys: bool = False, add_provider_specific_params: bool = False, @@ -3605,11 +3601,11 @@ def pre_process_non_default_params( passed_params: dict, special_params: dict, custom_llm_provider: str, - additional_drop_params: Optional[List[str]], + additional_drop_params: list[str] | None, model: str, remove_sensitive_keys: bool = False, add_provider_specific_params: bool = False, - provider_config: Optional[BaseConfig] = None, + provider_config: BaseConfig | None = None, ) -> dict: """ Pre-process non-default params to a standardized format @@ -3668,7 +3664,7 @@ def remove_sensitive_keys_from_dict(d: dict) -> dict: """ sensitive_key_phrases = ["key", "secret", "access", "credential"] remove_keys = [] - for key in d.keys(): + for key in d: if any(phrase in key.lower() for phrase in sensitive_key_phrases): remove_keys.append(key) for key in remove_keys: @@ -3678,7 +3674,7 @@ def remove_sensitive_keys_from_dict(d: dict) -> dict: def pre_process_optional_params(passed_params: dict, non_default_params: dict, custom_llm_provider: str) -> dict: """For .completion(), preprocess optional params""" - optional_params: Dict = {} + optional_params: dict = {} common_auth_dict = litellm.common_cloud_provider_auth_params if custom_llm_provider in common_auth_dict["providers"]: @@ -3787,15 +3783,15 @@ def get_optional_params( api_version=None, parallel_tool_calls=None, drop_params=None, - allowed_openai_params: Optional[List[str]] = None, + allowed_openai_params: list[str] | None = None, reasoning_effort=None, verbosity=None, additional_drop_params=None, - messages: Optional[List[AllMessageValues]] = None, - thinking: Optional[AnthropicThinkingParam] = None, - web_search_options: Optional[OpenAIWebSearchOptions] = None, - safety_identifier: Optional[str] = None, - base_model: Optional[str] = None, + messages: list[AllMessageValues] | None = None, + thinking: AnthropicThinkingParam | None = None, + web_search_options: OpenAIWebSearchOptions | None = None, + safety_identifier: str | None = None, + base_model: str | None = None, **kwargs, ): passed_params = locals().copy() @@ -3804,7 +3800,7 @@ def get_optional_params( # non_default_params / _check_valid_arg — it's a routing hint, not an # OpenAI param. passed_params.pop("base_model", None) - provider_config: Optional[BaseConfig] = None + provider_config: BaseConfig | None = None if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]: provider_config = ProviderConfigManager.get_provider_chat_config( model=model, @@ -3825,7 +3821,7 @@ def get_optional_params( custom_llm_provider=custom_llm_provider, ) - def _check_valid_arg(supported_params: List[str]): + def _check_valid_arg(supported_params: list[str]): """ Check if the params passed to completion() are supported by the provider @@ -3852,7 +3848,7 @@ def get_optional_params( if unsupported_params: if litellm.drop_params is True or (drop_params is not None and drop_params is True): - for k in unsupported_params.keys(): + for k in unsupported_params: non_default_params.pop(k, None) else: raise UnsupportedParamsError( @@ -4357,8 +4353,8 @@ def add_provider_specific_params_to_optional_params( optional_params: dict, passed_params: dict, custom_llm_provider: str, - openai_params: List[str], - additional_drop_params: Optional[list] = None, + openai_params: list[str], + additional_drop_params: list | None = None, ) -> dict: """ Add provider specific params to optional_params @@ -4368,7 +4364,7 @@ def add_provider_specific_params_to_optional_params( # for openai, azure we should pass the extra/passed params within `extra_body` https://github.com/openai/openai-python/blob/ac33853ba10d13ac149b1fa3ca6dba7d613065c9/src/openai/resources/models.py#L46 if _should_drop_param(k="extra_body", additional_drop_params=additional_drop_params) is False: extra_body = dict(passed_params.pop("extra_body", None) or {}) - for k in passed_params.keys(): + for k in passed_params: if k not in openai_params and passed_params[k] is not None: extra_body[k] = passed_params[k] if not isinstance(optional_params.get("extra_body"), dict): @@ -4386,7 +4382,7 @@ def add_provider_specific_params_to_optional_params( _ensure_extra_body_is_safe = getattr(sys.modules[__name__], "_ensure_extra_body_is_safe") optional_params["extra_body"] = _ensure_extra_body_is_safe(extra_body=processed_extra_body) else: - for k in passed_params.keys(): + for k in passed_params: if k not in openai_params and passed_params[k] is not None: if _should_drop_param(k=k, additional_drop_params=additional_drop_params): continue @@ -4434,11 +4430,11 @@ def get_non_default_params(passed_params: dict) -> dict: def calculate_max_parallel_requests( - max_parallel_requests: Optional[int], - rpm: Optional[int], - tpm: Optional[int], - default_max_parallel_requests: Optional[int], -) -> Optional[int]: + max_parallel_requests: int | None, + rpm: int | None, + tpm: int | None, + default_max_parallel_requests: int | None, +) -> int | None: """ Returns the max parallel requests to send to a deployment. @@ -4476,7 +4472,7 @@ def calculate_max_parallel_requests( return None -def _get_deployment_order(deployment: Union[Dict, Any]) -> Optional[int]: +def _get_deployment_order(deployment: dict | Any) -> int | None: """ Returns the routing order for a deployment. @@ -4489,7 +4485,7 @@ def _get_deployment_order(deployment: Union[Dict, Any]) -> Optional[int]: return order -def _get_order_filtered_deployments(healthy_deployments: List[Dict], target_order: Optional[int] = None) -> List: +def _get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list: if target_order is not None: filtered = [d for d in healthy_deployments if _get_deployment_order(d) == target_order] if filtered: @@ -4498,10 +4494,10 @@ def _get_order_filtered_deployments(healthy_deployments: List[Dict], target_orde return healthy_deployments # Default: pick min order group - _valid_orders: List[int] = [ + _valid_orders: list[int] = [ o for deployment in healthy_deployments for o in [_get_deployment_order(deployment)] if o is not None ] - min_order: Optional[int] = min(_valid_orders) if _valid_orders else None + min_order: int | None = min(_valid_orders) if _valid_orders else None if min_order is not None: filtered_deployments = [ @@ -4513,9 +4509,9 @@ def _get_order_filtered_deployments(healthy_deployments: List[Dict], target_orde def _get_excluded_filtered_deployments( - healthy_deployments: List[Dict], - excluded_deployment_ids: Optional[Iterable[str]] = None, -) -> List: + healthy_deployments: list[dict], + excluded_deployment_ids: Iterable[str] | None = None, +) -> list: """ Filter out deployments whose `model_info.id` appears in `excluded_deployment_ids`. @@ -4535,7 +4531,7 @@ def _get_excluded_filtered_deployments( return [d for d in healthy_deployments if (d.get("model_info") or {}).get("id") not in excluded_set] -def _get_model_region(custom_llm_provider: str, litellm_params: LiteLLM_Params) -> Optional[str]: +def _get_model_region(custom_llm_provider: str, litellm_params: LiteLLM_Params) -> str | None: """ Return the region for a model, for a given provider """ @@ -4560,7 +4556,7 @@ def _get_model_region(custom_llm_provider: str, litellm_params: LiteLLM_Params) return litellm_params.region_name -def _infer_model_region(litellm_params: LiteLLM_Params) -> Optional[AllowedModelRegion]: +def _infer_model_region(litellm_params: LiteLLM_Params) -> AllowedModelRegion | None: """ Infer if a model is in the EU or US region @@ -4575,7 +4571,7 @@ def _infer_model_region(litellm_params: LiteLLM_Params) -> Optional[AllowedModel model_region = _get_model_region(custom_llm_provider=custom_llm_provider, litellm_params=litellm_params) if model_region is None: - verbose_logger.debug("Cannot infer model region for model: {}".format(litellm_params.model)) + verbose_logger.debug(f"Cannot infer model region for model: {litellm_params.model}") return None if custom_llm_provider == "azure": @@ -4640,7 +4636,7 @@ def is_region_allowed(litellm_params: LiteLLM_Params, allowed_model_region: str) return False -def get_model_region(litellm_params: LiteLLM_Params, mode: Optional[str]) -> Optional[str]: +def get_model_region(litellm_params: LiteLLM_Params, mode: str | None) -> str | None: """ Pass the litellm params for an azure model, and get back the region """ @@ -4659,7 +4655,7 @@ def get_model_region(litellm_params: LiteLLM_Params, mode: Optional[str]) -> Opt mode=mode or "chat", ) - region: Optional[str] = response.get("x-ms-region", None) + region: str | None = response.get("x-ms-region", None) return region return None @@ -4679,7 +4675,7 @@ def _count_characters(text: str) -> int: return len(filtered_text) -def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) -> str: +def get_response_string(response_obj: ModelResponse | ModelResponseStream) -> str: # Handle Responses API streaming events if hasattr(response_obj, "type") and hasattr(response_obj, "response"): # This is a Responses API streaming event (e.g., ResponseCreatedEvent, ResponseCompletedEvent) @@ -4690,7 +4686,7 @@ def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) # Use list accumulation to avoid O(n^2) string concatenation: # repeatedly doing `response_str += part` copies the full string each time # because Python strings are immutable, so total work grows with n^2. - response_output_parts: List[str] = [] + response_output_parts: list[str] = [] for output_item in output_list: # Handle output items with content array if hasattr(output_item, "content"): @@ -4710,10 +4706,10 @@ def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) return delta if isinstance(delta, str) else "" # Handle standard ModelResponse and ModelResponseStream - _choices: Union[List[Choices], List[StreamingChoices]] = response_obj.choices + _choices: list[Choices] | list[StreamingChoices] = response_obj.choices # Use list accumulation to avoid O(n^2) string concatenation across choices - response_parts: List[str] = [] + response_parts: list[str] = [] for choice in _choices: if isinstance(choice, Choices): if choice.message.content is not None: @@ -4725,7 +4721,7 @@ def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) return "".join(response_parts) -def get_api_key(llm_provider: str, dynamic_api_key: Optional[str]): +def get_api_key(llm_provider: str, dynamic_api_key: str | None): api_key = dynamic_api_key or litellm.api_key # openai if llm_provider == "openai" or llm_provider == "text-completion-openai": @@ -4778,7 +4774,7 @@ def get_utc_datetime(): return datetime.utcnow() # type: ignore -def get_max_tokens(model: str) -> Optional[int]: +def get_max_tokens(model: str) -> int | None: """ Get the maximum number of output tokens allowed for a given model. @@ -4872,14 +4868,16 @@ def _strip_openai_finetune_model_name(model_name: str) -> str: return re.sub(r"(:[^:]+){3}$", "", model_name) -def _strip_model_name(model: str, custom_llm_provider: Optional[str]) -> str: +def _strip_model_name(model: str, custom_llm_provider: str | None) -> str: if custom_llm_provider and custom_llm_provider in ["bedrock", "bedrock_converse"]: stripped_bedrock_model = _get_base_bedrock_model(model_name=model) return stripped_bedrock_model - elif custom_llm_provider and (custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini"): - strip_version = _strip_stable_vertex_version(model_name=model) - return strip_version - elif custom_llm_provider and (custom_llm_provider == "databricks"): + elif ( + custom_llm_provider + and (custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini") + or custom_llm_provider + and (custom_llm_provider == "databricks") + ): strip_version = _strip_stable_vertex_version(model_name=model) return strip_version elif "ft:" in model: @@ -4890,7 +4888,7 @@ def _strip_model_name(model: str, custom_llm_provider: Optional[str]) -> str: # Global case-insensitive lookup map for model_cost (built eagerly at module import) -_model_cost_lowercase_map: Optional[Dict[str, str]] = None +_model_cost_lowercase_map: dict[str, str] | None = None # Monotonic counter bumped on every model_cost mutation. Consumers that # memoize derived state (e.g. provider-specific indices) can include this @@ -4918,7 +4916,7 @@ def _invalidate_model_cost_lowercase_map() -> None: _cached_get_model_info_helper.cache_clear() -def _rebuild_model_cost_lowercase_map() -> Dict[str, str]: +def _rebuild_model_cost_lowercase_map() -> dict[str, str]: """Rebuild the case-insensitive lookup map from the current model_cost. Returns: @@ -4931,7 +4929,7 @@ def _rebuild_model_cost_lowercase_map() -> Dict[str, str]: def _handle_stale_map_entry_rebuild( potential_key_lower: str, -) -> Optional[str]: +) -> str | None: """ Handle stale _model_cost_lowercase_map entry (key was popped). @@ -4950,7 +4948,7 @@ def _handle_stale_map_entry_rebuild( def _handle_new_key_with_scan( potential_key_lower: str, -) -> Optional[str]: +) -> str | None: """ Handle new key added to model_cost without invalidating _model_cost_lowercase_map. @@ -4967,7 +4965,7 @@ def _handle_new_key_with_scan( return None -def _get_model_cost_key(potential_key: str) -> Optional[str]: +def _get_model_cost_key(potential_key: str) -> str | None: """ Get the actual key from model_cost, with case-insensitive fallback. @@ -5012,7 +5010,7 @@ def _get_model_info_from_model_cost(key: str) -> dict: return litellm.model_cost[key] -def _check_provider_match(model_info: dict, custom_llm_provider: Optional[str]) -> bool: +def _check_provider_match(model_info: dict, custom_llm_provider: str | None) -> bool: """ Check if the model info provider matches the custom provider. @@ -5025,11 +5023,14 @@ def _check_provider_match(model_info: dict, custom_llm_provider: Optional[str]) if custom_llm_provider and ( model_info.get("litellm_provider") is not None and model_info["litellm_provider"] != custom_llm_provider ): - if custom_llm_provider == "vertex_ai" and model_info["litellm_provider"].startswith("vertex_ai"): - return True - elif custom_llm_provider == "fireworks_ai" and model_info["litellm_provider"].startswith("fireworks_ai"): - return True - elif custom_llm_provider.startswith("bedrock") and model_info["litellm_provider"].startswith("bedrock"): + if ( + custom_llm_provider == "vertex_ai" + and model_info["litellm_provider"].startswith("vertex_ai") + or custom_llm_provider == "fireworks_ai" + and model_info["litellm_provider"].startswith("fireworks_ai") + or custom_llm_provider.startswith("bedrock") + and model_info["litellm_provider"].startswith("bedrock") + ): return True elif ( custom_llm_provider == "litellm_proxy" @@ -5066,8 +5067,8 @@ class PotentialModelNamesAndCustomLLMProvider(TypedDict): def _get_model_info_from_generalization( model: str, potential_model_names: PotentialModelNamesAndCustomLLMProvider, - custom_llm_provider: Optional[str], -) -> Optional[tuple[str, dict]]: + custom_llm_provider: str | None, +) -> tuple[str, dict] | None: """Resolve an unmapped model via the declarative capability generalization rules. Tries the same name candidates as the exact lookups, in the same order, and @@ -5100,9 +5101,7 @@ def _get_model_info_from_generalization( return None -def _get_potential_model_names( - model: str, custom_llm_provider: Optional[str] -) -> PotentialModelNamesAndCustomLLMProvider: +def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> PotentialModelNamesAndCustomLLMProvider: if custom_llm_provider is None: # Get custom_llm_provider try: @@ -5119,15 +5118,12 @@ def _get_potential_model_names( split_model = model.split("/", 1)[1] combined_model_name = model stripped_model_name = _strip_model_name(model=split_model, custom_llm_provider=custom_llm_provider) - combined_stripped_model_name = "{}/{}".format(custom_llm_provider, stripped_model_name) + combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" else: split_model = model - combined_model_name = "{}/{}".format(custom_llm_provider, model) + combined_model_name = f"{custom_llm_provider}/{model}" stripped_model_name = _strip_model_name(model=model, custom_llm_provider=custom_llm_provider) - combined_stripped_model_name = "{}/{}".format( - custom_llm_provider, - stripped_model_name, - ) + combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" if custom_llm_provider in ("bedrock", "bedrock_converse"): from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix @@ -5143,7 +5139,7 @@ def _get_potential_model_names( ) -def _get_max_position_embeddings(model_name: str) -> Optional[int]: +def _get_max_position_embeddings(model_name: str) -> int | None: # Construct the URL for the config.json file config_url = f"https://huggingface.co/{model_name}/raw/main/config.json" @@ -5169,8 +5165,8 @@ def _get_max_position_embeddings(model_name: str) -> Optional[int]: @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) def _cached_get_model_info_helper( model: str, - custom_llm_provider: Optional[str], - api_base: Optional[str] = None, + custom_llm_provider: str | None, + api_base: str | None = None, ) -> ModelInfoBase: """ _get_model_info_helper wrapped with lru_cache @@ -5184,18 +5180,18 @@ def _cached_get_model_info_helper( ) -def get_provider_info(model: str, custom_llm_provider: Optional[str]) -> Optional[ProviderSpecificModelInfo]: +def get_provider_info(model: str, custom_llm_provider: str | None) -> ProviderSpecificModelInfo | None: ## PROVIDER-SPECIFIC INFORMATION # if custom_llm_provider == "predibase": # _model_info["supports_response_schema"] = True - provider_config: Optional[BaseLLMModelInfo] = None + provider_config: BaseLLMModelInfo | None = None if custom_llm_provider and custom_llm_provider in LlmProvidersSet: # Check if the provider string exists in LlmProviders enum provider_config = ProviderConfigManager.get_provider_model_info( model=model, provider=LlmProviders(custom_llm_provider) ) - model_info: Optional[ProviderSpecificModelInfo] = None + model_info: ProviderSpecificModelInfo | None = None if provider_config: model_info = provider_config.get_provider_info(model=model) @@ -5219,9 +5215,9 @@ _ABOVE_THRESHOLD_COST_KEY = re.compile(r"_above_\d+k?_tokens$") def _get_model_info_helper( model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, ) -> ModelInfoBase: """ Helper for 'get_model_info'. Separated out to avoid infinite loop caused by returning 'supported_openai_param's @@ -5235,9 +5231,9 @@ def _get_model_info_helper( if custom_llm_provider is not None and custom_llm_provider == "vertex_ai": if "meta/" + model in litellm.vertex_llama3_models: model = "meta/" + model - elif model + "@latest" in litellm.vertex_mistral_models: - model = model + "@latest" - elif model + "@latest" in litellm.vertex_ai_ai21_models: + elif ( + model + "@latest" in litellm.vertex_mistral_models or model + "@latest" in litellm.vertex_ai_ai21_models + ): model = model + "@latest" ########################## potential_model_names = _get_potential_model_names(model=model, custom_llm_provider=custom_llm_provider) @@ -5251,7 +5247,7 @@ def _get_model_info_helper( custom_llm_provider = potential_model_names["custom_llm_provider"] model_cost_custom_llm_provider = custom_llm_provider ######################### - provider_config: Optional[BaseLLMModelInfo] = None + provider_config: BaseLLMModelInfo | None = None if custom_llm_provider and custom_llm_provider in LlmProvidersSet: provider_config = ProviderConfigManager.get_provider_model_info( model=model, provider=LlmProviders(custom_llm_provider) @@ -5306,8 +5302,8 @@ def _get_model_info_helper( 5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given. """ - _model_info: Optional[Dict[str, Any]] = None - key: Optional[str] = None + _model_info: dict[str, Any] | None = None + key: str | None = None # Use case-insensitive lookup for all model name checks _matched_key = _get_model_cost_key(combined_model_name) @@ -5373,23 +5369,19 @@ def _get_model_info_helper( raise ValueError( "This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) - _input_cost_per_token: Optional[float] = _model_info.get("input_cost_per_token") + _input_cost_per_token: float | None = _model_info.get("input_cost_per_token") if _input_cost_per_token is None: # default value to 0, be noisy about this verbose_logger.debug( - "model={}, custom_llm_provider={} has no input_cost_per_token in model_cost_map. Defaulting to 0.".format( - model, custom_llm_provider - ) + f"model={model}, custom_llm_provider={custom_llm_provider} has no input_cost_per_token in model_cost_map. Defaulting to 0." ) _input_cost_per_token = 0 - _output_cost_per_token: Optional[float] = _model_info.get("output_cost_per_token") + _output_cost_per_token: float | None = _model_info.get("output_cost_per_token") if _output_cost_per_token is None: # default value to 0, be noisy about this verbose_logger.debug( - "model={}, custom_llm_provider={} has no output_cost_per_token in model_cost_map. Defaulting to 0.".format( - model, custom_llm_provider - ) + f"model={model}, custom_llm_provider={custom_llm_provider} has no output_cost_per_token in model_cost_map. Defaulting to 0." ) _output_cost_per_token = 0 @@ -5558,17 +5550,15 @@ def _get_model_info_helper( except Exception as e: verbose_logger.debug(f"Error getting model info: {e}") raise Exception( - "This model isn't mapped yet. model={}, custom_llm_provider={}. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json.".format( - model, custom_llm_provider - ) + f"This model isn't mapped yet. model={model}, custom_llm_provider={custom_llm_provider}. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json." ) def _build_model_info( model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, ) -> ModelInfo: supported_openai_params = litellm.get_supported_openai_params(model=model, custom_llm_provider=custom_llm_provider) @@ -5594,17 +5584,17 @@ def _build_model_info( @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) def _cached_get_model_info( model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, ) -> ModelInfo: return _build_model_info(model=model, custom_llm_provider=custom_llm_provider, api_base=api_base) def get_model_info( model: str, - custom_llm_provider: Optional[str] = None, - api_base: Optional[str] = None, - api_key: Optional[str] = None, + custom_llm_provider: str | None = None, + api_base: str | None = None, + api_key: str | None = None, ) -> ModelInfo: """ Get a dict for the maximum tokens (context window), input_cost_per_token, output_cost_per_token for a given model. @@ -5851,7 +5841,7 @@ def load_test_model( } -def get_provider_fields(custom_llm_provider: str) -> List[ProviderField]: +def get_provider_fields(custom_llm_provider: str) -> list[ProviderField]: """Return the fields required for each provider""" if custom_llm_provider == "databricks": @@ -5893,10 +5883,10 @@ def create_proxy_transport_and_mounts(): def validate_environment( - model: Optional[str] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, ) -> dict: """ Checks if the environment variables are valid for the given model. @@ -5911,7 +5901,7 @@ def validate_environment( - missing_keys (List[str]): A list of missing keys in the environment. """ keys_in_environment = False - missing_keys: List[str] = [] + missing_keys: list[str] = [] if model is None: return { @@ -6279,7 +6269,7 @@ def validate_environment( else: missing_keys.append("WANDB_API_KEY") - def filter_missing_keys(keys: List[str], exclude_pattern: str) -> List[str]: + def filter_missing_keys(keys: list[str], exclude_pattern: str) -> list[str]: """Filter out keys that contain the exclude_pattern (case insensitive).""" return [key for key in keys if exclude_pattern not in key.lower()] @@ -6383,7 +6373,7 @@ def _should_retry(status_code: int): def _get_retry_after_from_exception_header( - response_headers: Optional[httpx.Headers] = None, + response_headers: httpx.Headers | None = None, ): """ Reimplementation of openai's calculate retry after, since that one can't be imported. @@ -6419,9 +6409,9 @@ def _get_retry_after_from_exception_header( def _calculate_retry_after( remaining_retries: int, max_retries: int, - response_headers: Optional[httpx.Headers] = None, + response_headers: httpx.Headers | None = None, min_timeout: int = 0, -) -> Union[float, int]: +) -> float | int: retry_after = _get_retry_after_from_exception_header(response_headers) # Add some jitter (default JITTER is 0.75 - so upto 0.75s) @@ -6515,8 +6505,8 @@ class TextCompletionStreamWrapper: self, completion_stream, model, - stream_options: Optional[dict] = None, - custom_llm_provider: Optional[str] = None, + stream_options: dict | None = None, + custom_llm_provider: str | None = None, ): self.completion_stream = completion_stream self.model = model @@ -6552,7 +6542,7 @@ class TextCompletionStreamWrapper: return response except Exception as e: - raise Exception(f"Error occurred converting to text completion object - chunk: {chunk}; Error: {str(e)}") + raise Exception(f"Error occurred converting to text completion object - chunk: {chunk}; Error: {e!s}") def __next__(self): # model_response = ModelResponse(stream=True, model=self.model) @@ -6588,7 +6578,7 @@ class TextCompletionStreamWrapper: raise StopAsyncIteration -def mock_completion_streaming_obj(model_response, mock_response, model, n: Optional[int] = None): +def mock_completion_streaming_obj(model_response, mock_response, model, n: int | None = None): if isinstance(mock_response, litellm.MockException): raise mock_response if isinstance(mock_response, ModelResponseStream): @@ -6612,9 +6602,9 @@ def mock_completion_streaming_obj(model_response, mock_response, model, n: Optio async def async_mock_completion_streaming_obj( model_response, - mock_response: Union[str, "MockException", ModelResponseStream], + mock_response: str | MockException | ModelResponseStream, model, - n: Optional[int] = None, + n: int | None = None, ): if isinstance(mock_response, litellm.MockException): raise mock_response @@ -6725,7 +6715,7 @@ def get_token_count(messages, model): return token_counter(model=model, messages=messages) -def shorten_message_to_fit_limit(message, tokens_needed, model: Optional[str], raise_error_on_max_limit: bool = False): +def shorten_message_to_fit_limit(message, tokens_needed, model: str | None, raise_error_on_max_limit: bool = False): """ Shorten a message to fit within a token limit by removing characters from the middle. @@ -6783,7 +6773,7 @@ def shorten_message_to_fit_limit(message, tokens_needed, model: Optional[str], r # Credits for this code go to Killian Lucas def trim_messages( messages, - model: Optional[str] = None, + model: str | None = None, trim_ratio: float = DEFAULT_TRIM_RATIO, return_response_tokens: bool = False, max_tokens=None, @@ -6848,7 +6838,7 @@ def trim_messages( print_verbose( f"Need to trim input messages: {messages}, current_tokens{current_tokens}, max_tokens: {max_tokens}" ) - system_message_event: Optional[dict] = None + system_message_event: dict | None = None if system_message: system_message_event, max_tokens = process_system_message( system_message=system_message, max_tokens=max_tokens, model=model @@ -6878,7 +6868,7 @@ def trim_messages( return final_messages, response_tokens return final_messages except Exception as e: # [NON-Blocking, if error occurs just return final_messages - verbose_logger.exception("Got exception while token trimming - {}".format(str(e))) + verbose_logger.exception(f"Got exception while token trimming - {e!s}") return original_messages @@ -6888,7 +6878,7 @@ from litellm.caching.in_memory_cache import InMemoryCache class AvailableModelsCache(InMemoryCache): def __init__(self, ttl_seconds: int = 300, max_size: int = 1000): super().__init__(ttl_seconds, max_size) - self._env_hash: Optional[str] = None + self._env_hash: str | None = None def _get_env_hash(self) -> str: """Create a hash of relevant environment variables""" @@ -6905,8 +6895,8 @@ class AvailableModelsCache(InMemoryCache): def _get_cache_key( self, - custom_llm_provider: Optional[str], - litellm_params: Optional[LiteLLM_Params], + custom_llm_provider: str | None, + litellm_params: LiteLLM_Params | None, ) -> str: valid_str = "" @@ -6918,9 +6908,9 @@ class AvailableModelsCache(InMemoryCache): def get_cached_model_info( self, - custom_llm_provider: Optional[str] = None, - litellm_params: Optional[LiteLLM_Params] = None, - ) -> Optional[List[str]]: + custom_llm_provider: str | None = None, + litellm_params: LiteLLM_Params | None = None, + ) -> list[str] | None: """Get cached model info""" # Check if environment has changed if litellm_params is None and self._check_env_changed(): @@ -6929,7 +6919,7 @@ class AvailableModelsCache(InMemoryCache): cache_key = self._get_cache_key(custom_llm_provider, litellm_params) - result = cast(Optional[List[str]], self.get_cache(cache_key)) + result = cast(list[str] | None, self.get_cache(cache_key)) if result is not None: return copy.deepcopy(result) @@ -6938,8 +6928,8 @@ class AvailableModelsCache(InMemoryCache): def set_cached_model_info( self, custom_llm_provider: str, - litellm_params: Optional[LiteLLM_Params], - available_models: List[str], + litellm_params: LiteLLM_Params | None, + available_models: list[str], ): """Set cached model info""" cache_key = self._get_cache_key(custom_llm_provider, litellm_params) @@ -6951,9 +6941,9 @@ _model_cache = AvailableModelsCache() def _infer_valid_provider_from_env_vars( - custom_llm_provider: Optional[str] = None, -) -> List[str]: - valid_providers: List[str] = [] + custom_llm_provider: str | None = None, +) -> list[str]: + valid_providers: list[str] = [] environ_keys = os.environ.keys() for provider in litellm.provider_list: if custom_llm_provider and provider != custom_llm_provider: @@ -6977,8 +6967,8 @@ def _infer_valid_provider_from_env_vars( def _get_valid_models_from_provider_api( provider_config: BaseLLMModelInfo, custom_llm_provider: str, - litellm_params: Optional[LiteLLM_Params] = None, -) -> List[str]: + litellm_params: LiteLLM_Params | None = None, +) -> list[str]: try: cached_result = _model_cache.get_cached_model_info(custom_llm_provider, litellm_params) @@ -6997,12 +6987,12 @@ def _get_valid_models_from_provider_api( def get_valid_models( - check_provider_endpoint: Optional[bool] = None, - custom_llm_provider: Optional[str] = None, - litellm_params: Optional[LiteLLM_Params] = None, - api_key: Optional[str] = None, - api_base: Optional[str] = None, -) -> List[str]: + check_provider_endpoint: bool | None = None, + custom_llm_provider: str | None = None, + litellm_params: LiteLLM_Params | None = None, + api_key: str | None = None, + api_base: str | None = None, +) -> list[str]: """ Returns a list of valid LLMs based on the set environment variables @@ -7032,8 +7022,8 @@ def get_valid_models( check_provider_endpoint = check_provider_endpoint or litellm.check_provider_endpoint # get keys set in .env - valid_providers: List[str] = [] - valid_models: List[str] = [] + valid_providers: list[str] = [] + valid_models: list[str] = [] # for all valid providers, make a list of supported llms if custom_llm_provider: @@ -7075,19 +7065,23 @@ def print_args_passed_to_litellm(original_function, args, kwargs): return try: # we've already printed this for acompletion, don't print for completion - if "acompletion" in kwargs and kwargs["acompletion"] is True and original_function.__name__ == "completion": - return - elif "aembedding" in kwargs and kwargs["aembedding"] is True and original_function.__name__ == "embedding": - return - elif ( - "aimg_generation" in kwargs - and kwargs["aimg_generation"] is True - and original_function.__name__ == "img_generation" + if ( + "acompletion" in kwargs + and kwargs["acompletion"] is True + and original_function.__name__ == "completion" + or "aembedding" in kwargs + and kwargs["aembedding"] is True + and original_function.__name__ == "embedding" + or ( + "aimg_generation" in kwargs + and kwargs["aimg_generation"] is True + and original_function.__name__ == "img_generation" + ) ): return args_str = ", ".join(map(repr, args)) - kwargs_str = ", ".join(f"{key}={repr(value)}" for key, value in kwargs.items()) + kwargs_str = ", ".join(f"{key}={value!r}" for key, value in kwargs.items()) print_verbose( "\n", ) # new line before @@ -7147,7 +7141,7 @@ class ModelResponseIterator: if convert_to_delta is True: _stream_response = ModelResponseStream() _stream_response.choices[0].delta.content = model_response.choices[0].message.content # type: ignore - self.model_response: Union[ModelResponse, ModelResponseStream] = _stream_response + self.model_response: ModelResponse | ModelResponseStream = _stream_response else: self.model_response = model_response self.is_done = False @@ -7174,7 +7168,7 @@ class ModelResponseIterator: class ModelResponseListIterator: - def __init__(self, model_responses, delay: Optional[float] = None): + def __init__(self, model_responses, delay: float | None = None): self.model_responses = model_responses self.index = 0 self.delay = delay @@ -7288,7 +7282,7 @@ def get_base64_str(s: str) -> str: return s -def has_tool_call_blocks(messages: List[AllMessageValues]) -> bool: +def has_tool_call_blocks(messages: list[AllMessageValues]) -> bool: """ Returns true, if messages has tool call blocks. @@ -7301,7 +7295,7 @@ def has_tool_call_blocks(messages: List[AllMessageValues]) -> bool: def any_assistant_message_has_thinking_blocks( - messages: List[AllMessageValues], + messages: list[AllMessageValues], ) -> bool: """ Returns true if ANY assistant message has thinking_blocks. @@ -7322,7 +7316,7 @@ def any_assistant_message_has_thinking_blocks( def last_assistant_with_tool_calls_has_no_thinking_blocks( - messages: List[AllMessageValues], + messages: list[AllMessageValues], ) -> bool: """ Returns true if the last assistant message with tool_calls has no thinking_blocks. @@ -7353,7 +7347,7 @@ def last_assistant_with_tool_calls_has_no_thinking_blocks( return thinking_blocks is None or (hasattr(thinking_blocks, "__len__") and len(thinking_blocks) == 0) -def add_dummy_tool(custom_llm_provider: str) -> List[ChatCompletionToolParam]: +def add_dummy_tool(custom_llm_provider: str) -> list[ChatCompletionToolParam]: """ Prevent Anthropic from raising error when tool_use block exists but no tools are provided. @@ -7384,7 +7378,7 @@ from litellm.types.llms.openai import ( ) -def convert_to_dict(message: Union[BaseModel, dict]) -> dict: +def convert_to_dict(message: BaseModel | dict) -> dict: """ Converts a message to a dictionary if it's a Pydantic model. @@ -7402,7 +7396,7 @@ def convert_to_dict(message: Union[BaseModel, dict]) -> dict: raise TypeError(f"Invalid message type: {type(message)}. Expected dict or Pydantic model.") -def convert_list_message_to_dict(messages: List): +def convert_list_message_to_dict(messages: list): new_messages = [] for message in messages: convert_msg_to_dict = cast(AllMessageValues, convert_to_dict(message)) @@ -7411,7 +7405,7 @@ def convert_list_message_to_dict(messages: List): return new_messages -def validate_and_fix_openai_messages(messages: List): +def validate_and_fix_openai_messages(messages: list): """ Ensures all messages are valid OpenAI chat completion messages. @@ -7430,7 +7424,7 @@ def validate_and_fix_openai_messages(messages: List): return validate_chat_completion_user_messages(messages=new_messages) -def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]]: +def validate_and_fix_openai_tools(tools: list | None) -> list[dict] | None: """ Ensure tools is List[dict] and not List[BaseModel] """ @@ -7446,8 +7440,8 @@ def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]] def validate_and_fix_thinking_param( - thinking: Optional["AnthropicThinkingParam"], -) -> Optional["AnthropicThinkingParam"]: + thinking: AnthropicThinkingParam | None, +) -> AnthropicThinkingParam | None: """ Normalizes camelCase keys in the thinking param to snake_case. Handles clients that send budgetTokens instead of budget_tokens. @@ -7472,7 +7466,7 @@ def cleanup_none_field_in_message(message: AllMessageValues): return {k: v for k, v in new_message.items() if v is not None} -def validate_chat_completion_user_messages(messages: List[AllMessageValues]): +def validate_chat_completion_user_messages(messages: list[AllMessageValues]): """ Ensures all user messages are valid OpenAI chat completion messages. @@ -7514,8 +7508,8 @@ def validate_chat_completion_user_messages(messages: List[AllMessageValues]): def validate_chat_completion_tool_choice( - tool_choice: Optional[Union[dict, str]], -) -> Optional[Union[dict, str]]: + tool_choice: dict | str | None, +) -> dict | str | None: """ Confirm the tool choice is passed in the OpenAI format. @@ -7526,9 +7520,7 @@ def validate_chat_completion_tool_choice( ChatCompletionToolChoiceStringValues, ) - if tool_choice is None: - return tool_choice - elif isinstance(tool_choice, str): + if tool_choice is None or isinstance(tool_choice, str): return tool_choice elif isinstance(tool_choice, dict): # Handle Cursor IDE format: {"type": "auto"} -> return as-is @@ -7546,9 +7538,7 @@ def validate_chat_completion_tool_choice( ) -def validate_openai_optional_params( - stop: Optional[Union[str, List[str]]] = None, **kwargs -) -> Optional[Union[str, List[str]]]: +def validate_openai_optional_params(stop: str | list[str] | None = None, **kwargs) -> str | list[str] | None: """ Validates and fixes OpenAI optional parameters. @@ -7568,7 +7558,7 @@ def validate_openai_optional_params( @lru_cache(maxsize=1) -def _get_bundled_model_cost_map() -> Dict[str, Any]: +def _get_bundled_model_cost_map() -> dict[str, Any]: try: model_cost_path = resources.files("litellm").joinpath("model_prices_and_context_window_backup.json") return json.loads(model_cost_path.read_text()) @@ -7579,7 +7569,7 @@ def _get_bundled_model_cost_map() -> Dict[str, Any]: def _get_model_cost_entry_for_provider_config( model: str, provider: LlmProviders, -) -> Dict[str, Any]: +) -> dict[str, Any]: candidate_keys = (model, f"{provider.value}/{model}") for model_key in candidate_keys: model_info = litellm.model_cost.get(model_key) @@ -7598,7 +7588,7 @@ class ProviderConfigManager: # Dictionary mapping for O(1) provider lookup # Stores tuples of (factory_function, needs_model_parameter) # This is initialized lazily on first access to avoid circular imports - _PROVIDER_CONFIG_MAP: Optional[dict[LlmProviders, tuple[Callable, bool]]] = None + _PROVIDER_CONFIG_MAP: dict[LlmProviders, tuple[Callable, bool]] | None = None @staticmethod def _build_provider_config_map() -> dict[LlmProviders, tuple[Callable, bool]]: @@ -7768,7 +7758,7 @@ class ProviderConfigManager: } @staticmethod - def _get_azure_config(model: str, base_model: Optional[str] = None) -> BaseConfig: + def _get_azure_config(model: str, base_model: str | None = None) -> BaseConfig: """Get Azure config based on model type. When *base_model* is provided (e.g. ``"azure/gpt-5.2"``), it is used @@ -7847,8 +7837,8 @@ class ProviderConfigManager: def get_provider_chat_config( model: str, provider: LlmProviders, - base_model: Optional[str] = None, - ) -> Optional[BaseConfig]: + base_model: str | None = None, + ) -> BaseConfig | None: """ Returns the provider config for a given provider. @@ -7898,7 +7888,7 @@ class ProviderConfigManager: def get_provider_embedding_config( model: str, provider: LlmProviders, - ) -> Optional[BaseEmbeddingConfig]: + ) -> BaseEmbeddingConfig | None: if ( litellm.LlmProviders.VOYAGE == provider and litellm.VoyageContextualEmbeddingConfig.is_contextualized_embeddings(model) @@ -7985,8 +7975,8 @@ class ProviderConfigManager: def get_provider_rerank_config( model: str, provider: LlmProviders, - api_base: Optional[str], - present_version_params: List[str], + api_base: str | None, + present_version_params: list[str], ) -> BaseRerankConfig: if litellm.LlmProviders.COHERE == provider or litellm.LlmProviders.COHERE_CHAT == provider: if should_use_cohere_v1_client(api_base, present_version_params): @@ -8031,7 +8021,7 @@ class ProviderConfigManager: def get_provider_anthropic_messages_config( model: str, provider: LlmProviders, - ) -> Optional[BaseAnthropicMessagesConfig]: + ) -> BaseAnthropicMessagesConfig | None: return ProviderConfigManager._get_provider_anthropic_messages_config_cached(model=model, provider=provider) @staticmethod @@ -8039,7 +8029,7 @@ class ProviderConfigManager: def _get_provider_anthropic_messages_config_cached( model: str, provider: LlmProviders, - ) -> Optional[BaseAnthropicMessagesConfig]: + ) -> BaseAnthropicMessagesConfig | None: model_lower = model.lower() if litellm.LlmProviders.ANTHROPIC == provider: return litellm.AnthropicMessagesConfig() @@ -8104,7 +8094,7 @@ class ProviderConfigManager: def get_provider_audio_transcription_config( model: str, provider: LlmProviders, - ) -> Optional[BaseAudioTranscriptionConfig]: + ) -> BaseAudioTranscriptionConfig | None: model_cost_entry = _get_model_cost_entry_for_provider_config( model=model, provider=provider, @@ -8183,9 +8173,9 @@ class ProviderConfigManager: @staticmethod def get_provider_responses_api_config( - provider: Union[LlmProviders, str], - model: Optional[str] = None, - ) -> Optional[BaseResponsesAPIConfig]: + provider: LlmProviders | str, + model: str | None = None, + ) -> BaseResponsesAPIConfig | None: from litellm.llms.openai_like.dynamic_config import ( create_responses_config_class, ) @@ -8196,7 +8186,7 @@ class ProviderConfigManager: # Try to convert to enum for Python class lookup first. # Python classes take priority over JSON (they have custom overrides). - provider_enum: Optional[LlmProviders] = None + provider_enum: LlmProviders | None = None if isinstance(provider, LlmProviders): provider_enum = provider else: @@ -8220,9 +8210,9 @@ class ProviderConfigManager: @staticmethod def _get_python_responses_api_config( - provider: Optional[LlmProviders], - model: Optional[str] = None, - ) -> Optional[BaseResponsesAPIConfig]: + provider: LlmProviders | None, + model: str | None = None, + ) -> BaseResponsesAPIConfig | None: """Check for Python-class-based responses API configs (custom overrides).""" if provider is None: return None @@ -8294,7 +8284,7 @@ class ProviderConfigManager: @staticmethod def get_provider_skills_api_config( provider: LlmProviders, - ) -> Optional["BaseSkillsAPIConfig"]: + ) -> BaseSkillsAPIConfig | None: """ Get provider-specific Skills API configuration @@ -8311,7 +8301,7 @@ class ProviderConfigManager: @staticmethod def get_provider_evals_api_config( provider: LlmProviders, - ) -> Optional["BaseEvalsAPIConfig"]: + ) -> BaseEvalsAPIConfig | None: """ Get provider-specific Evals API configuration @@ -8342,9 +8332,9 @@ class ProviderConfigManager: @staticmethod def get_provider_model_info( - model: Optional[str], + model: str | None, provider: LlmProviders, - ) -> Optional[BaseLLMModelInfo]: + ) -> BaseLLMModelInfo | None: if LlmProviders.FIREWORKS_AI == provider: return litellm.FireworksAIConfig() elif LlmProviders.OPENAI == provider: @@ -8392,7 +8382,7 @@ class ProviderConfigManager: def get_provider_passthrough_config( model: str, provider: LlmProviders, - ) -> Optional[BasePassthroughConfig]: + ) -> BasePassthroughConfig | None: if LlmProviders.BEDROCK == provider: from litellm.llms.bedrock.passthrough.transformation import ( BedrockPassthroughConfig, @@ -8423,7 +8413,7 @@ class ProviderConfigManager: def get_provider_image_variation_config( model: str, provider: LlmProviders, - ) -> Optional[BaseImageVariationConfig]: + ) -> BaseImageVariationConfig | None: if LlmProviders.OPENAI == provider: return litellm.OpenAIImageVariationConfig() elif LlmProviders.TOPAZ == provider: @@ -8434,7 +8424,7 @@ class ProviderConfigManager: def get_provider_files_config( model: str, provider: LlmProviders, - ) -> Optional[BaseFilesConfig]: + ) -> BaseFilesConfig | None: if LlmProviders.GEMINI == provider: from litellm.llms.gemini.files.transformation import ( GoogleAIStudioFilesHandler, # experimental approach, to reduce bloat on __init__.py @@ -8463,7 +8453,7 @@ class ProviderConfigManager: def get_provider_batches_config( model: str, provider: LlmProviders, - ) -> Optional[BaseBatchesConfig]: + ) -> BaseBatchesConfig | None: if LlmProviders.BEDROCK == provider: from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig @@ -8473,7 +8463,7 @@ class ProviderConfigManager: @staticmethod def get_provider_vector_store_config( provider: LlmProviders, - ) -> Optional[CustomLogger]: + ) -> CustomLogger | None: from litellm.integrations.vector_store_integrations.bedrock_vector_store import ( BedrockVectorStore, ) @@ -8485,8 +8475,8 @@ class ProviderConfigManager: @staticmethod def get_provider_vector_stores_config( provider: LlmProviders, - api_type: Optional[str] = None, - ) -> Optional[BaseVectorStoreConfig]: + api_type: str | None = None, + ) -> BaseVectorStoreConfig | None: """ v2 vector store config, use this for new vector store integrations """ @@ -8562,7 +8552,7 @@ class ProviderConfigManager: @staticmethod def get_provider_vector_store_files_config( provider: LlmProviders, - ) -> Optional[BaseVectorStoreFilesConfig]: + ) -> BaseVectorStoreFilesConfig | None: if litellm.LlmProviders.OPENAI == provider: from litellm.llms.openai.vector_store_files.transformation import ( OpenAIVectorStoreFilesConfig, @@ -8575,7 +8565,7 @@ class ProviderConfigManager: def get_provider_image_generation_config( model: str, provider: LlmProviders, - ) -> Optional[BaseImageGenerationConfig]: + ) -> BaseImageGenerationConfig | None: if LlmProviders.OPENAI == provider: from litellm.llms.openai.image_generation import ( get_openai_image_generation_config, @@ -8682,9 +8672,9 @@ class ProviderConfigManager: @staticmethod def get_provider_video_config( - model: Optional[str], + model: str | None, provider: LlmProviders, - ) -> Optional[BaseVideoConfig]: + ) -> BaseVideoConfig | None: if LlmProviders.OPENAI == provider: from litellm.llms.openai.videos.transformation import OpenAIVideoConfig @@ -8710,7 +8700,7 @@ class ProviderConfigManager: @staticmethod def get_provider_container_config( provider: LlmProviders, - ) -> Optional[BaseContainerConfig]: + ) -> BaseContainerConfig | None: if LlmProviders.OPENAI == provider: from litellm.llms.openai.containers.transformation import ( OpenAIContainerConfig, @@ -8729,7 +8719,7 @@ class ProviderConfigManager: def get_provider_realtime_config( model: str, provider: LlmProviders, - ) -> Optional[BaseRealtimeConfig]: + ) -> BaseRealtimeConfig | None: if LlmProviders.GEMINI == provider: from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig @@ -8740,7 +8730,7 @@ class ProviderConfigManager: def get_provider_realtime_http_config( model: str, provider: LlmProviders, - ) -> Optional["BaseRealtimeHTTPConfig"]: + ) -> BaseRealtimeHTTPConfig | None: """ Return the HTTP transformation config for realtime HTTP endpoints (POST /realtime/client_secrets and POST /realtime/calls). @@ -8764,7 +8754,7 @@ class ProviderConfigManager: def get_provider_image_edit_config( model: str, provider: LlmProviders, - ) -> Optional[BaseImageEditConfig]: + ) -> BaseImageEditConfig | None: if LlmProviders.OPENAI == provider: from litellm.llms.openai.image_edit import get_openai_image_edit_config @@ -8831,7 +8821,7 @@ class ProviderConfigManager: def get_provider_ocr_config( model: str, provider: LlmProviders, - ) -> Optional["BaseOCRConfig"]: + ) -> BaseOCRConfig | None: """ Get OCR configuration for a given provider. """ @@ -8871,8 +8861,8 @@ class ProviderConfigManager: @staticmethod def get_provider_search_config( - provider: "SearchProviders", - ) -> Optional["BaseSearchConfig"]: + provider: SearchProviders, + ) -> BaseSearchConfig | None: """ Get Search configuration for a given provider. """ @@ -8944,7 +8934,7 @@ class ProviderConfigManager: def get_provider_text_to_speech_config( model: str, provider: LlmProviders, - ) -> Optional["BaseTextToSpeechConfig"]: + ) -> BaseTextToSpeechConfig | None: """ Get text-to-speech configuration for a given provider. """ @@ -8997,7 +8987,7 @@ class ProviderConfigManager: def get_provider_google_genai_generate_content_config( model: str, provider: LlmProviders, - ) -> Optional[BaseGoogleGenAIGenerateContentConfig]: + ) -> BaseGoogleGenAIGenerateContentConfig | None: if litellm.LlmProviders.GEMINI == provider: from litellm.llms.gemini.google_genai.transformation import ( GoogleGenAIConfig, @@ -9031,7 +9021,7 @@ class ProviderConfigManager: def get_end_user_id_for_cost_tracking( litellm_params: dict, service_type: Literal["litellm_logging", "prometheus"] = "litellm_logging", -) -> Optional[str]: +) -> str | None: """ Used for enforcing `disable_end_user_cost_tracking` param. @@ -9041,7 +9031,7 @@ def get_end_user_id_for_cost_tracking( _metadata = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params))) end_user_id = cast( - Optional[str], + str | None, litellm_params.get("user_api_key_end_user_id") or _metadata.get("user_api_key_end_user_id"), ) if litellm.disable_end_user_cost_tracking: @@ -9057,7 +9047,7 @@ def get_end_user_id_for_cost_tracking( return end_user_id -def should_use_cohere_v1_client(api_base: Optional[str], present_version_params: List[str]): +def should_use_cohere_v1_client(api_base: str | None, present_version_params: list[str]): if not api_base: return False uses_v1_params = ("max_chunks_per_doc" in present_version_params) and ( @@ -9092,9 +9082,9 @@ def get_prompt_cache_min_tokens(model: str) -> int: def is_prompt_caching_valid_prompt( model: str, - messages: Optional[List[AllMessageValues]], - tools: Optional[List[ChatCompletionToolParam]] = None, - custom_llm_provider: Optional[str] = None, + messages: list[AllMessageValues] | None, + tools: list[ChatCompletionToolParam] | None = None, + custom_llm_provider: str | None = None, min_token_count: int | None = None, ) -> bool: """ @@ -9126,7 +9116,7 @@ def is_prompt_caching_valid_prompt( return False -def extract_duration_from_srt_or_vtt(srt_or_vtt_content: str) -> Optional[float]: +def extract_duration_from_srt_or_vtt(srt_or_vtt_content: str) -> float | None: """ Extracts the total duration (in seconds) from SRT or VTT content. @@ -9206,7 +9196,7 @@ def get_non_default_completion_params(kwargs: dict) -> dict: return non_default_params -def peek_reasoning_summary_aliases(optional_params: dict) -> Optional[Any]: +def peek_reasoning_summary_aliases(optional_params: dict) -> Any | None: """Read AI-SDK-style reasoning summary from optional_params or nested extra_body. Uses key membership (not ``or`` chains) so falsy values like ``""`` are not skipped. @@ -9226,7 +9216,7 @@ def peek_reasoning_summary_aliases(optional_params: dict) -> Optional[Any]: def strip_reasoning_summary_aliases_from_optional_params( optional_params: dict, -) -> Tuple[dict, Optional[Any]]: +) -> tuple[dict, Any | None]: """Copy optional_params; remove reasoningSummary aliases from top-level and extra_body.""" op = dict(optional_params) rs_val = op.pop("reasoningSummary", None) @@ -9258,8 +9248,8 @@ def get_non_default_transcription_params(kwargs: dict) -> dict: def add_openai_metadata( - metadata: Optional[Mapping[str, Any]], -) -> Optional[Dict[str, str]]: + metadata: Mapping[str, Any] | None, +) -> dict[str, str] | None: """ Add metadata to openai optional parameters, excluding hidden params. @@ -9275,7 +9265,7 @@ def add_openai_metadata( if metadata is None: return None # Only include non-hidden parameters - visible_metadata: Dict[str, str] = { + visible_metadata: dict[str, str] = { str(k): v for k, v in metadata.items() if k != "hidden_params" and isinstance(v, str) } @@ -9352,13 +9342,13 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict ) -def jsonify_tools(tools: List[Any]) -> List[Dict]: +def jsonify_tools(tools: list[Any]) -> list[dict]: """ Fixes https://github.com/BerriAI/litellm/issues/9321 Where user passes in a pydantic base model """ - new_tools: List[Dict] = [] + new_tools: list[dict] = [] for tool in tools: if isinstance(tool, BaseModel): tool = tool.model_dump(exclude_none=True) @@ -9378,9 +9368,9 @@ def get_empty_usage() -> Usage: def should_run_mock_completion( - mock_response: Optional[Any], - mock_tool_calls: Optional[Any], - mock_timeout: Optional[Any], + mock_response: Any | None, + mock_tool_calls: Any | None, + mock_timeout: Any | None, ) -> bool: if mock_response or mock_tool_calls or mock_timeout: return True diff --git a/litellm/vector_store_files/__init__.py b/litellm/vector_store_files/__init__.py index 07c66de3678..eadcece6994 100644 --- a/litellm/vector_store_files/__init__.py +++ b/litellm/vector_store_files/__init__.py @@ -14,16 +14,16 @@ from .main import ( ) __all__ = [ - "create", "acreate", - "list", - "alist", - "retrieve", - "aretrieve", - "retrieve_content", - "aretrieve_content", - "update", - "aupdate", - "delete", "adelete", + "alist", + "aretrieve", + "aretrieve_content", + "aupdate", + "create", + "delete", + "list", + "retrieve", + "retrieve_content", + "update", ] diff --git a/litellm/vector_store_files/main.py b/litellm/vector_store_files/main.py index 41990de8ec2..987480075ff 100644 --- a/litellm/vector_store_files/main.py +++ b/litellm/vector_store_files/main.py @@ -4,7 +4,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, Optional, Union +from typing import Any, Union import httpx @@ -29,17 +29,17 @@ from litellm.vector_store_files.utils import VectorStoreFileRequestUtils base_llm_http_handler = BaseLLMHTTPHandler() VectorStoreFileAttributeValue = Union[str, int, float, bool] -VectorStoreFileAttributes = Dict[str, VectorStoreFileAttributeValue] +VectorStoreFileAttributes = dict[str, VectorStoreFileAttributeValue] -def _ensure_provider(custom_llm_provider: Optional[str]) -> str: +def _ensure_provider(custom_llm_provider: str | None) -> str: return custom_llm_provider or "openai" def _prepare_registry_credentials( *, vector_store_id: str, - kwargs: Dict[str, Any], + kwargs: dict[str, Any], ) -> None: if litellm.vector_store_registry is None: return @@ -56,13 +56,13 @@ async def acreate( *, vector_store_id: str, file_id: str, - attributes: Optional[VectorStoreFileAttributes] = None, - chunking_strategy: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + attributes: VectorStoreFileAttributes | None = None, + chunking_strategy: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileObject: local_vars = locals() @@ -108,19 +108,19 @@ def create( *, vector_store_id: str, file_id: str, - attributes: Optional[VectorStoreFileAttributes] = None, - chunking_strategy: Optional[Dict[str, Any]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + attributes: VectorStoreFileAttributes | None = None, + chunking_strategy: dict[str, Any] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileObject, Coroutine[Any, Any, VectorStoreFileObject]]: +) -> VectorStoreFileObject | Coroutine[Any, Any, VectorStoreFileObject]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("acreate", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) @@ -182,15 +182,15 @@ def create( async def alist( *, vector_store_id: str, - after: Optional[str] = None, - before: Optional[str] = None, - filter: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + after: str | None = None, + before: str | None = None, + filter: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileListResponse: local_vars = locals() @@ -235,21 +235,21 @@ async def alist( def list( *, vector_store_id: str, - after: Optional[str] = None, - before: Optional[str] = None, - filter: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + after: str | None = None, + before: str | None = None, + filter: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileListResponse, Coroutine[Any, Any, VectorStoreFileListResponse]]: +) -> VectorStoreFileListResponse | Coroutine[Any, Any, VectorStoreFileListResponse]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("alist", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) @@ -308,9 +308,9 @@ async def aretrieve( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileObject: local_vars = locals() @@ -351,15 +351,15 @@ def retrieve( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileObject, Coroutine[Any, Any, VectorStoreFileObject]]: +) -> VectorStoreFileObject | Coroutine[Any, Any, VectorStoreFileObject]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("aretrieve", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) @@ -417,9 +417,9 @@ async def aretrieve_content( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileContentResponse: local_vars = locals() @@ -459,15 +459,15 @@ def retrieve_content( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileContentResponse, Coroutine[Any, Any, VectorStoreFileContentResponse]]: +) -> VectorStoreFileContentResponse | Coroutine[Any, Any, VectorStoreFileContentResponse]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("aretrieve_content", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) @@ -526,10 +526,10 @@ async def aupdate( vector_store_id: str, file_id: str, attributes: VectorStoreFileAttributes, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileObject: local_vars = locals() @@ -572,16 +572,16 @@ def update( vector_store_id: str, file_id: str, attributes: VectorStoreFileAttributes, - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileObject, Coroutine[Any, Any, VectorStoreFileObject]]: +) -> VectorStoreFileObject | Coroutine[Any, Any, VectorStoreFileObject]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("aupdate", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) @@ -646,9 +646,9 @@ async def adelete( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreFileDeleteResponse: local_vars = locals() @@ -688,15 +688,15 @@ def delete( *, vector_store_id: str, file_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreFileDeleteResponse, Coroutine[Any, Any, VectorStoreFileDeleteResponse]]: +) -> VectorStoreFileDeleteResponse | Coroutine[Any, Any, VectorStoreFileDeleteResponse]: local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id") + litellm_call_id: str | None = kwargs.get("litellm_call_id") _is_async = kwargs.pop("adelete", False) is True custom_llm_provider = _ensure_provider(custom_llm_provider) diff --git a/litellm/vector_store_files/utils.py b/litellm/vector_store_files/utils.py index 0f97af0066f..4b93d11a959 100644 --- a/litellm/vector_store_files/utils.py +++ b/litellm/vector_store_files/utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, cast, get_type_hints +from typing import Any, cast, get_type_hints from litellm.types.vector_store_files import ( VectorStoreFileCreateRequest, @@ -11,25 +11,25 @@ class VectorStoreFileRequestUtils: """Helper utilities for constructing vector store file requests.""" @staticmethod - def _filter_params(params: Dict[str, Any], model: Any) -> Dict[str, Any]: + def _filter_params(params: dict[str, Any], model: Any) -> dict[str, Any]: valid_keys = get_type_hints(model).keys() return {key: value for key, value in params.items() if key in valid_keys and value is not None} @staticmethod def get_create_request_params( - params: Dict[str, Any], + params: dict[str, Any], ) -> VectorStoreFileCreateRequest: filtered = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileCreateRequest) return cast(VectorStoreFileCreateRequest, filtered) @staticmethod - def get_list_query_params(params: Dict[str, Any]) -> VectorStoreFileListQueryParams: + def get_list_query_params(params: dict[str, Any]) -> VectorStoreFileListQueryParams: filtered = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileListQueryParams) return cast(VectorStoreFileListQueryParams, filtered) @staticmethod def get_update_request_params( - params: Dict[str, Any], + params: dict[str, Any], ) -> VectorStoreFileUpdateRequest: filtered = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileUpdateRequest) return cast(VectorStoreFileUpdateRequest, filtered) diff --git a/litellm/vector_stores/__init__.py b/litellm/vector_stores/__init__.py index 011c620f133..10311d33b17 100644 --- a/litellm/vector_stores/__init__.py +++ b/litellm/vector_stores/__init__.py @@ -1,4 +1,4 @@ from .main import acreate, asearch, create, search from .vector_store_registry import VectorStoreRegistry -__all__ = ["search", "asearch", "create", "acreate", "VectorStoreRegistry"] +__all__ = ["VectorStoreRegistry", "acreate", "asearch", "create", "search"] diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index e7d930122b5..4ae445d5d75 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -7,7 +7,7 @@ import builtins import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Dict, List, Optional, Union +from typing import Any import httpx @@ -36,7 +36,7 @@ base_llm_http_handler = BaseLLMHTTPHandler() def mock_vector_store_search_response( - mock_results: Optional[List[VectorStoreSearchResult]] = None, + mock_results: list[VectorStoreSearchResult] | None = None, ): """Mock response for vector store search""" if mock_results is None: @@ -60,7 +60,7 @@ def mock_vector_store_search_response( def mock_vector_store_create_response( - mock_response: Optional[VectorStoreCreateResponse] = None, + mock_response: VectorStoreCreateResponse | None = None, ): """Mock response for vector store create""" if mock_response is None: @@ -89,19 +89,19 @@ def mock_vector_store_create_response( @client async def acreate( - name: Optional[str] = None, - file_ids: Optional[List[str]] = None, - expires_after: Optional[Dict] = None, - chunking_strategy: Optional[Dict] = None, - metadata: Optional[Dict[str, str]] = None, + name: str | None = None, + file_ids: list[str] | None = None, + expires_after: dict | None = None, + chunking_strategy: dict | None = None, + metadata: dict[str, str] | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreCreateResponse: """ @@ -153,21 +153,21 @@ async def acreate( @client def create( - name: Optional[str] = None, - file_ids: Optional[List[str]] = None, - expires_after: Optional[Dict] = None, - chunking_strategy: Optional[Dict] = None, - metadata: Optional[Dict[str, str]] = None, + name: str | None = None, + file_ids: list[str] | None = None, + expires_after: dict | None = None, + chunking_strategy: dict | None = None, + metadata: dict[str, str] | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreCreateResponse, Coroutine[Any, Any, VectorStoreCreateResponse]]: +) -> VectorStoreCreateResponse | Coroutine[Any, Any, VectorStoreCreateResponse]: """ Create a vector store. @@ -184,7 +184,7 @@ def create( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("acreate", False) is True # get llm provider logic @@ -267,19 +267,19 @@ def create( @client async def asearch( vector_store_id: str, - query: Union[str, List[str]], - filters: Optional[Dict] = None, - max_num_results: Optional[int] = None, - ranking_options: Optional[Dict] = None, - rewrite_query: Optional[bool] = None, + query: str | list[str], + filters: dict | None = None, + max_num_results: int | None = None, + ranking_options: dict | None = None, + rewrite_query: bool | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreSearchResponse: """ @@ -334,21 +334,21 @@ async def asearch( @client def search( vector_store_id: str, - query: Union[str, List[str]], - filters: Optional[Dict] = None, - max_num_results: Optional[int] = None, - ranking_options: Optional[Dict] = None, - rewrite_query: Optional[bool] = None, + query: str | list[str], + filters: dict | None = None, + max_num_results: int | None = None, + ranking_options: dict | None = None, + rewrite_query: bool | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreSearchResponse, Coroutine[Any, Any, VectorStoreSearchResponse]]: +) -> VectorStoreSearchResponse | Coroutine[Any, Any, VectorStoreSearchResponse]: """ Search a vector store for relevant chunks based on a query and file attributes filter. @@ -366,7 +366,7 @@ def search( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("asearch", False) is True # pull credentials from registry if available @@ -466,11 +466,11 @@ def search( @client async def aretrieve( vector_store_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreCreateResponse: """ @@ -518,13 +518,13 @@ async def aretrieve( @client def retrieve( vector_store_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreCreateResponse, Coroutine[Any, Any, VectorStoreCreateResponse]]: +) -> VectorStoreCreateResponse | Coroutine[Any, Any, VectorStoreCreateResponse]: """ Retrieve a vector store. @@ -537,7 +537,7 @@ def retrieve( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aretrieve", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -597,15 +597,15 @@ def retrieve( @client async def alist( - after: Optional[str] = None, - before: Optional[str] = None, - limit: Optional[int] = 20, - order: Optional[str] = "desc", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + after: str | None = None, + before: str | None = None, + limit: int | None = 20, + order: str | None = "desc", + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -655,15 +655,15 @@ async def alist( @client def list( - after: Optional[str] = None, - before: Optional[str] = None, - limit: Optional[int] = 20, - order: Optional[str] = "desc", - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + after: str | None = None, + before: str | None = None, + limit: int | None = 20, + order: str | None = "desc", + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -681,7 +681,7 @@ def list( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("alist", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -750,14 +750,14 @@ def list( @client async def aupdate( vector_store_id: str, - name: Optional[str] = None, - expires_after: Optional[Dict] = None, - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + name: str | None = None, + expires_after: dict | None = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ) -> VectorStoreCreateResponse: """ @@ -808,16 +808,16 @@ async def aupdate( @client def update( vector_store_id: str, - name: Optional[str] = None, - expires_after: Optional[Dict] = None, - metadata: Optional[Dict[str, str]] = None, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + name: str | None = None, + expires_after: dict | None = None, + metadata: dict[str, str] | None = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, -) -> Union[VectorStoreCreateResponse, Coroutine[Any, Any, VectorStoreCreateResponse]]: +) -> VectorStoreCreateResponse | Coroutine[Any, Any, VectorStoreCreateResponse]: """ Update a vector store. @@ -833,7 +833,7 @@ def update( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aupdate", False) is True litellm_params = GenericLiteLLMParams(**kwargs) @@ -905,11 +905,11 @@ def update( @client async def adelete( vector_store_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -957,11 +957,11 @@ async def adelete( @client def delete( vector_store_id: str, - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - custom_llm_provider: Optional[str] = None, + extra_headers: dict[str, Any] | None = None, + extra_query: dict[str, Any] | None = None, + extra_body: dict[str, Any] | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, **kwargs, ): """ @@ -976,7 +976,7 @@ def delete( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("adelete", False) is True litellm_params = GenericLiteLLMParams(**kwargs) diff --git a/litellm/vector_stores/utils.py b/litellm/vector_stores/utils.py index 27b8d546af1..9c3793901a1 100644 --- a/litellm/vector_stores/utils.py +++ b/litellm/vector_stores/utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, cast, get_type_hints +from typing import Any, cast, get_type_hints from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.types.vector_stores import ( @@ -12,7 +12,7 @@ class VectorStoreRequestUtils: @staticmethod def get_requested_vector_store_search_optional_param( - params: Dict[str, Any], + params: dict[str, Any], vector_store_provider_config: BaseVectorStoreConfig, ) -> VectorStoreSearchOptionalRequestParams: """ @@ -37,7 +37,7 @@ class VectorStoreRequestUtils: @staticmethod def get_requested_vector_store_create_optional_param( - params: Dict[str, Any], + params: dict[str, Any], ) -> VectorStoreCreateOptionalRequestParams: """ Filter parameters to only include those defined in VectorStoreCreateOptionalRequestParams. diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index 5070db1c89e..4abd587bce5 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -1,7 +1,7 @@ # litellm/proxy/vector_stores/vector_store_registry.py import json from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional, get_args +from typing import TYPE_CHECKING, Any, get_args from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import remove_items_at_indices @@ -25,16 +25,16 @@ else: class VectorStoreIndexRegistry: - def __init__(self, vector_store_indexes: List[LiteLLM_ManagedVectorStoreIndex] = []): - self.vector_store_indexes: List[LiteLLM_ManagedVectorStoreIndex] = vector_store_indexes + def __init__(self, vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = []): + self.vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = vector_store_indexes - def get_vector_store_indexes(self) -> List[LiteLLM_ManagedVectorStoreIndex]: + def get_vector_store_indexes(self) -> list[LiteLLM_ManagedVectorStoreIndex]: """ Returns the vector store indexes """ return self.vector_store_indexes - def get_vector_store_index_by_name(self, vector_store_index_name: str) -> Optional[LiteLLM_ManagedVectorStoreIndex]: + def get_vector_store_index_by_name(self, vector_store_index_name: str) -> LiteLLM_ManagedVectorStoreIndex | None: """ Returns the vector store index by name """ @@ -76,12 +76,12 @@ class VectorStoreIndexRegistry: @staticmethod async def _get_vector_store_indexes_from_db( - prisma_client: Optional[PrismaClient], - ) -> List[LiteLLM_ManagedVectorStoreIndex]: + prisma_client: PrismaClient | None, + ) -> list[LiteLLM_ManagedVectorStoreIndex]: """ Get vector stores from the database """ - vector_stores_from_db: List[LiteLLM_ManagedVectorStoreIndex] = [] + vector_stores_from_db: list[LiteLLM_ManagedVectorStoreIndex] = [] if prisma_client is not None: _vector_stores_from_db = await ManagedVectorStoreIndexRepository(prisma_client).table.find_many( order={"created_at": "desc"}, @@ -94,11 +94,11 @@ class VectorStoreIndexRegistry: class VectorStoreRegistry: - def __init__(self, vector_stores: List[LiteLLM_ManagedVectorStore] = []): - self.vector_stores: List[LiteLLM_ManagedVectorStore] = vector_stores - self.vector_store_ids_to_vector_store_map: Dict[str, LiteLLM_ManagedVectorStore] = {} + def __init__(self, vector_stores: list[LiteLLM_ManagedVectorStore] = []): + self.vector_stores: list[LiteLLM_ManagedVectorStore] = vector_stores + self.vector_store_ids_to_vector_store_map: dict[str, LiteLLM_ManagedVectorStore] = {} - def _extract_tool_params(self, tool: Dict) -> VectorStoreToolParams: + def _extract_tool_params(self, tool: dict) -> VectorStoreToolParams: """ Extract supported parameters from a tool definition. @@ -112,13 +112,13 @@ class VectorStoreRegistry: return VectorStoreToolParams(**kwargs) - def get_vector_store_ids_to_run(self, non_default_params: Dict, tools: Optional[List[Dict]] = None) -> List[str]: + def get_vector_store_ids_to_run(self, non_default_params: dict, tools: list[dict] | None = None) -> list[str]: """ Returns the vector store ids to run vector_store_ids can be provided in two ways: """ - vector_store_ids: List[str] = [] + vector_store_ids: list[str] = [] # 1. check if vector_store_ids is provided in the non_default_params vector_store_ids_param = non_default_params.get("vector_store_ids") @@ -132,9 +132,9 @@ class VectorStoreRegistry: def get_and_pop_recognised_vector_store_tools( self, - tools: Optional[List[Dict]] = None, - vector_store_ids: Optional[List[str]] = None, - ) -> Dict[str, VectorStoreToolParams]: + tools: list[dict] | None = None, + vector_store_ids: list[str] | None = None, + ) -> dict[str, VectorStoreToolParams]: """ Returns and pops recognized vector store tools from the tools list. @@ -145,7 +145,7 @@ class VectorStoreRegistry: Returns: Dict mapping vector_store_id to its extracted tool parameters """ - params_by_id: Dict[str, VectorStoreToolParams] = {} + params_by_id: dict[str, VectorStoreToolParams] = {} if not tools: return params_by_id @@ -153,7 +153,7 @@ class VectorStoreRegistry: if vector_store_ids is None: vector_store_ids = [] - tools_to_remove: List[int] = [] + tools_to_remove: list[int] = [] for i, tool in enumerate(tools): tool_vector_store_ids = tool.get("vector_store_ids", []) @@ -180,8 +180,8 @@ class VectorStoreRegistry: return params_by_id def get_vector_store_to_run( - self, non_default_params: Dict, tools: Optional[List[Dict]] = None - ) -> Optional[LiteLLM_ManagedVectorStore]: + self, non_default_params: dict, tools: list[dict] | None = None + ) -> LiteLLM_ManagedVectorStore | None: """ Returns the vector store to run @@ -204,9 +204,7 @@ class VectorStoreRegistry: return vector_store return None - def get_litellm_managed_vector_store_from_registry( - self, vector_store_id: str - ) -> Optional[LiteLLM_ManagedVectorStore]: + def get_litellm_managed_vector_store_from_registry(self, vector_store_id: str) -> LiteLLM_ManagedVectorStore | None: """ Returns the vector store from the registry """ @@ -216,8 +214,8 @@ class VectorStoreRegistry: return None async def get_litellm_managed_vector_store_from_registry_or_db( - self, vector_store_id: str, prisma_client: Optional[PrismaClient] = None - ) -> Optional[LiteLLM_ManagedVectorStore]: + self, vector_store_id: str, prisma_client: PrismaClient | None = None + ) -> LiteLLM_ManagedVectorStore | None: """ Returns the vector store from the registry, falling back to database if not found. This ensures synchronization across multiple instances. @@ -237,13 +235,13 @@ class VectorStoreRegistry: self.add_vector_store_to_registry(vector_store=db_vector_store) return db_vector_store except Exception as e: - verbose_logger.debug(f"Error fetching vector store from database: {str(e)}") + verbose_logger.debug(f"Error fetching vector store from database: {e!s}") return None def get_litellm_managed_vector_store_from_registry_by_name( self, vector_store_name: str - ) -> Optional[LiteLLM_ManagedVectorStore]: + ) -> LiteLLM_ManagedVectorStore | None: """ Returns the vector store from the registry by name """ @@ -253,8 +251,8 @@ class VectorStoreRegistry: return None def pop_vector_stores_to_run( - self, non_default_params: Dict, tools: Optional[List[Dict]] = None - ) -> List[LiteLLM_ManagedVectorStore]: + self, non_default_params: dict, tools: list[dict] | None = None + ) -> list[LiteLLM_ManagedVectorStore]: """ Pops the vector stores to run with their tool parameters merged. @@ -268,12 +266,12 @@ class VectorStoreRegistry: List of vector stores with tool parameters merged into litellm_params """ # Pop vector_store_ids from params - vector_store_ids: List[str] = non_default_params.pop("vector_store_ids", None) or [] + vector_store_ids: list[str] = non_default_params.pop("vector_store_ids", None) or [] # Extract params from tools and collect IDs params_by_id = self.get_and_pop_recognised_vector_store_tools(tools=tools, vector_store_ids=vector_store_ids) - vector_stores_to_run: List[LiteLLM_ManagedVectorStore] = [] + vector_stores_to_run: list[LiteLLM_ManagedVectorStore] = [] for vector_store_id in vector_store_ids: for vector_store in self.vector_stores: @@ -296,10 +294,10 @@ class VectorStoreRegistry: async def pop_vector_stores_to_run_with_db_fallback( self, - non_default_params: Dict, - tools: Optional[List[Dict]] = None, - prisma_client: Optional[PrismaClient] = None, - ) -> List[LiteLLM_ManagedVectorStore]: + non_default_params: dict, + tools: list[dict] | None = None, + prisma_client: PrismaClient | None = None, + ) -> list[LiteLLM_ManagedVectorStore]: """ Pops the vector stores to run with their tool parameters merged. Falls back to database if vector stores are not found in memory. @@ -316,12 +314,12 @@ class VectorStoreRegistry: List of vector stores with tool parameters merged into litellm_params """ # Pop vector_store_ids from params - vector_store_ids: List[str] = non_default_params.pop("vector_store_ids", None) or [] + vector_store_ids: list[str] = non_default_params.pop("vector_store_ids", None) or [] # Extract params from tools and collect IDs params_by_id = self.get_and_pop_recognised_vector_store_tools(tools=tools, vector_store_ids=vector_store_ids) - vector_stores_to_run: List[LiteLLM_ManagedVectorStore] = [] + vector_stores_to_run: list[LiteLLM_ManagedVectorStore] = [] for vector_store_id in vector_store_ids: vector_store = None @@ -348,7 +346,7 @@ class VectorStoreRegistry: self.delete_vector_store_from_registry(vector_store_id=vector_store_id) vector_store = None except Exception as e: - verbose_logger.debug(f"Error verifying vector store {vector_store_id} in database: {str(e)}") + verbose_logger.debug(f"Error verifying vector store {vector_store_id} in database: {e!s}") # Fall back to database if not found in memory (or was deleted) if vector_store is None and prisma_client is not None: @@ -357,7 +355,7 @@ class VectorStoreRegistry: vector_store_id=vector_store_id, prisma_client=prisma_client ) except Exception as e: - verbose_logger.debug(f"Error fetching vector store {vector_store_id} from database: {str(e)}") + verbose_logger.debug(f"Error fetching vector store {vector_store_id} from database: {e!s}") if vector_store is not None: # Create a copy to avoid modifying the registry @@ -376,8 +374,8 @@ class VectorStoreRegistry: return vector_stores_to_run def _get_vector_store_ids_from_tool_calls( - self, tools: Optional[List[Dict]] = None, vector_store_ids: List[str] = [] - ) -> List[str]: + self, tools: list[dict] | None = None, vector_store_ids: list[str] = [] + ) -> list[str]: """ Returns the vector store ids from the tool calls """ @@ -387,7 +385,7 @@ class VectorStoreRegistry: vector_store_ids.extend(tool["vector_store_ids"]) return vector_store_ids - def load_vector_stores_from_config(self, vector_stores_config: List[Dict]): + def load_vector_stores_from_config(self, vector_stores_config: list[dict]): """ Loads vector stores from the litellm proxy config.yaml """ @@ -395,7 +393,7 @@ class VectorStoreRegistry: # cast to VectorStoreConfig litellm_vector_store_config = LiteLLM_VectorStoreConfig(**vector_store_config) vector_store_name = litellm_vector_store_config.get("vector_store_name") - vector_store_litellm_params: Dict[str, Any] = litellm_vector_store_config.get("litellm_params") or {} + vector_store_litellm_params: dict[str, Any] = litellm_vector_store_config.get("litellm_params") or {} vector_store_id = vector_store_litellm_params.get("vector_store_id") if vector_store_id is None: @@ -479,12 +477,12 @@ class VectorStoreRegistry: @staticmethod async def _get_vector_stores_from_db( - prisma_client: Optional[PrismaClient], - ) -> List[LiteLLM_ManagedVectorStore]: + prisma_client: PrismaClient | None, + ) -> list[LiteLLM_ManagedVectorStore]: """ Get vector stores from the database """ - vector_stores_from_db: List[LiteLLM_ManagedVectorStore] = [] + vector_stores_from_db: list[LiteLLM_ManagedVectorStore] = [] if prisma_client is not None: _vector_stores_from_db = await ManagedVectorStoresRepository(prisma_client).table.find_many( order={"created_at": "desc"}, @@ -495,7 +493,7 @@ class VectorStoreRegistry: vector_stores_from_db.append(_litellm_managed_vector_store) return vector_stores_from_db - def get_credentials_for_vector_store(self, vector_store_id: str) -> Dict[str, Any]: + def get_credentials_for_vector_store(self, vector_store_id: str) -> dict[str, Any]: """ Get the credentials for a vector store diff --git a/litellm/videos/__init__.py b/litellm/videos/__init__.py index 9fb66d7557a..787eb7419cb 100644 --- a/litellm/videos/__init__.py +++ b/litellm/videos/__init__.py @@ -22,22 +22,22 @@ from .main import ( ) __all__ = [ - "avideo_generation", - "video_generation", - "avideo_list", - "video_list", - "avideo_status", - "video_status", "avideo_content", - "video_content", - "avideo_remix", - "video_remix", "avideo_create_character", - "video_create_character", - "avideo_get_character", - "video_get_character", "avideo_edit", - "video_edit", "avideo_extension", + "avideo_generation", + "avideo_get_character", + "avideo_list", + "avideo_remix", + "avideo_status", + "video_content", + "video_create_character", + "video_edit", "video_extension", + "video_generation", + "video_get_character", + "video_list", + "video_remix", + "video_status", ] diff --git a/litellm/videos/main.py b/litellm/videos/main.py index 1d17b566b5c..f1eaa09546a 100644 --- a/litellm/videos/main.py +++ b/litellm/videos/main.py @@ -3,7 +3,7 @@ import contextvars import json from collections.abc import Coroutine from functools import partial -from typing import Dict, List, Literal, Optional, Union, overload +from typing import Literal, overload import litellm from litellm.constants import DEFAULT_VIDEO_ENDPOINT_MODEL @@ -32,18 +32,18 @@ llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() @client async def avideo_generation( prompt: str, - model: Optional[str] = None, - input_reference: Optional[FileTypes] = None, - seconds: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + input_reference: FileTypes | None = None, + seconds: str | None = None, + size: str | None = None, + user: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> VideoObject: """ @@ -120,16 +120,16 @@ async def avideo_generation( @overload def video_generation( prompt: str, - model: Optional[str] = None, - input_reference: Optional[FileTypes] = None, - seconds: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + input_reference: FileTypes | None = None, + seconds: str | None = None, + size: str | None = None, + user: str | None = None, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_generation: Literal[True], **kwargs: object, @@ -140,16 +140,16 @@ def video_generation( @overload def video_generation( prompt: str, - model: Optional[str] = None, - input_reference: Optional[FileTypes] = None, - seconds: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + input_reference: FileTypes | None = None, + seconds: str | None = None, + size: str | None = None, + user: str | None = None, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_generation: Literal[False] = False, **kwargs: object, @@ -162,23 +162,20 @@ def video_generation( @client def video_generation( prompt: str, - model: Optional[str] = None, - input_reference: Optional[FileTypes] = None, - seconds: Optional[str] = None, - size: Optional[str] = None, - user: Optional[str] = None, + model: str | None = None, + input_reference: FileTypes | None = None, + seconds: str | None = None, + size: str | None = None, + user: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[ - VideoObject, - Coroutine[object, object, VideoObject], -]: +) -> VideoObject | Coroutine[object, object, VideoObject]: """ Maps the https://api.openai.com/v1/videos endpoint. @@ -187,7 +184,7 @@ def video_generation( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -207,7 +204,7 @@ def video_generation( ) # get provider config - video_generation_provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + video_generation_provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=model, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -222,7 +219,7 @@ def video_generation( ) # Get optional parameters for the video generation API - video_generation_request_params: Dict = VideoGenerationRequestUtils.get_optional_params_video_generation( + video_generation_request_params: dict = VideoGenerationRequestUtils.get_optional_params_video_generation( model=model, video_generation_provider_config=video_generation_provider_config, video_generation_optional_params=video_generation_optional_params, @@ -273,19 +270,16 @@ def video_generation( @client def video_content( video_id: str, - timeout: Optional[float] = None, - custom_llm_provider: Optional[str] = None, - variant: Optional[str] = None, + timeout: float | None = None, + custom_llm_provider: str | None = None, + variant: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[ - bytes, - Coroutine[object, object, bytes], -]: +) -> bytes | Coroutine[object, object, bytes]: """ Download video content from OpenAI's video API. @@ -318,7 +312,7 @@ def video_content( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True # Try to decode provider from video_id if not explicitly provided @@ -330,7 +324,7 @@ def video_content( litellm_params = GenericLiteLLMParams(**kwargs) # get provider config - video_provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + video_provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -341,7 +335,7 @@ def video_content( local_vars.update(kwargs) # For video content download, we don't need complex optional parameter handling # Just pass the basic parameters that are relevant for content download - video_content_request_params: Dict = { + video_content_request_params: dict = { "video_id": video_id, } @@ -386,14 +380,14 @@ def video_content( @client async def avideo_content( video_id: str, - timeout: Optional[float] = None, - custom_llm_provider: Optional[str] = None, - variant: Optional[str] = None, + timeout: float | None = None, + custom_llm_provider: str | None = None, + variant: str | None = None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> bytes: """ @@ -462,9 +456,9 @@ async def avideo_remix( custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> VideoObject: """ @@ -528,10 +522,10 @@ def video_remix( video_id: str, prompt: str, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_remix: Literal[True], **kwargs: object, @@ -544,10 +538,10 @@ def video_remix( video_id: str, prompt: str, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_remix: Literal[False] = False, **kwargs: object, @@ -565,14 +559,11 @@ def video_remix( custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[ - VideoObject, - Coroutine[object, object, VideoObject], -]: +) -> VideoObject | Coroutine[object, object, VideoObject]: """ Maps the https://api.openai.com/v1/videos/{video_id}/remix endpoint. @@ -581,7 +572,7 @@ def video_remix( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -602,7 +593,7 @@ def video_remix( litellm_params = GenericLiteLLMParams(**kwargs) # get provider config - video_remix_provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + video_remix_provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -612,7 +603,7 @@ def video_remix( local_vars.update(kwargs) # For video remix, we need the video_id and prompt - video_remix_request_params: Dict = { + video_remix_request_params: dict = { "video_id": video_id, "prompt": prompt, } @@ -661,19 +652,19 @@ def video_remix( ##### Video List ####################### @client async def avideo_list( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, - api_key: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, + api_key: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> List[VideoObject]: +) -> list[VideoObject]: """ Asynchronously calls the `video_list` function with the given arguments and keyword arguments. @@ -740,35 +731,35 @@ async def avideo_list( # Overload for when avideo_list=True (returns Coroutine) @overload def video_list( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_list: Literal[True], **kwargs: object, -) -> Coroutine[object, object, List[VideoObject]]: +) -> Coroutine[object, object, list[VideoObject]]: ... @overload def video_list( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_list: Literal[False] = False, **kwargs: object, -) -> List[VideoObject]: +) -> list[VideoObject]: ... # fmt: on @@ -776,21 +767,18 @@ def video_list( @client def video_list( - after: Optional[str] = None, - limit: Optional[int] = None, - order: Optional[str] = None, + after: str | None = None, + limit: int | None = None, + order: str | None = None, timeout=600, # default to 10 minutes custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[ - List[VideoObject], - Coroutine[object, object, List[VideoObject]], -]: +) -> list[VideoObject] | Coroutine[object, object, list[VideoObject]]: """ Maps the https://api.openai.com/v1/videos endpoint. @@ -799,7 +787,7 @@ def video_list( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -817,7 +805,7 @@ def video_list( litellm_params = GenericLiteLLMParams(**kwargs) # get provider config - video_list_provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + video_list_provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -827,7 +815,7 @@ def video_list( local_vars.update(kwargs) # For video list, we need the query parameters - video_list_request_params: Dict = { + video_list_request_params: dict = { "after": after, "limit": limit, "order": order, @@ -883,9 +871,9 @@ async def avideo_status( custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> VideoObject: """ @@ -947,10 +935,10 @@ async def avideo_status( def video_status( video_id: str, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_status: Literal[True], **kwargs: object, @@ -962,10 +950,10 @@ def video_status( def video_status( video_id: str, timeout: int = 600, - custom_llm_provider: Optional[str] = None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, *, avideo_status: Literal[False] = False, **kwargs: object, @@ -982,14 +970,11 @@ def video_status( custom_llm_provider=None, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[ - VideoObject, - Coroutine[object, object, VideoObject], -]: +) -> VideoObject | Coroutine[object, object, VideoObject]: """ Retrieve video status from OpenAI's video API. @@ -1020,7 +1005,7 @@ def video_status( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True # Check for mock response first @@ -1041,7 +1026,7 @@ def video_status( litellm_params = GenericLiteLLMParams(**kwargs) # get provider config - video_status_provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + video_status_provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1051,7 +1036,7 @@ def video_status( local_vars.update(kwargs) # For video status, we need the video_id - video_status_request_params: Dict = { + video_status_request_params: dict = { "video_id": video_id, } @@ -1101,9 +1086,9 @@ async def avideo_create_character( video: FileTypes, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> CharacterObject: """ @@ -1156,11 +1141,11 @@ def video_create_character( video: FileTypes, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[CharacterObject, Coroutine[object, object, CharacterObject]]: +) -> CharacterObject | Coroutine[object, object, CharacterObject]: """ Create a character from an uploaded video file. Maps to POST /v1/videos/characters @@ -1168,7 +1153,7 @@ def video_create_character( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True mock_response = kwargs.get("mock_response", None) @@ -1182,7 +1167,7 @@ def video_create_character( litellm_params = GenericLiteLLMParams(**kwargs) - provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1191,7 +1176,7 @@ def video_create_character( raise ValueError(f"video create character is not supported for {custom_llm_provider}") local_vars.update(kwargs) - request_params: Dict = {"name": name} + request_params: dict = {"name": name} litellm_logging_obj.update_environment_variables( model="", @@ -1231,9 +1216,9 @@ async def avideo_get_character( character_id: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> CharacterObject: """ @@ -1281,11 +1266,11 @@ def video_get_character( character_id: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[CharacterObject, Coroutine[object, object, CharacterObject]]: +) -> CharacterObject | Coroutine[object, object, CharacterObject]: """ Retrieve a character by ID. Maps to GET /v1/videos/characters/{character_id} @@ -1293,7 +1278,7 @@ def video_get_character( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True mock_response = kwargs.get("mock_response", None) @@ -1307,7 +1292,7 @@ def video_get_character( litellm_params = GenericLiteLLMParams(**kwargs) - provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1316,7 +1301,7 @@ def video_get_character( raise ValueError(f"video get character is not supported for {custom_llm_provider}") local_vars.update(kwargs) - request_params: Dict = {"character_id": character_id} + request_params: dict = {"character_id": character_id} litellm_logging_obj.update_environment_variables( model="", @@ -1356,9 +1341,9 @@ async def avideo_edit( prompt: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> VideoObject: """ @@ -1408,11 +1393,11 @@ def video_edit( prompt: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[VideoObject, Coroutine[object, object, VideoObject]]: +) -> VideoObject | Coroutine[object, object, VideoObject]: """ Create a video edit job. Maps to POST /v1/videos/edits @@ -1420,7 +1405,7 @@ def video_edit( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True mock_response = kwargs.get("mock_response", None) @@ -1435,7 +1420,7 @@ def video_edit( litellm_params = GenericLiteLLMParams(**kwargs) - provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1444,7 +1429,7 @@ def video_edit( raise ValueError(f"video edit is not supported for {custom_llm_provider}") local_vars.update(kwargs) - request_params: Dict = {"video_id": video_id, "prompt": prompt} + request_params: dict = {"video_id": video_id, "prompt": prompt} litellm_logging_obj.update_environment_variables( model="", @@ -1487,9 +1472,9 @@ async def avideo_extension( seconds: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, ) -> VideoObject: """ @@ -1541,11 +1526,11 @@ def video_extension( seconds: str, timeout=600, custom_llm_provider=None, - extra_headers: Optional[Dict[str, object]] = None, - extra_query: Optional[Dict[str, object]] = None, - extra_body: Optional[Dict[str, object]] = None, + extra_headers: dict[str, object] | None = None, + extra_query: dict[str, object] | None = None, + extra_body: dict[str, object] | None = None, **kwargs, -) -> Union[VideoObject, Coroutine[object, object, VideoObject]]: +) -> VideoObject | Coroutine[object, object, VideoObject]: """ Create a video extension. Maps to POST /v1/videos/extensions @@ -1553,7 +1538,7 @@ def video_extension( local_vars = locals() try: litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True mock_response = kwargs.get("mock_response", None) @@ -1568,7 +1553,7 @@ def video_extension( litellm_params = GenericLiteLLMParams(**kwargs) - provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config( + provider_config: BaseVideoConfig | None = ProviderConfigManager.get_provider_video_config( model=None, provider=litellm.LlmProviders(custom_llm_provider), ) @@ -1577,7 +1562,7 @@ def video_extension( raise ValueError(f"video extension is not supported for {custom_llm_provider}") local_vars.update(kwargs) - request_params: Dict = { + request_params: dict = { "video_id": video_id, "prompt": prompt, "seconds": seconds, diff --git a/litellm/videos/utils.py b/litellm/videos/utils.py index 42e0d7f4a27..766057b7191 100644 --- a/litellm/videos/utils.py +++ b/litellm/videos/utils.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, cast +from typing import Any, cast import litellm from litellm.llms.base_llm.videos.transformation import BaseVideoConfig @@ -14,7 +14,7 @@ class VideoGenerationRequestUtils: model: str, video_generation_provider_config: BaseVideoConfig, video_generation_optional_params: VideoCreateOptionalRequestParams, - ) -> Dict: + ) -> dict: """ Get optional parameters for the video generation API. @@ -46,7 +46,7 @@ class VideoGenerationRequestUtils: @staticmethod def get_requested_video_generation_optional_param( - params: Dict[str, Any], + params: dict[str, Any], ) -> VideoCreateOptionalRequestParams: """ Filter parameters to only include those defined in VideoCreateOptionalRequestParams. @@ -75,12 +75,12 @@ class VideoGenerationRequestUtils: cleaned_kwargs = filter_out_litellm_params(kwargs={k: v for k, v in raw_kwargs.items() if v is not None}) - optional_params: Dict[str, Any] = { + optional_params: dict[str, Any] = { **base_params, **cleaned_kwargs, } - merged_extra_body: Dict[str, Any] = {} + merged_extra_body: dict[str, Any] = {} for extra_body_candidate in (top_level_extra_body, kwargs_extra_body): if isinstance(extra_body_candidate, dict): for key, value in extra_body_candidate.items(): diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 537e019af1e..d3ef01940bb 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -57,7 +57,7 @@ "limit": 6 }, "B033": { - "limit": 4 + "limit": 0 }, "BLE001": { "limit": 2899 @@ -81,7 +81,7 @@ "limit": 4 }, "C901": { - "limit": 314 + "limit": 312 }, "D419": { "limit": 9 @@ -114,16 +114,16 @@ "limit": 23 }, "FURB136": { - "limit": 4 + "limit": 0 }, "FURB168": { - "limit": 4 + "limit": 0 }, "FURB188": { - "limit": 52 + "limit": 0 }, "I001": { - "limit": 196 + "limit": 0 }, "LOG015": { "limit": 8 @@ -144,10 +144,10 @@ "limit": 74 }, "PIE790": { - "limit": 278 + "limit": 0 }, "PIE800": { - "limit": 4 + "limit": 0 }, "PIE804": { "limit": 24 @@ -159,7 +159,7 @@ "limit": 31 }, "PLC0208": { - "limit": 4 + "limit": 0 }, "PLC0414": { "limit": 38 @@ -171,22 +171,22 @@ "limit": 4 }, "PLR0402": { - "limit": 9 + "limit": 0 }, "PLR1704": { "limit": 6 }, "PLR1711": { - "limit": 34 + "limit": 0 }, "PLR1714": { "limit": 261 }, "PLR1730": { - "limit": 10 + "limit": 0 }, "PLR2044": { - "limit": 4 + "limit": 0 }, "PLW0127": { "limit": 43 @@ -207,25 +207,25 @@ "limit": 5 }, "PYI030": { - "limit": 5 + "limit": 0 }, "PYI036": { "limit": 5 }, "PYI041": { - "limit": 12 + "limit": 0 }, "PYI064": { - "limit": 5 + "limit": 0 }, "RET501": { - "limit": 38 + "limit": 0 }, "RET504": { "limit": 702 }, "RUF010": { - "limit": 874 + "limit": 0 }, "RUF012": { "limit": 168 @@ -237,16 +237,16 @@ "limit": 41 }, "RUF022": { - "limit": 85 + "limit": 0 }, "RUF023": { - "limit": 5 + "limit": 0 }, "RUF046": { - "limit": 8 + "limit": 6 }, "RUF051": { - "limit": 6 + "limit": 0 }, "RUF059": { "limit": 73 @@ -273,7 +273,7 @@ "limit": 6 }, "SIM114": { - "limit": 111 + "limit": 0 }, "SIM115": { "limit": 5 @@ -282,7 +282,7 @@ "limit": 10 }, "SIM118": { - "limit": 114 + "limit": 0 }, "SIM201": { "limit": 4 @@ -303,10 +303,10 @@ "limit": 8 }, "TC005": { - "limit": 9 + "limit": 0 }, "TID251": { - "limit": 2648 + "limit": 0 }, "TRY002": { "limit": 547 @@ -324,22 +324,22 @@ "limit": 879 }, "UP006": { - "limit": 12050 + "limit": 0 }, "UP007": { - "limit": 2526 + "limit": 0 }, "UP008": { - "limit": 5 + "limit": 0 }, "UP012": { - "limit": 7 + "limit": 0 }, "UP018": { - "limit": 21 + "limit": 0 }, "UP024": { - "limit": 15 + "limit": 0 }, "UP028": { "limit": 5 @@ -348,21 +348,21 @@ "limit": 5 }, "UP032": { - "limit": 626 + "limit": 0 }, "UP034": { - "limit": 4 + "limit": 0 }, "UP035": { - "limit": 1909 + "limit": 0 }, "UP036": { "limit": 4 }, "UP037": { - "limit": 104 + "limit": 0 }, "UP045": { - "limit": 17793 + "limit": 0 } } diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 5fbc3c4869b..0af5ad6cd9b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -1,4 +1,3 @@ -from typing import List from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -307,7 +306,7 @@ EXPECTED_RESPONSE_MODELS = { "/customer/update": CustomerResponse, "/customer/delete": DeleteCustomersResponse, "/customer/info": CustomerResponse, - "/customer/list": List[CustomerResponse], + "/customer/list": list[CustomerResponse], "/customer/daily/activity": SpendAnalyticsPaginatedResponse, } diff --git a/type-discipline-budget.json b/type-discipline-budget.json index b790d9acb0f..f071c381916 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23191 + "limit": 23350 }, "LIT002": { - "limit": 27275 + "limit": 27256 }, "LIT003": { "limit": 292 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1106 + "limit": 1105 }, "LIT007": { "limit": 0 @@ -24,6 +24,6 @@ "limit": 1004 }, "LIT009": { - "limit": 2467 + "limit": 2465 } } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 109d638fb9c..90552b6eae1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -23578,7 +23578,7 @@ export interface components { * @description Default role assigned to new users created * @default internal_user_viewer */ - user_role: ("internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer") | null; + user_role: ("proxy_admin" | "proxy_admin_viewer" | "internal_user" | "internal_user_viewer") | null; }; /** * DefaultTeamSSOParams @@ -26172,7 +26172,7 @@ export interface components { [key: string]: unknown; } | null; /** Stream Timeout */ - stream_timeout?: number | string | null; + stream_timeout?: string | number | null; /** Tag Regex */ tag_regex?: string[] | null; /** Tags */ @@ -34304,7 +34304,7 @@ export interface components { [key: string]: unknown; } | null; /** Stream Timeout */ - stream_timeout?: number | string | null; + stream_timeout?: string | number | null; /** Tag Regex */ tag_regex?: string[] | null; /** Tags */