From deaade232d352379e4cca16f93417524b37efd42 Mon Sep 17 00:00:00 2001 From: mateo Date: Sun, 2 Aug 2026 00:57:44 +0000 Subject: [PATCH 01/11] feat(gemini): add gemini-robotics-er-2-preview and gemini-robotics-er-1.6-preview pricing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 93 +++++++++++++++++++ model_prices_and_context_window.json | 93 +++++++++++++++++++ 2 files changed, 186 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 07f04136bd0..158f7e3b8ba 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18883,6 +18883,99 @@ "search_context_size_high": 0.035 } }, + "gemini/gemini-robotics-er-2-preview": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_audio_token": 2e-06, + "input_cost_per_token": 2e-06, + "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 1e-05, + "output_cost_per_token": 1e-05, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-robotics-er-2", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini/gemini-robotics-er-1.6-preview": { + "input_cost_per_audio_token": 2e-06, + "input_cost_per_token": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 5e-06, + "output_cost_per_token": 5e-06, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-robotics-er", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, "gemini-2.5-computer-use-preview-10-2025": { "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 346f613ea3e..5e5e741705a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -18961,6 +18961,99 @@ "search_context_size_high": 0.035 } }, + "gemini/gemini-robotics-er-2-preview": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_audio_token": 2e-06, + "input_cost_per_token": 2e-06, + "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 1e-05, + "output_cost_per_token": 1e-05, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-robotics-er-2", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini/gemini-robotics-er-1.6-preview": { + "input_cost_per_audio_token": 2e-06, + "input_cost_per_token": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 5e-06, + "output_cost_per_token": 5e-06, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-robotics-er", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, "gemini-2.5-computer-use-preview-10-2025": { "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, From 2708620d6a599cc73c1950a942d26ac26a7ed3d4 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 4 Aug 2026 12:54:39 -0700 Subject: [PATCH 02/11] feat(lint): enforce Final on locals and freeze function parameters (LIT010, LIT011) --- litellm/__init__.py | 641 ++--- litellm/_internal_context.py | 3 +- litellm/_lazy_imports.py | 28 +- litellm/_lazy_imports_registry.py | 54 +- litellm/_logging.py | 58 +- litellm/_redis.py | 141 +- litellm/_redis_credential_provider.py | 24 +- litellm/_service_logger.py | 14 +- litellm/a2a_protocol/card_resolver.py | 22 +- litellm/a2a_protocol/client.py | 6 +- litellm/a2a_protocol/cost_calculator.py | 22 +- .../a2a_protocol/exception_mapping_utils.py | 16 +- .../litellm_completion_bridge/handler.py | 46 +- .../transformation.py | 30 +- litellm/a2a_protocol/main.py | 103 +- .../providers/bedrock_agentcore/config.py | 6 +- .../providers/bedrock_agentcore/handler.py | 18 +- .../bedrock_agentcore/transformation.py | 28 +- .../providers/pydantic_ai_agents/handler.py | 6 +- .../pydantic_ai_agents/transformation.py | 48 +- .../providers/watsonx_orchestrate/config.py | 6 +- .../providers/watsonx_orchestrate/handler.py | 82 +- .../watsonx_orchestrate/transformation.py | 24 +- litellm/a2a_protocol/streaming_iterator.py | 38 +- litellm/a2a_protocol/utils.py | 20 +- litellm/anthropic_beta_headers_manager.py | 43 +- .../exceptions/exception_mapping_utils.py | 10 +- litellm/assistants/main.py | 188 +- litellm/assistants/utils.py | 20 +- litellm/batch_completion/main.py | 19 +- litellm/batches/batch_utils.py | 48 +- litellm/batches/main.py | 166 +- litellm/budget_manager.py | 26 +- litellm/caching/_embedding_router.py | 6 +- litellm/caching/_internal_lru_cache.py | 4 +- litellm/caching/azure_blob_cache.py | 21 +- litellm/caching/base_cache.py | 4 +- litellm/caching/caching.py | 110 +- litellm/caching/caching_handler.py | 86 +- litellm/caching/disk_cache.py | 14 +- litellm/caching/dual_cache.py | 28 +- litellm/caching/evicted_client_closer.py | 33 +- litellm/caching/gcs_cache.py | 39 +- litellm/caching/in_memory_cache.py | 22 +- litellm/caching/llm_caching_handler.py | 7 +- litellm/caching/qdrant_semantic_cache.py | 80 +- litellm/caching/redis_cache.py | 228 +- litellm/caching/redis_cluster_cache.py | 16 +- litellm/caching/redis_semantic_cache.py | 84 +- litellm/caching/s3_cache.py | 23 +- litellm/caching/valkey_semantic_cache.py | 48 +- .../handler.py | 98 +- .../transformation.py | 130 +- litellm/compression/compress.py | 72 +- litellm/compression/content_detection.py | 13 +- litellm/compression/message_stubbing.py | 25 +- litellm/compression/scoring/bm25.py | 23 +- .../compression/scoring/embedding_scorer.py | 26 +- litellm/constants.py | 726 +++--- litellm/containers/endpoint_factory.py | 66 +- litellm/containers/main.py | 148 +- litellm/containers/utils.py | 28 +- litellm/cost_calculator.py | 204 +- .../speech_to_completion_bridge/handler.py | 22 +- .../transformation.py | 24 +- litellm/evals/main.py | 304 +-- litellm/exceptions.py | 36 +- litellm/experimental_mcp_client/client.py | 113 +- litellm/experimental_mcp_client/tools.py | 16 +- litellm/files/main.py | 164 +- litellm/files/streaming.py | 33 +- litellm/files/utils.py | 4 +- litellm/fine_tuning/main.py | 94 +- litellm/google_genai/adapters/handler.py | 20 +- .../google_genai/adapters/transformation.py | 62 +- litellm/google_genai/main.py | 56 +- litellm/google_genai/streaming_iterator.py | 14 +- litellm/images/main.py | 153 +- litellm/images/utils.py | 22 +- .../SlackAlerting/batching_handler.py | 8 +- .../SlackAlerting/budget_alert_types.py | 4 +- .../SlackAlerting/hanging_request_check.py | 12 +- .../SlackAlerting/slack_alerting.py | 208 +- litellm/integrations/SlackAlerting/utils.py | 8 +- litellm/integrations/agentops/agentops.py | 16 +- .../anthropic_cache_control_hook.py | 46 +- litellm/integrations/argilla.py | 68 +- litellm/integrations/arize/__init__.py | 14 +- litellm/integrations/arize/_utils.py | 140 +- litellm/integrations/arize/arize.py | 20 +- litellm/integrations/arize/arize_phoenix.py | 70 +- .../arize/arize_phoenix_client.py | 12 +- .../arize/arize_phoenix_prompt_manager.py | 30 +- litellm/integrations/athina.py | 13 +- .../azure_sentinel/azure_sentinel.py | 35 +- .../azure_storage/azure_storage.py | 47 +- litellm/integrations/bitbucket/__init__.py | 12 +- .../bitbucket/bitbucket_client.py | 42 +- .../bitbucket/bitbucket_prompt_manager.py | 30 +- litellm/integrations/braintrust_logging.py | 77 +- .../integrations/braintrust_mock_client.py | 11 +- litellm/integrations/cloudzero/cloudzero.py | 54 +- .../cloudzero/cz_resource_names.py | 26 +- .../integrations/cloudzero/cz_stream_api.py | 24 +- litellm/integrations/cloudzero/database.py | 8 +- litellm/integrations/cloudzero/transform.py | 48 +- .../code_interpreter_interception/handler.py | 130 +- .../compression_interception/handler.py | 64 +- litellm/integrations/custom_batch_logger.py | 5 +- litellm/integrations/custom_guardrail.py | 135 +- litellm/integrations/custom_logger.py | 67 +- litellm/integrations/custom_sso_handler.py | 4 +- litellm/integrations/datadog/datadog.py | 114 +- .../datadog/datadog_cost_management.py | 26 +- .../integrations/datadog/datadog_handler.py | 11 +- .../integrations/datadog/datadog_llm_obs.py | 100 +- .../integrations/datadog/datadog_metrics.py | 61 +- .../datadog/datadog_mock_client.py | 4 +- .../datadog/datadog_team_handler.py | 10 +- litellm/integrations/deepeval/api.py | 21 +- litellm/integrations/deepeval/deepeval.py | 23 +- litellm/integrations/deepeval/utils.py | 3 +- litellm/integrations/dotprompt/__init__.py | 20 +- .../dotprompt/dotprompt_manager.py | 28 +- .../integrations/dotprompt/prompt_manager.py | 46 +- litellm/integrations/dynamodb.py | 22 +- litellm/integrations/email_alerting.py | 31 +- .../email_templates/email_footer.py | 4 +- .../email_templates/key_created_email.py | 4 +- .../email_templates/key_rotated_email.py | 4 +- .../integrations/email_templates/templates.py | 12 +- .../email_templates/user_invitation_email.py | 4 +- litellm/integrations/focus/database.py | 20 +- .../focus/destinations/factory.py | 4 +- .../focus/destinations/gcs_destination.py | 24 +- .../focus/destinations/mavvrik_destination.py | 50 +- .../focus/destinations/s3_destination.py | 24 +- .../focus/destinations/vantage_destination.py | 44 +- litellm/integrations/focus/export_engine.py | 34 +- litellm/integrations/focus/focus_logger.py | 40 +- litellm/integrations/focus/schema.py | 4 +- litellm/integrations/focus/serializers/csv.py | 5 +- .../integrations/focus/serializers/parquet.py | 5 +- litellm/integrations/focus/transformer.py | 13 +- litellm/integrations/galileo.py | 94 +- litellm/integrations/gcs_bucket/gcs_bucket.py | 56 +- .../gcs_bucket/gcs_bucket_base.py | 42 +- .../gcs_bucket/gcs_bucket_mock_client.py | 7 +- litellm/integrations/gcs_pubsub/pub_sub.py | 18 +- .../generic_api/generic_api_callback.py | 62 +- .../generic_prompt_management/__init__.py | 16 +- .../generic_prompt_manager.py | 38 +- litellm/integrations/gitlab/__init__.py | 16 +- litellm/integrations/gitlab/gitlab_client.py | 68 +- .../gitlab/gitlab_prompt_manager.py | 68 +- litellm/integrations/greenscale.py | 11 +- litellm/integrations/helicone.py | 33 +- litellm/integrations/helicone_mock_client.py | 4 +- litellm/integrations/humanloop.py | 28 +- litellm/integrations/lago.py | 36 +- litellm/integrations/langfuse/langfuse.py | 132 +- .../integrations/langfuse/langfuse_handler.py | 8 +- .../langfuse/langfuse_mock_client.py | 4 +- .../integrations/langfuse/langfuse_otel.py | 66 +- .../langfuse/langfuse_otel_attributes.py | 14 +- .../langfuse/langfuse_prompt_management.py | 40 +- litellm/integrations/langsmith.py | 98 +- litellm/integrations/langsmith_mock_client.py | 4 +- litellm/integrations/langtrace.py | 20 +- litellm/integrations/levo/levo.py | 20 +- .../litellm_agent_model_resolver.py | 6 +- litellm/integrations/literal_ai.py | 49 +- litellm/integrations/logfire_logger.py | 22 +- litellm/integrations/lunary.py | 15 +- .../mavvrik_focus/mavvrik_focus_logger.py | 40 +- litellm/integrations/mlflow.py | 42 +- litellm/integrations/mock_client_factory.py | 14 +- litellm/integrations/newrelic/newrelic.py | 92 +- litellm/integrations/openmeter.py | 33 +- litellm/integrations/opentelemetry.py | 457 ++-- .../base_otel_llm_obs_attributes.py | 4 +- .../opentelemetry_utils/gen_ai_semconv.py | 30 +- litellm/integrations/opik/opik.py | 14 +- .../opik/opik_payload_builder/api.py | 32 +- .../opik/opik_payload_builder/extractors.py | 20 +- .../opik_payload_builder/payload_builders.py | 14 +- .../opik/opik_payload_builder/types.py | 4 +- litellm/integrations/opik/utils.py | 22 +- litellm/integrations/otel/emitter.py | 19 +- litellm/integrations/otel/logger.py | 82 +- litellm/integrations/otel/mappers/__init__.py | 9 +- litellm/integrations/otel/mappers/base.py | 3 +- litellm/integrations/otel/mappers/genai.py | 7 +- litellm/integrations/otel/mappers/legacy.py | 4 +- .../otel/mappers/openinference.py | 5 +- litellm/integrations/otel/mappers/utils.py | 6 +- litellm/integrations/otel/model/baggage.py | 4 +- litellm/integrations/otel/model/config.py | 4 +- litellm/integrations/otel/model/metadata.py | 38 +- litellm/integrations/otel/model/payloads.py | 54 +- litellm/integrations/otel/model/semconv.py | 8 +- litellm/integrations/otel/model/spans.py | 16 +- litellm/integrations/otel/mount.py | 16 +- litellm/integrations/otel/plumbing/context.py | 31 +- litellm/integrations/otel/plumbing/events.py | 3 +- litellm/integrations/otel/plumbing/metrics.py | 62 +- .../integrations/otel/plumbing/providers.py | 40 +- litellm/integrations/otel/plumbing/routing.py | 16 +- litellm/integrations/otel/presets/__init__.py | 9 +- litellm/integrations/otel/presets/agentops.py | 18 +- litellm/integrations/otel/presets/arize.py | 16 +- litellm/integrations/otel/presets/langfuse.py | 12 +- .../integrations/otel/presets/langtrace.py | 4 +- litellm/integrations/otel/presets/levo.py | 6 +- litellm/integrations/otel/presets/phoenix.py | 10 +- litellm/integrations/otel/presets/utils.py | 3 +- litellm/integrations/otel/presets/weave.py | 12 +- litellm/integrations/otel/runtime.py | 6 +- litellm/integrations/posthog.py | 66 +- litellm/integrations/posthog_mock_client.py | 4 +- litellm/integrations/prometheus.py | 551 ++-- .../prometheus_helpers/__init__.py | 6 +- .../bounded_prometheus_series_tracker.py | 10 +- .../prometheus_helpers/prometheus_api.py | 41 +- litellm/integrations/prometheus_services.py | 36 +- litellm/integrations/prompt_layer.py | 9 +- .../integrations/prompt_management_base.py | 20 +- litellm/integrations/rubrik.py | 144 +- litellm/integrations/s3.py | 32 +- litellm/integrations/s3_v2.py | 116 +- litellm/integrations/sqs.py | 29 +- litellm/integrations/supabase.py | 5 +- litellm/integrations/traceloop.py | 9 +- .../integrations/vantage/vantage_logger.py | 26 +- .../vector_store_pre_call_hook.py | 22 +- litellm/integrations/weave/weave_otel.py | 54 +- .../websearch_interception/handler.py | 184 +- .../websearch_interception/tools.py | 22 +- .../websearch_interception/transformation.py | 24 +- litellm/integrations/weights_biases.py | 34 +- litellm/interactions/agents/http_handler.py | 72 +- litellm/interactions/agents/main.py | 108 +- litellm/interactions/agents/utils.py | 3 +- litellm/interactions/http_handler.py | 44 +- .../handler.py | 13 +- .../streaming_iterator.py | 29 +- .../transformation.py | 24 +- litellm/interactions/main.py | 104 +- litellm/interactions/streaming_iterator.py | 16 +- litellm/interactions/utils.py | 12 +- .../api_route_to_call_types.py | 10 +- litellm/litellm_core_utils/app_crypto.py | 19 +- litellm/litellm_core_utils/asyncify.py | 11 +- .../litellm_core_utils/audio_utils/utils.py | 27 +- .../chat_completion_agentic_loop.py | 36 +- litellm/litellm_core_utils/cli_token_utils.py | 15 +- .../cloud_storage_security.py | 30 +- .../litellm_core_utils/completion_timeout.py | 3 +- litellm/litellm_core_utils/core_helpers.py | 42 +- .../litellm_core_utils/coroutine_checker.py | 8 +- .../litellm_core_utils/credential_accessor.py | 4 +- .../custom_logger_registry.py | 4 +- litellm/litellm_core_utils/dd_tracing.py | 12 +- .../litellm_core_utils/default_encoding.py | 7 +- .../dot_notation_indexing.py | 22 +- litellm/litellm_core_utils/duration_parser.py | 70 +- litellm/litellm_core_utils/env_utils.py | 3 +- .../exception_mapping_utils.py | 66 +- .../fallback_generalizations.py | 34 +- litellm/litellm_core_utils/fallback_utils.py | 14 +- litellm/litellm_core_utils/get_blog_posts.py | 21 +- .../litellm_core_utils/get_litellm_params.py | 16 +- .../get_llm_provider_logic.py | 32 +- .../litellm_core_utils/get_model_cost_map.py | 42 +- .../get_provider_specific_headers.py | 6 +- .../get_supported_openai_params.py | 6 +- .../health_check_helpers.py | 10 +- .../litellm_core_utils/health_check_utils.py | 4 +- .../initialize_dynamic_callback_params.py | 16 +- .../json_validation_rule.py | 10 +- litellm/litellm_core_utils/litellm_logging.py | 537 ++-- .../llm_cost_calc/tiered_pricing.py | 16 +- .../llm_cost_calc/tool_call_cost_tracking.py | 74 +- .../litellm_core_utils/llm_cost_calc/utils.py | 122 +- .../litellm_core_utils/llm_request_utils.py | 10 +- .../convert_dict_to_response.py | 74 +- .../llm_response_utils/get_api_base.py | 4 +- .../llm_response_utils/get_headers.py | 9 +- .../llm_response_utils/response_metadata.py | 26 +- .../logging_callback_manager.py | 52 +- litellm/litellm_core_utils/logging_utils.py | 66 +- litellm/litellm_core_utils/logging_worker.py | 43 +- .../litellm_core_utils/model_param_helper.py | 37 +- .../model_response_utils.py | 4 +- .../prompt_templates/common_utils.py | 145 +- .../prompt_templates/factory.py | 430 +-- .../huggingface_template_handler.py | 26 +- .../prompt_templates/image_handling.py | 27 +- .../litellm_core_utils/realtime_streaming.py | 122 +- litellm/litellm_core_utils/redact_messages.py | 34 +- .../request_timeout_resolver.py | 4 +- litellm/litellm_core_utils/safe_json_dumps.py | 6 +- .../litellm_core_utils/secret_redaction.py | 7 +- .../sensitive_data_masker.py | 26 +- .../litellm_core_utils/service_tier_utils.py | 16 +- .../specialty_caches/dynamic_logging_cache.py | 14 +- .../streaming_chunk_builder_utils.py | 72 +- .../litellm_core_utils/streaming_handler.py | 199 +- .../thread_pool_executor.py | 5 +- litellm/litellm_core_utils/token_counter.py | 90 +- litellm/litellm_core_utils/url_utils.py | 88 +- litellm/llms/__init__.py | 10 +- .../chat/guardrail_translation/__init__.py | 4 +- .../a2a/chat/guardrail_translation/handler.py | 64 +- litellm/llms/a2a/chat/streaming_iterator.py | 12 +- litellm/llms/a2a/chat/transformation.py | 36 +- litellm/llms/a2a/common_utils.py | 22 +- litellm/llms/ai21/chat/transformation.py | 4 +- litellm/llms/aiml/chat/transformation.py | 4 +- .../aiml/image_generation/cost_calculator.py | 6 +- .../aiml/image_generation/transformation.py | 14 +- .../aiohttp_openai/chat/transformation.py | 4 +- .../llms/amazon_nova/chat/transformation.py | 6 +- litellm/llms/anthropic/batches/handler.py | 10 +- .../llms/anthropic/batches/transformation.py | 40 +- .../chat/guardrail_translation/__init__.py | 4 +- .../chat/guardrail_translation/handler.py | 116 +- litellm/llms/anthropic/chat/handler.py | 106 +- litellm/llms/anthropic/chat/transformation.py | 297 ++- litellm/llms/anthropic/common_utils.py | 136 +- .../anthropic/completion/transformation.py | 37 +- litellm/llms/anthropic/cost_calculation.py | 20 +- .../llms/anthropic/count_tokens/handler.py | 18 +- .../anthropic/count_tokens/token_counter.py | 8 +- .../anthropic/count_tokens/transformation.py | 4 +- .../adapters/handler.py | 89 +- .../adapters/streaming_iterator.py | 85 +- .../adapters/transformation.py | 141 +- .../context_management/constants.py | 30 +- .../context_management/dispatcher.py | 10 +- .../editors/clear_tool_uses.py | 32 +- .../context_management/editors/compact.py | 136 +- .../context_management/result.py | 4 +- .../messages/agentic_streaming_iterator.py | 52 +- .../messages/fake_stream_iterator.py | 38 +- .../messages/handler.py | 45 +- .../messages/interceptors/__init__.py | 4 +- .../messages/interceptors/advisor.py | 56 +- .../messages/mcp_handler.py | 20 +- .../messages/streaming_iterator.py | 18 +- .../messages/transformation.py | 92 +- .../messages/utils.py | 6 +- .../responses_adapters/handler.py | 30 +- .../responses_adapters/streaming_iterator.py | 16 +- .../responses_adapters/transformation.py | 62 +- litellm/llms/anthropic/files/handler.py | 38 +- .../llms/anthropic/files/transformation.py | 56 +- .../llms/anthropic/skills/transformation.py | 26 +- litellm/llms/apiserpent/search/defaults.py | 18 +- .../llms/apiserpent/search/transformation.py | 32 +- .../text_to_speech/transformation.py | 40 +- litellm/llms/azure/assistants.py | 80 +- .../audio_transcription/transformation.py | 30 +- litellm/llms/azure/audio_transcriptions.py | 22 +- litellm/llms/azure/azure.py | 164 +- litellm/llms/azure/batches/handler.py | 16 +- .../llms/azure/chat/gpt_5_transformation.py | 14 +- litellm/llms/azure/chat/gpt_transformation.py | 18 +- .../azure/chat/o_series_transformation.py | 12 +- litellm/llms/azure/common_utils.py | 110 +- litellm/llms/azure/completion/handler.py | 52 +- .../llms/azure/containers/transformation.py | 11 +- litellm/llms/azure/cost_calculation.py | 8 +- litellm/llms/azure/exception_mapping.py | 14 +- litellm/llms/azure/files/handler.py | 18 +- litellm/llms/azure/fine_tuning/handler.py | 20 +- .../llms/azure/image_edit/transformation.py | 12 +- .../llms/azure/passthrough/transformation.py | 14 +- litellm/llms/azure/realtime/handler.py | 16 +- .../azure/realtime/http_transformation.py | 14 +- .../responses/o_series_transformation.py | 10 +- .../llms/azure/responses/transformation.py | 52 +- .../azure/text_to_speech/transformation.py | 52 +- litellm/llms/azure_ai/agents/handler.py | 59 +- .../llms/azure_ai/agents/transformation.py | 14 +- .../anthropic/count_tokens/handler.py | 18 +- .../anthropic/count_tokens/token_counter.py | 8 +- .../anthropic/count_tokens/transformation.py | 6 +- litellm/llms/azure_ai/anthropic/handler.py | 19 +- .../anthropic/messages_transformation.py | 8 +- .../llms/azure_ai/anthropic/transformation.py | 22 +- .../azure_model_router/transformation.py | 8 +- litellm/llms/azure_ai/chat/transformation.py | 44 +- litellm/llms/azure_ai/common_utils.py | 6 +- litellm/llms/azure_ai/cost_calculator.py | 14 +- .../azure_ai/embed/cohere_transformation.py | 20 +- litellm/llms/azure_ai/embed/handler.py | 34 +- litellm/llms/azure_ai/image_edit/__init__.py | 4 +- .../image_edit/flux2_transformation.py | 12 +- .../azure_ai/image_edit/mai_transformation.py | 18 +- .../azure_ai/image_edit/transformation.py | 4 +- .../image_generation/cost_calculator.py | 8 +- .../image_generation/flux_transformation.py | 4 +- .../image_generation/mai_transformation.py | 18 +- litellm/llms/azure_ai/ocr/common_utils.py | 4 +- .../document_intelligence/transformation.py | 80 +- litellm/llms/azure_ai/ocr/transformation.py | 24 +- .../llms/azure_ai/rerank/transformation.py | 12 +- .../azure_ai/vector_stores/transformation.py | 32 +- .../audio_transcription/transformation.py | 6 +- litellm/llms/base_llm/base_model_iterator.py | 12 +- litellm/llms/base_llm/base_utils.py | 16 +- litellm/llms/base_llm/chat/transformation.py | 15 +- .../base_llm/containers/transformation.py | 4 +- .../files/azure_blob_storage_backend.py | 49 +- .../guardrail_translation/base_translation.py | 4 +- .../base_llm/guardrail_translation/utils.py | 10 +- .../base_managed_resource.py | 59 +- .../base_llm/managed_resources/isolation.py | 10 +- .../llms/base_llm/managed_resources/utils.py | 36 +- .../base_llm/passthrough/transformation.py | 10 +- .../base_llm/realtime/http_transformation.py | 3 +- .../llms/base_llm/rerank/transformation.py | 6 +- .../llms/base_llm/responses/transformation.py | 6 +- .../llms/base_llm/sandbox/transformation.py | 6 +- .../llms/base_llm/search/transformation.py | 8 +- litellm/llms/baseten/chat.py | 14 +- .../bedrock/audio_transcription/__init__.py | 11 +- litellm/llms/bedrock/base_aws_llm.py | 241 +- litellm/llms/bedrock/batches/handler.py | 48 +- .../llms/bedrock/batches/transformation.py | 104 +- .../bedrock/chat/agentcore/transformation.py | 108 +- litellm/llms/bedrock/chat/converse_handler.py | 80 +- .../bedrock/chat/converse_transformation.py | 206 +- .../chat/invoke_agent/transformation.py | 74 +- litellm/llms/bedrock/chat/invoke_handler.py | 66 +- .../amazon_ai21_transformation.py | 3 +- .../amazon_cohere_transformation.py | 5 +- .../amazon_deepseek_transformation.py | 22 +- .../amazon_llama_transformation.py | 3 +- .../amazon_mistral_transformation.py | 4 +- .../amazon_moonshot_transformation.py | 18 +- .../amazon_nova_transformation.py | 12 +- .../amazon_openai_transformation.py | 8 +- .../amazon_qwen2_transformation.py | 8 +- .../amazon_qwen3_transformation.py | 16 +- .../amazon_titan_transformation.py | 3 +- ...mazon_twelvelabs_pegasus_transformation.py | 44 +- .../anthropic_claude2_transformation.py | 3 +- .../anthropic_claude3_transformation.py | 48 +- .../base_invoke_transformation.py | 62 +- .../bedrock/chat/mantle/transformation.py | 16 +- .../bedrock/claude_platform/common_utils.py | 10 +- .../messages_transformation.py | 6 +- .../bedrock/claude_platform/transformation.py | 6 +- litellm/llms/bedrock/common_utils.py | 141 +- .../count_tokens/bedrock_token_counter.py | 12 +- litellm/llms/bedrock/count_tokens/handler.py | 26 +- .../bedrock/count_tokens/transformation.py | 38 +- .../embed/amazon_nova_transformation.py | 24 +- .../embed/amazon_titan_g1_transformation.py | 7 +- .../amazon_titan_multimodal_transformation.py | 10 +- .../embed/amazon_titan_v2_transformation.py | 7 +- .../bedrock/embed/cohere_transformation.py | 6 +- litellm/llms/bedrock/embed/embedding.py | 82 +- .../twelvelabs_marengo_transformation.py | 22 +- litellm/llms/bedrock/files/handler.py | 20 +- litellm/llms/bedrock/files/transformation.py | 176 +- ...n_nova_canvas_image_edit_transformation.py | 78 +- litellm/llms/bedrock/image_edit/handler.py | 36 +- .../image_edit/stability_transformation.py | 28 +- .../amazon_nova_canvas_transformation.py | 24 +- .../amazon_stability1_transformation.py | 23 +- .../amazon_stability3_transformation.py | 13 +- .../amazon_titan_transformation.py | 21 +- .../image_generation/cost_calculator.py | 4 +- .../bedrock/image_generation/image_handler.py | 42 +- .../anthropic_claude3_transformation.py | 104 +- .../bedrock/messages/mantle_transformation.py | 18 +- .../guardrail_translation/handler.py | 88 +- .../bedrock/passthrough/transformation.py | 34 +- litellm/llms/bedrock/realtime/handler.py | 26 +- .../llms/bedrock/realtime/transformation.py | 154 +- .../llms/bedrock/realtime/trigger_audio.py | 3 +- litellm/llms/bedrock/rerank/handler.py | 28 +- litellm/llms/bedrock/rerank/transformation.py | 12 +- .../bedrock/vector_stores/transformation.py | 40 +- .../bedrock_mantle/chat/transformation.py | 10 +- litellm/llms/bedrock_mantle/common_utils.py | 19 +- .../responses/transformation.py | 36 +- .../llms/black_forest_labs/common_utils.py | 17 +- .../black_forest_labs/image_edit/handler.py | 36 +- .../image_edit/transformation.py | 34 +- .../image_generation/handler.py | 40 +- .../image_generation/transformation.py | 18 +- litellm/llms/brave/search/transformation.py | 38 +- litellm/llms/bytez/chat/transformation.py | 58 +- litellm/llms/bytez/common_utils.py | 4 +- litellm/llms/cerebras/chat.py | 8 +- litellm/llms/chatgpt/authenticator.py | 118 +- litellm/llms/chatgpt/chat/streaming_utils.py | 6 +- litellm/llms/chatgpt/chat/transformation.py | 14 +- litellm/llms/chatgpt/common_utils.py | 58 +- .../llms/chatgpt/responses/transformation.py | 36 +- litellm/llms/clarifai/chat/transformation.py | 8 +- .../llms/cloudflare/chat/transformation.py | 6 +- litellm/llms/codestral/completion/handler.py | 33 +- .../codestral/completion/transformation.py | 11 +- litellm/llms/cohere/chat/transformation.py | 28 +- litellm/llms/cohere/chat/v2_transformation.py | 30 +- litellm/llms/cohere/common_utils.py | 72 +- litellm/llms/cohere/embed/handler.py | 10 +- litellm/llms/cohere/embed/transformation.py | 14 +- .../llms/cohere/embed/v1_transformation.py | 12 +- .../rerank/guardrail_translation/__init__.py | 4 +- .../rerank/guardrail_translation/handler.py | 12 +- litellm/llms/cohere/rerank/transformation.py | 8 +- .../llms/cohere/rerank_v2/transformation.py | 4 +- litellm/llms/cometapi/chat/transformation.py | 18 +- litellm/llms/cometapi/embed/transformation.py | 12 +- .../image_generation/cost_calculator.py | 6 +- .../image_generation/transformation.py | 8 +- .../llms/compactifai/chat/transformation.py | 8 +- litellm/llms/custom_httpx/aiohttp_handler.py | 36 +- .../llms/custom_httpx/aiohttp_transport.py | 46 +- .../llms/custom_httpx/async_client_cleanup.py | 7 +- .../llms/custom_httpx/container_handler.py | 78 +- litellm/llms/custom_httpx/http_handler.py | 152 +- litellm/llms/custom_httpx/httpx_handler.py | 9 +- litellm/llms/custom_httpx/llm_http_handler.py | 959 ++++--- litellm/llms/custom_httpx/mock_transport.py | 11 +- litellm/llms/dashscope/chat/transformation.py | 4 +- litellm/llms/dashscope/cost_calculator.py | 23 +- .../llms/dashscope/embed/transformation.py | 22 +- .../image_generation/transformation.py | 18 +- .../llms/dashscope/rerank/transformation.py | 28 +- .../llms/databricks/chat/transformation.py | 62 +- litellm/llms/databricks/common_utils.py | 26 +- litellm/llms/databricks/cost_calculator.py | 8 +- litellm/llms/databricks/embed/handler.py | 3 +- .../llms/databricks/embed/transformation.py | 3 +- .../databricks/responses/transformation.py | 6 +- litellm/llms/databricks/streaming_utils.py | 11 +- .../llms/dataforseo/search/transformation.py | 20 +- litellm/llms/datarobot/chat/transformation.py | 9 +- .../audio_transcription/transformation.py | 31 +- litellm/llms/deepinfra/chat/transformation.py | 16 +- .../llms/deepinfra/rerank/transformation.py | 36 +- litellm/llms/deepseek/chat/transformation.py | 28 +- .../llms/deepseek/messages/transformation.py | 8 +- .../llms/deprecated_providers/aleph_alpha.py | 25 +- litellm/llms/deprecated_providers/palm.py | 17 +- .../chat/transformation.py | 8 +- .../llms/duckduckgo/search/transformation.py | 22 +- litellm/llms/e2b/sandbox/transformation.py | 58 +- .../audio_transcription/transformation.py | 22 +- .../text_to_speech/transformation.py | 40 +- litellm/llms/exa_ai/search/transformation.py | 10 +- litellm/llms/fal_ai/cost_calculator.py | 6 +- .../llms/fal_ai/image_generation/__init__.py | 4 +- .../image_generation/bria_transformation.py | 20 +- .../bytedance_transformation.py | 10 +- .../flux_pro_v11_transformation.py | 10 +- .../flux_pro_v11_ultra_transformation.py | 20 +- .../flux_schnell_transformation.py | 10 +- .../ideogram_v3_transformation.py | 10 +- .../imagen4_transformation.py | 20 +- .../nano_banana_transformation.py | 10 +- .../recraft_v3_transformation.py | 14 +- .../stable_diffusion_transformation.py | 14 +- .../fal_ai/image_generation/transformation.py | 12 +- litellm/llms/fastcrw/search/transformation.py | 12 +- .../featherless_ai/chat/transformation.py | 6 +- .../llms/firecrawl/search/transformation.py | 16 +- .../llms/fireworks_ai/chat/transformation.py | 80 +- litellm/llms/fireworks_ai/common_utils.py | 14 +- .../fireworks_ai/completion/transformation.py | 8 +- litellm/llms/fireworks_ai/cost_calculator.py | 30 +- .../embed/fireworks_ai_transformation.py | 4 +- .../fireworks_ai/rerank/transformation.py | 24 +- litellm/llms/gdc/chat/transformation.py | 30 +- litellm/llms/gemini/agents/transformation.py | 22 +- litellm/llms/gemini/chat/transformation.py | 6 +- litellm/llms/gemini/common_utils.py | 78 +- litellm/llms/gemini/cost_calculator.py | 10 +- litellm/llms/gemini/count_tokens/handler.py | 18 +- litellm/llms/gemini/files/transformation.py | 56 +- .../gemini/google_genai/transformation.py | 32 +- .../llms/gemini/image_edit/transformation.py | 24 +- .../image_generation/cost_calculator.py | 12 +- .../gemini/image_generation/transformation.py | 18 +- .../llms/gemini/image_usage_transformation.py | 14 +- .../gemini/interactions/transformation.py | 42 +- .../llms/gemini/realtime/transformation.py | 186 +- .../gemini/vector_stores/transformation.py | 44 +- litellm/llms/gemini/videos/transformation.py | 94 +- litellm/llms/gigachat/authenticator.py | 41 +- litellm/llms/gigachat/chat/streaming.py | 10 +- litellm/llms/gigachat/chat/transformation.py | 32 +- .../llms/gigachat/embedding/transformation.py | 11 +- litellm/llms/gigachat/file_handler.py | 71 +- litellm/llms/github_copilot/authenticator.py | 56 +- .../github_copilot/chat/transformation.py | 30 +- litellm/llms/github_copilot/common_utils.py | 11 +- .../embedding/transformation.py | 8 +- .../github_copilot/messages/transformation.py | 10 +- .../responses/transformation.py | 32 +- .../llms/google_pse/search/transformation.py | 20 +- .../llms/gradient_ai/chat/transformation.py | 16 +- litellm/llms/groq/chat/transformation.py | 39 +- litellm/llms/groq/cost_calculator.py | 12 +- litellm/llms/groq/stt/transformation.py | 5 +- .../llms/hosted_vllm/chat/transformation.py | 29 +- .../hosted_vllm/embedding/transformation.py | 6 +- .../llms/hosted_vllm/rerank/transformation.py | 22 +- .../hosted_vllm/responses/transformation.py | 4 +- .../transcriptions/transformation.py | 4 +- .../llms/huggingface/chat/transformation.py | 14 +- litellm/llms/huggingface/common_utils.py | 16 +- litellm/llms/huggingface/embedding/handler.py | 52 +- .../huggingface/embedding/transformation.py | 38 +- .../llms/huggingface/rerank/transformation.py | 28 +- .../llms/hyperbolic/chat/transformation.py | 4 +- litellm/llms/inception/chat/transformation.py | 4 +- .../inception/completion/transformation.py | 4 +- .../llms/infinity/embedding/transformation.py | 8 +- .../llms/infinity/rerank/transformation.py | 14 +- .../llms/jina_ai/embedding/transformation.py | 14 +- litellm/llms/jina_ai/rerank/transformation.py | 28 +- litellm/llms/lambda_ai/chat/transformation.py | 4 +- litellm/llms/langflow/a2a.py | 10 +- litellm/llms/langflow/chat/transformation.py | 40 +- litellm/llms/langgraph/chat/sse_iterator.py | 17 +- litellm/llms/langgraph/chat/transformation.py | 50 +- litellm/llms/lemonade/chat/transformation.py | 30 +- litellm/llms/lemonade/cost_calculator.py | 6 +- litellm/llms/linkup/search/transformation.py | 12 +- .../llms/litellm_proxy/chat/transformation.py | 12 +- .../litellm_proxy/skills/code_execution.py | 14 +- .../llms/litellm_proxy/skills/constants.py | 8 +- litellm/llms/litellm_proxy/skills/handler.py | 38 +- .../litellm_proxy/skills/prompt_injection.py | 32 +- .../litellm_proxy/skills/sandbox_executor.py | 32 +- .../litellm_proxy/skills/transformation.py | 16 +- litellm/llms/llamafile/chat/transformation.py | 4 +- litellm/llms/lm_studio/chat/transformation.py | 4 +- .../llms/lm_studio/embed/transformation.py | 3 +- litellm/llms/manus/files/transformation.py | 46 +- .../llms/manus/responses/transformation.py | 30 +- litellm/llms/maritalk.py | 4 +- .../llms/meta_llama/chat/transformation.py | 6 +- .../milvus/vector_stores/transformation.py | 32 +- litellm/llms/minimax/chat/transformation.py | 8 +- .../llms/minimax/messages/transformation.py | 4 +- .../minimax/text_to_speech/transformation.py | 66 +- .../audio_transcription/transformation.py | 20 +- litellm/llms/mistral/chat/transformation.py | 40 +- .../ocr/guardrail_translation/__init__.py | 4 +- .../ocr/guardrail_translation/handler.py | 26 +- litellm/llms/mistral/ocr/transformation.py | 12 +- .../llms/modelscope/chat/transformation.py | 6 +- .../image_generation/transformation.py | 16 +- litellm/llms/moonshot/chat/transformation.py | 14 +- litellm/llms/morph/chat/transformation.py | 4 +- litellm/llms/nebius/chat/transformation.py | 4 +- litellm/llms/nlp_cloud/chat/handler.py | 15 +- litellm/llms/nlp_cloud/chat/transformation.py | 16 +- litellm/llms/nscale/chat/transformation.py | 6 +- .../llms/nvidia_nim/chat/transformation.py | 4 +- litellm/llms/nvidia_nim/embed.py | 3 +- .../rerank/ranking_transformation.py | 4 +- .../llms/nvidia_nim/rerank/transformation.py | 42 +- .../audio_transcription/audio_utils.py | 30 +- .../audio_transcription/handler.py | 64 +- .../audio_transcription/transformation.py | 30 +- litellm/llms/nvidia_riva/common_utils.py | 16 +- litellm/llms/oci/chat/cohere.py | 32 +- litellm/llms/oci/chat/generic.py | 32 +- litellm/llms/oci/chat/transformation.py | 91 +- litellm/llms/oci/common_utils.py | 118 +- litellm/llms/oci/embed/transformation.py | 30 +- litellm/llms/ollama/chat/transformation.py | 46 +- litellm/llms/ollama/common_utils.py | 44 +- litellm/llms/ollama/completion/handler.py | 22 +- .../llms/ollama/completion/transformation.py | 42 +- litellm/llms/oobabooga/chat/oobabooga.py | 22 +- litellm/llms/oobabooga/chat/transformation.py | 6 +- .../llms/openai/chat/gpt_5_transformation.py | 32 +- .../openai/chat/gpt_audio_transformation.py | 6 +- .../llms/openai/chat/gpt_transformation.py | 93 +- .../chat/guardrail_translation/__init__.py | 4 +- .../chat/guardrail_translation/handler.py | 114 +- .../openai/chat/o_series_transformation.py | 18 +- litellm/llms/openai/common_utils.py | 46 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 20 +- litellm/llms/openai/completion/handler.py | 53 +- .../llms/openai/completion/transformation.py | 8 +- litellm/llms/openai/completion/utils.py | 6 +- .../llms/openai/containers/transformation.py | 56 +- litellm/llms/openai/containers/utils.py | 6 +- litellm/llms/openai/cost_calculation.py | 22 +- litellm/llms/openai/data_residency.py | 5 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 14 +- litellm/llms/openai/evals/transformation.py | 77 +- litellm/llms/openai/fine_tuning/handler.py | 36 +- litellm/llms/openai/image_edit/__init__.py | 4 +- .../image_edit/dalle2_transformation.py | 18 +- .../llms/openai/image_edit/transformation.py | 22 +- .../image_generation/cost_calculator.py | 8 +- .../dall_e_2_transformation.py | 10 +- .../dall_e_3_transformation.py | 10 +- .../image_generation/gpt_transformation.py | 10 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 12 +- .../llms/openai/image_variations/handler.py | 37 +- litellm/llms/openai/openai.py | 270 +- litellm/llms/openai/realtime/handler.py | 16 +- .../openai/responses/count_tokens/handler.py | 18 +- .../responses/count_tokens/token_counter.py | 10 +- .../responses/count_tokens/transformation.py | 12 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 78 +- .../llms/openai/responses/transformation.py | 102 +- .../speech/guardrail_translation/__init__.py | 4 +- .../speech/guardrail_translation/handler.py | 12 +- .../transcriptions/gpt_transformation.py | 4 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 12 +- litellm/llms/openai/transcriptions/handler.py | 24 +- .../transcriptions/whisper_transformation.py | 11 +- .../vector_store_files/transformation.py | 30 +- .../openai/vector_stores/transformation.py | 26 +- litellm/llms/openai/videos/transformation.py | 96 +- litellm/llms/openai_like/chat/handler.py | 32 +- .../llms/openai_like/chat/transformation.py | 10 +- litellm/llms/openai_like/dynamic_config.py | 24 +- litellm/llms/openai_like/embedding/handler.py | 11 +- litellm/llms/openai_like/json_loader.py | 7 +- .../openai_like/messages/transformation.py | 18 +- .../openai_like/responses/transformation.py | 4 +- .../llms/openrouter/chat/transformation.py | 32 +- .../openrouter/embedding/transformation.py | 12 +- .../openrouter/image_edit/transformation.py | 40 +- .../image_generation/transformation.py | 28 +- .../openrouter/responses/transformation.py | 4 +- .../opensandbox/sandbox/transformation.py | 106 +- .../audio_transcription/transformation.py | 22 +- litellm/llms/ovhcloud/chat/transformation.py | 16 +- .../llms/ovhcloud/embedding/transformation.py | 12 +- .../llms/parallel_ai/search/transformation.py | 20 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 64 +- .../llms/perplexity/chat/transformation.py | 34 +- litellm/llms/perplexity/cost_calculator.py | 26 +- .../perplexity/embedding/transformation.py | 18 +- .../perplexity/responses/transformation.py | 10 +- .../llms/perplexity/search/transformation.py | 16 +- litellm/llms/petals/completion/handler.py | 23 +- .../llms/petals/completion/transformation.py | 4 +- .../pg_vector/vector_stores/transformation.py | 8 +- litellm/llms/predibase/chat/handler.py | 25 +- litellm/llms/predibase/chat/transformation.py | 34 +- litellm/llms/ragflow/chat/transformation.py | 16 +- .../ragflow/vector_stores/transformation.py | 28 +- litellm/llms/recraft/cost_calculator.py | 6 +- .../llms/recraft/image_edit/transformation.py | 28 +- .../image_generation/transformation.py | 10 +- litellm/llms/reducto/common.py | 26 +- litellm/llms/reducto/ocr/transformation.py | 38 +- litellm/llms/replicate/chat/handler.py | 15 +- litellm/llms/replicate/chat/transformation.py | 30 +- litellm/llms/runwayml/cost_calculator.py | 6 +- .../image_generation/transformation.py | 40 +- .../runwayml/text_to_speech/transformation.py | 68 +- .../llms/runwayml/videos/transformation.py | 88 +- .../vector_stores/transformation.py | 32 +- litellm/llms/sagemaker/chat/handler.py | 45 +- litellm/llms/sagemaker/chat/transformation.py | 22 +- litellm/llms/sagemaker/common_utils.py | 27 +- litellm/llms/sagemaker/completion/handler.py | 124 +- .../sagemaker/completion/transformation.py | 26 +- .../sagemaker/embedding/transformation.py | 12 +- litellm/llms/sagemaker/nova/transformation.py | 8 +- litellm/llms/sambanova/chat.py | 8 +- .../sambanova/embedding/transformation.py | 10 +- litellm/llms/sap/chat/handler.py | 5 +- litellm/llms/sap/chat/models.py | 8 +- litellm/llms/sap/chat/transformation.py | 60 +- litellm/llms/sap/credentials.py | 82 +- litellm/llms/sap/embed/transformation.py | 22 +- .../audio_transcription/transformation.py | 20 +- .../llms/searchapi/search/transformation.py | 22 +- litellm/llms/searxng/search/transformation.py | 16 +- litellm/llms/serper/search/transformation.py | 14 +- litellm/llms/snowflake/chat/transformation.py | 92 +- .../snowflake/embedding/transformation.py | 6 +- litellm/llms/snowflake/utils.py | 8 +- .../soniox/audio_transcription/handler.py | 85 +- .../audio_transcription/transformation.py | 28 +- litellm/llms/soniox/common_utils.py | 54 +- .../stability/image_edit/transformations.py | 28 +- .../image_generation/transformation.py | 16 +- litellm/llms/tavily/search/transformation.py | 10 +- litellm/llms/tencent/chat/transformation.py | 10 +- .../llms/tinyfish/search/transformation.py | 50 +- litellm/llms/together_ai/chat.py | 6 +- .../together_ai/completion/transformation.py | 6 +- litellm/llms/together_ai/cost_calculator.py | 9 +- litellm/llms/together_ai/rerank/handler.py | 18 +- .../llms/together_ai/rerank/transformation.py | 12 +- .../topaz/image_variations/transformation.py | 14 +- .../llms/triton/completion/transformation.py | 36 +- .../llms/triton/embedding/transformation.py | 12 +- litellm/llms/v0/chat/transformation.py | 4 +- .../vercel_ai_gateway/chat/transformation.py | 16 +- .../embedding/transformation.py | 8 +- .../vertex_ai/agent_engine/sse_iterator.py | 10 +- .../vertex_ai/agent_engine/transformation.py | 64 +- .../audio_transcription/transformation.py | 41 +- .../vertex_ai/aws_credentials_supplier.py | 3 +- litellm/llms/vertex_ai/batches/handler.py | 88 +- .../llms/vertex_ai/batches/transformation.py | 54 +- litellm/llms/vertex_ai/common_utils.py | 102 +- .../context_caching/transformation.py | 26 +- .../vertex_ai_context_caching.py | 48 +- litellm/llms/vertex_ai/cost_calculator.py | 18 +- .../llms/vertex_ai/count_tokens/handler.py | 8 +- litellm/llms/vertex_ai/files/handler.py | 26 +- .../llms/vertex_ai/files/transformation.py | 130 +- litellm/llms/vertex_ai/fine_tuning/handler.py | 62 +- .../llms/vertex_ai/gemini/transformation.py | 124 +- .../vertex_and_google_ai_studio_gemini.py | 303 ++- .../batch_embed_content_handler.py | 40 +- .../batch_embed_content_transformation.py | 85 +- .../vertex_ai/google_genai/transformation.py | 8 +- litellm/llms/vertex_ai/image_edit/__init__.py | 4 +- .../vertex_ai/image_edit/cost_calculator.py | 8 +- .../vertex_gemini_transformation.py | 44 +- .../vertex_imagen_transformation.py | 54 +- .../vertex_ai/image_generation/__init__.py | 4 +- .../image_generation/cost_calculator.py | 12 +- .../image_generation_handler.py | 34 +- .../vertex_gemini_transformation.py | 34 +- .../vertex_imagen_transformation.py | 30 +- .../embedding_handler.py | 18 +- .../multimodal_embeddings/transformation.py | 24 +- .../vertex_ai/ocr/deepseek_transformation.py | 26 +- litellm/llms/vertex_ai/ocr/transformation.py | 30 +- .../llms/vertex_ai/rag_engine/ingestion.py | 26 +- .../vertex_ai/rag_engine/transformation.py | 16 +- .../llms/vertex_ai/realtime/transformation.py | 31 +- .../llms/vertex_ai/rerank/transformation.py | 42 +- .../text_to_speech/text_to_speech_handler.py | 34 +- .../text_to_speech/transformation.py | 48 +- .../vector_stores/rag_api/transformation.py | 52 +- .../search_api/transformation.py | 60 +- litellm/llms/vertex_ai/vertex_ai_aws_wif.py | 16 +- .../llms/vertex_ai/vertex_ai_non_gemini.py | 34 +- .../ai21/transformation.py | 3 +- .../transformation.py | 22 +- .../anthropic/output_params_utils.py | 10 +- .../anthropic/transformation.py | 24 +- .../count_tokens/handler.py | 26 +- .../gpt_oss/transformation.py | 6 +- .../llama3/transformation.py | 20 +- .../vertex_ai_partner_models/main.py | 21 +- .../llms/vertex_ai/vertex_embeddings/bge.py | 22 +- .../vertex_embeddings/embedding_handler.py | 34 +- .../vertex_embeddings/transformation.py | 34 +- .../vertex_ai/vertex_gemma_models/main.py | 5 +- .../vertex_gemma_models/transformation.py | 32 +- litellm/llms/vertex_ai/vertex_llm_base.py | 72 +- .../vertex_ai/vertex_model_garden/main.py | 7 +- .../llms/vertex_ai/videos/transformation.py | 126 +- litellm/llms/vllm/common_utils.py | 10 +- litellm/llms/vllm/completion/handler.py | 25 +- .../llms/vllm/passthrough/transformation.py | 4 +- litellm/llms/volcengine/__init__.py | 4 +- .../llms/volcengine/chat/transformation.py | 6 +- litellm/llms/volcengine/common_utils.py | 4 +- .../volcengine/embedding/transformation.py | 18 +- .../volcengine/responses/transformation.py | 110 +- .../llms/voyage/embedding/transformation.py | 6 +- .../embedding/transformation_contextual.py | 6 +- .../embedding/transformation_multimodal.py | 14 +- litellm/llms/voyage/rerank/transformation.py | 22 +- litellm/llms/wandb/chat/transformation.py | 4 +- .../audio_transcription/transformation.py | 30 +- litellm/llms/watsonx/chat/handler.py | 7 +- litellm/llms/watsonx/chat/transformation.py | 8 +- litellm/llms/watsonx/common_utils.py | 32 +- .../llms/watsonx/completion/transformation.py | 47 +- litellm/llms/watsonx/embed/transformation.py | 16 +- .../watsonx/passthrough/transformation.py | 6 +- litellm/llms/watsonx/rerank/transformation.py | 34 +- litellm/llms/xai/chat/transformation.py | 22 +- litellm/llms/xai/common_utils.py | 8 +- litellm/llms/xai/cost_calculator.py | 18 +- litellm/llms/xai/oauth.py | 102 +- litellm/llms/xai/realtime/transformation.py | 42 +- litellm/llms/xai/responses/transformation.py | 24 +- .../image_generation/transformation.py | 4 +- litellm/llms/you_com/search/transformation.py | 18 +- litellm/llms/zai/chat/transformation.py | 8 +- litellm/main.py | 2302 ++++++++--------- litellm/models/team.py | 4 +- litellm/ocr/main.py | 84 +- litellm/passthrough/main.py | 51 +- litellm/passthrough/timeout_utils.py | 7 +- litellm/passthrough/utils.py | 17 +- .../mcp_server/auth/token_endpoint_auth.py | 5 +- .../mcp_server/auth/user_api_key_auth_mcp.py | 298 +-- .../mcp_server/bridge_token_flow.py | 58 +- .../mcp_server/byok_oauth_endpoints.py | 66 +- .../mcp_server/cost_calculator.py | 14 +- litellm/proxy/_experimental/mcp_server/db.py | 241 +- .../mcp_server/discoverable_endpoints.py | 318 +-- .../mcp_server/elicitation_handler.py | 12 +- .../_experimental/mcp_server/exceptions.py | 6 +- .../mcp_server/faults/classify.py | 14 +- .../mcp_server/faults/list_outcomes.py | 12 +- .../mcp_server/faults/render_oauth.py | 6 +- .../mcp_server/faults/traversal.py | 5 +- .../_experimental/mcp_server/faults/types.py | 10 +- .../mcp_server/gateway_dcr_flow.py | 122 +- .../guardrail_translation/__init__.py | 4 +- .../guardrail_translation/handler.py | 38 +- .../_experimental/mcp_server/mcp_context.py | 7 +- .../_experimental/mcp_server/mcp_debug.py | 24 +- .../mcp_server/mcp_server_manager.py | 675 +++-- .../mcp_server/oauth2_flow_backfill.py | 18 +- .../mcp_server/oauth2_token_cache.py | 44 +- .../mcp_server/oauth_issuer_stamp_backfill.py | 18 +- .../_experimental/mcp_server/oauth_utils.py | 118 +- .../mcp_server/openapi_to_mcp_generator.py | 90 +- .../outbound_credentials/adapter.py | 48 +- .../authz_code_refresher.py | 20 +- .../bridge_credentials.py | 36 +- .../client_credentials.py | 36 +- .../dual_cache_token_backend.py | 4 +- .../outbound_credentials/envelope.py | 46 +- .../outbound_credentials/oauth_token_store.py | 22 +- .../per_user_oauth_store.py | 34 +- .../redis_distributed_lock.py | 8 +- .../redis_refresh_coordinator.py | 14 +- .../outbound_credentials/resolver.py | 33 +- .../session_credentials.py | 28 +- .../outbound_credentials/session_token.py | 26 +- .../sso_assertion_store.py | 34 +- .../outbound_credentials/token_cache_codec.py | 3 +- .../outbound_credentials/token_endpoint.py | 17 +- .../token_exchange_provider.py | 20 +- .../outbound_credentials/token_exchanger.py | 56 +- .../mcp_server/outbound_credentials/types.py | 4 +- .../outbound_credentials/v2_token_store.py | 9 +- .../mcp_server/rest_endpoints.py | 162 +- .../mcp_server/sampling_handler.py | 132 +- .../mcp_server/semantic_tool_filter.py | 34 +- .../proxy/_experimental/mcp_server/server.py | 454 ++-- .../_experimental/mcp_server/sse_transport.py | 20 +- .../_experimental/mcp_server/tool_registry.py | 4 +- .../_experimental/mcp_server/tool_search.py | 22 +- .../_experimental/mcp_server/toolset_db.py | 19 +- .../mcp_server/ui_session_utils.py | 18 +- .../proxy/_experimental/mcp_server/utils.py | 98 +- litellm/proxy/_lazy_features.py | 24 +- litellm/proxy/_lazy_openapi_snapshot.py | 11 +- litellm/proxy/_logging.py | 13 +- litellm/proxy/_types.py | 86 +- litellm/proxy/a2a/agent_card.py | 32 +- litellm/proxy/a2a/discovery.py | 16 +- litellm/proxy/a2a/endpoints.py | 6 +- litellm/proxy/a2a/version_convert.py | 84 +- .../proxy/agent_endpoints/a2a_endpoints.py | 124 +- litellm/proxy/agent_endpoints/a2a_routing.py | 12 +- .../proxy/agent_endpoints/agent_registry.py | 86 +- .../auth/agent_permission_handler.py | 50 +- .../proxy/agent_endpoints/databricks_oauth.py | 46 +- litellm/proxy/agent_endpoints/endpoints.py | 92 +- .../agent_endpoints/model_list_helpers.py | 6 +- .../analytics_endpoints.py | 4 +- .../analytics_endpoints/cache_activity.py | 26 +- .../claude_code_marketplace.py | 48 +- .../proxy/anthropic_endpoints/endpoints.py | 46 +- .../anthropic_endpoints/skills_endpoints.py | 36 +- litellm/proxy/auth/auth_checks.py | 416 +-- .../proxy/auth/auth_checks_organization.py | 25 +- litellm/proxy/auth/auth_exception_handler.py | 8 +- litellm/proxy/auth/auth_utils.py | 180 +- litellm/proxy/auth/budget_throttle.py | 3 +- litellm/proxy/auth/handle_jwt.py | 220 +- litellm/proxy/auth/ip_address_utils.py | 32 +- litellm/proxy/auth/litellm_license.py | 24 +- litellm/proxy/auth/login_utils.py | 14 +- litellm/proxy/auth/model_checks.py | 40 +- litellm/proxy/auth/network.py | 14 +- litellm/proxy/auth/oauth2_check.py | 30 +- litellm/proxy/auth/oauth2_proxy_hook.py | 11 +- litellm/proxy/auth/rds_iam_token.py | 12 +- litellm/proxy/auth/resolvers/store.py | 22 +- litellm/proxy/auth/roles.py | 3 +- litellm/proxy/auth/route_checks.py | 45 +- litellm/proxy/auth/trusted_proxy_utils.py | 12 +- litellm/proxy/auth/user_api_key_auth.py | 220 +- litellm/proxy/batches_endpoints/endpoints.py | 106 +- litellm/proxy/caching_routes.py | 22 +- litellm/proxy/client/chat.py | 22 +- litellm/proxy/client/cli/commands/agents.py | 61 +- litellm/proxy/client/cli/commands/auth.py | 94 +- .../client/cli/commands/autoroute/commands.py | 41 +- .../client/cli/commands/autoroute/config.py | 28 +- .../client/cli/commands/autoroute/process.py | 27 +- .../client/cli/commands/autoroute/settings.py | 22 +- .../client/cli/commands/autoroute/wizard.py | 49 +- litellm/proxy/client/cli/commands/chat.py | 34 +- litellm/proxy/client/cli/commands/config.py | 21 +- .../proxy/client/cli/commands/credentials.py | 28 +- .../proxy/client/cli/commands/encryption.py | 8 +- litellm/proxy/client/cli/commands/http.py | 9 +- litellm/proxy/client/cli/commands/keys.py | 48 +- .../proxy/client/cli/commands/model_groups.py | 8 +- litellm/proxy/client/cli/commands/models.py | 94 +- .../proxy/client/cli/commands/private_json.py | 3 +- litellm/proxy/client/cli/commands/teams.py | 28 +- litellm/proxy/client/cli/commands/up.py | 64 +- litellm/proxy/client/cli/commands/users.py | 22 +- litellm/proxy/client/cli/interface.py | 43 +- litellm/proxy/client/cli/main.py | 10 +- litellm/proxy/client/credentials.py | 38 +- litellm/proxy/client/health.py | 4 +- litellm/proxy/client/http_client.py | 8 +- litellm/proxy/client/keys.py | 54 +- litellm/proxy/client/model_groups.py | 12 +- litellm/proxy/client/models.py | 54 +- litellm/proxy/client/teams.py | 20 +- litellm/proxy/client/users.py | 28 +- litellm/proxy/common_request_processing.py | 338 ++- litellm/proxy/common_utils/admin_ui_utils.py | 7 +- litellm/proxy/common_utils/banner.py | 4 +- .../proxy/common_utils/cache_coordinator.py | 14 +- litellm/proxy/common_utils/callback_utils.py | 78 +- .../proxy/common_utils/config_sync_pubsub.py | 38 +- .../proxy/common_utils/custom_openapi_spec.py | 8 +- litellm/proxy/common_utils/debug_utils.py | 128 +- .../common_utils/encrypt_decrypt_utils.py | 38 +- .../expired_ui_session_key_cleanup_manager.py | 18 +- litellm/proxy/common_utils/get_routes.py | 12 +- .../html_forms/cli_sso_success.py | 4 +- .../html_forms/jwt_display_template.py | 4 +- .../proxy/common_utils/html_forms/ui_login.py | 9 +- .../proxy/common_utils/http_parsing_utils.py | 60 +- .../proxy/common_utils/json_merge_patch.py | 10 +- .../common_utils/key_rotation_manager.py | 23 +- .../proxy/common_utils/load_config_utils.py | 27 +- .../proxy/common_utils/model_listing_utils.py | 22 +- .../common_utils/openai_endpoint_utils.py | 10 +- .../common_utils/openapi_schema_compat.py | 12 +- litellm/proxy/common_utils/path_utils.py | 7 +- .../proxy/common_utils/performance_utils.py | 36 +- .../common_utils/proxy_rate_limit_error.py | 6 +- litellm/proxy/common_utils/rbac_utils.py | 8 +- .../proxy/common_utils/reset_budget_job.py | 62 +- .../proxy/common_utils/resource_ownership.py | 4 +- .../proxy/common_utils/static_asset_utils.py | 9 +- litellm/proxy/common_utils/swagger_utils.py | 4 +- .../proxy/common_utils/user_api_key_cache.py | 24 +- litellm/proxy/compliance_checks.py | 50 +- .../pass_through_endpoints.py | 4 +- .../proxy/config_resolvers/_descriptors.py | 12 +- litellm/proxy/config_resolvers/alerting.py | 6 +- litellm/proxy/config_resolvers/sso.py | 17 +- .../proxy/container_endpoints/endpoints.py | 30 +- .../container_endpoints/handler_factory.py | 36 +- .../proxy/container_endpoints/ownership.py | 96 +- .../proxy/credential_endpoints/endpoints.py | 50 +- .../proxy/custom_hooks/custom_ui_sso_hook.py | 6 +- litellm/proxy/custom_prompt_management.py | 4 +- litellm/proxy/custom_sso.py | 4 +- litellm/proxy/db/check_migration.py | 11 +- litellm/proxy/db/create_views.py | 8 +- litellm/proxy/db/db_spend_update_writer.py | 154 +- .../db_transaction_queue/base_update_queue.py | 4 +- .../daily_spend_update_queue.py | 11 +- .../db_transaction_queue/pod_lock_manager.py | 14 +- .../redis_update_buffer.py | 60 +- .../db_transaction_queue/spend_log_cleanup.py | 13 +- .../spend_logs_partition_manager.py | 25 +- .../spend_update_queue.py | 15 +- .../tool_discovery_queue.py | 4 +- litellm/proxy/db/db_url_settings.py | 40 +- litellm/proxy/db/dynamo_db.py | 12 +- litellm/proxy/db/exception_handler.py | 20 +- litellm/proxy/db/log_db_metrics.py | 11 +- litellm/proxy/db/prisma_client.py | 96 +- litellm/proxy/db/query_engine_reaper.py | 41 +- litellm/proxy/db/routing_prisma_wrapper.py | 22 +- litellm/proxy/db/spend_counter_reseed.py | 30 +- litellm/proxy/db/spend_log_batching.py | 7 +- litellm/proxy/db/spend_log_tool_index.py | 14 +- litellm/proxy/db/tool_registry_writer.py | 44 +- .../ui_discovery_endpoints.py | 13 +- .../enterprise_billing/billing_metrics.py | 78 +- .../proxy/example_config_yaml/custom_auth.py | 7 +- .../example_config_yaml/custom_callbacks.py | 11 +- .../example_config_yaml/custom_callbacks1.py | 4 +- .../example_config_yaml/custom_guardrail.py | 8 +- .../example_config_yaml/custom_handler.py | 4 +- .../custom_team_metadata_validate.py | 5 +- .../team_metadata_validator_e2e.py | 23 +- .../proxy/fine_tuning_endpoints/endpoints.py | 64 +- .../google_endpoints/agents_endpoints.py | 39 +- litellm/proxy/google_endpoints/endpoints.py | 38 +- litellm/proxy/guardrails/_content_utils.py | 26 +- .../proxy/guardrails/guardrail_endpoints.py | 267 +- litellm/proxy/guardrails/guardrail_helpers.py | 7 +- .../guardrail_hooks/aim/__init__.py | 8 +- .../guardrails/guardrail_hooks/aim/aim.py | 54 +- .../guardrail_hooks/akto/__init__.py | 8 +- .../guardrails/guardrail_hooks/akto/akto.py | 76 +- .../guardrail_hooks/aporia_ai/__init__.py | 8 +- .../guardrail_hooks/aporia_ai/aporia_ai.py | 26 +- .../guardrail_hooks/azure/__init__.py | 10 +- .../guardrails/guardrail_hooks/azure/base.py | 16 +- .../guardrail_hooks/azure/prompt_shield.py | 8 +- .../guardrail_hooks/azure/text_moderation.py | 12 +- .../guardrail_hooks/bedrock_guardrails.py | 266 +- .../block_code_execution/__init__.py | 28 +- .../block_code_execution.py | 52 +- .../guardrail_hooks/cato_networks/__init__.py | 8 +- .../cato_networks/cato_networks.py | 98 +- .../cisco_ai_defense/__init__.py | 18 +- .../cisco_ai_defense/cisco_ai_defense.py | 252 +- .../cisco_ai_defense/cisco_ai_defense_mcp.py | 92 +- .../guardrail_hooks/compresr/__init__.py | 12 +- .../guardrail_hooks/compresr/compresr.py | 178 +- .../guardrail_hooks/content_text.py | 23 +- .../crowdstrike_aidr/__init__.py | 10 +- .../crowdstrike_aidr/crowdstrike_aidr.py | 85 +- .../guardrail_hooks/custom_code/__init__.py | 12 +- .../custom_code/custom_code_guardrail.py | 24 +- .../guardrail_hooks/custom_code/primitives.py | 34 +- .../custom_code/response_rejection_code.py | 6 +- .../guardrail_hooks/custom_code/sandbox.py | 8 +- .../guardrail_hooks/custom_guardrail.py | 6 +- .../guardrail_hooks/deepkeep/__init__.py | 8 +- .../guardrail_hooks/deepkeep/deepkeep.py | 64 +- .../guardrail_hooks/dynamoai/__init__.py | 8 +- .../guardrail_hooks/dynamoai/dynamoai.py | 66 +- .../guardrail_hooks/enkryptai/__init__.py | 8 +- .../guardrail_hooks/enkryptai/enkryptai.py | 47 +- .../generic_guardrail_api/__init__.py | 12 +- .../generic_guardrail_api.py | 76 +- .../guardrail_hooks/grayswan/__init__.py | 14 +- .../guardrail_hooks/grayswan/grayswan.py | 126 +- .../guardrail_hooks/guardrails_ai/__init__.py | 8 +- .../guardrails_ai/guardrails_ai.py | 29 +- .../guardrail_hooks/headroom/__init__.py | 8 +- .../guardrail_hooks/headroom/headroom.py | 110 +- .../guardrail_hooks/hiddenlayer/__init__.py | 12 +- .../hiddenlayer/hiddenlayer.py | 38 +- .../ibm_guardrails/__init__.py | 22 +- .../ibm_guardrails/ibm_detector.py | 50 +- .../guardrail_hooks/javelin/__init__.py | 8 +- .../guardrail_hooks/javelin/javelin.py | 24 +- .../guardrails/guardrail_hooks/lakera_ai.py | 28 +- .../guardrail_hooks/lakera_ai_v2.py | 45 +- .../guardrail_hooks/lasso/__init__.py | 8 +- .../guardrails/guardrail_hooks/lasso/lasso.py | 119 +- .../litellm_content_filter/__init__.py | 10 +- .../competitor_intent/airline.py | 38 +- .../competitor_intent/base.py | 38 +- .../litellm_content_filter/content_filter.py | 196 +- .../guardrail_benchmarks/test_eval.py | 60 +- .../litellm_content_filter/patterns.py | 24 +- .../llm_as_a_judge/__init__.py | 66 +- .../mcp_end_user_permission/__init__.py | 12 +- .../mcp_end_user_permission.py | 28 +- .../mcp_jwt_signer/__init__.py | 16 +- .../mcp_jwt_signer/mcp_jwt_signer.py | 126 +- .../guardrail_hooks/mcp_security/__init__.py | 12 +- .../mcp_security/mcp_security_guardrail.py | 16 +- .../microsoft_purview/__init__.py | 16 +- .../guardrail_hooks/microsoft_purview/base.py | 80 +- .../microsoft_purview/purview_dlp.py | 65 +- .../guardrail_hooks/model_armor/__init__.py | 8 +- .../model_armor/file_scanning.py | 30 +- .../model_armor/model_armor.py | 93 +- .../guardrail_hooks/noma/__init__.py | 10 +- .../guardrails/guardrail_hooks/noma/noma.py | 82 +- .../guardrail_hooks/noma/noma_v2.py | 46 +- .../guardrail_hooks/onyx/__init__.py | 8 +- .../guardrails/guardrail_hooks/onyx/onyx.py | 12 +- .../guardrail_hooks/openai/__init__.py | 14 +- .../guardrail_hooks/openai/moderations.py | 30 +- .../guardrail_hooks/ovalix/__init__.py | 18 +- .../guardrail_hooks/ovalix/ovalix.py | 42 +- .../guardrail_hooks/pangea/__init__.py | 10 +- .../guardrail_hooks/pangea/pangea.py | 26 +- .../panw_prisma_airs/__init__.py | 10 +- .../panw_prisma_airs/panw_prisma_airs.py | 220 +- .../guardrail_hooks/pillar/__init__.py | 14 +- .../guardrail_hooks/pillar/pillar.py | 80 +- .../guardrails/guardrail_hooks/presidio.py | 130 +- .../prompt_security/__init__.py | 8 +- .../prompt_security/prompt_security.py | 116 +- .../guardrail_hooks/promptguard/__init__.py | 8 +- .../promptguard/promptguard.py | 45 +- .../guardrail_hooks/qohash/__init__.py | 8 +- .../guardrail_hooks/qohash/qohash.py | 8 +- .../guardrail_hooks/qualifire/__init__.py | 8 +- .../guardrail_hooks/qualifire/qualifire.py | 38 +- .../guardrail_hooks/repelloai/__init__.py | 8 +- .../guardrail_hooks/repelloai/repelloai.py | 96 +- .../guardrail_hooks/rubrik/__init__.py | 8 +- .../semantic_guard/__init__.py | 10 +- .../semantic_guard/route_loader.py | 16 +- .../semantic_guard/semantic_guard.py | 20 +- .../guardrail_hooks/singulr/__init__.py | 8 +- .../guardrail_hooks/singulr/singulr.py | 28 +- .../guardrail_hooks/straiker/__init__.py | 22 +- .../guardrail_hooks/straiker/straiker.py | 106 +- .../guardrail_hooks/tool_permission.py | 100 +- .../guardrail_hooks/tool_policy/__init__.py | 4 +- .../tool_policy/tool_policy_guardrail.py | 30 +- .../unified_guardrail/__init__.py | 4 +- .../unified_guardrail/unified_guardrail.py | 86 +- .../guardrail_hooks/vigil_guard/__init__.py | 8 +- .../vigil_guard/vigil_guard.py | 81 +- .../guardrail_hooks/xecguard/__init__.py | 8 +- .../guardrail_hooks/xecguard/xecguard.py | 115 +- .../zscaler_ai_guard/__init__.py | 8 +- .../zscaler_ai_guard/zscaler_ai_guard.py | 62 +- .../guardrails/guardrail_initializers.py | 28 +- .../proxy/guardrails/guardrail_registry.py | 104 +- litellm/proxy/guardrails/init_guardrails.py | 12 +- .../proxy/guardrails/tool_name_extraction.py | 16 +- litellm/proxy/guardrails/usage_endpoints.py | 154 +- litellm/proxy/guardrails/usage_tracking.py | 12 +- litellm/proxy/health_check.py | 105 +- .../shared_health_check_manager.py | 40 +- .../health_endpoints/_health_endpoints.py | 178 +- litellm/proxy/hooks/__init__.py | 4 +- litellm/proxy/hooks/azure_content_safety.py | 13 +- litellm/proxy/hooks/batch_rate_limiter.py | 82 +- litellm/proxy/hooks/batch_redis_get.py | 10 +- litellm/proxy/hooks/cache_control_check.py | 6 +- litellm/proxy/hooks/dynamic_rate_limiter.py | 27 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 98 +- .../proxy/hooks/key_management_event_hooks.py | 40 +- litellm/proxy/hooks/litellm_skills/main.py | 86 +- litellm/proxy/hooks/max_budget_limiter.py | 10 +- .../hooks/max_budget_per_session_limiter.py | 48 +- litellm/proxy/hooks/max_iterations_limiter.py | 32 +- .../proxy/hooks/mcp_semantic_filter/hook.py | 64 +- .../proxy/hooks/model_max_budget_limiter.py | 45 +- .../proxy/hooks/parallel_request_limiter.py | 120 +- .../hooks/parallel_request_limiter_v3.py | 383 ++- .../proxy/hooks/prompt_injection_detection.py | 20 +- .../proxy/hooks/proxy_track_cost_callback.py | 64 +- litellm/proxy/hooks/rate_limiter_utils.py | 12 +- litellm/proxy/hooks/responses_id_security.py | 38 +- litellm/proxy/hooks/sensitive_data_routing.py | 22 +- .../hooks/user_management_event_hooks.py | 7 +- litellm/proxy/image_endpoints/endpoints.py | 41 +- litellm/proxy/lambda.py | 4 +- litellm/proxy/litellm_pre_call_utils.py | 342 +-- .../callback_logs_endpoints.py | 30 +- .../access_group_endpoints.py | 58 +- .../budget_management_endpoints.py | 25 +- .../cache_settings_endpoints.py | 76 +- .../callback_management_endpoints.py | 9 +- .../common_daily_activity.py | 102 +- .../management_endpoints/common_utils.py | 46 +- .../compliance_endpoints.py | 12 +- .../config_override_endpoints.py | 72 +- .../coordination_redis_endpoints.py | 63 +- .../cost_tracking_settings.py | 52 +- .../credential_migration.py | 58 +- .../customer_endpoints.py | 59 +- .../fallback_management_endpoints.py | 28 +- .../internal_user_endpoints.py | 212 +- .../jwt_key_mapping_endpoints.py | 34 +- .../key_management_endpoints.py | 488 ++-- .../management_v1/__init__.py | 4 +- .../management_v1/budgets.py | 26 +- .../management_v1/common.py | 19 +- .../management_v1/list_framework.py | 74 +- .../management_v1/spend_logs.py | 34 +- .../mcp_management_endpoints.py | 430 +-- ...model_access_group_management_endpoints.py | 80 +- .../model_management_endpoints.py | 186 +- .../organization_endpoints.py | 151 +- .../policy_endpoints/ai_policy_suggester.py | 21 +- .../policy_endpoints/endpoints.py | 154 +- .../router_settings_endpoints.py | 24 +- .../scim/scim_transformations.py | 44 +- .../management_endpoints/scim/scim_v2.py | 331 ++- .../sso/custom_microsoft_sso.py | 11 +- .../management_endpoints/sso/saml_sso.py | 136 +- .../tag_management_endpoints.py | 91 +- .../team_callback_endpoints.py | 44 +- .../management_endpoints/team_endpoints.py | 419 ++- .../tool_management_endpoints.py | 76 +- litellm/proxy/management_endpoints/types.py | 4 +- litellm/proxy/management_endpoints/ui_sso.py | 684 +++-- .../usage_endpoints/ai_usage_chat.py | 112 +- .../usage_endpoints/endpoints.py | 8 +- .../user_agent_analytics_endpoints.py | 112 +- .../workflow_management_endpoints.py | 38 +- .../proxy/management_helpers/audit_logs.py | 19 +- .../object_permission_utils.py | 114 +- .../team_member_permission_checks.py | 26 +- .../team_metadata_validation.py | 34 +- .../management_helpers/user_invitation.py | 7 +- litellm/proxy/management_helpers/utils.py | 82 +- litellm/proxy/memory/memory_endpoints.py | 60 +- .../billable_request_metrics_middleware.py | 38 +- .../in_flight_requests_middleware.py | 4 +- .../middleware/prometheus_auth_middleware.py | 16 +- .../request_size_limit_middleware.py | 15 +- .../middleware/security_headers_middleware.py | 9 +- litellm/proxy/ocr_endpoints/endpoints.py | 18 +- .../proxy/openai_evals_endpoints/endpoints.py | 88 +- .../openai_files_endpoints/common_utils.py | 114 +- .../file_content_streaming_handler.py | 12 +- .../openai_files_endpoints/files_endpoints.py | 160 +- .../storage_backend_service.py | 28 +- .../pass_through_endpoints/common_utils.py | 4 +- .../jsonpath_extractor.py | 6 +- .../llm_passthrough_endpoints.py | 294 +-- .../anthropic_passthrough_logging_handler.py | 92 +- .../assembly_passthrough_logging_handler.py | 30 +- .../base_passthrough_logging_handler.py | 24 +- .../cohere_passthrough_logging_handler.py | 21 +- .../cursor_passthrough_logging_handler.py | 25 +- .../gemini_passthrough_logging_handler.py | 28 +- .../openai_passthrough_logging_handler.py | 87 +- ...tex_ai_live_passthrough_logging_handler.py | 54 +- .../vertex_passthrough_logging_handler.py | 88 +- .../managed_id_codec.py | 15 +- .../managed_id_rewriter.py | 118 +- .../pass_through_endpoints.py | 326 +-- .../passthrough_endpoint_router.py | 16 +- .../passthrough_guardrails.py | 24 +- .../streaming_handler.py | 21 +- .../pass_through_endpoints/success_handler.py | 22 +- .../upstream_usage_headers.py | 17 +- litellm/proxy/plugin_routes.py | 57 +- .../policy_engine/attachment_registry.py | 30 +- litellm/proxy/policy_engine/init_policies.py | 28 +- .../proxy/policy_engine/pipeline_executor.py | 12 +- .../proxy/policy_engine/policy_endpoints.py | 66 +- litellm/proxy/policy_engine/policy_matcher.py | 10 +- .../proxy/policy_engine/policy_registry.py | 114 +- .../policy_engine/policy_resolve_endpoints.py | 49 +- .../proxy/policy_engine/policy_resolver.py | 30 +- .../proxy/policy_engine/policy_validator.py | 44 +- litellm/proxy/prisma_migration.py | 6 +- litellm/proxy/prometheus_cleanup.py | 3 +- litellm/proxy/prompts/init_prompts.py | 4 +- litellm/proxy/prompts/prompt_endpoints.py | 172 +- litellm/proxy/prompts/prompt_registry.py | 21 +- litellm/proxy/proxy_cli.py | 126 +- litellm/proxy/proxy_server.py | 1683 ++++++------ .../public_endpoints/public_endpoints.py | 50 +- litellm/proxy/rag_endpoints/endpoints.py | 74 +- litellm/proxy/read_model_list.py | 4 +- litellm/proxy/realtime_endpoints/endpoints.py | 102 +- litellm/proxy/rerank_endpoints/endpoints.py | 21 +- .../proxy/response_api_endpoints/endpoints.py | 158 +- .../response_polling/background_streaming.py | 20 +- .../proxy/response_polling/polling_handler.py | 26 +- litellm/proxy/route_llm_request.py | 54 +- litellm/proxy/search_endpoints/endpoints.py | 22 +- .../search_tool_management.py | 50 +- .../search_endpoints/search_tool_registry.py | 35 +- .../shutdown/graceful_shutdown_manager.py | 15 +- .../spend_tracking/budget_reservation.py | 168 +- .../spend_tracking/cloudzero_endpoints.py | 39 +- .../spend_tracking/cold_storage_handler.py | 8 +- .../spend_tracking/compression_savings.py | 5 +- litellm/proxy/spend_tracking/savings.py | 52 +- .../spend_tracking/spend_log_error_logger.py | 4 +- .../spend_management_endpoints.py | 330 ++- .../spend_tracking/spend_tracking_utils.py | 142 +- .../proxy/spend_tracking/vantage_endpoints.py | 61 +- litellm/proxy/types_utils/utils.py | 40 +- .../proxy_setting_endpoints.py | 166 +- .../user_banner_endpoints.py | 16 +- litellm/proxy/utils.py | 681 +++-- .../proxy/vector_store_endpoints/endpoints.py | 36 +- .../management_endpoints.py | 96 +- litellm/proxy/vector_store_endpoints/utils.py | 52 +- .../vector_store_files_endpoints/endpoints.py | 84 +- .../vertex_ai_endpoints/langfuse_endpoints.py | 39 +- litellm/proxy/video_endpoints/endpoints.py | 142 +- litellm/proxy/video_endpoints/utils.py | 10 +- litellm/proxy_auth/credentials.py | 10 +- litellm/rag/ingestion/base_ingestion.py | 36 +- litellm/rag/ingestion/bedrock_ingestion.py | 114 +- litellm/rag/ingestion/gemini_ingestion.py | 64 +- litellm/rag/ingestion/openai_ingestion.py | 14 +- litellm/rag/ingestion/s3_vectors_ingestion.py | 70 +- litellm/rag/ingestion/vertex_ai_ingestion.py | 80 +- litellm/rag/main.py | 89 +- litellm/rag/rag_query.py | 12 +- .../recursive_character_text_splitter.py | 10 +- litellm/rag/utils.py | 6 +- litellm/realtime_api/main.py | 90 +- litellm/repositories/base_repository.py | 18 +- litellm/repositories/budget_repository.py | 6 +- litellm/repositories/config_repository.py | 24 +- .../repositories/credentials_repository.py | 6 +- litellm/repositories/model_repository.py | 30 +- .../object_permission_repository.py | 6 +- .../repositories/organization_repository.py | 8 +- litellm/repositories/project_repository.py | 8 +- litellm/repositories/team_repository.py | 50 +- .../repositories/user_banner_repository.py | 8 +- litellm/repositories/user_repository.py | 20 +- .../verification_token_repository.py | 46 +- litellm/rerank_api/main.py | 52 +- litellm/rerank_api/rerank_utils.py | 4 +- .../responses/file_search/emulated_handler.py | 76 +- .../custom_tools.py | 28 +- .../handler.py | 18 +- .../session_handler.py | 28 +- .../streaming_iterator.py | 70 +- .../transformation.py | 160 +- litellm/responses/main.py | 266 +- .../responses/mcp/chat_completions_handler.py | 61 +- .../mcp/litellm_proxy_mcp_handler.py | 103 +- .../responses/mcp/mcp_streaming_iterator.py | 82 +- litellm/responses/mcp/request_context.py | 8 +- litellm/responses/sse_output_recovery.py | 20 +- litellm/responses/streaming_iterator.py | 244 +- litellm/responses/utils.py | 179 +- litellm/router.py | 1294 +++++---- .../adaptive_router/adaptive_router.py | 72 +- .../router_strategy/adaptive_router/bandit.py | 25 +- .../adaptive_router/classifier.py | 5 +- .../router_strategy/adaptive_router/config.py | 36 +- .../router_strategy/adaptive_router/hooks.py | 42 +- .../adaptive_router/signals.py | 26 +- .../adaptive_router/update_queue.py | 12 +- .../auto_router/auto_router.py | 10 +- .../auto_router/litellm_encoder.py | 10 +- .../router_strategy/base_routing_strategy.py | 27 +- litellm/router_strategy/budget_limiter.py | 74 +- .../complexity_router/complexity_router.py | 260 +- .../complexity_router/config.py | 36 +- .../evals/eval_complexity_router.py | 14 +- litellm/router_strategy/lar1_routing.py | 34 +- litellm/router_strategy/least_busy.py | 47 +- litellm/router_strategy/lowest_cost.py | 61 +- litellm/router_strategy/lowest_latency.py | 100 +- litellm/router_strategy/lowest_tpm_rpm.py | 33 +- litellm/router_strategy/lowest_tpm_rpm_v2.py | 98 +- .../router_strategy/quality_router/config.py | 4 +- .../quality_router/quality_router.py | 48 +- litellm/router_strategy/simple_shuffle.py | 4 +- litellm/router_strategy/tag_based_routing.py | 62 +- .../add_retry_fallback_headers.py | 46 +- .../router_utils/auto_router_model_naming.py | 22 +- litellm/router_utils/batch_utils.py | 9 +- .../client_initalization_utils.py | 18 +- .../clientside_credential_handler.py | 10 +- litellm/router_utils/common_utils.py | 44 +- litellm/router_utils/cooldown_cache.py | 24 +- litellm/router_utils/cooldown_callbacks.py | 18 +- litellm/router_utils/cooldown_handlers.py | 38 +- .../router_utils/fallback_event_handlers.py | 6 +- litellm/router_utils/handle_error.py | 12 +- litellm/router_utils/health_state_cache.py | 8 +- .../router_utils/pattern_match_deployments.py | 19 +- .../deployment_affinity_check.py | 52 +- .../encrypted_content_affinity_check.py | 32 +- .../io_token_rate_limit_check.py | 102 +- .../pre_call_checks/model_rate_limit_check.py | 80 +- .../prompt_caching_deployment_check.py | 20 +- .../responses_api_deployment_check.py | 7 +- litellm/router_utils/prompt_caching_cache.py | 24 +- .../track_deployment_metrics.py | 10 +- litellm/router_utils/search_api_router.py | 28 +- litellm/rust_bridge/loader.py | 3 +- litellm/rust_bridge/messages.py | 8 +- litellm/rust_bridge/ocr.py | 16 +- litellm/rust_bridge/responses_websocket.py | 8 +- litellm/rust_bridge/transcription.py | 10 +- litellm/sandbox/main.py | 14 +- litellm/scheduler.py | 17 +- litellm/search/cost_calculator.py | 10 +- litellm/search/main.py | 38 +- litellm/secret_managers/aws_secret_manager.py | 20 +- .../secret_managers/aws_secret_manager_v2.py | 54 +- .../secret_managers/base_secret_manager.py | 10 +- .../custom_secret_manager_loader.py | 15 +- .../cyberark_secret_manager.py | 52 +- .../get_azure_ad_token_provider.py | 6 +- litellm/secret_managers/google_kms.py | 3 +- .../secret_managers/google_secret_manager.py | 19 +- .../hashicorp_secret_manager.py | 112 +- litellm/secret_managers/main.py | 59 +- .../secret_managers/secret_manager_handler.py | 14 +- litellm/setup_wizard.py | 74 +- litellm/skills/main.py | 116 +- litellm/timeout.py | 13 +- litellm/types/agents.py | 4 +- litellm/types/caching.py | 2 +- litellm/types/completion.py | 11 +- litellm/types/files.py | 18 +- litellm/types/guardrails.py | 8 +- .../anthropic_cache_control_hook.py | 2 +- litellm/types/integrations/argilla.py | 4 +- litellm/types/integrations/cloudzero.py | 2 +- litellm/types/integrations/custom_logger.py | 12 +- litellm/types/integrations/datadog.py | 6 +- litellm/types/integrations/gcs_bucket.py | 8 +- litellm/types/integrations/posthog.py | 4 +- litellm/types/integrations/prometheus.py | 48 +- litellm/types/integrations/slack_alerting.py | 18 +- litellm/types/interactions/generated.py | 4 +- litellm/types/llms/anthropic.py | 18 +- litellm/types/llms/anthropic_tool_search.py | 10 +- litellm/types/llms/azure.py | 6 +- litellm/types/llms/azure_ai.py | 2 +- litellm/types/llms/base.py | 4 +- litellm/types/llms/bedrock.py | 8 +- litellm/types/llms/bedrock_invoke_agents.py | 6 +- litellm/types/llms/cohere.py | 2 +- litellm/types/llms/custom_http.py | 2 +- litellm/types/llms/databricks.py | 2 +- litellm/types/llms/gemini.py | 2 +- litellm/types/llms/oci.py | 2 +- litellm/types/llms/openai.py | 23 +- litellm/types/llms/openai_evals.py | 2 +- litellm/types/llms/stability.py | 10 +- litellm/types/llms/vertex_ai.py | 4 +- .../cache_settings_endpoints.py | 6 +- .../coordination_redis_endpoints.py | 4 +- .../router_settings_endpoints.py | 6 +- litellm/types/mcp.py | 6 +- .../types/mcp_server/mcp_server_manager.py | 4 +- .../passthrough_endpoints/assembly_ai.py | 6 +- .../pass_through_endpoints.py | 8 +- .../azure/azure_text_moderation.py | 4 +- .../proxy/guardrails/guardrail_hooks/base.py | 2 +- .../guardrail_hooks/bedrock_guardrails.py | 2 +- .../guardrail_hooks/block_code_execution.py | 4 +- .../guardrail_hooks/cisco_ai_defense.py | 2 +- .../guardrail_hooks/generic_guardrail_api.py | 6 +- .../guardrail_hooks/litellm_content_filter.py | 2 +- .../guardrails/guardrail_hooks/straiker.py | 4 +- .../guardrail_hooks/tool_permission.py | 4 +- .../guardrails/guardrail_hooks/xecguard.py | 4 +- .../guardrail_hooks/zscaler_ai_guard.py | 10 +- .../internal_user_endpoints.py | 10 +- .../key_management_endpoints.py | 4 +- .../management_endpoints/management_v1.py | 2 +- .../proxy/management_endpoints/scim_v2.py | 28 +- .../management_endpoints/team_endpoints.py | 2 +- .../proxy/policy_engine/pipeline_types.py | 6 +- litellm/types/rag.py | 2 +- litellm/types/realtime.py | 2 +- litellm/types/responses/main.py | 6 +- litellm/types/router.py | 23 +- litellm/types/search.py | 2 +- litellm/types/services.py | 4 +- litellm/types/tool_management.py | 2 +- litellm/types/utils.py | 55 +- litellm/types/vector_stores.py | 2 +- litellm/types/videos/utils.py | 44 +- litellm/utils.py | 760 +++--- litellm/vector_store_files/main.py | 160 +- litellm/vector_store_files/utils.py | 8 +- litellm/vector_stores/main.py | 164 +- litellm/vector_stores/utils.py | 12 +- .../vector_stores/vector_store_registry.py | 38 +- litellm/videos/main.py | 254 +- litellm/videos/utils.py | 18 +- scripts/check_type_discipline.py | 349 ++- scripts/type_discipline_gate.py | 50 +- .../test_check_type_discipline.py | 270 +- .../test_litellm/test_type_discipline_gate.py | 13 + type-discipline-budget.json | 6 + 1587 files changed, 37473 insertions(+), 36659 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index 62c41c2959b..7cc37a3d864 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -27,18 +27,19 @@ if os.getenv("LITELLM_MODE", "DEV") == "DEV": _dotenv.load_dotenv(override=_dev_env_hot_reload_enabled()) from typing import ( - Callable, - List, - Optional, - Dict, - Union, Any, - Literal, + Callable, + Dict, + Final, get_args, - TYPE_CHECKING, - Tuple, + List, + Literal, + Optional, overload, + Tuple, Type, + TYPE_CHECKING, + Union, ) from litellm.types.integrations.datadog import DatadogInitParams from litellm.types.integrations.newrelic import NewRelicInitParams @@ -94,7 +95,7 @@ import httpx # register_async_client_cleanup is lazy-loaded and called on first access -litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV" +litellm_mode: Final = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV" #################################################### @@ -106,7 +107,7 @@ CALLBACK_TYPES = Union[str, Callable, "CustomLogger"] # CustomLogger is lazy-lo input_callback: List[CALLBACK_TYPES] = [] success_callback: List[CALLBACK_TYPES] = [] failure_callback: List[CALLBACK_TYPES] = [] -service_callback: List[CALLBACK_TYPES] = [] +service_callback: Final[List[CALLBACK_TYPES]] = [] audit_log_callbacks: List[CALLBACK_TYPES] = [] # logging_callback_manager is lazy-loaded via __getattr__ _custom_logger_compatible_callbacks_literal = Literal[ @@ -162,25 +163,25 @@ _custom_logger_compatible_callbacks_literal = Literal[ "compression_interception", "newrelic", ] -cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None -logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None -_known_custom_logger_compatible_callbacks: List = list(get_args(_custom_logger_compatible_callbacks_literal)) +cold_storage_custom_logger: Final[Optional[_custom_logger_compatible_callbacks_literal]] = None +logged_real_time_event_types: Final[Optional[Union[List[str], Literal["*"]]]] = None +_known_custom_logger_compatible_callbacks: Final[List] = list(get_args(_custom_logger_compatible_callbacks_literal)) callbacks: List[ Union[Callable, _custom_logger_compatible_callbacks_literal, "CustomLogger"] # CustomLogger is lazy-loaded ] = [] callback_settings: Dict[str, Dict[str, Any]] = {} initialized_langfuse_clients: int = 0 -langfuse_default_tags: Optional[List[str]] = None -langsmith_batch_size: Optional[int] = None -prometheus_initialize_budget_metrics: Optional[bool] = False -prometheus_latency_buckets: Optional[List[float]] = None -require_auth_for_metrics_endpoint: Optional[bool] = True -argilla_batch_size: Optional[int] = None +langfuse_default_tags: Final[Optional[List[str]]] = None +langsmith_batch_size: Final[Optional[int]] = None +prometheus_initialize_budget_metrics: Final[Optional[bool]] = False +prometheus_latency_buckets: Final[Optional[List[float]]] = None +require_auth_for_metrics_endpoint: Final[Optional[bool]] = True +argilla_batch_size: Final[Optional[int]] = None datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload. -gcs_pub_sub_use_v1: Optional[bool] = False # if you want to use v1 gcs pubsub logged payload -generic_api_use_v1: Optional[bool] = False # if you want to use v1 generic api logged payload -argilla_transformation_object: Optional[Dict[str, Any]] = None -_async_input_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded +gcs_pub_sub_use_v1: Final[Optional[bool]] = False # if you want to use v1 gcs pubsub logged payload +generic_api_use_v1: Final[Optional[bool]] = False # if you want to use v1 generic api logged payload +argilla_transformation_object: Final[Optional[Dict[str, Any]]] = None +_async_input_callback: Final[List[Union[str, Callable, "CustomLogger"]]] = ( # CustomLogger is lazy-loaded [] ) # internal variable - async custom callbacks are routed here. _async_success_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded @@ -189,13 +190,13 @@ _async_success_callback: List[Union[str, Callable, "CustomLogger"]] = ( # Custo _async_failure_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded [] ) # internal variable - async custom callbacks are routed here. -pre_call_rules: List[Callable] = [] +pre_call_rules: Final[List[Callable]] = [] post_call_rules: List[Callable] = [] turn_off_message_logging: Optional[bool] = False -standard_logging_payload_excluded_fields: Optional[List[str]] = ( +standard_logging_payload_excluded_fields: Final[Optional[List[str]]] = ( None # Fields to exclude from StandardLoggingPayload before callbacks receive it ) -log_raw_request_response: bool = False +log_raw_request_response: Final[bool] = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False # When True (default — preserves historical behavior), the Router appends @@ -206,27 +207,27 @@ redact_user_api_key_info: Optional[bool] = False # Deprecation: planned to flip to False (redact by default) in a future # major release; opt in early with `litellm.expose_router_debug_in_errors # = False`. -expose_router_debug_in_errors: bool = True -filter_invalid_headers: Optional[bool] = False -add_user_information_to_llm_headers: Optional[bool] = ( +expose_router_debug_in_errors: Final[bool] = True +filter_invalid_headers: Final[Optional[bool]] = False +add_user_information_to_llm_headers: Final[Optional[bool]] = ( None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers ) -overwrite_user_with_key_hash: bool = ( +overwrite_user_with_key_hash: Final[bool] = ( False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id ) store_audit_logs = False # Enterprise feature, allow users to see audit logs -skip_system_message_in_guardrail: bool = False -skip_tool_message_in_guardrail: bool = False +skip_system_message_in_guardrail: Final[bool] = False +skip_tool_message_in_guardrail: Final[bool] = False ### end of callbacks ############# -email: Optional[str] = ( +email: Final[Optional[str]] = ( None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 ) -token: Optional[str] = ( +token: Final[Optional[str]] = ( None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 ) -telemetry = True -max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults +telemetry: Final = True +max_tokens: Final[int] = DEFAULT_MAX_TOKENS # OpenAI Defaults drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False)) modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False)) use_chat_completions_url_for_anthropic_messages: bool = bool( @@ -243,7 +244,7 @@ use_chat_completions_url_for_anthropic_messages: bool = bool( # Or via `litellm_settings.strip_anthropic_total_tokens: true` in # config.yaml. strip_anthropic_total_tokens: bool = False -route_all_chat_openai_to_responses: bool = ( +route_all_chat_openai_to_responses: Final[bool] = ( os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true" ) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge # When True, Gemini/Vertex Live setup is deferred until client `session.update`. @@ -253,126 +254,126 @@ use_legacy_interactions_schema: bool = ( os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true" ) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs` # schema instead of the new `steps` schema. Remove this flag after June 8, 2026. -retry = True +retry: Final = True ### AUTH ### -api_key: Optional[str] = None -openai_key: Optional[str] = None -groq_key: Optional[str] = None -gigachat_key: Optional[str] = None -xai_key: Optional[str] = None -databricks_key: Optional[str] = None -openai_like_key: Optional[str] = None -azure_key: Optional[str] = None -anthropic_key: Optional[str] = None -autorouter_savings_baseline_model: Optional[str] = None -replicate_key: Optional[str] = None -bytez_key: Optional[str] = None -gdc_key: Optional[str] = None -gdc_api_base: Optional[str] = None -cohere_key: Optional[str] = None -infinity_key: Optional[str] = None -clarifai_key: Optional[str] = None -maritalk_key: Optional[str] = None -ai21_key: Optional[str] = None -ollama_key: Optional[str] = None -openrouter_key: Optional[str] = None -datarobot_key: Optional[str] = None -predibase_key: Optional[str] = None -huggingface_key: Optional[str] = None -vertex_project: Optional[str] = None -vertex_location: Optional[str] = None -predibase_tenant_id: Optional[str] = None -togetherai_api_key: Optional[str] = None -cloudflare_api_key: Optional[str] = None -vercel_ai_gateway_key: Optional[str] = None -baseten_key: Optional[str] = None -llama_api_key: Optional[str] = None -aleph_alpha_key: Optional[str] = None -nlp_cloud_key: Optional[str] = None -novita_api_key: Optional[str] = None -snowflake_key: Optional[str] = None -gradient_ai_api_key: Optional[str] = None -nebius_key: Optional[str] = None -wandb_key: Optional[str] = None -heroku_key: Optional[str] = None -cometapi_key: Optional[str] = None -ovhcloud_key: Optional[str] = None -lemonade_key: Optional[str] = None -sap_service_key: Optional[str] = None -amazon_nova_api_key: Optional[str] = None -inception_key: Optional[str] = None -common_cloud_provider_auth_params: dict = { +api_key: Final[Optional[str]] = None +openai_key: Final[Optional[str]] = None +groq_key: Final[Optional[str]] = None +gigachat_key: Final[Optional[str]] = None +xai_key: Final[Optional[str]] = None +databricks_key: Final[Optional[str]] = None +openai_like_key: Final[Optional[str]] = None +azure_key: Final[Optional[str]] = None +anthropic_key: Final[Optional[str]] = None +autorouter_savings_baseline_model: Final[Optional[str]] = None +replicate_key: Final[Optional[str]] = None +bytez_key: Final[Optional[str]] = None +gdc_key: Final[Optional[str]] = None +gdc_api_base: Final[Optional[str]] = None +cohere_key: Final[Optional[str]] = None +infinity_key: Final[Optional[str]] = None +clarifai_key: Final[Optional[str]] = None +maritalk_key: Final[Optional[str]] = None +ai21_key: Final[Optional[str]] = None +ollama_key: Final[Optional[str]] = None +openrouter_key: Final[Optional[str]] = None +datarobot_key: Final[Optional[str]] = None +predibase_key: Final[Optional[str]] = None +huggingface_key: Final[Optional[str]] = None +vertex_project: Final[Optional[str]] = None +vertex_location: Final[Optional[str]] = None +predibase_tenant_id: Final[Optional[str]] = None +togetherai_api_key: Final[Optional[str]] = None +cloudflare_api_key: Final[Optional[str]] = None +vercel_ai_gateway_key: Final[Optional[str]] = None +baseten_key: Final[Optional[str]] = None +llama_api_key: Final[Optional[str]] = None +aleph_alpha_key: Final[Optional[str]] = None +nlp_cloud_key: Final[Optional[str]] = None +novita_api_key: Final[Optional[str]] = None +snowflake_key: Final[Optional[str]] = None +gradient_ai_api_key: Final[Optional[str]] = None +nebius_key: Final[Optional[str]] = None +wandb_key: Final[Optional[str]] = None +heroku_key: Final[Optional[str]] = None +cometapi_key: Final[Optional[str]] = None +ovhcloud_key: Final[Optional[str]] = None +lemonade_key: Final[Optional[str]] = None +sap_service_key: Final[Optional[str]] = None +amazon_nova_api_key: Final[Optional[str]] = None +inception_key: Final[Optional[str]] = None +common_cloud_provider_auth_params: Final[dict] = { "params": ["project", "region_name", "token"], "providers": ["vertex_ai", "bedrock", "watsonx", "azure", "vertex_ai_beta"], } -use_litellm_proxy: bool = False # when True, requests will be sent to the specified litellm proxy endpoint -use_client: bool = False +use_litellm_proxy: Final[bool] = False # when True, requests will be sent to the specified litellm proxy endpoint +use_client: Final[bool] = False ssl_verify: Union[str, bool] = True -ssl_security_level: Optional[str] = None -ssl_certificate: Optional[str] = None +ssl_security_level: Final[Optional[str]] = None +ssl_certificate: Final[Optional[str]] = None user_url_validation: bool = True user_url_allowed_hosts: List[str] = [] provider_url_destination_allowed_hosts: List[str] = [] -ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance +ssl_ecdh_curve: Final[Optional[str]] = None # Set to 'X25519' to disable PQC and improve performance disable_streaming_logging: bool = False -disable_token_counter: bool = False -disable_add_transform_inline_image_block: bool = False -disable_add_user_agent_to_request_tags: bool = False +disable_token_counter: Final[bool] = False +disable_add_transform_inline_image_block: Final[bool] = False +disable_add_user_agent_to_request_tags: Final[bool] = False disable_anthropic_gemini_context_caching_transform: bool = False enable_anthropic_prompt_caching: bool = os.getenv("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", "false").lower() == "true" -_anthropic_prompt_caching_ttl_env: Optional[str] = os.getenv("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL") -anthropic_prompt_caching_ttl: Optional[Literal["5m", "1h"]] = ( +_anthropic_prompt_caching_ttl_env: Final[Optional[str]] = os.getenv("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL") +anthropic_prompt_caching_ttl: Final[Optional[Literal["5m", "1h"]]] = ( "1h" if _anthropic_prompt_caching_ttl_env == "1h" else "5m" if _anthropic_prompt_caching_ttl_env == "5m" else None ) disable_vertex_batch_output_transformation: bool = False -extra_spend_tag_headers: Optional[List[str]] = None +extra_spend_tag_headers: Final[Optional[List[str]]] = None in_memory_llm_clients_cache: "LLMClientCache" -safe_memory_mode: bool = False -enable_azure_ad_token_refresh: Optional[bool] = False +safe_memory_mode: Final[bool] = False +enable_azure_ad_token_refresh: Final[Optional[bool]] = False # Proxy Authentication - auto-obtain/refresh OAuth2/JWT tokens for LiteLLM Proxy proxy_auth: Optional[Any] = None ### DEFAULT AZURE API VERSION ### -AZURE_DEFAULT_API_VERSION = "2025-02-01-preview" # this is updated to the latest +AZURE_DEFAULT_API_VERSION: Final = "2025-02-01-preview" # this is updated to the latest ### DEFAULT WATSONX API VERSION ### -WATSONX_DEFAULT_API_VERSION = "2024-03-13" +WATSONX_DEFAULT_API_VERSION: Final = "2024-03-13" ### COHERE EMBEDDINGS DEFAULT TYPE ### -COHERE_DEFAULT_EMBEDDING_INPUT_TYPE: "COHERE_EMBEDDING_INPUT_TYPES" = "search_document" +COHERE_DEFAULT_EMBEDDING_INPUT_TYPE: Final["COHERE_EMBEDDING_INPUT_TYPES"] = "search_document" ### CREDENTIALS ### credential_list: List["CredentialItem"] = [] ### GUARDRAILS ### -llamaguard_model_name: Optional[str] = None -openai_moderations_model_name: Optional[str] = None -presidio_ad_hoc_recognizers: Optional[str] = None -google_moderation_confidence_threshold: Optional[float] = None -llamaguard_unsafe_content_categories: Optional[str] = None -blocked_user_list: Optional[Union[str, List]] = None -banned_keywords_list: Optional[Union[str, List]] = None -llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all" +llamaguard_model_name: Final[Optional[str]] = None +openai_moderations_model_name: Final[Optional[str]] = None +presidio_ad_hoc_recognizers: Final[Optional[str]] = None +google_moderation_confidence_threshold: Final[Optional[float]] = None +llamaguard_unsafe_content_categories: Final[Optional[str]] = None +blocked_user_list: Final[Optional[Union[str, List]]] = None +banned_keywords_list: Final[Optional[Union[str, List]]] = None +llm_guard_mode: Final[Literal["all", "key-specific", "request-specific"]] = "all" guardrail_name_config_map: Dict[str, GuardrailItem] = {} -include_cost_in_streaming_usage: bool = False -reasoning_auto_summary: bool = False +include_cost_in_streaming_usage: Final[bool] = False +reasoning_auto_summary: Final[bool] = False ### PROMPTS #### from litellm.types.prompts.init_prompts import PromptSpec -prompt_name_config_map: Dict[str, PromptSpec] = {} +prompt_name_config_map: Final[Dict[str, PromptSpec]] = {} ################## ### PREVIEW FEATURES ### -enable_preview_features: bool = False +enable_preview_features: Final[bool] = False return_response_headers: bool = False # get response headers from LLM Api providers - example x-remaining-requests, -enable_json_schema_validation: bool = False +enable_json_schema_validation: Final[bool] = False enable_model_config_credential_overrides: bool = False -enable_key_alias_format_validation: bool = ( +enable_key_alias_format_validation: Final[bool] = ( False # opt-in validation of key_alias format on /key/generate and /key/update ) -enable_gemini_default_thinking_level_low: bool = ( +enable_gemini_default_thinking_level_low: Final[bool] = ( False # opt-in: force thinkingLevel low/minimal for Gemini 3 thinking param mapping ) #################### -logging: bool = True -enable_loadbalancing_on_batch_endpoints: Optional[bool] = None -require_managed_files: bool = False # proxy only - require target_model_names on POST /v1/files -enable_caching_on_provider_specific_optional_params: bool = ( +logging: Final[bool] = True +enable_loadbalancing_on_batch_endpoints: Final[Optional[bool]] = None +require_managed_files: Final[bool] = False # proxy only - require target_model_names on POST /v1/files +enable_caching_on_provider_specific_optional_params: Final[bool] = ( False # feature-flag for caching on optional params - e.g. 'top_k' ) caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 @@ -380,91 +381,91 @@ caching_with_models: bool = False # # Not used anymore, will be removed in next cache: Optional["Cache"] = None # cache object <- use this - https://docs.litellm.ai/docs/caching default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None -default_redis_batch_cache_expiry: Optional[float] = None -model_alias_map: Dict[str, str] = {} +default_redis_batch_cache_expiry: Final[Optional[float]] = None +model_alias_map: Final[Dict[str, str]] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers -budget_duration: Optional[str] = ( +budget_duration: Final[Optional[str]] = ( None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). ) -default_soft_budget: float = DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 -budget_exceeded_throttle_percentage: Optional[float] = None -forward_traceparent_to_llm_provider: bool = False +default_soft_budget: Final[float] = DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 +budget_exceeded_throttle_percentage: Final[Optional[float]] = None +forward_traceparent_to_llm_provider: Final[bool] = False _current_cost = 0.0 # private variable, used if max budget is set -error_logs: Dict = {} +error_logs: Final[Dict] = {} add_function_to_prompt: bool = ( False # if function calling not supported by api, append function call details to system prompt ) -client_session: Optional[httpx.Client] = None -aclient_session: Optional[httpx.AsyncClient] = None -model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks' -model_cost_map_url: str = os.getenv( +client_session: Final[Optional[httpx.Client]] = None +aclient_session: Final[Optional[httpx.AsyncClient]] = None +model_fallbacks: Final[Optional[List]] = None # Deprecated for 'litellm.fallbacks' +model_cost_map_url: Final[str] = os.getenv( "LITELLM_MODEL_COST_MAP_URL", "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json", ) -blog_posts_url: str = os.getenv( +blog_posts_url: Final[str] = os.getenv( "LITELLM_BLOG_POSTS_URL", "https://docs.litellm.ai/blog/rss.xml", ) -anthropic_beta_headers_url: str = os.getenv( +anthropic_beta_headers_url: Final[str] = os.getenv( "LITELLM_ANTHROPIC_BETA_HEADERS_URL", "https://raw.githubusercontent.com/BerriAI/litellm/main/litellm/anthropic_beta_headers_config.json", ) suppress_debug_info: bool = False dynamodb_table_name: Optional[str] = None -s3_callback_params: Optional[Dict] = None -s3_audit_callback_params: Optional[Dict] = None -datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None -datadog_params: Optional[Union[DatadogInitParams, Dict]] = None -newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None +s3_callback_params: Final[Optional[Dict]] = None +s3_audit_callback_params: Final[Optional[Dict]] = None +datadog_llm_observability_params: Final[Optional[Union[DatadogLLMObsInitParams, Dict]]] = None +datadog_params: Final[Optional[Union[DatadogInitParams, Dict]]] = None +newrelic_params: Final[Optional[Union[NewRelicInitParams, Dict]]] = None aws_sqs_callback_params: Optional[Dict] = None -generic_logger_headers: Optional[Dict] = None -default_key_generate_params: Optional[Dict] = None -default_key_max_budget_alert_emails: Optional[Dict[str, list]] = None +generic_logger_headers: Final[Optional[Dict]] = None +default_key_generate_params: Final[Optional[Dict]] = None +default_key_max_budget_alert_emails: Final[Optional[Dict[str, list]]] = None upperbound_key_generate_params: Optional[LiteLLM_UpperboundKeyGenerateParams] = None -key_generation_settings: Optional["StandardKeyGenerationConfig"] = None +key_generation_settings: Final[Optional["StandardKeyGenerationConfig"]] = None default_internal_user_params: Optional[Dict] = None -default_team_params: Optional[Union[DefaultTeamSSOParams, Dict]] = None -default_team_settings: Optional[List] = None -max_user_budget: Optional[float] = None +default_team_params: Final[Optional[Union[DefaultTeamSSOParams, Dict]]] = None +default_team_settings: Final[Optional[List]] = None +max_user_budget: Final[Optional[float]] = None default_max_internal_user_budget: Optional[float] = None max_internal_user_budget: Optional[float] = None max_ui_session_budget: Optional[float] = ( 1.0 # USD budget for each dashboard login session (playground, test connection) ) -internal_user_budget_duration: Optional[str] = None -tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None -max_end_user_budget: Optional[float] = None -max_end_user_budget_id: Optional[str] = None +internal_user_budget_duration: Final[Optional[str]] = None +tag_budget_config: Final[Optional[Dict[str, "BudgetConfig"]]] = None +max_end_user_budget: Final[Optional[float]] = None +max_end_user_budget_id: Final[Optional[str]] = None # When True, end-user IDs extracted from requests are validated against # LiteLLM_EndUserTable / LiteLLM_UserTable. Values that do not resolve to a # known row are dropped before reaching spend logs. Defaults to False for # backwards compatibility — arbitrary client-supplied identifiers still # pass through unchanged. -validate_end_user_id_in_db: bool = False -disable_end_user_cost_tracking: Optional[bool] = None -disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None -enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None -custom_prometheus_metadata_labels: List[str] = [] -custom_prometheus_tags: List[str] = [] -prometheus_metrics_config: Optional[List] = None -prometheus_exclude_metrics: Optional[List[str]] = None -prometheus_exclude_labels: Optional[List[str]] = None -prometheus_emit_stream_label: bool = False +validate_end_user_id_in_db: Final[bool] = False +disable_end_user_cost_tracking: Final[Optional[bool]] = None +disable_end_user_cost_tracking_prometheus_only: Final[Optional[bool]] = None +enable_end_user_cost_tracking_prometheus_only: Final[Optional[bool]] = None +custom_prometheus_metadata_labels: Final[List[str]] = [] +custom_prometheus_tags: Final[List[str]] = [] +prometheus_metrics_config: Final[Optional[List]] = None +prometheus_exclude_metrics: Final[Optional[List[str]]] = None +prometheus_exclude_labels: Final[Optional[List[str]]] = None +prometheus_emit_stream_label: Final[bool] = False # Opt-in: emit `rate_limit_category` and `rate_limit_type` labels on # `litellm_proxy_failed_requests_metric`. Off by default to preserve the # pre-unification label set so existing dashboards / recording rules keyed on # that metric keep matching after upgrade. Enable when downstream consumers # are ready to split 429s by source (vendor vs. litellm) and dimension # (RPM/TPM/concurrent/budget). -prometheus_emit_rate_limit_labels: bool = False -prometheus_user_budget_label_include_email_alias: bool = False -prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000 -prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0 -prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0 -disable_add_prefix_to_prompt: bool = False # used by anthropic, to disable adding prefix to prompt +prometheus_emit_rate_limit_labels: Final[bool] = False +prometheus_user_budget_label_include_email_alias: Final[bool] = False +prometheus_end_user_metrics_max_series_per_metric: Final[Optional[int]] = 10000 +prometheus_end_user_metrics_ttl_seconds: Final[Optional[float]] = 3600.0 +prometheus_end_user_metrics_cleanup_interval_seconds: Final[Optional[float]] = 60.0 +disable_add_prefix_to_prompt: Final[bool] = False # used by anthropic, to disable adding prefix to prompt disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. public_mcp_servers: Optional[List[str]] = None public_mcp_hub_strict_whitelist: bool = True @@ -475,7 +476,7 @@ public_agent_groups: Optional[List[str]] = None # Old format: { "displayName": "url" } (for backward compatibility) public_model_groups_links: Dict[str, Union[str, Dict[str, Any]]] = {} #### REQUEST PRIORITIZATION ####### -priority_reservation: Optional[Dict[str, Union[float, "PriorityReservationDict"]]] = None +priority_reservation: Final[Optional[Dict[str, Union[float, "PriorityReservationDict"]]]] = None # priority_reservation_settings is lazy-loaded via __getattr__ # Only declare for type checking - at runtime __getattr__ handles it if TYPE_CHECKING: @@ -484,25 +485,25 @@ if TYPE_CHECKING: ######## Networking Settings ######## use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. -aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings +aiohttp_trust_env: Final[bool] = False # set to true to use HTTP_ Proxy settings disable_aiohttp_transport: bool = False # Set this to true to use httpx instead -disable_aiohttp_trust_env: bool = False # When False, aiohttp will respect HTTP(S)_PROXY env vars +disable_aiohttp_trust_env: Final[bool] = False # When False, aiohttp will respect HTTP(S)_PROXY env vars force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. -network_mock: bool = False # When True, use mock transport — no real network calls +network_mock: Final[bool] = False # When True, use mock transport — no real network calls ####### STOP SEQUENCE LIMIT ####### -disable_stop_sequence_limit: bool = False # when True, stop sequence limit is disabled +disable_stop_sequence_limit: Final[bool] = False # when True, stop sequence limit is disabled #### RETRIES #### num_retries: Optional[int] = None # per model endpoint -max_fallbacks: Optional[int] = None -default_fallbacks: Optional[List] = None -fallbacks: Optional[List] = None -context_window_fallbacks: Optional[List] = None -content_policy_fallbacks: Optional[List] = None -allowed_fails: int = 3 -allow_dynamic_callback_disabling: bool = True -num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries) +max_fallbacks: Final[Optional[int]] = None +default_fallbacks: Final[Optional[List]] = None +fallbacks: Final[Optional[List]] = None +context_window_fallbacks: Final[Optional[List]] = None +content_policy_fallbacks: Final[Optional[List]] = None +allowed_fails: Final[int] = 3 +allow_dynamic_callback_disabling: Final[bool] = True +num_retries_per_request: Final[Optional[int]] = None # for the request overall (incl. fallbacks + model retries) ####### SECRET MANAGERS ##################### secret_manager_client: Optional[Any] = ( None # list of instantiated key management clients - e.g. azure kv, infisical, etc. @@ -514,7 +515,7 @@ _key_management_system: Optional["KeyManagementSystem"] = None # We'll import it after the lazy import system is set up # We can't define it here because KeyManagementSettings is lazy-loaded #### PII MASKING #### -output_parse_pii: bool = False +output_parse_pii: Final[bool] = False ############################################# from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map @@ -527,8 +528,8 @@ cost_margin_config: Dict[ # Fixed: {"openai": {"fixed_amount": 0.001}} = $0.001 per request # Global: {"global": 0.05} = 5% global margin on all providers # Combined: {"vertex_ai": {"percentage": 0.08, "fixed_amount": 0.0005}} -custom_prompt_dict: Dict[str, dict] = {} -check_provider_endpoint = False +custom_prompt_dict: Final[Dict[str, dict]] = {} +check_provider_endpoint: Final = False ####### THREAD-SPECIFIC DATA #################### @@ -537,7 +538,7 @@ class MyLocal(threading.local): self.user = "Hello World" -_thread_context = MyLocal() +_thread_context: Final = MyLocal() def identify(event_details): @@ -548,126 +549,126 @@ def identify(event_details): ####### ADDITIONAL PARAMS ################### configurable params if you use proxy models like Helicone, map spend to org id, etc. api_base: Optional[str] = None -headers = None +headers: Final = None api_version: Optional[str] = None -organization = None -project = None +organization: Final = None +project: Final = None config_path = None -vertex_ai_safety_settings: Optional[dict] = None +vertex_ai_safety_settings: Final[Optional[dict]] = None ####### COMPLETION MODELS ################### from typing import Set -open_ai_chat_completion_models: Set = set() -open_ai_text_completion_models: Set = set() -cohere_models: Set = set() -cohere_chat_models: Set = set() -mistral_chat_models: Set = set() -text_completion_codestral_models: Set = set() -text_completion_inception_models: Set = set() -anthropic_models: Set = set() -openrouter_models: Set = set() -datarobot_models: Set = set() -vertex_language_models: Set = set() -vertex_vision_models: Set = set() -vertex_chat_models: Set = set() -vertex_code_chat_models: Set = set() -vertex_ai_image_models: Set = set() -vertex_ai_video_models: Set = set() -vertex_text_models: Set = set() -vertex_code_text_models: Set = set() -vertex_embedding_models: Set = set() -vertex_anthropic_models: Set = set() -vertex_llama3_models: Set = set() -vertex_deepseek_models: Set = set() -vertex_ai_ai21_models: Set = set() -vertex_mistral_models: Set = set() -vertex_openai_models: Set = set() -vertex_minimax_models: Set = set() -vertex_moonshot_models: Set = set() -vertex_zai_models: Set = set() -ai21_models: Set = set() -ai21_chat_models: Set = set() -nlp_cloud_models: Set = set() -aleph_alpha_models: Set = set() -bedrock_models: Set = set() -bedrock_converse_models: Set = set(BEDROCK_CONVERSE_MODELS) -fal_ai_models: Set = set() -fireworks_ai_models: Set = set() -fireworks_ai_embedding_models: Set = set() -deepinfra_models: Set = set() -perplexity_models: Set = set() -watsonx_models: Set = set() -gemini_models: Set = set() -xai_models: Set = set() -zai_models: Set = set() -deepseek_models: Set = set() -tencent_models: Set = set() -runwayml_models: Set = set() -azure_ai_models: Set = set() -jina_ai_models: Set = set() -voyage_models: Set = set() -infinity_models: Set = set() -heroku_models: Set = set() -databricks_models: Set = set() -cloudflare_models: Set = set() -codestral_models: Set = set() -friendliai_models: Set = set() -featherless_ai_models: Set = set() -palm_models: Set = set() -groq_models: Set = set() -azure_models: Set = set() -azure_anthropic_models: Set = set() -azure_text_models: Set = set() -anyscale_models: Set = set() -cerebras_models: Set = set() -galadriel_models: Set = set() -nvidia_nim_models: Set = set() -nvidia_riva_models: Set = set() -soniox_models: Set = set() -sambanova_models: Set = set() -sambanova_embedding_models: Set = set() -novita_models: Set = set() -assemblyai_models: Set = set() -snowflake_models: Set = set() -gradient_ai_models: Set = set() -llama_models: Set = set() -nscale_models: Set = set() -nebius_models: Set = set() -nebius_embedding_models: Set = set() -aiml_models: Set = set() -deepgram_models: Set = set() -elevenlabs_models: Set = set() -dashscope_models: Set = set() -moonshot_models: Set = set() -publicai_models: Set = set() -darkbloom_models: Set = set() -v0_models: Set = set() -morph_models: Set = set() -lambda_ai_models: Set = set() -inception_models: Set = set() -hyperbolic_models: Set = set() -black_forest_labs_models: Set = set() -recraft_models: Set = set() -cometapi_models: Set = set() -oci_models: Set = set() -vercel_ai_gateway_models: Set = set() -volcengine_models: Set = set() -wandb_models: Set = set(WANDB_MODELS) -ovhcloud_models: Set = set() -ovhcloud_embedding_models: Set = set() -lemonade_models: Set = set() -docker_model_runner_models: Set = set() -amazon_nova_models: Set = set() -stability_models: Set = set() -github_copilot_models: Set = set() -chatgpt_models: Set = set() -minimax_models: Set = set() -aws_polly_models: Set = set() -gigachat_models: Set = set() -llamagate_models: Set = set() -reducto_models: Set = set() -bedrock_mantle_models: Set = set() +open_ai_chat_completion_models: Final[Set] = set() +open_ai_text_completion_models: Final[Set] = set() +cohere_models: Final[Set] = set() +cohere_chat_models: Final[Set] = set() +mistral_chat_models: Final[Set] = set() +text_completion_codestral_models: Final[Set] = set() +text_completion_inception_models: Final[Set] = set() +anthropic_models: Final[Set] = set() +openrouter_models: Final[Set] = set() +datarobot_models: Final[Set] = set() +vertex_language_models: Final[Set] = set() +vertex_vision_models: Final[Set] = set() +vertex_chat_models: Final[Set] = set() +vertex_code_chat_models: Final[Set] = set() +vertex_ai_image_models: Final[Set] = set() +vertex_ai_video_models: Final[Set] = set() +vertex_text_models: Final[Set] = set() +vertex_code_text_models: Final[Set] = set() +vertex_embedding_models: Final[Set] = set() +vertex_anthropic_models: Final[Set] = set() +vertex_llama3_models: Final[Set] = set() +vertex_deepseek_models: Final[Set] = set() +vertex_ai_ai21_models: Final[Set] = set() +vertex_mistral_models: Final[Set] = set() +vertex_openai_models: Final[Set] = set() +vertex_minimax_models: Final[Set] = set() +vertex_moonshot_models: Final[Set] = set() +vertex_zai_models: Final[Set] = set() +ai21_models: Final[Set] = set() +ai21_chat_models: Final[Set] = set() +nlp_cloud_models: Final[Set] = set() +aleph_alpha_models: Final[Set] = set() +bedrock_models: Final[Set] = set() +bedrock_converse_models: Final[Set] = set(BEDROCK_CONVERSE_MODELS) +fal_ai_models: Final[Set] = set() +fireworks_ai_models: Final[Set] = set() +fireworks_ai_embedding_models: Final[Set] = set() +deepinfra_models: Final[Set] = set() +perplexity_models: Final[Set] = set() +watsonx_models: Final[Set] = set() +gemini_models: Final[Set] = set() +xai_models: Final[Set] = set() +zai_models: Final[Set] = set() +deepseek_models: Final[Set] = set() +tencent_models: Final[Set] = set() +runwayml_models: Final[Set] = set() +azure_ai_models: Final[Set] = set() +jina_ai_models: Final[Set] = set() +voyage_models: Final[Set] = set() +infinity_models: Final[Set] = set() +heroku_models: Final[Set] = set() +databricks_models: Final[Set] = set() +cloudflare_models: Final[Set] = set() +codestral_models: Final[Set] = set() +friendliai_models: Final[Set] = set() +featherless_ai_models: Final[Set] = set() +palm_models: Final[Set] = set() +groq_models: Final[Set] = set() +azure_models: Final[Set] = set() +azure_anthropic_models: Final[Set] = set() +azure_text_models: Final[Set] = set() +anyscale_models: Final[Set] = set() +cerebras_models: Final[Set] = set() +galadriel_models: Final[Set] = set() +nvidia_nim_models: Final[Set] = set() +nvidia_riva_models: Final[Set] = set() +soniox_models: Final[Set] = set() +sambanova_models: Final[Set] = set() +sambanova_embedding_models: Final[Set] = set() +novita_models: Final[Set] = set() +assemblyai_models: Final[Set] = set() +snowflake_models: Final[Set] = set() +gradient_ai_models: Final[Set] = set() +llama_models: Final[Set] = set() +nscale_models: Final[Set] = set() +nebius_models: Final[Set] = set() +nebius_embedding_models: Final[Set] = set() +aiml_models: Final[Set] = set() +deepgram_models: Final[Set] = set() +elevenlabs_models: Final[Set] = set() +dashscope_models: Final[Set] = set() +moonshot_models: Final[Set] = set() +publicai_models: Final[Set] = set() +darkbloom_models: Final[Set] = set() +v0_models: Final[Set] = set() +morph_models: Final[Set] = set() +lambda_ai_models: Final[Set] = set() +inception_models: Final[Set] = set() +hyperbolic_models: Final[Set] = set() +black_forest_labs_models: Final[Set] = set() +recraft_models: Final[Set] = set() +cometapi_models: Final[Set] = set() +oci_models: Final[Set] = set() +vercel_ai_gateway_models: Final[Set] = set() +volcengine_models: Final[Set] = set() +wandb_models: Final[Set] = set(WANDB_MODELS) +ovhcloud_models: Final[Set] = set() +ovhcloud_embedding_models: Final[Set] = set() +lemonade_models: Final[Set] = set() +docker_model_runner_models: Final[Set] = set() +amazon_nova_models: Final[Set] = set() +stability_models: Final[Set] = set() +github_copilot_models: Final[Set] = set() +chatgpt_models: Final[Set] = set() +minimax_models: Final[Set] = set() +aws_polly_models: Final[Set] = set() +gigachat_models: Final[Set] = set() +llamagate_models: Final[Set] = set() +reducto_models: Final[Set] = set() +bedrock_mantle_models: Final[Set] = set() def is_bedrock_pricing_only_model(key: str) -> bool: @@ -681,12 +682,12 @@ def is_bedrock_pricing_only_model(key: str) -> bool: bool: True if the key matches the Bedrock pattern, False otherwise. """ # Regex to match 'bedrock//' - bedrock_pattern = re.compile(r"^bedrock/[a-zA-Z0-9_-]+/.+$") + bedrock_pattern: Final = re.compile(r"^bedrock/[a-zA-Z0-9_-]+/.+$") if "month-commitment" in key: return True - is_match = bedrock_pattern.match(key) + is_match: Final = bedrock_pattern.match(key) return is_match is not None @@ -704,7 +705,7 @@ def is_openai_finetune_model(key: str) -> bool: def add_known_models(model_cost_map: Optional[Dict] = None): - _map = model_cost_map if model_cost_map is not None else model_cost + _map: Final = model_cost_map if model_cost_map is not None else model_cost for key, value in _map.items(): if value.get("litellm_provider") == "openai" and not is_openai_finetune_model(key): open_ai_chat_completion_models.add(key) @@ -957,7 +958,7 @@ add_known_models() # used for Cost Tracking & Token counting # https://azure.microsoft.com/en-in/pricing/details/cognitive-services/openai-service/ # Azure returns gpt-35-turbo in their responses, we need to map this to azure/gpt-3.5-turbo for token counting -azure_llms = { +azure_llms: Final = { "gpt-35-turbo": "azure/gpt-35-turbo", "gpt-35-turbo-16k": "azure/gpt-35-turbo-16k", "gpt-35-turbo-instruct": "azure/gpt-35-turbo-instruct", @@ -966,19 +967,19 @@ azure_llms = { "azure/gpt-41-nano": "gpt-4.1-nano", } -azure_embedding_models = { +azure_embedding_models: Final = { "ada": "azure/ada", } -petals_models = [ +petals_models: Final = [ "petals-team/StableBeluga2", ] -ollama_models = ["llama2"] +ollama_models: Final = ["llama2"] -maritalk_models = ["maritalk"] +maritalk_models: Final = ["maritalk"] -model_list = list( +model_list: Final = list( open_ai_chat_completion_models | open_ai_text_completion_models | cohere_models @@ -1065,12 +1066,12 @@ model_list = list( | set(clarifai_models) ) -model_list_set = set(model_list) +model_list_set: Final = set(model_list) # provider_list is lazy-loaded via __getattr__ to avoid importing LlmProviders at import time -models_by_provider: dict = { +models_by_provider: Final[dict] = { "openai": open_ai_chat_completion_models | open_ai_text_completion_models, "text-completion-openai": open_ai_text_completion_models, "cohere": cohere_models | cohere_chat_models, @@ -1178,7 +1179,7 @@ models_by_provider: dict = { } # mapping for those models which have larger equivalents -longer_context_model_fallback_dict: dict = { +longer_context_model_fallback_dict: Final[dict] = { # openai chat completion models "gpt-3.5-turbo": "gpt-3.5-turbo-16k", "gpt-3.5-turbo-0301": "gpt-3.5-turbo-16k-0301", @@ -1201,7 +1202,7 @@ longer_context_model_fallback_dict: dict = { ####### EMBEDDING MODELS ################### -all_embedding_models = ( +all_embedding_models: Final = ( open_ai_embedding_models | set(cohere_embedding_models) | set(bedrock_embedding_models) @@ -1213,10 +1214,10 @@ all_embedding_models = ( ) ####### IMAGE GENERATION MODELS ################### -openai_image_generation_models = ["dall-e-2", "dall-e-3"] +openai_image_generation_models: Final = ["dall-e-2", "dall-e-3"] ####### VIDEO GENERATION MODELS ################### -openai_video_generation_models = ["sora-2"] +openai_video_generation_models: Final = ["sora-2"] # timeout is lazy-loaded via __getattr__ # get_llm_provider is lazy-loaded via __getattr__ @@ -1248,7 +1249,7 @@ from .llms.vertex_ai.vertex_embeddings.transformation import ( VertexAITextEmbeddingConfig, ) -vertexAITextEmbeddingConfig = VertexAITextEmbeddingConfig() +vertexAITextEmbeddingConfig: Final = VertexAITextEmbeddingConfig() from .llms.bedrock.embed.amazon_titan_v2_transformation import ( @@ -1423,11 +1424,11 @@ from . import rag from .types.llms.custom_llm import CustomLLMItem custom_provider_map: List[CustomLLMItem] = [] -_custom_providers: List[str] = [] # internal helper util, used to track names of custom providers -disable_hf_tokenizer_download: Optional[bool] = ( +_custom_providers: Final[List[str]] = [] # internal helper util, used to track names of custom providers +disable_hf_tokenizer_download: Final[Optional[bool]] = ( None # disable huggingface tokenizer download. Defaults to openai clk100 ) -global_disable_no_log_param: bool = False +global_disable_no_log_param: Final[bool] = False ### CLI UTILITIES ### from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key @@ -2140,11 +2141,11 @@ def __getattr__(name: str) -> Any: # Use cached registry from _lazy_imports instead of importing tuples every time from ._lazy_imports import _get_lazy_import_registry - registry = _get_lazy_import_registry() + registry: Final = _get_lazy_import_registry() # Check if name is in registry and call the cached handler function if name in registry: - handler_func = registry[name] + handler_func: Final = registry[name] return handler_func(name) # Lazy load encoding from main.py to avoid heavy tiktoken import @@ -2197,7 +2198,7 @@ def __getattr__(name: str) -> Any: return _globals["openaiOSeriesConfig"] # Lazy load other config instances - _config_instances = { + _config_instances: Final = { "openAIGPTConfig": "OpenAIGPTConfig", "openAIGPTAudioConfig": "OpenAIGPTAudioConfig", "openAIGPT5Config": "OpenAIGPT5Config", @@ -2239,7 +2240,7 @@ def __getattr__(name: str) -> Any: # Check if already cached if "priority_reservation_settings" not in _globals: # Import the class and instantiate it - PriorityReservationSettings = __getattr__("PriorityReservationSettings") + PriorityReservationSettings: Final = __getattr__("PriorityReservationSettings") _globals["priority_reservation_settings"] = PriorityReservationSettings() return _globals["priority_reservation_settings"] @@ -2251,7 +2252,7 @@ def __getattr__(name: str) -> Any: # Check if already cached if "logging_callback_manager" not in _globals: # Import the class and instantiate it - LoggingCallbackManager = __getattr__("LoggingCallbackManager") + LoggingCallbackManager: Final = __getattr__("LoggingCallbackManager") _globals["logging_callback_manager"] = LoggingCallbackManager() return _globals["logging_callback_manager"] diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py index 727ca87d3e8..f856fe0f2b3 100644 --- a/litellm/_internal_context.py +++ b/litellm/_internal_context.py @@ -7,7 +7,8 @@ asyncio task and cannot be injected via HTTP request bodies. """ from contextvars import ContextVar +from typing import Final # When True, suppresses async logging and billing for internal sub-calls # (e.g., emulated file-search steps that make nested LLM calls). -is_internal_call: ContextVar[bool] = ContextVar("is_internal_call", default=False) +is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False) diff --git a/litellm/_lazy_imports.py b/litellm/_lazy_imports.py index 4eee525f6a4..63142ee4f2f 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, cast +from typing import Any, Final, 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 @@ -233,7 +233,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate raise AttributeError(f"{category} lazy import: unknown attribute {name!r}") # Step 2: Get the cache (where we store imported things) - _globals = _get_litellm_globals() + _globals: Final = _get_litellm_globals() # Step 3: If we've already imported it, just return the cached version if name in _globals: @@ -255,7 +255,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate # Step 6: Get the actual attribute from the module # Example: getattr(utils_module, "ModelResponse") returns the ModelResponse class - value = getattr(module, attr_name) + value: Final = getattr(module, attr_name) # Step 7: Cache it so we don't have to import again next time _globals[name] = value @@ -339,7 +339,7 @@ def _lazy_import_utils_module(name: str) -> Any: raise AttributeError(f"Utils module lazy import: unknown attribute {name!r}") # Get the cache (where we store imported things) - use utils globals - _globals = _get_utils_globals() + _globals: Final = _get_utils_globals() # If we've already imported it, just return the cached version if name in _globals: @@ -355,7 +355,7 @@ def _lazy_import_utils_module(name: str) -> Any: module = importlib.import_module(module_path) # Get the actual attribute from the module - value = getattr(module, attr_name) + value: Final = getattr(module, attr_name) # Cache it so we don't have to import again next time _globals[name] = value @@ -379,15 +379,15 @@ def _lazy_import_llm_client_cache(name: str) -> Any: - "in_memory_llm_clients_cache" is a singleton instance of that class So we need custom logic to handle both cases. """ - _globals = _get_litellm_globals() + _globals: Final = _get_litellm_globals() # If already cached, return it if name in _globals: return _globals[name] # Import the class - module = importlib.import_module("litellm.caching.llm_caching_handler") - LLMClientCache = getattr(module, "LLMClientCache") + module: Final = importlib.import_module("litellm.caching.llm_caching_handler") + LLMClientCache: Final = getattr(module, "LLMClientCache") # If they want the class itself, return it if name == "LLMClientCache": @@ -396,7 +396,7 @@ def _lazy_import_llm_client_cache(name: str) -> Any: # If they want the singleton instance, create it (only once) if name == "in_memory_llm_clients_cache": - instance = LLMClientCache() + instance: Final = LLMClientCache() _globals["in_memory_llm_clients_cache"] = instance return instance @@ -412,7 +412,7 @@ def _lazy_import_http_handlers(name: str) -> Any: - They need configuration (timeout, etc.) from the module globals - They use factory functions instead of direct instantiation """ - _globals = _get_litellm_globals() + _globals: Final = _get_litellm_globals() if name == "module_level_aclient": # Create an async HTTP client using the factory function @@ -420,11 +420,11 @@ def _lazy_import_http_handlers(name: str) -> Any: # Get timeout from module config (if set) timeout = _globals.get("request_timeout") - params = {"timeout": timeout, "client_alias": "module level aclient"} + params: Final = {"timeout": timeout, "client_alias": "module level aclient"} # Create the client instance - provider_id = cast(Any, "litellm_module_level_client") - async_client = get_async_httpx_client( + provider_id: Final = cast(Any, "litellm_module_level_client") + async_client: Final = get_async_httpx_client( llm_provider=provider_id, params=params, ) @@ -438,7 +438,7 @@ def _lazy_import_http_handlers(name: str) -> Any: from litellm.llms.custom_httpx.http_handler import HTTPHandler timeout = _globals.get("request_timeout") - sync_client = HTTPHandler(timeout=timeout) + sync_client: Final = HTTPHandler(timeout=timeout) # Cache it _globals["module_level_client"] = sync_client diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 488331e3895..37f111c2324 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -5,21 +5,23 @@ This module contains all the name tuples and import maps used by the lazy import Separated from the handler functions for better organization. """ +from typing import Final + # Cost calculator names that support lazy loading via _lazy_import_cost_calculator -COST_CALCULATOR_NAMES = ( +COST_CALCULATOR_NAMES: Final = ( "completion_cost", "cost_per_token", "response_cost_calculator", ) # Litellm logging names that support lazy loading via _lazy_import_litellm_logging -LITELLM_LOGGING_NAMES = ( +LITELLM_LOGGING_NAMES: Final = ( "Logging", "modify_integration", ) # Utils names that support lazy loading via _lazy_import_utils -UTILS_NAMES = ( +UTILS_NAMES: Final = ( "exception_type", "get_optional_params", "get_response_string", @@ -66,20 +68,20 @@ UTILS_NAMES = ( ) # Token counter names that support lazy loading via _lazy_import_token_counter -TOKEN_COUNTER_NAMES = ("get_modified_max_tokens",) +TOKEN_COUNTER_NAMES: Final = ("get_modified_max_tokens",) # LLM client cache names that support lazy loading via _lazy_import_llm_client_cache -LLM_CLIENT_CACHE_NAMES = ( +LLM_CLIENT_CACHE_NAMES: Final = ( "LLMClientCache", "in_memory_llm_clients_cache", ) # Bedrock type names that support lazy loading via _lazy_import_bedrock_types -BEDROCK_TYPES_NAMES = ("COHERE_EMBEDDING_INPUT_TYPES",) +BEDROCK_TYPES_NAMES: Final = ("COHERE_EMBEDDING_INPUT_TYPES",) # Common types from litellm.types.utils that support lazy loading via # _lazy_import_types_utils -TYPES_UTILS_NAMES = ( +TYPES_UTILS_NAMES: Final = ( "ImageObject", "BudgetConfig", "all_litellm_params", @@ -92,7 +94,7 @@ TYPES_UTILS_NAMES = ( ) # Caching / cache classes that support lazy loading via _lazy_import_caching -CACHING_NAMES = ( +CACHING_NAMES: Final = ( "Cache", "DualCache", "RedisCache", @@ -100,20 +102,20 @@ CACHING_NAMES = ( ) # HTTP handler names that support lazy loading via _lazy_import_http_handlers -HTTP_HANDLER_NAMES = ( +HTTP_HANDLER_NAMES: Final = ( "module_level_aclient", "module_level_client", ) # Dotprompt integration names that support lazy loading via _lazy_import_dotprompt -DOTPROMPT_NAMES = ( +DOTPROMPT_NAMES: Final = ( "global_prompt_manager", "global_prompt_directory", "set_global_prompt_directory", ) # LLM config classes that support lazy loading via _lazy_import_llm_configs -LLM_CONFIG_NAMES = ( +LLM_CONFIG_NAMES: Final = ( "AmazonConverseConfig", "OpenAILikeChatConfig", "GaladrielChatConfig", @@ -328,7 +330,7 @@ LLM_CONFIG_NAMES = ( ) # Types that support lazy loading via _lazy_import_types -TYPES_NAMES = ( +TYPES_NAMES: Final = ( "GuardrailItem", "DefaultTeamSSOParams", "LiteLLM_UpperboundKeyGenerateParams", @@ -344,14 +346,14 @@ TYPES_NAMES = ( ) # LLM provider logic names that support lazy loading via _lazy_import_llm_provider_logic -LLM_PROVIDER_LOGIC_NAMES = ( +LLM_PROVIDER_LOGIC_NAMES: Final = ( "get_llm_provider", "remove_index_from_tool_calls", ) # Utils module names that support lazy loading via _lazy_import_utils_module # These are attributes accessed from litellm.utils module -UTILS_MODULE_NAMES = ( +UTILS_MODULE_NAMES: Final = ( "encoding", "BaseVectorStore", "CredentialAccessor", @@ -423,7 +425,7 @@ UTILS_MODULE_NAMES = ( ) # Import maps for registry pattern - reduces repetition -_UTILS_IMPORT_MAP = { +_UTILS_IMPORT_MAP: Final = { "exception_type": (".utils", "exception_type"), "get_optional_params": (".utils", "get_optional_params"), "get_response_string": (".utils", "get_response_string"), @@ -478,13 +480,13 @@ _UTILS_IMPORT_MAP = { ), } -_COST_CALCULATOR_IMPORT_MAP = { +_COST_CALCULATOR_IMPORT_MAP: Final = { "completion_cost": (".cost_calculator", "completion_cost"), "cost_per_token": (".cost_calculator", "cost_per_token"), "response_cost_calculator": (".cost_calculator", "response_cost_calculator"), } -_TYPES_UTILS_IMPORT_MAP = { +_TYPES_UTILS_IMPORT_MAP: Final = { "ImageObject": (".types.utils", "ImageObject"), "BudgetConfig": (".types.utils", "BudgetConfig"), "all_litellm_params": (".types.utils", "all_litellm_params"), @@ -496,28 +498,28 @@ _TYPES_UTILS_IMPORT_MAP = { "GenericStreamingChunk": (".types.utils", "GenericStreamingChunk"), } -_TOKEN_COUNTER_IMPORT_MAP = { +_TOKEN_COUNTER_IMPORT_MAP: Final = { "get_modified_max_tokens": ( "litellm.litellm_core_utils.token_counter", "get_modified_max_tokens", ), } -_BEDROCK_TYPES_IMPORT_MAP = { +_BEDROCK_TYPES_IMPORT_MAP: Final = { "COHERE_EMBEDDING_INPUT_TYPES": ( "litellm.types.llms.bedrock", "COHERE_EMBEDDING_INPUT_TYPES", ), } -_CACHING_IMPORT_MAP = { +_CACHING_IMPORT_MAP: Final = { "Cache": ("litellm.caching.caching", "Cache"), "DualCache": ("litellm.caching.caching", "DualCache"), "RedisCache": ("litellm.caching.caching", "RedisCache"), "InMemoryCache": ("litellm.caching.caching", "InMemoryCache"), } -_LITELLM_LOGGING_IMPORT_MAP = { +_LITELLM_LOGGING_IMPORT_MAP: Final = { "Logging": ("litellm.litellm_core_utils.litellm_logging", "Logging"), "modify_integration": ( "litellm.litellm_core_utils.litellm_logging", @@ -525,7 +527,7 @@ _LITELLM_LOGGING_IMPORT_MAP = { ), } -_DOTPROMPT_IMPORT_MAP = { +_DOTPROMPT_IMPORT_MAP: Final = { "global_prompt_manager": ( "litellm.integrations.dotprompt", "global_prompt_manager", @@ -540,7 +542,7 @@ _DOTPROMPT_IMPORT_MAP = { ), } -_TYPES_IMPORT_MAP = { +_TYPES_IMPORT_MAP: Final = { "GuardrailItem": ("litellm.types.guardrails", "GuardrailItem"), "DefaultTeamSSOParams": ( "litellm.types.proxy.management_endpoints.ui_sso", @@ -569,7 +571,7 @@ _TYPES_IMPORT_MAP = { ), } -_LLM_PROVIDER_LOGIC_IMPORT_MAP = { +_LLM_PROVIDER_LOGIC_IMPORT_MAP: Final = { "get_llm_provider": ( "litellm.litellm_core_utils.get_llm_provider_logic", "get_llm_provider", @@ -580,7 +582,7 @@ _LLM_PROVIDER_LOGIC_IMPORT_MAP = { ), } -_LLM_CONFIGS_IMPORT_MAP = { +_LLM_CONFIGS_IMPORT_MAP: Final = { "AmazonConverseConfig": ( ".llms.bedrock.chat.converse_transformation", "AmazonConverseConfig", @@ -1215,7 +1217,7 @@ _LLM_CONFIGS_IMPORT_MAP = { } # Import map for utils module lazy imports -_UTILS_MODULE_IMPORT_MAP = { +_UTILS_MODULE_IMPORT_MAP: Final = { "encoding": ("litellm.main", "encoding"), "BaseVectorStore": ( "litellm.integrations.vector_store_integrations.base_vector_store", diff --git a/litellm/_logging.py b/litellm/_logging.py index a41784e9170..c5a9fd0f8c7 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -4,7 +4,7 @@ import os import sys from datetime import datetime from logging import Formatter -from typing import Any +from typing import Any, Final from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads @@ -17,7 +17,7 @@ if set_verbose is True: "`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs." ) -_ENABLE_SECRET_REDACTION = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true" +_ENABLE_SECRET_REDACTION: Final = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true" def _redact_string(value: str) -> str: @@ -74,14 +74,14 @@ class SecretRedactionFilter(logging.Filter): return True -_secret_filter = SecretRedactionFilter() +_secret_filter: Final = SecretRedactionFilter() json_logs = bool(os.getenv("JSON_LOGS", False)) # Create a handler for the logger (you may need to adapt this based on your needs) -log_level = os.getenv("LITELLM_LOG", "DEBUG") -numeric_level: str = getattr(logging, log_level.upper()) -handler = logging.StreamHandler() +log_level: Final = os.getenv("LITELLM_LOG", "DEBUG") +numeric_level: Final[str] = getattr(logging, log_level.upper()) +handler: Final = logging.StreamHandler() handler.setLevel(numeric_level) handler.addFilter(_secret_filter) @@ -94,10 +94,10 @@ def _try_parse_json_message(message: str) -> dict[str, Any] | None: """ if not message or not isinstance(message, str): return None - msg_stripped = message.strip() + msg_stripped: Final = message.strip() if not (msg_stripped.startswith("{") or msg_stripped.startswith("[")): return None - parsed = safe_json_loads(message, default=None) + parsed: Final = safe_json_loads(message, default=None) if parsed is None or not isinstance(parsed, dict): return None return parsed @@ -144,7 +144,7 @@ def _get_standard_record_attrs() -> frozenset: return frozenset(logging.LogRecord("", 0, "", 0, "", (), None).__dict__.keys()) -_STANDARD_RECORD_ATTRS = _get_standard_record_attrs() +_STANDARD_RECORD_ATTRS: Final = _get_standard_record_attrs() class JsonFormatter(Formatter): @@ -153,12 +153,12 @@ class JsonFormatter(Formatter): def formatTime(self, record, datefmt=None): # Use datetime to format the timestamp in ISO 8601 format - dt = datetime.fromtimestamp(record.created) + dt: Final = datetime.fromtimestamp(record.created) return dt.isoformat() def format(self, record): - message_str = record.getMessage() - json_record: dict[str, Any] = { + message_str: Final = record.getMessage() + json_record: Final[dict[str, Any]] = { "message": message_str, "level": record.levelname, "timestamp": self.formatTime(record), @@ -193,13 +193,13 @@ class JsonFormatter(Formatter): # Function to set up exception handlers for JSON logging def _setup_json_exception_handlers(formatter): # Create a handler with JSON formatting for exceptions - error_handler = logging.StreamHandler() + error_handler: Final = logging.StreamHandler() error_handler.setFormatter(formatter) error_handler.addFilter(_secret_filter) # Setup excepthook for uncaught exceptions def json_excepthook(exc_type, exc_value, exc_traceback): - record = logging.LogRecord( + record: Final = logging.LogRecord( name="LiteLLM", level=logging.ERROR, pathname="", @@ -217,10 +217,10 @@ def _setup_json_exception_handlers(formatter): import asyncio def async_json_exception_handler(loop, context): - exception = context.get("exception") + exception: Final = context.get("exception") if exception: - exc_type = type(exception) - record = logging.LogRecord( + exc_type: Final = type(exception) + record: Final = logging.LogRecord( name="LiteLLM", level=logging.ERROR, pathname="", @@ -243,7 +243,7 @@ if json_logs: handler.setFormatter(JsonFormatter()) _setup_json_exception_handlers(JsonFormatter()) else: - formatter = logging.Formatter( + formatter: Final = logging.Formatter( "\033[92m%(asctime)s - %(name)s:%(levelname)s\033[0m: %(filename)s:%(lineno)s - %(message)s", datefmt="%H:%M:%S", ) @@ -263,20 +263,20 @@ verbose_logger.addHandler(handler) def _suppress_loggers(): """Suppress noisy loggers at INFO level""" # Suppress httpx request logging at INFO level - httpx_logger = logging.getLogger("httpx") + httpx_logger: Final = logging.getLogger("httpx") httpx_logger.setLevel(logging.WARNING) # Suppress APScheduler logging at INFO level - apscheduler_executors_logger = logging.getLogger("apscheduler.executors.default") + apscheduler_executors_logger: Final = logging.getLogger("apscheduler.executors.default") apscheduler_executors_logger.setLevel(logging.WARNING) - apscheduler_scheduler_logger = logging.getLogger("apscheduler.scheduler") + apscheduler_scheduler_logger: Final = logging.getLogger("apscheduler.scheduler") apscheduler_scheduler_logger.setLevel(logging.WARNING) # Call the suppression function _suppress_loggers() -ALL_LOGGERS = [ +ALL_LOGGERS: Final = [ logging.getLogger(), verbose_logger, verbose_router_logger, @@ -293,11 +293,11 @@ def _get_loggers_to_initialize(): """ import litellm - loggers = list(ALL_LOGGERS) + loggers: Final = list(ALL_LOGGERS) # Add langfuse logger if langfuse is being used as a callback - langfuse_callbacks = {"langfuse", "langfuse_otel"} - all_callbacks = set(litellm.success_callback + litellm.failure_callback) + langfuse_callbacks: Final = {"langfuse", "langfuse_otel"} + all_callbacks: Final = set(litellm.success_callback + litellm.failure_callback) if langfuse_callbacks & all_callbacks: loggers.append(logging.getLogger("langfuse")) @@ -325,12 +325,12 @@ def _get_uvicorn_json_log_config(): This ensures that uvicorn's access logs, error logs, and all application logs are formatted as JSON when json_logs is enabled. """ - json_formatter_class = "litellm._logging.JsonFormatter" + json_formatter_class: Final = "litellm._logging.JsonFormatter" # Use the module-level log_level variable for consistency - uvicorn_log_level = log_level.upper() + uvicorn_log_level: Final = log_level.upper() - log_config = { + log_config: Final = { "version": 1, "disable_existing_loggers": False, "formatters": { @@ -384,7 +384,7 @@ def _turn_on_json(): - Adds a JSON formatter to all loggers """ - handler = logging.StreamHandler() + handler: Final = logging.StreamHandler() handler.setFormatter(JsonFormatter()) _initialize_loggers_with_handler(handler) # Set up exception handlers diff --git a/litellm/_redis.py b/litellm/_redis.py index e05e9d4eb20..ed014a83c25 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -13,6 +13,7 @@ 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 Final import redis # type: ignore import redis.asyncio as async_redis # type: ignore @@ -32,20 +33,20 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from ._logging import verbose_logger -AZURE_REDIS_SCOPE = "https://redis.azure.com/.default" +AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" def _get_redis_kwargs(): - arg_spec = inspect.getfullargspec(redis.Redis) + arg_spec: Final = inspect.getfullargspec(redis.Redis) # Only allow primitive arguments - exclude_args = { + exclude_args: Final = { "self", "connection_pool", "retry", } - include_args = { + include_args: Final = { "url", "redis_connect_func", "gcp_service_account", @@ -56,7 +57,7 @@ def _get_redis_kwargs(): "azure_client_secret", } - available_args = {x for x in arg_spec.args if x not in exclude_args} | include_args + available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args return available_args @@ -92,9 +93,9 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]: """ if client is None: client = redis.Redis - connection_cls = async_redis.Connection if client is async_redis.Redis else redis.Connection + connection_cls: Final = async_redis.Connection if client is async_redis.Redis else redis.Connection - exclude_args = frozenset( + exclude_args: Final = frozenset( { "self", "connection_pool", @@ -103,7 +104,7 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]: ) # Only allow primitive arguments - include_args = ("url", "max_connections") + include_args: Final = ("url", "max_connections") return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args @@ -111,10 +112,10 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]: def _get_redis_cluster_kwargs(client=None): if client is None: client = redis.Redis.from_url - arg_spec = inspect.getfullargspec(redis.RedisCluster) + arg_spec: Final = inspect.getfullargspec(redis.RedisCluster) # Only allow primitive arguments - exclude_args = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"} + exclude_args: Final = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"} available_args = {x for x in arg_spec.args if x not in exclude_args} available_args |= { @@ -142,15 +143,15 @@ def _get_redis_cluster_kwargs(client=None): def _get_redis_env_kwarg_mapping(): - PREFIX = "REDIS_" + PREFIX: Final = "REDIS_" return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs()} def _redis_kwargs_from_environment(): - mapping = _get_redis_env_kwarg_mapping() + mapping: Final = _get_redis_env_kwarg_mapping() - return_dict = {} + return_dict: Final = {} for k, v in mapping.items(): value = get_secret(k, default_value=None) # type: ignore if value is not None: @@ -183,7 +184,7 @@ def create_gcp_iam_redis_connect_func( self._parser.on_connect(self) - auth_args = (_generate_gcp_iam_access_token(service_account),) + auth_args: Final = (_generate_gcp_iam_access_token(service_account),) self.send_command("AUTH", *auth_args, check_health=False) try: @@ -224,9 +225,9 @@ def _build_azure_credential( "azure-identity is required for Azure AD Redis authentication. Install it with: pip install azure-identity" ) - _client_id = azure_client_id or os.environ.get("AZURE_CLIENT_ID") - _tenant_id = azure_tenant_id or os.environ.get("AZURE_TENANT_ID") - _client_secret = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET") + _client_id: Final = azure_client_id or os.environ.get("AZURE_CLIENT_ID") + _tenant_id: Final = azure_tenant_id or os.environ.get("AZURE_TENANT_ID") + _client_secret: Final = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET") if _client_id and _tenant_id and _client_secret: return ClientSecretCredential( @@ -253,12 +254,12 @@ def _generate_azure_ad_redis_token( (``AzureADCredentialProvider``) keep the credential alive across connections so the Azure SDK's internal cache + silent refresh apply. """ - credential = _build_azure_credential( + credential: Final = _build_azure_credential( azure_client_id=azure_client_id, azure_tenant_id=azure_tenant_id, azure_client_secret=azure_client_secret, ) - token = credential.get_token(AZURE_REDIS_SCOPE) + token: Final = credential.get_token(AZURE_REDIS_SCOPE) return token.token @@ -274,7 +275,7 @@ def create_azure_ad_redis_connect_func( closure) and reused across connections — the Azure SDK handles token caching and silent renewal internally. Only ``get_token`` is called per connection. """ - credential = _build_azure_credential( + credential: Final = _build_azure_credential( azure_client_id=azure_client_id, azure_tenant_id=azure_tenant_id, azure_client_secret=azure_client_secret, @@ -290,11 +291,11 @@ def create_azure_ad_redis_connect_func( self._parser.on_connect(self) - access_token = credential.get_token(AZURE_REDIS_SCOPE).token + access_token: Final = credential.get_token(AZURE_REDIS_SCOPE).token # Only include username when explicitly set — sending AUTH "" # is invalid for most ACL-configured Azure Redis instances. - username = os.environ.get("REDIS_USERNAME", "") + username: Final = os.environ.get("REDIS_USERNAME", "") if username: auth_args = (username, access_token) else: @@ -353,23 +354,23 @@ def _get_redis_client_logic(**env_overrides): value = get_secret(v) # type: ignore env_overrides[k] = value - environment_kwargs = _redis_kwargs_from_environment() + environment_kwargs: Final = _redis_kwargs_from_environment() # An explicitly configured connection target outranks REDIS_URL from the # environment. Without this, the url branch below strips the caller's # host/port/password and silently connects to whatever REDIS_URL names. - caller_named_a_target = any( + caller_named_a_target: Final = any( env_overrides.get(key) is not None for key in ("host", "startup_nodes", "sentinel_nodes") ) if caller_named_a_target and env_overrides.get("url") is None: environment_kwargs.pop("url", None) - redis_kwargs = { + redis_kwargs: Final = { **environment_kwargs, **env_overrides, } - _startup_nodes: str | list | None = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore + _startup_nodes: Final[str | list | None] = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore "REDIS_CLUSTER_NODES" ) @@ -380,21 +381,21 @@ def _get_redis_client_logic(**env_overrides): elif _startup_nodes is None: redis_kwargs.pop("startup_nodes", None) - _sentinel_nodes: str | list | None = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore + _sentinel_nodes: Final[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: str | None = redis_kwargs.get("sentinel_password", None) or get_secret_str( + _sentinel_password: Final[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: str | None = redis_kwargs.get("service_name", None) or get_secret( # type: ignore + _service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret( # type: ignore "REDIS_SERVICE_NAME" ) @@ -402,8 +403,8 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs["service_name"] = _service_name # Handle GCP IAM authentication - _gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT") - _gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS") + _gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT") + _gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS") if _gcp_service_account is not None: verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.") @@ -422,9 +423,9 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs # Handle Azure AD authentication (after GCP IAM block) - _azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") + _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") - _azure_ad_enabled = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" + _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" if _azure_ad_enabled and _gcp_service_account is not None: verbose_logger.warning( @@ -433,9 +434,9 @@ def _get_redis_client_logic(**env_overrides): ) if _azure_ad_enabled and _gcp_service_account is None: - _azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID") - _azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID") - _azure_client_secret = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET") + _azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID") + _azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID") + _azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET") verbose_logger.debug("Setting up Azure AD authentication for Redis.") redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func( @@ -480,7 +481,7 @@ def _get_redis_client_logic(**env_overrides): def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: - _redis_cluster_nodes_in_env: str | None = get_secret("REDIS_CLUSTER_NODES") # type: ignore + _redis_cluster_nodes_in_env: Final[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) @@ -492,13 +493,13 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: verbose_logger.debug("init_redis_cluster: startup nodes are being initialized.") from redis.cluster import ClusterNode - args = _get_redis_cluster_kwargs() - cluster_kwargs = {} + args: Final = _get_redis_cluster_kwargs() + cluster_kwargs: Final = {} for arg in redis_kwargs: if arg in args: cluster_kwargs[arg] = redis_kwargs[arg] - new_startup_nodes: list[ClusterNode] = [] + new_startup_nodes: Final[list[ClusterNode]] = [] for item in redis_kwargs["startup_nodes"]: new_startup_nodes.append(ClusterNode(**item)) @@ -508,8 +509,8 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict: - connection_kwargs = {} - args = _get_redis_kwargs() + connection_kwargs: Final = {} + args: Final = _get_redis_kwargs() for arg in redis_kwargs: if arg in args: connection_kwargs[arg] = redis_kwargs[arg] @@ -518,12 +519,12 @@ def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict: def _init_redis_sentinel(redis_kwargs) -> redis.Redis: - sentinel_nodes = redis_kwargs.get("sentinel_nodes") - sentinel_password = redis_kwargs.get("sentinel_password") - service_name = redis_kwargs.get("service_name") - connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + sentinel_nodes: Final = redis_kwargs.get("sentinel_nodes") + sentinel_password: Final = redis_kwargs.get("sentinel_password") + service_name: Final = redis_kwargs.get("service_name") + connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs) connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) - sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs: Final = dict(connection_kwargs) sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: @@ -532,7 +533,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.") # Set up the Sentinel client - sentinel = redis.Sentinel( + sentinel: Final = redis.Sentinel( sentinel_nodes, sentinel_kwargs=sentinel_kwargs, ) @@ -543,12 +544,12 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: - sentinel_nodes = redis_kwargs.get("sentinel_nodes") - sentinel_password = redis_kwargs.get("sentinel_password") - service_name = redis_kwargs.get("service_name") - connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + sentinel_nodes: Final = redis_kwargs.get("sentinel_nodes") + sentinel_password: Final = redis_kwargs.get("sentinel_password") + service_name: Final = redis_kwargs.get("service_name") + connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs) connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) - sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs: Final = dict(connection_kwargs) sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: @@ -557,7 +558,7 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.") # Set up the Sentinel client - sentinel = async_redis.Sentinel( + sentinel: Final = async_redis.Sentinel( sentinel_nodes, sentinel_kwargs=sentinel_kwargs, ) @@ -568,14 +569,14 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: def get_redis_client(**env_overrides): - redis_kwargs = _get_redis_client_logic(**env_overrides) + redis_kwargs: Final = _get_redis_client_logic(**env_overrides) if "startup_nodes" in redis_kwargs: return init_redis_cluster(redis_kwargs) if "url" in redis_kwargs and redis_kwargs["url"] is not None: - args = _get_redis_url_kwargs() - url_kwargs = {} + args: Final = _get_redis_url_kwargs() + url_kwargs: Final = {} for arg in redis_kwargs: if arg in args: url_kwargs[arg] = redis_kwargs[arg] @@ -593,13 +594,13 @@ def get_redis_async_client( connection_pool: async_redis.BlockingConnectionPool | None = None, **env_overrides, ) -> async_redis.Redis | async_redis.RedisCluster: - redis_kwargs = _get_redis_client_logic(**env_overrides) + redis_kwargs: Final = _get_redis_client_logic(**env_overrides) if "startup_nodes" in redis_kwargs: from redis.cluster import ClusterNode args = _get_redis_cluster_kwargs() - cluster_kwargs = {} + cluster_kwargs: Final = {} for arg in redis_kwargs: if arg in args: cluster_kwargs[arg] = redis_kwargs[arg] @@ -621,7 +622,7 @@ def get_redis_async_client( username=os.environ.get("REDIS_USERNAME") or None, ) - new_startup_nodes: list[ClusterNode] = [] + new_startup_nodes: Final[list[ClusterNode]] = [] for item in redis_kwargs["startup_nodes"]: new_startup_nodes.append(ClusterNode(**item)) @@ -635,7 +636,7 @@ def get_redis_async_client( cluster_kwargs.setdefault("socket_keepalive", True) # Create async RedisCluster with IAM token as password if available - cluster_client = async_redis.RedisCluster( + cluster_client: Final = async_redis.RedisCluster( startup_nodes=new_startup_nodes, **cluster_kwargs, # type: ignore ) @@ -646,7 +647,7 @@ def get_redis_async_client( if connection_pool is not None: return async_redis.Redis(connection_pool=connection_pool) args = _get_redis_url_kwargs(client=async_redis.Redis) - url_kwargs = {} + url_kwargs: Final = {} for arg in redis_kwargs: if arg in args: url_kwargs[arg] = redis_kwargs[arg] @@ -686,15 +687,15 @@ def get_redis_async_client( def get_redis_connection_pool( **env_overrides, ) -> async_redis.BlockingConnectionPool | None: - redis_kwargs = _get_redis_client_logic(**env_overrides) + redis_kwargs: Final = _get_redis_client_logic(**env_overrides) verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs) if "startup_nodes" in redis_kwargs: return None if "url" in redis_kwargs and redis_kwargs["url"] is not None: - allowed_args = _get_redis_url_kwargs(client=async_redis.Redis) - pool_kwargs = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"} + allowed_args: Final = _get_redis_url_kwargs(client=async_redis.Redis) + pool_kwargs: Final = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"} pool_kwargs["timeout"] = REDIS_CONNECTION_POOL_TIMEOUT pool_kwargs["url"] = redis_kwargs["url"] if "max_connections" in redis_kwargs: @@ -710,7 +711,7 @@ def get_redis_connection_pool( # Wrap GCP / Azure AD auth in a CredentialProvider so pool-managed # connections re-fetch tokens via the SDK's internal cache + silent refresh # rather than reusing a single token captured at pool creation. - redis_connect_func = redis_kwargs.pop("redis_connect_func", None) + redis_connect_func: Final = redis_kwargs.pop("redis_connect_func", None) if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"): redis_kwargs["credential_provider"] = AzureADCredentialProvider( redis_connect_func._azure_credential, @@ -737,7 +738,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: if not verbose_logger.isEnabledFor(logging.DEBUG): return - console = Console() + console: Final = Console() # Initialize the sensitive data masker masker = SensitiveDataMasker() @@ -746,10 +747,10 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: masked_redis_kwargs = masker.mask_dict(redis_kwargs) # Create main panel title - title = Text("Redis Configuration", style="bold blue") + title: Final = Text("Redis Configuration", style="bold blue") # Create configuration table - config_table = Table( + config_table: Final = Table( title="🔧 Redis Connection Parameters", show_header=True, header_style="bold magenta", @@ -786,7 +787,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: connection_type = "Redis (URL-based)" # Create connection type info - info_table = Table( + info_table: Final = Table( title="📊 Connection Info", show_header=True, header_style="bold green", diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 762e8bcd928..8b8bbf9366f 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,21 +1,21 @@ import asyncio import threading import time -from typing import Any +from typing import Any, Final from redis.credentials import CredentialProvider # type: ignore[attr-defined] # Azure AD scope for Redis Cache for Azure. -AZURE_REDIS_SCOPE = "https://redis.azure.com/.default" +AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" # GCP IAM tokens are valid for 1 hour. Cache for 55 minutes to refresh before expiry. -_GCP_IAM_TOKEN_TTL_SECONDS = 3300 +_GCP_IAM_TOKEN_TTL_SECONDS: Final = 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_lock = threading.Lock() +_token_cache: Final[dict[str, tuple[str, float]]] = {} +_token_cache_lock: Final = threading.Lock() def _generate_gcp_iam_access_token(service_account: str) -> str: @@ -36,12 +36,12 @@ def _generate_gcp_iam_access_token(service_account: str) -> str: "Install it with: pip install google-cloud-iam" ) - client = iam_credentials_v1.IAMCredentialsClient() - request = iam_credentials_v1.GenerateAccessTokenRequest( + client: Final = iam_credentials_v1.IAMCredentialsClient() + request: Final = iam_credentials_v1.GenerateAccessTokenRequest( name=service_account, scope=["https://www.googleapis.com/auth/cloud-platform"], ) - response = client.generate_access_token(request=request) + response: Final = client.generate_access_token(request=request) return str(response.access_token) @@ -96,11 +96,11 @@ class GCPIAMCredentialProvider(CredentialProvider): self._gcp_service_account = gcp_service_account def get_credentials(self) -> tuple[str]: - token = _get_cached_gcp_iam_token(self._gcp_service_account) + token: Final = _get_cached_gcp_iam_token(self._gcp_service_account) return (token,) async def get_credentials_async(self) -> tuple[str]: - token = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account) + token: Final = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account) return (token,) @@ -120,13 +120,13 @@ class AzureADCredentialProvider(CredentialProvider): self._username = username def get_credentials(self) -> tuple[str] | tuple[str, str]: - token = self._credential.get_token(AZURE_REDIS_SCOPE).token + token: Final = self._credential.get_token(AZURE_REDIS_SCOPE).token if self._username: return (self._username, token) return (token,) async def get_credentials_async(self) -> tuple[str] | tuple[str, str]: - token_obj = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE) + token_obj: Final = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE) if self._username: return (self._username, token_obj.token) return (token_obj.token,) diff --git a/litellm/_service_logger.py b/litellm/_service_logger.py index 7eea82b4e74..06ae3f41c19 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, Union +from typing import TYPE_CHECKING, Any, Final, Union import litellm from litellm._logging import verbose_logger @@ -67,7 +67,7 @@ class ServiceLogging(CustomLogger): whether the callback is the logger instance itself or the ``"otel"`` string (which routes to the proxy's registered ``open_telemetry_logger``). """ - otel_v2_cls = _get_otel_v2_class() + otel_v2_cls: Final = _get_otel_v2_class() def _is_otel_logger(obj: Any) -> bool: if isinstance(obj, OpenTelemetry): @@ -101,7 +101,7 @@ class ServiceLogging(CustomLogger): try: # Try to get the current event loop - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() # Check if the loop is running if loop.is_running(): # If we're in a running loop, create a task @@ -163,7 +163,7 @@ class ServiceLogging(CustomLogger): if self.mock_testing: self.mock_testing_async_success_hook += 1 - payload = ServiceLoggerPayload( + payload: Final = ServiceLoggerPayload( is_error=False, error=None, service=service, @@ -178,7 +178,7 @@ class ServiceLogging(CustomLogger): # (the V2 logger self-registers its instance even when the string is # present, unlike V1). Without this guard each such reference emits its own # span, so a single DB call shows up as duplicate ``postgres ...`` spans. - emitted_otel_logger_ids: set = set() + emitted_otel_logger_ids: Final[set] = set() for callback in litellm.service_callback: if callback == "prometheus_system": await self.init_prometheus_services_logger_if_none() @@ -267,7 +267,7 @@ class ServiceLogging(CustomLogger): elif isinstance(error, str): error_message = error - payload = ServiceLoggerPayload( + payload: Final = ServiceLoggerPayload( is_error=True, error=error_message, service=service, @@ -278,7 +278,7 @@ class ServiceLogging(CustomLogger): # Dedupe OTel loggers per event — see ``async_service_success_hook`` for why # the same logger can be referenced twice in ``service_callback``. - emitted_otel_logger_ids: set = set() + emitted_otel_logger_ids: Final[set] = set() for callback in litellm.service_callback: if callback == "prometheus_system": await self.init_prometheus_services_logger_if_none() diff --git a/litellm/a2a_protocol/card_resolver.py b/litellm/a2a_protocol/card_resolver.py index 81a7813d14c..f6b74bcbb42 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 +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_logger from litellm.constants import LOCALHOST_URL_PATTERNS @@ -43,18 +43,18 @@ def is_localhost_or_internal_url(url: str | None) -> bool: if not url: return False - url_lower = url.lower() + url_lower: Final = url.lower() return any(pattern in url_lower for pattern in LOCALHOST_URL_PATTERNS) def get_agent_card_url(agent_card: "AgentCard") -> str | None: """Return the agent endpoint URL from the resolved SDK card.""" - url = getattr(agent_card, "url", None) + url: Final = getattr(agent_card, "url", None) if url: return url - interfaces = getattr(agent_card, "supported_interfaces", None) + interfaces: Final = getattr(agent_card, "supported_interfaces", None) if interfaces: return getattr(interfaces[0], "url", None) return None @@ -62,11 +62,11 @@ def get_agent_card_url(agent_card: "AgentCard") -> str | None: def set_agent_card_url(agent_card: "AgentCard", url: str) -> None: """Set the agent endpoint URL on the resolved SDK card.""" - normalized = url.rstrip("/") + "/" + normalized: Final = url.rstrip("/") + "/" if hasattr(agent_card, "url"): agent_card.url = normalized - interfaces = getattr(agent_card, "supported_interfaces", None) + interfaces: Final = getattr(agent_card, "supported_interfaces", None) if interfaces: interfaces[0].url = normalized @@ -86,16 +86,16 @@ def fix_agent_card_url(agent_card: "AgentCard", base_url: str) -> "AgentCard": Returns: The agent card with the URL fixed if necessary """ - card_url = getattr(agent_card, "url", None) + card_url: Final = getattr(agent_card, "url", None) if card_url and is_localhost_or_internal_url(card_url): # Normalize base_url to ensure it ends with / - fixed_url = base_url.rstrip("/") + "/" + fixed_url: Final = base_url.rstrip("/") + "/" agent_card.url = fixed_url - interfaces = getattr(agent_card, "supported_interfaces", None) + interfaces: Final = getattr(agent_card, "supported_interfaces", None) if interfaces: - interface_url = getattr(interfaces[0], "url", None) + interface_url: Final = getattr(interfaces[0], "url", None) if interface_url and is_localhost_or_internal_url(interface_url): interfaces[0].url = base_url.rstrip("/") + "/" @@ -140,7 +140,7 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc] ) # Try both well-known paths - paths = [ + paths: Final = [ AGENT_CARD_WELL_KNOWN_PATH, PREV_AGENT_CARD_WELL_KNOWN_PATH, ] diff --git a/litellm/a2a_protocol/client.py b/litellm/a2a_protocol/client.py index 8fbd4b8b81c..0ded50d25b7 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 +from typing import TYPE_CHECKING, Final from litellm.types.agents import LiteLLMSendMessageResponse @@ -92,7 +92,7 @@ class A2AClient: """Send a message to the A2A agent.""" from litellm.a2a_protocol.main import asend_message - a2a_client = await self._get_client() + a2a_client: Final = await self._get_client() return await asend_message(a2a_client=a2a_client, request=request) async def send_message_streaming( @@ -101,6 +101,6 @@ class A2AClient: """Send a streaming message to the A2A agent.""" from litellm.a2a_protocol.main import asend_message_streaming - a2a_client = await self._get_client() + a2a_client: Final = await self._get_client() async for chunk in asend_message_streaming(a2a_client=a2a_client, request=request): yield chunk diff --git a/litellm/a2a_protocol/cost_calculator.py b/litellm/a2a_protocol/cost_calculator.py index 7e6c20a31e5..31ca81c44dd 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 +from typing import TYPE_CHECKING, Any, Final if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( @@ -42,23 +42,23 @@ class A2ACostCalculator: if litellm_logging_obj is None: return 0.0 - model_call_details = litellm_logging_obj.model_call_details + model_call_details: Final = litellm_logging_obj.model_call_details # Check if user set a custom response cost (backward compatibility) - response_cost = model_call_details.get("response_cost", None) + response_cost: Final = model_call_details.get("response_cost", None) if response_cost is not None: return float(response_cost) # Get litellm_params for cost parameters - litellm_params = model_call_details.get("litellm_params", {}) or {} + litellm_params: Final = model_call_details.get("litellm_params", {}) or {} # Check for cost_per_query (fixed cost per query) if litellm_params.get("cost_per_query") is not None: return float(litellm_params["cost_per_query"]) # Check for token-based pricing - input_cost_per_token = litellm_params.get("input_cost_per_token") - output_cost_per_token = litellm_params.get("output_cost_per_token") + input_cost_per_token: Final = litellm_params.get("input_cost_per_token") + output_cost_per_token: Final = litellm_params.get("output_cost_per_token") if input_cost_per_token is not None or output_cost_per_token is not None: return A2ACostCalculator._calculate_token_based_cost( @@ -88,16 +88,16 @@ class A2ACostCalculator: float: The calculated cost """ # Get usage from model_call_details - usage = model_call_details.get("usage") + usage: Final = model_call_details.get("usage") if usage is None: return 0.0 # Get token counts - prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0 - completion_tokens = getattr(usage, "completion_tokens", 0) or 0 + prompt_tokens: Final = getattr(usage, "prompt_tokens", 0) or 0 + completion_tokens: Final = getattr(usage, "completion_tokens", 0) or 0 # Calculate costs - input_cost = prompt_tokens * (float(input_cost_per_token) if input_cost_per_token else 0.0) - output_cost = completion_tokens * (float(output_cost_per_token) if output_cost_per_token else 0.0) + input_cost: Final = prompt_tokens * (float(input_cost_per_token) if input_cost_per_token else 0.0) + output_cost: Final = completion_tokens * (float(output_cost_per_token) if output_cost_per_token else 0.0) return input_cost + output_cost diff --git a/litellm/a2a_protocol/exception_mapping_utils.py b/litellm/a2a_protocol/exception_mapping_utils.py index 16979667fe5..d2c4cdf7a65 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 +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_logger from litellm.a2a_protocol.card_resolver import ( @@ -53,7 +53,7 @@ class A2AExceptionCheckers: if not isinstance(error_str, str): return False - error_str_lower = error_str.lower() + error_str_lower: Final = error_str.lower() return any(pattern in error_str_lower for pattern in CONNECTION_ERROR_PATTERNS) @staticmethod @@ -83,8 +83,8 @@ class A2AExceptionCheckers: if not isinstance(error_str, str): return False - error_str_lower = error_str.lower() - agent_card_patterns = [ + error_str_lower: Final = error_str.lower() + agent_card_patterns: Final = [ "agent card", "agent-card", ".well-known", @@ -118,7 +118,7 @@ def map_a2a_exception( A2AAgentCardError: If the error is related to agent card issues A2AError: For other A2A-related errors """ - error_str = str(original_exception) + error_str: Final = str(original_exception) # Check for localhost URL connection error (special case - retryable) if ( @@ -190,7 +190,7 @@ async def handle_a2a_localhost_retry( "rewrite, so the upstream URL cannot be corrected." ) - request_type = "streaming " if is_streaming else "" + request_type: Final = "streaming " if is_streaming else "" verbose_logger.warning( "A2A %srequest to '%s' failed: %s. Agent card contains localhost/internal URL. Retrying with base_url '%s'.", request_type, @@ -205,14 +205,14 @@ async def handle_a2a_localhost_retry( # Reuse the httpx client LiteLLM attached at creation. It carries this agent's # trace-id and auth headers, so a fresh client would drop them. Only clients built # by ``create_a2a_client`` have it; an externally-supplied client cannot be retried. - httpx_client = getattr(a2a_client, "_litellm_httpx_client", None) + httpx_client: Final = getattr(a2a_client, "_litellm_httpx_client", None) if httpx_client is None: raise RuntimeError( "Cannot retry A2A localhost URL fix: the client was not created by " "create_a2a_client, so no LiteLLM httpx client is attached." ) - new_client = await create_client( # pyright: ignore[reportOptionalCall] + new_client: Final = await create_client( # pyright: ignore[reportOptionalCall] agent_card, client_config=ClientConfig( # pyright: ignore[reportOptionalCall] httpx_client=httpx_client, diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 21366602d1a..9c0564ca594 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 +from typing import Any, Final import litellm from litellm._logging import verbose_logger @@ -25,10 +25,10 @@ from litellm.interactions.agents.utils import merge_agent_headers # litellm_params key carrying the authenticated principal (hashed virtual key) so # A2A provider configs can scope provider-side state (e.g. LangFlow session memory) # per key instead of trusting the client-supplied A2A contextId. -A2A_USER_API_KEY_HASH_PARAM = "litellm_a2a_user_api_key_hash" +A2A_USER_API_KEY_HASH_PARAM: Final = "litellm_a2a_user_api_key_hash" # Agent metadata fields stored in litellm_params that are not valid litellm.acompletion() kwargs -_AGENT_ONLY_PARAMS = frozenset( +_AGENT_ONLY_PARAMS: Final = frozenset( { "is_public", "agent_name", @@ -70,7 +70,7 @@ class A2ACompletionBridgeHandler: """ custom_llm_provider = litellm_params.get("custom_llm_provider") if not _skip_a2a_provider_routing: - a2a_provider_config = A2AProviderConfigManager.get_provider_config( + a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config( custom_llm_provider=custom_llm_provider, model=litellm_params.get("model"), ) @@ -87,14 +87,14 @@ class A2ACompletionBridgeHandler: ) # Extract message from params - message = params.get("message", {}) + message: Final = params.get("message", {}) # Transform A2A message to OpenAI format - openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) + openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) # Get completion params custom_llm_provider = litellm_params.get("custom_llm_provider") - model = litellm_params.get("model", "agent") + model: Final = litellm_params.get("model", "agent") # Build full model string if provider specified # Skip prepending if model already starts with the provider prefix @@ -106,14 +106,14 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A completion bridge: model=%s, api_base=%s", full_model, api_base) # Build completion params dict - completion_params: dict[str, Any] = { + completion_params: Final[dict[str, Any]] = { "model": full_model, "messages": openai_messages, "api_base": api_base, "stream": False, } # Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.) - litellm_params_to_add = { + litellm_params_to_add: Final = { k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS @@ -135,10 +135,10 @@ class A2ACompletionBridgeHandler: ) # Call litellm.acompletion - response = await litellm.acompletion(**completion_params) + response: Final = await litellm.acompletion(**completion_params) # Transform response to A2A format - a2a_response = A2ACompletionBridgeTransformation.openai_response_to_a2a_response( + a2a_response: Final = A2ACompletionBridgeTransformation.openai_response_to_a2a_response( response=response, request_id=request_id, ) @@ -179,7 +179,7 @@ class A2ACompletionBridgeHandler: """ custom_llm_provider = litellm_params.get("custom_llm_provider") if not _skip_a2a_provider_routing: - a2a_provider_config = A2AProviderConfigManager.get_provider_config( + a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config( custom_llm_provider=custom_llm_provider, model=litellm_params.get("model"), ) @@ -199,20 +199,20 @@ class A2ACompletionBridgeHandler: return # Extract message from params - message = params.get("message", {}) + message: Final = params.get("message", {}) # Create streaming context - ctx = A2AStreamingContext( + ctx: Final = A2AStreamingContext( request_id=request_id, input_message=message, ) # Transform A2A message to OpenAI format - openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) + openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) # Get completion params custom_llm_provider = litellm_params.get("custom_llm_provider") - model = litellm_params.get("model", "agent") + model: Final = litellm_params.get("model", "agent") # Build full model string if provider specified # Skip prepending if model already starts with the provider prefix @@ -224,14 +224,14 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A completion bridge streaming: model=%s, api_base=%s", full_model, api_base) # Build completion params dict - completion_params: dict[str, Any] = { + completion_params: Final[dict[str, Any]] = { "model": full_model, "messages": openai_messages, "api_base": api_base, "stream": True, } # Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.) - litellm_params_to_add = { + litellm_params_to_add: Final = { k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS @@ -253,11 +253,11 @@ class A2ACompletionBridgeHandler: ) # 1. Emit initial task event (kind: "task", status: "submitted") - task_event = A2ACompletionBridgeTransformation.create_task_event(ctx) + task_event: Final = A2ACompletionBridgeTransformation.create_task_event(ctx) yield task_event # 2. Emit status update (kind: "status-update", status: "working") - working_event = A2ACompletionBridgeTransformation.create_status_update_event( + working_event: Final = A2ACompletionBridgeTransformation.create_status_update_event( ctx=ctx, state="working", final=False, @@ -266,7 +266,7 @@ class A2ACompletionBridgeHandler: yield working_event # Call litellm.acompletion with streaming - response = await litellm.acompletion(**completion_params) + response: Final = await litellm.acompletion(**completion_params) # 3. Accumulate content and emit artifact update accumulated_text = "" @@ -286,14 +286,14 @@ class A2ACompletionBridgeHandler: # Emit artifact update with accumulated content if accumulated_text: - artifact_event = A2ACompletionBridgeTransformation.create_artifact_update_event( + artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( ctx=ctx, text=accumulated_text, ) yield artifact_event # 4. Emit final status update (kind: "status-update", status: "completed", final: true) - completed_event = A2ACompletionBridgeTransformation.create_status_update_event( + completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event( ctx=ctx, state="completed", final=True, diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index e216abb6c6d..15cf77708f9 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 +from typing import Any, Final from uuid import uuid4 from litellm._logging import verbose_logger @@ -48,7 +48,7 @@ class A2ACompletionBridgeTransformation: @staticmethod 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: Final[list[str]] = [] for part in parts: if not isinstance(part, dict): continue @@ -71,10 +71,10 @@ class A2ACompletionBridgeTransformation: 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: Final[dict[str, Any]] = {} if params and isinstance(params.get("metadata"), dict): merged.update(params["metadata"]) - message_metadata = a2a_message.get("metadata") + message_metadata: Final = a2a_message.get("metadata") if isinstance(message_metadata, dict): merged.update(message_metadata) return merged or None @@ -90,7 +90,7 @@ class A2ACompletionBridgeTransformation: Uses ``extra_body`` so we do not collide with LiteLLM's spend-log ``metadata`` kwarg. """ - forward_metadata = A2ACompletionBridgeTransformation.get_forward_metadata( + forward_metadata: Final = A2ACompletionBridgeTransformation.get_forward_metadata( a2a_message=a2a_message, params=params, ) @@ -103,9 +103,9 @@ class A2ACompletionBridgeTransformation: # Layer client-supplied A2A metadata under any agent-owner-configured # ``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_metadata: Final = extra_body.get("metadata") + existing_dict: Final[dict[str, Any]] = existing_metadata if isinstance(existing_metadata, dict) else {} + merged_metadata: Final[dict[str, Any]] = {**forward_metadata, **existing_dict} extra_body = {**extra_body, "metadata": merged_metadata} completion_params["extra_body"] = extra_body @@ -124,7 +124,7 @@ class A2ACompletionBridgeTransformation: Returns: List of OpenAI-format messages """ - role = a2a_message.get("role", "user") + role: Final = a2a_message.get("role", "user") parts = a2a_message.get("parts", []) # Map A2A roles to OpenAI roles @@ -139,11 +139,11 @@ class A2ACompletionBridgeTransformation: if not isinstance(parts, list): parts = [] - content = A2ACompletionBridgeTransformation._extract_text_from_a2a_parts(parts) + content: Final = A2ACompletionBridgeTransformation._extract_text_from_a2a_parts(parts) # 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: Final[dict[str, Any]] = {"role": openai_role, "content": content} verbose_logger.debug( "A2A -> OpenAI transform: role=%s -> %s, content_length=%s", role, openai_role, len(content) @@ -169,12 +169,12 @@ class A2ACompletionBridgeTransformation: # Extract content from response content = "" if hasattr(response, "choices") and response.choices: - choice = response.choices[0] + choice: Final = response.choices[0] if hasattr(choice, "message") and choice.message: content = choice.message.content or "" # Build A2A message - a2a_message = { + a2a_message: Final = { "kind": "message", "role": "agent", "parts": [{"kind": "text", "text": content}], @@ -182,7 +182,7 @@ class A2ACompletionBridgeTransformation: } # Build A2A response - a2a_response = { + a2a_response: Final = { "jsonrpc": "2.0", "id": request_id, "result": a2a_message, @@ -245,7 +245,7 @@ class A2ACompletionBridgeTransformation: final: Whether this is the final event message_text: Optional message text for 'working' status """ - status: dict[str, Any] = { + status: Final[dict[str, Any]] = { "state": state, "timestamp": A2ACompletionBridgeTransformation._get_timestamp(), } diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index ec2d3ccf1f7..6b2541bc8a9 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -13,12 +13,7 @@ import asyncio import datetime import uuid from collections.abc import AsyncIterator, Coroutine -from typing import ( - TYPE_CHECKING, - Any, - Optional, - cast, -) +from typing import TYPE_CHECKING, Any, Final, Optional, cast import litellm from litellm._logging import verbose_logger, verbose_proxy_logger @@ -80,7 +75,7 @@ from litellm.a2a_protocol.exception_mapping_utils import ( from litellm.a2a_protocol.exceptions import A2ALocalhostURLError # Use our custom resolver instead of the default A2A SDK resolver -A2ACardResolver = LiteLLMA2ACardResolver +A2ACardResolver: Final = LiteLLMA2ACardResolver def _set_usage_on_logging_obj( @@ -96,9 +91,9 @@ def _set_usage_on_logging_obj( prompt_tokens: Number of input tokens completion_tokens: Number of output tokens """ - litellm_logging_obj = kwargs.get("litellm_logging_obj") + litellm_logging_obj: Final = kwargs.get("litellm_logging_obj") if litellm_logging_obj is not None: - usage = litellm.Usage( + usage: Final = litellm.Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens, @@ -120,13 +115,13 @@ def _set_agent_id_on_logging_obj( if agent_id is None: return - litellm_logging_obj = kwargs.get("litellm_logging_obj") + litellm_logging_obj: Final = kwargs.get("litellm_logging_obj") if litellm_logging_obj is not None: # Set agent_id directly on model_call_details (same pattern as custom_llm_provider) litellm_logging_obj.model_call_details["agent_id"] = agent_id -_A2A_COST_PARAM_KEYS = ("cost_per_query", "input_cost_per_token", "output_cost_per_token") +_A2A_COST_PARAM_KEYS: Final = ("cost_per_query", "input_cost_per_token", "output_cost_per_token") def _set_litellm_params_on_logging_obj( @@ -141,7 +136,7 @@ def _set_litellm_params_on_logging_obj( litellm_params already carries metadata / proxy_server_request / user-key context, so merge the pricing keys in rather than replacing the dict. """ - logging_obj = kwargs.get("litellm_logging_obj") + logging_obj: Final = kwargs.get("litellm_logging_obj") if logging_obj is None: return @@ -149,7 +144,7 @@ def _set_litellm_params_on_logging_obj( if not cost_params: return - existing = logging_obj.model_call_details.get("litellm_params") or {} + existing: Final = logging_obj.model_call_details.get("litellm_params") or {} logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params} @@ -162,17 +157,17 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: dict[str, Any]) -> str: """ agent_name = "unknown" - agent_card = _get_a2a_client_agent_card(a2a_client) + agent_card: Final = _get_a2a_client_agent_card(a2a_client) if agent_card is not None: agent_name = getattr(agent_card, "name", "unknown") or "unknown" # Build model string - model = f"a2a_agent/{agent_name}" - custom_llm_provider = "a2a_agent" + model: Final = f"a2a_agent/{agent_name}" + custom_llm_provider: Final = "a2a_agent" # Set on litellm_logging_obj if available (for standard logging payload) - litellm_logging_obj = kwargs.get("litellm_logging_obj") + litellm_logging_obj: Final = kwargs.get("litellm_logging_obj") if litellm_logging_obj is not None: litellm_logging_obj.model = model litellm_logging_obj.custom_llm_provider = custom_llm_provider @@ -212,7 +207,7 @@ async def _send_message_via_completion_bridge( params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params) - response_dict = await A2ACompletionBridgeHandler.handle_non_streaming( + response_dict: Final = await A2ACompletionBridgeHandler.handle_non_streaming( request_id=str(request.id), params=params, litellm_params=litellm_params, @@ -230,18 +225,18 @@ async def _send_message(a2a_client: "A2AClientType", request: "SendMessageReques "The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk" ) - pb_request = _a2a_conversions.to_core_send_message_request(request) + pb_request: Final = _a2a_conversions.to_core_send_message_request(request) last_event = None async for event in a2a_client.send_message(pb_request): last_event = event if last_event is None: raise RuntimeError("A2A send_message failed: no response received from agent.") - stream_compat = _a2a_conversions.to_compat_stream_response( + stream_compat: Final = _a2a_conversions.to_compat_stream_response( last_event, request_id=request.id, ) - result = stream_compat.result + result: Final = stream_compat.result if not isinstance(result, (Message, Task)): raise RuntimeError( "A2A send_message failed: non-streaming message/send expects the " @@ -305,7 +300,7 @@ async def _stream_messages( "The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk" ) - pb_request = _a2a_conversions.to_core_send_message_request(request) + pb_request: Final = _a2a_conversions.to_core_send_message_request(request) async for event in a2a_client.send_message(pb_request): compat_chunk = _a2a_conversions.to_compat_stream_response( event, @@ -425,9 +420,9 @@ async def asend_message( ``` """ litellm_params = litellm_params or {} - logging_obj = kwargs.get("litellm_logging_obj") + logging_obj: Final = kwargs.get("litellm_logging_obj") trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None - custom_llm_provider = litellm_params.get("custom_llm_provider") + custom_llm_provider: Final = litellm_params.get("custom_llm_provider") # Route through completion bridge if custom_llm_provider is set if custom_llm_provider: @@ -450,7 +445,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: Final[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) @@ -461,15 +456,15 @@ async def asend_message( # Type assertion: a2a_client is guaranteed to be non-None here assert a2a_client is not None - agent_name = _get_a2a_model_info(a2a_client, kwargs) + agent_name: Final = _get_a2a_model_info(a2a_client, kwargs) verbose_logger.info("A2A send_message request_id=%s, agent=%s", request.id, agent_name) # Get agent card URL for localhost retry logic - agent_card = _get_a2a_client_agent_card(a2a_client) - card_url = get_agent_card_url(agent_card) if agent_card else None + agent_card: Final = _get_a2a_client_agent_card(a2a_client) + card_url: Final = get_agent_card_url(agent_card) if agent_card else None - a2a_response = await _execute_a2a_send_with_retry( + a2a_response: Final = await _execute_a2a_send_with_retry( a2a_client=a2a_client, request=request, agent_card=agent_card, @@ -481,10 +476,10 @@ async def asend_message( verbose_logger.info("A2A send_message completed, request_id=%s", request.id) # Wrap in LiteLLM response type for _hidden_params support - response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id)) + response: Final = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id)) # Calculate token usage from request and response - response_dict = a2a_response.model_dump(mode="json", exclude_none=True) + response_dict: Final = a2a_response.model_dump(mode="json", exclude_none=True) ( prompt_tokens, completion_tokens, @@ -549,10 +544,10 @@ def _build_streaming_logging_obj( proxy_server_request: dict[str, Any] | None, ) -> Logging: """Build logging object for streaming A2A requests.""" - start_time = datetime.datetime.now() - model = f"a2a_agent/{agent_name}" + start_time: Final = datetime.datetime.now() + model: Final = f"a2a_agent/{agent_name}" - logging_obj = Logging( + logging_obj: Final = Logging( model=model, messages=[{"role": "user", "content": "streaming-request"}], stream=False, @@ -569,7 +564,7 @@ def _build_streaming_logging_obj( if agent_id: logging_obj.model_call_details["agent_id"] = agent_id - _litellm_params = litellm_params.copy() if litellm_params else {} + _litellm_params: Final = litellm_params.copy() if litellm_params else {} if metadata: _litellm_params["metadata"] = metadata if proxy_server_request: @@ -632,7 +627,7 @@ async def asend_message_streaming( ``` """ litellm_params = litellm_params or {} - custom_llm_provider = litellm_params.get("custom_llm_provider") + custom_llm_provider: Final = litellm_params.get("custom_llm_provider") # Route through completion bridge if custom_llm_provider is set if custom_llm_provider: @@ -647,7 +642,7 @@ async def asend_message_streaming( ) # Extract params from request - params = ( + params: Final = ( request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params) ) @@ -664,15 +659,15 @@ async def asend_message_streaming( if request is None: raise ValueError("request is required") - _raw_logging_obj = kwargs.get("litellm_logging_obj") + _raw_logging_obj: Final = kwargs.get("litellm_logging_obj") logging_obj: Logging | None = _raw_logging_obj if isinstance(_raw_logging_obj, Logging) else None if a2a_client is None: if api_base is None: raise ValueError("Either a2a_client or api_base is required for standard A2A flow") - logging_trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None - trace_id = logging_trace_id or (str(request.id) if request.id else str(uuid.uuid4())) - extra_headers: dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id} + logging_trace_id: Final = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None + trace_id: Final = logging_trace_id or (str(request.id) if request.id else str(uuid.uuid4())) + extra_headers: Final[dict[str, str]] = {"X-LiteLLM-Trace-Id": trace_id} if agent_id: extra_headers["X-LiteLLM-Agent-Id"] = agent_id if agent_extra_headers: @@ -685,7 +680,7 @@ async def asend_message_streaming( assert a2a_client is not None - agent_name = _get_a2a_model_info(a2a_client, kwargs) + agent_name: Final = _get_a2a_model_info(a2a_client, kwargs) if logging_obj is None: logging_obj = _build_streaming_logging_obj( @@ -699,10 +694,10 @@ async def asend_message_streaming( verbose_logger.info("A2A send_message_streaming request_id=%s, agent=%s", request.id, agent_name) - agent_card = _get_a2a_client_agent_card(a2a_client) - card_url = get_agent_card_url(agent_card) if agent_card else None + agent_card: Final = _get_a2a_client_agent_card(a2a_client) + card_url: Final = get_agent_card_url(agent_card) if agent_card else None - stream = _execute_a2a_stream_with_retry( + stream: Final = _execute_a2a_stream_with_retry( a2a_client=a2a_client, request=request, agent_card=agent_card, @@ -769,21 +764,21 @@ async def create_a2a_client( # Only pass params that AsyncHTTPHandler.__init__ accepts (e.g. timeout). # Use "disable_aiohttp_transport" key for cache-key-only data (it's # filtered out before reaching the constructor). - _client_params: dict = {"timeout": timeout} + _client_params: Final[dict] = {"timeout": timeout} if extra_headers: # Encode headers into a cache-key-only param so each unique header # set produces a distinct cache key. _client_params["disable_aiohttp_transport"] = str(sorted(extra_headers.items())) - _async_handler = get_async_httpx_client( + _async_handler: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.A2AProvider, params=_client_params, ) - httpx_client = _async_handler.client + httpx_client: Final = _async_handler.client if extra_headers: httpx_client.headers.update(extra_headers) verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys())) - a2a_client = await create_client( # pyright: ignore[reportOptionalCall] + a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall] base_url, client_config=ClientConfig( # pyright: ignore[reportOptionalCall] httpx_client=httpx_client, @@ -794,7 +789,7 @@ async def create_a2a_client( # the configured httpx client (with this agent's trace-id/auth headers) without # excavating a2a-sdk private internals. a2a_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined] - agent_card = getattr(a2a_client, "_card", None) + agent_card: Final = getattr(a2a_client, "_card", None) if agent_card is not None: a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined] @@ -827,17 +822,17 @@ async def aget_agent_card( verbose_logger.info("Fetching agent card from %s", base_url) # Use LiteLLM's cached httpx client - http_handler = get_async_httpx_client( + http_handler: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.A2A, params={"timeout": timeout}, ) - httpx_client = http_handler.client + httpx_client: Final = http_handler.client - resolver = A2ACardResolver( + resolver: Final = A2ACardResolver( httpx_client=httpx_client, base_url=base_url, ) - agent_card = await resolver.get_agent_card() + agent_card: Final = await resolver.get_agent_card() verbose_logger.info("Fetched agent card: %s", agent_card.name if hasattr(agent_card, "name") else "unknown") return agent_card diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py index 9390b0a94e2..2b37c0c4906 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 +from typing import Any, Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.bedrock_agentcore.handler import ( @@ -28,7 +28,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): **kwargs, ) -> dict[str, Any]: """Handle non-streaming request to AgentCore A2A agent.""" - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)" @@ -48,7 +48,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): **kwargs, ) -> AsyncIterator[dict[str, Any]]: """Handle streaming request to AgentCore A2A agent.""" - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)" diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py index d19137ef4a9..db57072ca38 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, cast +from typing import Any, Final, cast from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( @@ -55,16 +55,16 @@ class BedrockAgentCoreA2AHandler: verbose_logger.info("BedrockAgentCore A2A: Sending non-streaming request to %s", url) - client = get_async_httpx_client( + client: Final = get_async_httpx_client( llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), ) - response = await client.post( + response: Final = await client.post( url, headers=headers, data=body, ) response.raise_for_status() - response_data = response.json() + response_data: Final = response.json() if "error" in response_data: verbose_logger.warning("BedrockAgentCore A2A: Agent returned error: %s", response_data["error"]) @@ -102,10 +102,10 @@ class BedrockAgentCoreA2AHandler: verbose_logger.info("BedrockAgentCore A2A: Sending streaming request to %s", url) - client = get_async_httpx_client( + client: Final = get_async_httpx_client( llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), ) - response = await client.post( + response: Final = await client.post( url, headers=headers, data=body, @@ -114,15 +114,15 @@ class BedrockAgentCoreA2AHandler: response.raise_for_status() # Check content type — AgentCore may return JSON instead of SSE - content_type = response.headers.get("content-type", "").lower() + content_type: Final = response.headers.get("content-type", "").lower() if "application/json" in content_type: # Single JSON response fallback (not SSE) verbose_logger.debug( "BedrockAgentCore A2A streaming: received JSON instead of SSE, yielding as single event" ) - response_body = await response.aread() - response_data = json.loads(response_body) + response_body: Final = await response.aread() + response_data: Final = json.loads(response_body) yield response_data else: # SSE stream — parse data: lines diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py index 0c1e01e7f9f..32252711997 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 +from typing import Any, Final from litellm._logging import verbose_logger from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig @@ -23,13 +23,13 @@ from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreCo # ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``; # ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and # the ``x-amz-*`` family are owned by SigV4 itself. -_RESERVED_EXACT_HEADERS = frozenset( +_RESERVED_EXACT_HEADERS: Final = frozenset( { "authorization", "host", } ) -_RESERVED_PREFIX_HEADERS: tuple[str, ...] = ( +_RESERVED_PREFIX_HEADERS: Final[tuple[str, ...]] = ( "x-amzn-bedrock-agentcore-runtime-", "x-amz-", ) @@ -47,8 +47,8 @@ def _filter_reserved_headers( if not agent_extra_headers: return None - filtered: dict[str, str] = {} - dropped: list = [] + filtered: Final[dict[str, str]] = {} + dropped: Final[list] = [] for k, v in agent_extra_headers.items(): k_lower = k.lower() if k_lower in _RESERVED_EXACT_HEADERS or any(k_lower.startswith(prefix) for prefix in _RESERVED_PREFIX_HEADERS): @@ -107,19 +107,19 @@ class BedrockAgentCoreA2ATransformation: """ # Extract model and strip the "bedrock/" prefix # "bedrock/agentcore/arn:aws:..." → "agentcore/arn:aws:..." - model = litellm_params.get("model", "") + model: Final = litellm_params.get("model", "") if model.startswith("bedrock/"): agentcore_model = model[len("bedrock/") :] else: agentcore_model = model # Build optional_params from litellm_params (everything except model and custom_llm_provider) - optional_params = {k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider")} + optional_params: Final = {k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider")} - agentcore_config = AmazonAgentCoreConfig() + agentcore_config: Final = AmazonAgentCoreConfig() # Derive URL from ARN - url = agentcore_config.get_complete_url( + url: Final = agentcore_config.get_complete_url( api_base=optional_params.get("api_base"), api_key=optional_params.get("api_key"), model=agentcore_model, @@ -129,7 +129,7 @@ class BedrockAgentCoreA2ATransformation: ) # Construct JSON-RPC 2.0 envelope - json_rpc_body = { + json_rpc_body: Final = { "jsonrpc": "2.0", "method": method, "id": request_id, @@ -138,17 +138,17 @@ class BedrockAgentCoreA2ATransformation: # Set required AgentCore session headers (normally set by transform_request, # which we skip because it also builds {"prompt": "..."}) - headers: dict = {} - session_id = agentcore_config._get_runtime_session_id(optional_params) + headers: Final[dict] = {} + session_id: Final = agentcore_config._get_runtime_session_id(optional_params) headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id - runtime_user_id = agentcore_config._get_runtime_user_id(optional_params) + runtime_user_id: Final = agentcore_config._get_runtime_user_id(optional_params) if runtime_user_id: headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id # Merge per-request agent headers before signing so SigV4 covers them. # Reserved headers are stripped first to prevent client-controlled values # from spoofing the AgentCore runtime identity / SigV4 metadata. - safe_extra_headers = _filter_reserved_headers(agent_extra_headers) + safe_extra_headers: Final = _filter_reserved_headers(agent_extra_headers) if safe_extra_headers: headers.update(safe_extra_headers) diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py index da1fc1eb657..c083c0267f7 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 +from typing import Any, Final from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( @@ -50,7 +50,7 @@ class PydanticAIHandler: verbose_logger.info("Pydantic AI: Routing to Pydantic AI agent at %s", api_base) # Send request directly to Pydantic AI agent - response_data = await PydanticAITransformation.send_non_streaming_request( + response_data: Final = await PydanticAITransformation.send_non_streaming_request( api_base=api_base, request_id=request_id, params=params, @@ -95,7 +95,7 @@ class PydanticAIHandler: verbose_logger.info("Pydantic AI: Faking streaming for Pydantic AI agent at %s", api_base) # Get raw task response first (not the transformed A2A format) - raw_response = await PydanticAITransformation.send_and_get_raw_response( + raw_response: Final = await PydanticAITransformation.send_and_get_raw_response( api_base=api_base, request_id=request_id, params=params, diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py index d8b22282d3a..339da998d56 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, cast +from typing import Any, Final, cast from uuid import uuid4 from litellm._logging import verbose_logger @@ -163,7 +163,7 @@ class PydanticAITransformation: params_dict["message"]["kind"] = "message" # Build A2A JSON-RPC request using message/send method for FastA2A compatibility - a2a_request = { + a2a_request: Final = { "jsonrpc": "2.0", "id": request_id, "method": "message/send", @@ -171,16 +171,16 @@ class PydanticAITransformation: } # FastA2A uses root endpoint (/) not /messages - endpoint = api_base.rstrip("/") + endpoint: Final = api_base.rstrip("/") verbose_logger.info("Pydantic AI: Sending non-streaming request to %s", endpoint) # Send request to Pydantic AI agent using shared async HTTP client - client = get_async_httpx_client( + client: Final = get_async_httpx_client( llm_provider=cast(Any, "pydantic_ai_agent"), params={"timeout": timeout}, ) - response = await client.post( + response: Final = await client.post( endpoint, json=a2a_request, headers={ @@ -192,13 +192,13 @@ class PydanticAITransformation: response_data = response.json() # Check if task is already completed - result = response_data.get("result", {}) - status = result.get("status", {}) - state = status.get("state", "") + result: Final = response_data.get("result", {}) + status: Final = result.get("status", {}) + state: Final = status.get("state", "") if state != "completed": # Need to poll for completion - task_id = result.get("id") + task_id: Final = result.get("id") if task_id: verbose_logger.info("Pydantic AI: Task %s submitted, polling for completion...", task_id) response_data = await PydanticAITransformation._poll_for_completion( @@ -235,7 +235,7 @@ class PydanticAITransformation: Standard A2A non-streaming response format with message """ # Get raw task response - raw_response = await PydanticAITransformation._send_and_poll_raw( + raw_response: Final = await PydanticAITransformation._send_and_poll_raw( api_base=api_base, request_id=request_id, params=params, @@ -313,7 +313,7 @@ class PydanticAITransformation: full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data) # Build standard A2A message - a2a_message = { + a2a_message: Final = { "kind": "message", "role": "agent", "parts": parts if parts else [{"kind": "text", "text": full_text}], @@ -342,10 +342,10 @@ class PydanticAITransformation: Returns: Tuple of (full_text, message_id, parts) """ - result = response_data.get("result", {}) + result: Final = response_data.get("result", {}) # Try to extract from artifacts first (preferred for results) - artifacts = result.get("artifacts", []) + artifacts: Final = result.get("artifacts", []) if artifacts: for artifact in artifacts: parts = artifact.get("parts", []) @@ -356,7 +356,7 @@ class PydanticAITransformation: return text, str(uuid4()), parts # Fall back to history - get the last agent message - history = result.get("history", []) + history: Final = result.get("history", []) for msg in reversed(history): if msg.get("role") == "agent": parts = msg.get("parts", []) @@ -369,7 +369,7 @@ class PydanticAITransformation: return full_text, message_id, parts # Fall back to message field (original format) - message = result.get("message", {}) + message: Final = result.get("message", {}) if message: parts = message.get("parts", []) message_id = message.get("messageId", str(uuid4())) @@ -410,8 +410,8 @@ class PydanticAITransformation: full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data) # Extract input message from raw response for history - result = response_data.get("result", {}) - history = result.get("history", []) + result: Final = response_data.get("result", {}) + history: Final = result.get("history", []) input_message = {} for msg in history: if msg.get("role") == "user": @@ -419,14 +419,14 @@ class PydanticAITransformation: break # Generate IDs for streaming events - task_id = str(uuid4()) - context_id = str(uuid4()) - artifact_id = str(uuid4()) - input_message_id = input_message.get("messageId", str(uuid4())) + task_id: Final = str(uuid4()) + context_id: Final = str(uuid4()) + artifact_id: Final = str(uuid4()) + input_message_id: Final = input_message.get("messageId", str(uuid4())) # 1. Emit initial task event (kind: "task", status: "submitted") # Format matches A2ACompletionBridgeTransformation.create_task_event - task_event = { + task_event: Final = { "jsonrpc": "2.0", "id": request_id, "result": { @@ -452,7 +452,7 @@ class PydanticAITransformation: # 2. Emit status update (kind: "status-update", status: "working") # Format matches A2ACompletionBridgeTransformation.create_status_update_event - working_event = { + working_event: Final = { "jsonrpc": "2.0", "id": request_id, "result": { @@ -503,7 +503,7 @@ class PydanticAITransformation: await asyncio.sleep(delay_ms / 1000.0) # 4. Emit final status update (kind: "status-update", status: "completed", final: true) - completed_event = { + completed_event: Final = { "jsonrpc": "2.0", "id": request_id, "result": { diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index 7c526b89c35..ca84d3e07b4 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 +from typing import Any, Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import ( @@ -22,7 +22,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): **kwargs: Any, ) -> dict[str, Any]: """Handle a non-streaming A2A request via WXO runs API.""" - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " @@ -42,7 +42,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): **kwargs: Any, ) -> AsyncIterator[dict[str, Any]]: """Handle a streaming A2A request via WXO streaming runs API.""" - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index 2c7c04cec0b..bb29700cd46 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, NamedTuple, cast +from typing import Any, Final, NamedTuple, cast import httpx @@ -21,11 +21,11 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.types.llms.custom_http import httpxSpecialProvider -_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]] = {} +_IBM_CLOUD_IAM_URL: Final = "https://iam.cloud.ibm.com/identity/token" +_POLL_INTERVAL_S: Final = 2.0 +_MAX_POLL_ATTEMPTS: Final = 90 +_TOKEN_CACHE_TTL_BUFFER_S: Final = 60 +_token_cache: Final[dict[str, tuple[str, float]]] = {} class WXORequestParams(NamedTuple): @@ -53,14 +53,14 @@ class WatsonxOrchestrateHandler: api_key: str, username: str | None, ) -> str: - material = f"{auth_mode}:{cp4d_host}:{username or ''}:{api_key}" + material: Final = 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: 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() + expires_at: Final = int(expiration) + wall: Final = now_wall if now_wall is not None else time.time() return max(expires_at - int(wall), 0) @staticmethod @@ -71,9 +71,9 @@ class WatsonxOrchestrateHandler: 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() - cached = _token_cache.get(cache_key) + cache_key: Final = WatsonxOrchestrateHandler._token_cache_key(auth_mode, cp4d_host, api_key, username) + now: Final = time.monotonic() + cached: Final = _token_cache.get(cache_key) if cached and cached[1] > now: return cached[0] @@ -96,7 +96,7 @@ class WatsonxOrchestrateHandler: else: if not username: raise ValueError("'username' is required in litellm_params when auth_mode='cp4d'") - token_url = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize" + token_url: Final = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize" response = await client.post( token_url, json={"username": username, "api_key": api_key}, @@ -105,13 +105,13 @@ class WatsonxOrchestrateHandler: response.raise_for_status() payload = response.json() token = str(payload["token"]) - expiration = payload.get("expiration") + expiration: Final = payload.get("expiration") if expiration is None: ttl_s = 3600 else: ttl_s = WatsonxOrchestrateHandler._cp4d_token_ttl_seconds(expiration) - expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0) + expires_at: Final = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0) _token_cache[cache_key] = (token, expires_at) for stale_key, (_, stale_expires_at) in list(_token_cache.items()): if stale_expires_at <= now: @@ -127,7 +127,7 @@ class WatsonxOrchestrateHandler: max_attempts: int = _MAX_POLL_ATTEMPTS, interval_s: float = _POLL_INTERVAL_S, ) -> dict[str, Any]: - url = f"{base_url}/v1/orchestrate/runs/{run_id}" + url: Final = f"{base_url}/v1/orchestrate/runs/{run_id}" for attempt in range(max_attempts): await asyncio.sleep(interval_s) @@ -152,7 +152,7 @@ class WatsonxOrchestrateHandler: ) -> 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 "" + run_id: Final = run_data.get("run_id") or run_data.get("id") or "" if not run_id: raise ValueError(f"WXO: No run_id in response: {run_data}") run_data = await WatsonxOrchestrateHandler._poll_run( @@ -188,10 +188,10 @@ class WatsonxOrchestrateHandler: @staticmethod 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 "" - api_key = litellm_params.get("api_key") or "" + cp4d_host: Final = litellm_params.get("cp4d_host") or "" + instance_id: Final = litellm_params.get("instance_id") or "" + wxo_agent_id: Final = litellm_params.get("wxo_agent_id") or "" + api_key: Final = litellm_params.get("api_key") or "" if not cp4d_host: raise ValueError("'cp4d_host' is required in litellm_params for WXO agents") @@ -218,29 +218,29 @@ class WatsonxOrchestrateHandler: params: dict[str, Any], litellm_params: dict[str, Any], ) -> dict[str, Any]: - wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) + wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) - client = WatsonxOrchestrateHandler._http_client(timeout=90.0) - token = await WatsonxOrchestrateHandler._get_bearer_token( + client: Final = WatsonxOrchestrateHandler._http_client(timeout=90.0) + token: Final = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=wxo.cp4d_host, auth_mode=wxo.auth_mode, api_key=wxo.api_key, username=wxo.username, client=client, ) - base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id) - auth_headers = { + base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id) + auth_headers: Final = { "Authorization": f"Bearer {token}", "Content-Type": "application/json", "Accept": "application/json", } - text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params) - body = WatsonxOrchestrateTransformation.build_wxo_run_body( + text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params) + body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body( wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id ) - run_response = await client.post( + run_response: Final = await client.post( f"{base_url}/v1/orchestrate/runs", json=body, headers=auth_headers, @@ -255,7 +255,7 @@ class WatsonxOrchestrateHandler: client=client, ) - response_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(run_data) + response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(run_data) return WatsonxOrchestrateTransformation.build_a2a_message_response(request_id=request_id, text=response_text) @staticmethod @@ -266,29 +266,29 @@ class WatsonxOrchestrateHandler: chunk_size: int = 50, delay_ms: int = 10, ) -> AsyncIterator[dict[str, Any]]: - wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) + wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) - client = WatsonxOrchestrateHandler._http_client(timeout=120.0) - token = await WatsonxOrchestrateHandler._get_bearer_token( + client: Final = WatsonxOrchestrateHandler._http_client(timeout=120.0) + token: Final = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=wxo.cp4d_host, auth_mode=wxo.auth_mode, api_key=wxo.api_key, username=wxo.username, client=client, ) - base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id) - auth_headers = { + base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id) + auth_headers: Final = { "Authorization": f"Bearer {token}", "Content-Type": "application/json", "Accept": "text/event-stream, application/json", } - text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params) - body = WatsonxOrchestrateTransformation.build_wxo_run_body( + text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params) + body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body( wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id ) try: - response = await client.post( + response: Final = await client.post( f"{base_url}/v1/orchestrate/runs/stream", json=body, headers=auth_headers, @@ -306,7 +306,7 @@ class WatsonxOrchestrateHandler: params=params, litellm_params=litellm_params, ) - response_text = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result) + response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result) async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text( text=response_text, request_id=request_id, @@ -316,9 +316,9 @@ class WatsonxOrchestrateHandler: yield chunk return - content_type = response.headers.get("content-type", "").lower() + content_type: Final = response.headers.get("content-type", "").lower() if "text/event-stream" not in content_type: - response_body = await response.aread() + response_body: Final = await response.aread() result = json.loads(response_body) result = await WatsonxOrchestrateHandler._get_successful_run_data( run_data=result, diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py index 18b0795aa8a..3748d8043cc 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 +from typing import Any, Final from uuid import uuid4 from litellm._logging import verbose_logger @@ -35,9 +35,9 @@ class WatsonxOrchestrateTransformation: A2A format: params.message.parts[*] where part.kind == "text" """ - message = params.get("message", {}) - parts = message.get("parts", []) - texts = [] + message: Final = params.get("message", {}) + parts: Final = message.get("parts", []) + texts: Final = [] for part in parts: if not isinstance(part, dict): continue @@ -53,7 +53,7 @@ class WatsonxOrchestrateTransformation: thread_id: str | None = None, ) -> dict[str, Any]: """Build the WXO POST /v1/orchestrate/runs request body.""" - body: dict[str, Any] = { + body: Final[dict[str, Any]] = { "agent_id": wxo_agent_id, "message": { "role": "user", @@ -96,7 +96,7 @@ class WatsonxOrchestrateTransformation: pass # Tertiary: results as a raw string - results = result.get("results") + results: Final = result.get("results") if results and isinstance(results, str): return results @@ -104,11 +104,11 @@ class WatsonxOrchestrateTransformation: @staticmethod def extract_text_from_a2a_message_response(a2a_response: dict[str, Any]) -> str: - result = a2a_response.get("result") + result: Final = a2a_response.get("result") if not isinstance(result, dict): verbose_logger.warning("WXO: A2A response missing result object") return "" - parts = result.get("parts") + parts: Final = result.get("parts") if not isinstance(parts, list): verbose_logger.warning("WXO: A2A result has no parts list") return "" @@ -150,9 +150,9 @@ class WatsonxOrchestrateTransformation: 3. artifact-update chunks 4. status-update (kind="status-update", state="completed", final=True) """ - task_id = str(uuid4()) - context_id = str(uuid4()) - artifact_id = str(uuid4()) + task_id: Final = str(uuid4()) + context_id: Final = str(uuid4()) + artifact_id: Final = str(uuid4()) # 1. Task submitted yield { @@ -181,7 +181,7 @@ class WatsonxOrchestrateTransformation: await asyncio.sleep(delay_ms / 1000.0) # 3. Artifact chunks (always emit at least one chunk, even for empty text) - text_to_chunk = text or "" + text_to_chunk: Final = text or "" for i in range(0, max(len(text_to_chunk), 1), chunk_size): chunk_text = text_to_chunk[i : i + chunk_size] is_last = (i + chunk_size) >= max(len(text_to_chunk), 1) diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index f954cb187b5..413691f233d 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 +from typing import TYPE_CHECKING, Any, Final import litellm from litellm._logging import verbose_logger @@ -47,7 +47,7 @@ class A2AStreamingIterator: async def __anext__(self) -> "SendStreamingMessageResponse": try: - chunk = await self.stream.__anext__() + chunk: Final = await self.stream.__anext__() # Store chunk self.chunks.append(chunk) @@ -71,8 +71,8 @@ class A2AStreamingIterator: def _collect_text_from_chunk(self, chunk: Any) -> None: """Extract text from a streaming chunk and add to collected parts.""" try: - chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} - text = A2ARequestUtils.extract_text_from_response(chunk_dict) + chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} + text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict) if text: self.collected_text_parts.append(text) except Exception: @@ -81,10 +81,10 @@ class A2AStreamingIterator: def _is_completed_chunk(self, chunk: Any) -> bool: """Check if chunk indicates stream completion.""" try: - chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} - result = chunk_dict.get("result", {}) + chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} + result: Final = chunk_dict.get("result", {}) if isinstance(result, dict): - status = result.get("status", {}) + status: Final = result.get("status", {}) if isinstance(status, dict): return status.get("state") == "completed" except Exception: @@ -94,21 +94,21 @@ class A2AStreamingIterator: async def _handle_stream_complete(self) -> None: """Handle logging and token counting when stream completes.""" try: - end_time = datetime.now() + end_time: Final = datetime.now() # Calculate tokens from collected text - input_message = A2ARequestUtils.get_input_message_from_request(self.request) - input_text = A2ARequestUtils.extract_text_from_message(input_message) - prompt_tokens = A2ARequestUtils.count_tokens(input_text) + input_message: Final = A2ARequestUtils.get_input_message_from_request(self.request) + input_text: Final = A2ARequestUtils.extract_text_from_message(input_message) + prompt_tokens: Final = A2ARequestUtils.count_tokens(input_text) # Use the last (most complete) text from chunks - output_text = self.collected_text_parts[-1] if self.collected_text_parts else "" - completion_tokens = A2ARequestUtils.count_tokens(output_text) + output_text: Final = self.collected_text_parts[-1] if self.collected_text_parts else "" + completion_tokens: Final = A2ARequestUtils.count_tokens(output_text) - total_tokens = prompt_tokens + completion_tokens + total_tokens: Final = prompt_tokens + completion_tokens # Create usage object - usage = litellm.Usage( + usage: Final = litellm.Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=total_tokens, @@ -120,11 +120,11 @@ class A2AStreamingIterator: self.logging_obj.model_call_details["stream"] = False # Calculate cost using A2ACostCalculator - response_cost = A2ACostCalculator.calculate_a2a_cost(self.logging_obj) + response_cost: Final = A2ACostCalculator.calculate_a2a_cost(self.logging_obj) self.logging_obj.model_call_details["response_cost"] = response_cost # Build result for logging - result = self._build_logging_result(usage) + result: Final = self._build_logging_result(usage) # Call success handlers - they will build standard_logging_object asyncio.create_task( @@ -150,7 +150,7 @@ class A2AStreamingIterator: def _build_logging_result(self, usage: litellm.Usage) -> dict[str, Any]: """Build a result dict for logging.""" - result: dict[str, Any] = { + result: Final[dict[str, Any]] = { "id": getattr(self.request, "id", "unknown"), "jsonrpc": "2.0", "usage": (usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)), @@ -159,7 +159,7 @@ class A2AStreamingIterator: # Add final chunk result if available if self.final_chunk: try: - chunk_dict = self.final_chunk.model_dump(mode="json", exclude_none=True) + chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True) result["result"] = chunk_dict.get("result", {}) except Exception: pass diff --git a/litellm/a2a_protocol/utils.py b/litellm/a2a_protocol/utils.py index d6a45e39a02..f2e61f66105 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 +from typing import TYPE_CHECKING, Any, Final 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: Final[list[str]] = [] for part in parts: if isinstance(part, dict): if part.get("kind") == "text": @@ -56,7 +56,7 @@ class A2ARequestUtils: Returns: Text from response message parts """ - result = response_dict.get("result", {}) + result: Final = response_dict.get("result", {}) if not isinstance(result, dict): return "" @@ -66,7 +66,7 @@ class A2ARequestUtils: if result.get("kind") == "message": return A2ARequestUtils.extract_text_from_message(result) - message = result.get("message", {}) + message: Final = result.get("message", {}) return A2ARequestUtils.extract_text_from_message(message) @staticmethod @@ -82,7 +82,7 @@ class A2ARequestUtils: Returns: The message object/dict or None """ - params = getattr(request, "params", None) + params: Final = getattr(request, "params", None) if params is None: return None return getattr(params, "message", None) @@ -128,14 +128,14 @@ class A2ARequestUtils: input_message = A2ARequestUtils.get_input_message_from_request(request) if input_message is not None and hasattr(input_message, "model_dump"): input_message = input_message.model_dump(mode="json") - input_text = A2ARequestUtils.extract_text_from_message(input_message) - prompt_tokens = A2ARequestUtils.count_tokens(input_text) + input_text: Final = A2ARequestUtils.extract_text_from_message(input_message) + prompt_tokens: Final = A2ARequestUtils.count_tokens(input_text) # Count output tokens - output_text = A2ARequestUtils.extract_text_from_response(response_dict) - completion_tokens = A2ARequestUtils.count_tokens(output_text) + output_text: Final = A2ARequestUtils.extract_text_from_response(response_dict) + completion_tokens: Final = A2ARequestUtils.count_tokens(output_text) - total_tokens = prompt_tokens + completion_tokens + total_tokens: Final = prompt_tokens + completion_tokens return prompt_tokens, completion_tokens, total_tokens diff --git a/litellm/anthropic_beta_headers_manager.py b/litellm/anthropic_beta_headers_manager.py index 063a83e6f38..abce47c191e 100644 --- a/litellm/anthropic_beta_headers_manager.py +++ b/litellm/anthropic_beta_headers_manager.py @@ -25,6 +25,7 @@ Environment Variables: import json import os from importlib.resources import files +from typing import Final import httpx @@ -46,7 +47,7 @@ class GetAnthropicBetaHeadersConfig: def load_local_beta_headers_config() -> dict: """Load the local backup beta headers config bundled with the package.""" try: - content = json.loads( + content: Final = json.loads( files("litellm").joinpath("anthropic_beta_headers_config.json").read_text(encoding="utf-8") ) return content @@ -79,14 +80,14 @@ class GetAnthropicBetaHeadersConfig: return False # Check for at least one provider key - provider_keys = [ + provider_keys: Final = [ "anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai", ] - has_provider = any(key in fetched_config for key in provider_keys) + has_provider: Final = any(key in fetched_config for key in provider_keys) if not has_provider: verbose_logger.warning( @@ -113,7 +114,7 @@ class GetAnthropicBetaHeadersConfig: Returns the parsed JSON dict. Raises on network/parse errors (caller is expected to handle). """ - response = httpx.get(url, timeout=timeout) + response: Final = httpx.get(url, timeout=timeout) response.raise_for_status() return response.json() @@ -138,7 +139,7 @@ def get_beta_headers_config(url: str) -> dict: return GetAnthropicBetaHeadersConfig.load_local_beta_headers_config() try: - content = GetAnthropicBetaHeadersConfig.fetch_remote_beta_headers_config(url) + content: Final = GetAnthropicBetaHeadersConfig.fetch_remote_beta_headers_config(url) except Exception as e: verbose_logger.warning( "LiteLLM: Failed to fetch remote beta headers config from %s: %s. Falling back to local backup.", @@ -206,8 +207,8 @@ def get_provider_name(provider: str) -> str: Returns: Canonical provider name """ - config = _load_beta_headers_config() - aliases = config.get("provider_aliases", {}) + config: Final = _load_beta_headers_config() + aliases: Final = config.get("provider_aliases", {}) return aliases.get(provider, provider) @@ -233,13 +234,13 @@ def filter_and_transform_beta_headers( if not beta_headers: return [] - config = _load_beta_headers_config() + config: Final = _load_beta_headers_config() provider = get_provider_name(provider) # Get the header mapping for this provider - provider_mapping = config.get(provider, {}) + provider_mapping: Final = config.get(provider, {}) - filtered_headers: set[str] = set() + filtered_headers: Final[set[str]] = set() for header in beta_headers: header = header.strip() @@ -279,9 +280,9 @@ def is_beta_header_supported( Returns: True if the header is in the mapping with a non-null value, False otherwise """ - config = _load_beta_headers_config() + config: Final = _load_beta_headers_config() provider = get_provider_name(provider) - provider_mapping = config.get(provider, {}) + provider_mapping: Final = config.get(provider, {}) # Header is supported if it's in the mapping and has a non-null value return beta_header in provider_mapping and provider_mapping[beta_header] is not None @@ -303,11 +304,11 @@ def get_provider_beta_header( Returns: The provider-specific header name if supported, or None if unsupported/unknown """ - config = _load_beta_headers_config() + config: Final = _load_beta_headers_config() provider = get_provider_name(provider) # Get the header mapping for this provider - provider_mapping = config.get(provider, {}) + provider_mapping: Final = config.get(provider, {}) # Check if header is in the mapping if anthropic_beta_header not in provider_mapping: @@ -332,15 +333,15 @@ def update_headers_with_filtered_beta( Returns: Updated headers dict """ - existing_beta = headers.get("anthropic-beta") + existing_beta: Final = headers.get("anthropic-beta") if not existing_beta: return headers # Parse existing beta headers - beta_values = [b.strip() for b in existing_beta.split(",") if b.strip()] + beta_values: Final = [b.strip() for b in existing_beta.split(",") if b.strip()] # Filter and transform based on provider - filtered_beta_values = filter_and_transform_beta_headers( + filtered_beta_values: Final = filter_and_transform_beta_headers( beta_headers=beta_values, provider=provider, ) @@ -374,11 +375,11 @@ def update_request_with_filtered_beta( """ headers = update_headers_with_filtered_beta(headers=headers, provider=provider) - existing_body_betas = request_data.get("anthropic_beta") + existing_body_betas: Final = request_data.get("anthropic_beta") if not existing_body_betas: return headers, request_data - filtered_body_betas = filter_and_transform_beta_headers( + filtered_body_betas: Final = filter_and_transform_beta_headers( beta_headers=existing_body_betas, provider=provider, ) @@ -401,9 +402,9 @@ def get_unsupported_headers(provider: str) -> list[str]: Returns: List of unsupported Anthropic beta header names """ - config = _load_beta_headers_config() + config: Final = _load_beta_headers_config() provider = get_provider_name(provider) - provider_mapping = config.get(provider, {}) + provider_mapping: Final = config.get(provider, {}) # Return headers with null values return [header for header, value in provider_mapping.items() if value is None] diff --git a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py index c0038dd7d83..ad8be8ec40e 100644 --- a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py +++ b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py @@ -4,13 +4,15 @@ Utilities for mapping exceptions to Anthropic error format. Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format. """ +from typing import Final + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads 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: Final[dict[int, AnthropicErrorType]] = { 400: "invalid_request_error", 401: "authentication_error", 403: "permission_error", @@ -50,9 +52,9 @@ class AnthropicExceptionMapping: "request_id": "req_..." } """ - error_type = AnthropicExceptionMapping.get_error_type(status_code) + error_type: Final = AnthropicExceptionMapping.get_error_type(status_code) - response: AnthropicErrorResponse = { + response: Final[AnthropicErrorResponse] = { "type": "error", "error": { "type": error_type, @@ -76,7 +78,7 @@ class AnthropicExceptionMapping: - Generic: {"message": "..."} - Plain strings """ - parsed = safe_json_loads(raw_message) + parsed: Final = safe_json_loads(raw_message) if isinstance(parsed, dict): # Bedrock format if "detail" in parsed and isinstance(parsed["detail"], dict): diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index b476b3993d4..237e35fdd5e 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, Literal +from typing import Any, Final, Literal import httpx from openai import AsyncOpenAI, OpenAI @@ -29,8 +29,8 @@ from ..types.router import * from .utils import get_optional_params_add_message ####### ENVIRONMENT VARIABLES ################### -openai_assistants_api = OpenAIAssistantsAPI() -azure_assistants_api = AzureAssistantsAPI() +openai_assistants_api: Final = OpenAIAssistantsAPI() +azure_assistants_api: Final = AzureAssistantsAPI() ### ASSISTANTS ### @@ -40,23 +40,23 @@ async def aget_assistants( client: AsyncOpenAI | None = None, **kwargs, ) -> AsyncCursorPage[Assistant]: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["aget_assistants"] = True try: # Use a partial function to pass your keyword arguments - func = partial(get_assistants, custom_llm_provider, client, **kwargs) + func: Final = partial(get_assistants, custom_llm_provider, client, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -80,11 +80,11 @@ def get_assistants( api_version: str | None = None, **kwargs, ) -> SyncCursorPage[Assistant]: - aget_assistants: bool | None = kwargs.pop("aget_assistants", None) + aget_assistants: Final[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) - litellm_params_dict = get_litellm_params(**kwargs) + optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -95,7 +95,7 @@ def get_assistants( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -111,7 +111,7 @@ def get_assistants( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -147,7 +147,7 @@ def get_assistants( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -197,25 +197,25 @@ async def acreate_assistants( client: AsyncOpenAI | None = None, **kwargs, ) -> Assistant: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["async_create_assistants"] = True - model = kwargs.pop("model", None) + model: Final = kwargs.pop("model", None) try: kwargs["client"] = client # Use a partial function to pass your keyword arguments - func = partial(create_assistants, custom_llm_provider, model, **kwargs) + func: Final = partial(create_assistants, custom_llm_provider, model, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model=model, custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -249,11 +249,11 @@ def create_assistants( api_version: str | None = None, **kwargs, ) -> Assistant | Coroutine[Any, Any, Assistant]: - async_create_assistants: bool | None = kwargs.pop("async_create_assistants", None) + async_create_assistants: Final[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) - litellm_params_dict = get_litellm_params(**kwargs) + optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -264,7 +264,7 @@ def create_assistants( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -296,7 +296,7 @@ def create_assistants( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -333,7 +333,7 @@ def create_assistants( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -380,24 +380,24 @@ async def adelete_assistant( client: AsyncOpenAI | None = None, **kwargs, ) -> AssistantDeleted: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["async_delete_assistants"] = True try: kwargs["client"] = client # Use a partial function to pass your keyword arguments - func = partial(delete_assistant, custom_llm_provider, **kwargs) + func: Final = partial(delete_assistant, custom_llm_provider, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -422,11 +422,11 @@ def delete_assistant( api_version: str | None = None, **kwargs, ) -> AssistantDeleted | Coroutine[Any, Any, AssistantDeleted]: - optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) + optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) - async_delete_assistants: bool | None = kwargs.pop("async_delete_assistants", None) + async_delete_assistants: Final[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") @@ -439,7 +439,7 @@ def delete_assistant( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -455,7 +455,7 @@ def delete_assistant( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None ) # set API KEY @@ -484,7 +484,7 @@ def delete_assistant( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -530,23 +530,23 @@ def delete_assistant( async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwargs) -> Thread: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["acreate_thread"] = True try: # Use a partial function to pass your keyword arguments - func = partial(create_thread, custom_llm_provider, **kwargs) + func: Final = partial(create_thread, custom_llm_provider, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -592,9 +592,9 @@ def create_thread( ) ``` """ - acreate_thread = kwargs.get("acreate_thread", None) - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + acreate_thread: Final = kwargs.get("acreate_thread", None) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -605,7 +605,7 @@ def create_thread( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -624,7 +624,7 @@ def create_thread( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -661,7 +661,7 @@ def create_thread( 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", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -704,23 +704,23 @@ async def aget_thread( client: AsyncOpenAI | None = None, **kwargs, ) -> Thread: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["aget_thread"] = True try: # Use a partial function to pass your keyword arguments - func = partial(get_thread, custom_llm_provider, thread_id, client, **kwargs) + func: Final = partial(get_thread, custom_llm_provider, thread_id, client, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -743,9 +743,9 @@ def get_thread( **kwargs, ) -> Thread: """Get the thread object, given a thread_id""" - aget_thread = kwargs.pop("aget_thread", None) - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + aget_thread: Final = kwargs.pop("aget_thread", None) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -755,7 +755,7 @@ def get_thread( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -772,7 +772,7 @@ def get_thread( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -809,7 +809,7 @@ def get_thread( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -858,12 +858,12 @@ async def a_add_message( client=None, **kwargs, ) -> OpenAIMessage: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["a_add_message"] = True try: # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( add_message, custom_llm_provider, thread_id, @@ -876,15 +876,15 @@ async def a_add_message( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -912,12 +912,12 @@ def add_message( **kwargs, ) -> OpenAIMessage: ### COMMON OBJECTS ### - a_add_message = kwargs.pop("a_add_message", None) - _message_data = MessageData(role=role, content=content, attachments=attachments, metadata=metadata) - litellm_params_dict = get_litellm_params(**kwargs) - optional_params = GenericLiteLLMParams(**kwargs) + a_add_message: Final = kwargs.pop("a_add_message", None) + _message_data: Final = MessageData(role=role, content=content, attachments=attachments, metadata=metadata) + litellm_params_dict: Final = get_litellm_params(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) - message_data = get_optional_params_add_message( + message_data: Final = get_optional_params_add_message( role=_message_data["role"], content=_message_data["content"], attachments=_message_data["attachments"], @@ -934,7 +934,7 @@ def add_message( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -951,7 +951,7 @@ def add_message( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -988,7 +988,7 @@ def add_message( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -1029,12 +1029,12 @@ async def aget_messages( client: AsyncOpenAI | None = None, **kwargs, ) -> AsyncCursorPage[OpenAIMessage]: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["aget_messages"] = True try: # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( get_messages, custom_llm_provider, thread_id, @@ -1043,15 +1043,15 @@ async def aget_messages( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -1074,9 +1074,9 @@ def get_messages( client: Any | None = None, **kwargs, ) -> SyncCursorPage[OpenAIMessage]: - aget_messages = kwargs.pop("aget_messages", None) - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + aget_messages: Final = kwargs.pop("aget_messages", None) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -1087,7 +1087,7 @@ def get_messages( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -1105,7 +1105,7 @@ def get_messages( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -1141,7 +1141,7 @@ def get_messages( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token: str | None = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) @@ -1198,12 +1198,12 @@ async def arun_thread( client: Any | None = None, **kwargs, ) -> Run: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() ### PASS ARGS TO GET ASSISTANTS ### kwargs["arun_thread"] = True try: # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( run_thread, custom_llm_provider, thread_id, @@ -1219,15 +1219,15 @@ async def arun_thread( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model="", custom_llm_provider=custom_llm_provider ) # type: ignore # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -1267,9 +1267,9 @@ def run_thread( **kwargs, ) -> Run: """Run a given thread + assistant.""" - arun_thread = kwargs.pop("arun_thread", None) - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + arun_thread: Final = kwargs.pop("arun_thread", None) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -1280,7 +1280,7 @@ def run_thread( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -1296,7 +1296,7 @@ def run_thread( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -1341,7 +1341,7 @@ def run_thread( or get_secret("AZURE_API_KEY") ) # type: ignore - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) azure_ad_token = None if extra_body is not None: azure_ad_token = extra_body.pop("azure_ad_token", None) diff --git a/litellm/assistants/utils.py b/litellm/assistants/utils.py index d23dcd973b4..e80131a5011 100644 --- a/litellm/assistants/utils.py +++ b/litellm/assistants/utils.py @@ -1,3 +1,5 @@ +from typing import Final + import litellm from ..exceptions import UnsupportedParamsError @@ -17,13 +19,13 @@ def get_optional_params_add_message( Reference - https://learn.microsoft.com/en-us/azure/ai-services/openai/assistants-reference-messages?tabs=python#create-message """ - passed_params = locals() + passed_params: Final = locals() custom_llm_provider = passed_params.pop("custom_llm_provider") - special_params = passed_params.pop("kwargs") + special_params: Final = passed_params.pop("kwargs") for k, v in special_params.items(): passed_params[k] = v - default_params = { + default_params: Final = { "role": None, "content": None, "attachments": None, @@ -36,7 +38,7 @@ def get_optional_params_add_message( ## raise exception if non-default value passed for non-openai/azure embedding calls def _check_valid_arg(supported_params): if len(non_default_params.keys()) > 0: - keys = list(non_default_params.keys()) + keys: Final = list(non_default_params.keys()) for k in keys: if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values non_default_params.pop(k, None) @@ -50,7 +52,7 @@ def get_optional_params_add_message( if custom_llm_provider == "openai": optional_params = non_default_params elif custom_llm_provider == "azure": - supported_params = litellm.AzureOpenAIAssistantsAPIConfig().get_supported_openai_create_message_params() + supported_params: Final = litellm.AzureOpenAIAssistantsAPIConfig().get_supported_openai_create_message_params() _check_valid_arg(supported_params=supported_params) optional_params = litellm.AzureOpenAIAssistantsAPIConfig().map_openai_params_create_message_params( non_default_params=non_default_params, optional_params=optional_params @@ -72,13 +74,13 @@ def get_optional_params_image_gen( **kwargs, ): # retrieve all parameters passed to the function - passed_params = locals() + passed_params: Final = locals() custom_llm_provider = passed_params.pop("custom_llm_provider") - special_params = passed_params.pop("kwargs") + special_params: Final = passed_params.pop("kwargs") for k, v in special_params.items(): passed_params[k] = v - default_params = { + default_params: Final = { "n": None, "quality": None, "response_format": None, @@ -93,7 +95,7 @@ def get_optional_params_image_gen( ## raise exception if non-default value passed for non-openai/azure embedding calls def _check_valid_arg(supported_params): if len(non_default_params.keys()) > 0: - keys = list(non_default_params.keys()) + keys: Final = list(non_default_params.keys()) for k in keys: if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values non_default_params.pop(k, None) diff --git a/litellm/batch_completion/main.py b/litellm/batch_completion/main.py index fb892789b15..702dd194fda 100644 --- a/litellm/batch_completion/main.py +++ b/litellm/batch_completion/main.py @@ -1,4 +1,5 @@ from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait +from typing import Final import litellm from litellm._logging import print_verbose @@ -55,17 +56,17 @@ def batch_completion( Returns: list: A list of completion results. """ - args = locals() + args: Final = locals() - batch_messages = messages - completions = [] + batch_messages: Final = messages + completions: Final = [] model = model custom_llm_provider = None if model.split("/", 1)[0] in litellm.provider_list: custom_llm_provider = model.split("/", 1)[0] model = model.split("/", 1)[1] if custom_llm_provider == "vllm": - optional_params = get_optional_params( + optional_params: Final = get_optional_params( functions=functions, function_call=function_call, temperature=temperature, @@ -145,7 +146,7 @@ def batch_completion_models(*args, **kwargs): if "model" in kwargs: kwargs.pop("model") if "models" in kwargs: - models = kwargs["models"] + models: Final = kwargs["models"] kwargs.pop("models") futures = {} with ThreadPoolExecutor(max_workers=len(models)) as executor: @@ -156,10 +157,10 @@ def batch_completion_models(*args, **kwargs): if future.result() is not None: return future.result() elif "deployments" in kwargs: - deployments = kwargs["deployments"] + deployments: Final = kwargs["deployments"] kwargs.pop("deployments") kwargs.pop("model_list") - nested_kwargs = kwargs.pop("kwargs", {}) + nested_kwargs: Final = kwargs.pop("kwargs", {}) futures = {} with ThreadPoolExecutor(max_workers=len(deployments)) as executor: for deployment in deployments: @@ -238,10 +239,10 @@ def batch_completion_models_all_responses(*args, **kwargs): if len(models) == 0: return [] - responses = [] + responses: Final = [] with concurrent.futures.ThreadPoolExecutor(max_workers=len(models)) as executor: - futures = [executor.submit(litellm.completion, *args, model=model, **kwargs) for model in models] + futures: Final = [executor.submit(litellm.completion, *args, model=model, **kwargs) for model in models] for future in futures: try: diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 4f90b50eaa2..4835fd722bc 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, Literal +from typing import Any, Final, Literal import litellm from litellm._logging import verbose_logger @@ -141,9 +141,9 @@ def _aggregate_batch_cost_usage_models( ) -> 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)) + line_stats: Final = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info)) - cache_token_params = { + cache_token_params: Final = { key: tokens for key, tokens in ( ("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)), @@ -151,14 +151,14 @@ def _aggregate_batch_cost_usage_models( ) if tokens > 0 } - batch_usage = Usage( + batch_usage: Final = Usage( total_tokens=sum(stats.total_tokens for stats in line_stats), prompt_tokens=sum(stats.prompt_tokens for stats in line_stats), completion_tokens=sum(stats.completion_tokens for stats in line_stats), **cache_token_params, ) - batch_models = [model_name] if model_name else [stats.model for stats in line_stats if stats.model] - total_cost = sum((stats.cost for stats in line_stats), 0.0) + batch_models: Final = [model_name] if model_name else [stats.model for stats in line_stats if stats.model] + total_cost: Final = sum((stats.cost for stats in line_stats), 0.0) verbose_logger.debug("batch output aggregate: cost=%s usage=%s models=%s", total_cost, batch_usage, batch_models) return total_cost, batch_usage, batch_models @@ -184,7 +184,7 @@ def calculate_vertex_ai_batch_cost_and_usage( total_tokens = 0 prompt_tokens = 0 completion_tokens = 0 - actual_model_name = model_name or "gemini-2.0-flash-001" + actual_model_name: Final = model_name or "gemini-2.0-flash-001" for response in vertex_ai_batch_responses: response_body = response.get("response") @@ -254,7 +254,7 @@ async def _fetch_batch_output_file_content( raise ValueError("Output file id is None cannot retrieve file content") file_id = batch.output_file_id - is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: try: file_id = is_base64_unified_file_id.split("llm_output_file_id,")[1].split(";")[0] @@ -265,16 +265,16 @@ async def _fetch_batch_output_file_content( ) # Build kwargs for afile_content with credentials from litellm_params - file_content_kwargs = { + file_content_kwargs: Final = { "file_id": file_id, "custom_llm_provider": custom_llm_provider, } # Extract and add credentials for file access - credentials = _extract_file_access_credentials(litellm_params) + credentials: Final = _extract_file_access_credentials(litellm_params) file_content_kwargs.update(credentials) - _file_content = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType] + _file_content: Final = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType] return _file_content.content @@ -291,11 +291,11 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: Returns: Dictionary containing only the credentials needed for file access """ - credentials = {} + credentials: Final = {} if litellm_params: # List of credential keys that should be passed to file operations - credential_keys = [ + credential_keys: Final = [ "api_key", "api_base", "api_version", @@ -355,7 +355,7 @@ def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]: # A batch request's input tokens scale roughly with its serialized size, so this # is a conservative per-row fallback when the token counter cannot measure a row. -_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN = 4 +_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN: Final = 4 def _estimate_batch_entry_tokens(raw_line: bytes) -> int: @@ -370,18 +370,18 @@ def _count_entry_tokens( model_name: str | None = None, ) -> int: """Token-count a single batch input entry's body (chat / text / embedding).""" - body = entry.get("body", {}) or {} - model = body.get("model", model_name or "") + body: Final = entry.get("body", {}) or {} + model: Final = body.get("model", model_name or "") - messages = body.get("messages") + messages: Final = body.get("messages") if messages: return token_counter(model=model, messages=messages) - prompt = body.get("prompt") + prompt: Final = body.get("prompt") if prompt: return _count_prompt_or_input_tokens(model=model, value=prompt) - input_data = body.get("input") + input_data: Final = body.get("input") if input_data: return _count_prompt_or_input_tokens(model=model, value=input_data) @@ -432,8 +432,8 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov usage_object=response_body.get("usage", None) or {}, reasoning_content=None, ) - _usage_dict = response_body.get("usage", None) or {} - usage: Usage = Usage(**_usage_dict) + _usage_dict: Final = response_body.get("usage", None) or {} + usage: Final[Usage] = Usage(**_usage_dict) return usage @@ -455,8 +455,8 @@ def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("message", None) or {} if custom_llm_provider == "bedrock": return batch_job_output_file.get("modelOutput", None) or {} - _response: dict = batch_job_output_file.get("response", None) or {} - _response_body = _response.get("body", None) or {} + _response: Final[dict] = batch_job_output_file.get("response", None) or {} + _response_body: Final = _response.get("body", None) or {} return _response_body @@ -472,5 +472,5 @@ def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provi return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("type") == "succeeded" if custom_llm_provider == "bedrock": return batch_job_output_file.get("modelOutput") is not None and batch_job_output_file.get("error") is None - _response: dict = batch_job_output_file.get("response", None) or {} + _response: Final[dict] = batch_job_output_file.get("response", None) or {} return _response.get("status_code", None) == 200 diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 33f2f4613bf..d6c5f0a509f 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, Literal, cast +from typing import Any, Final, Literal, cast import httpx from openai.types.batch import BatchRequestCounts @@ -54,10 +54,10 @@ from litellm.utils import ( ) ####### ENVIRONMENT VARIABLES ################### -openai_batches_instance = OpenAIBatchesAPI() -azure_batches_instance = AzureBatchesAPI() -vertex_ai_batches_instance = VertexAIBatchPrediction(gcs_bucket_name="") -anthropic_batches_instance = AnthropicBatchesHandler() +openai_batches_instance: Final = OpenAIBatchesAPI() +azure_batches_instance: Final = AzureBatchesAPI() +vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="") +anthropic_batches_instance: Final = AnthropicBatchesHandler() base_llm_http_handler = BaseLLMHTTPHandler() ################################################# @@ -80,13 +80,13 @@ def _resolve_timeout( Returns: Resolved timeout as float """ - timeout = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout + timeout: Final = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout # Handle httpx.Timeout objects if isinstance(timeout, httpx.Timeout): if supports_httpx_timeout(custom_llm_provider) is False: # Extract read timeout for providers that don't support httpx.Timeout - read_timeout = timeout.read or default_timeout + read_timeout: Final = timeout.read or default_timeout return float(read_timeout) else: # For providers that support httpx.Timeout, we still need to return a float @@ -119,11 +119,11 @@ async def acreate_batch( LiteLLM Equivalent of POST: https://api.openai.com/v1/batches """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acreate_batch"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( create_batch, completion_window, endpoint, @@ -137,9 +137,9 @@ async def acreate_batch( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -169,10 +169,10 @@ def create_batch( LiteLLM Equivalent of POST: https://api.openai.com/v1/batches """ try: - optional_params = GenericLiteLLMParams(**kwargs) - 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) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_call_id: Final = kwargs.get("litellm_call_id", None) + proxy_server_request: Final = kwargs.get("proxy_server_request", None) + model_info: Final = kwargs.get("model_info", None) model: str | None = kwargs.get("model", None) try: if model is not None: @@ -185,11 +185,11 @@ def create_batch( "litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - %s", e ) - _is_async = kwargs.pop("acreate_batch", False) is True - litellm_params = dict(GenericLiteLLMParams(**kwargs)) - litellm_logging_obj: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None)) + _is_async: Final = kwargs.pop("acreate_batch", False) is True + litellm_params: Final = dict(GenericLiteLLMParams(**kwargs)) + litellm_logging_obj: Final[LiteLLMLoggingObj] = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None)) ### TIMEOUT LOGIC ### - timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider) + timeout: Final = _resolve_timeout(optional_params, kwargs, custom_llm_provider) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model, @@ -206,7 +206,7 @@ def create_batch( custom_llm_provider=custom_llm_provider, ) - _create_batch_request = CreateBatchRequest( + _create_batch_request: Final = CreateBatchRequest( completion_window=completion_window, endpoint=endpoint, input_file_id=input_file_id, @@ -248,7 +248,7 @@ def create_batch( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -301,13 +301,13 @@ def create_batch( ) elif custom_llm_provider == "vertex_ai": api_base = optional_params.api_base or "" - vertex_ai_project = ( + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_ai_batches_instance.create_batch( _is_async=_is_async, @@ -350,11 +350,11 @@ async def aretrieve_batch( LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id} """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["aretrieve_batch"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( retrieve_batch, batch_id, custom_llm_provider, @@ -364,9 +364,9 @@ async def aretrieve_batch( **kwargs, ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -397,7 +397,7 @@ def _handle_retrieve_batch_providers_without_provider_config( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -422,7 +422,7 @@ def _handle_retrieve_batch_providers_without_provider_config( ) elif custom_llm_provider == "azure": api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") - api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") + api_version: Final = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") api_key = ( optional_params.api_key @@ -432,7 +432,7 @@ def _handle_retrieve_batch_providers_without_provider_config( or get_secret_str("AZURE_API_KEY") ) - extra_body = optional_params.get("extra_body", {}) + extra_body: Final = optional_params.get("extra_body", {}) if extra_body is not None: extra_body.pop("azure_ad_token", None) else: @@ -450,13 +450,13 @@ def _handle_retrieve_batch_providers_without_provider_config( ) elif custom_llm_provider == "vertex_ai": api_base = optional_params.api_base or "" - vertex_ai_project = ( + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_ai_batches_instance.retrieve_batch( _is_async=_is_async, @@ -519,11 +519,11 @@ def retrieve_batch( LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id} """ try: - optional_params = GenericLiteLLMParams(**kwargs) - litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj", None) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_logging_obj: Final[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( + litellm_params: Final = get_litellm_params( custom_llm_provider=custom_llm_provider, **kwargs, ) @@ -542,21 +542,21 @@ def retrieve_batch( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _retrieve_batch_request = RetrieveBatchRequest( + _retrieve_batch_request: Final = RetrieveBatchRequest( batch_id=batch_id, extra_headers=extra_headers, extra_body=extra_body, ) - _is_async = kwargs.pop("aretrieve_batch", False) is True - client = kwargs.get("client", None) + _is_async: Final = kwargs.pop("aretrieve_batch", False) is True + client: Final = kwargs.get("client", None) # Bedrock has two distinct ARN families that need different APIs: # * async-invoke ARNs (Twelve Labs Marengo embeddings) -> bedrock-runtime data plane @@ -568,7 +568,7 @@ def retrieve_batch( if batch_id.startswith("arn:aws") and ":bedrock:" in batch_id: if ":async-invoke/" in batch_id: # Remove aws_region_name from kwargs to avoid duplicate parameter - async_kwargs = kwargs.copy() + async_kwargs: Final = kwargs.copy() async_kwargs.pop("aws_region_name", None) return BedrockBatchesHandler._handle_async_invoke_status( @@ -578,7 +578,7 @@ def retrieve_batch( **async_kwargs, ) if ":model-invocation-job/" in batch_id: - mij_kwargs = kwargs.copy() + mij_kwargs: Final = kwargs.copy() mij_kwargs.pop("aws_region_name", None) return BedrockBatchesHandler._handle_model_invocation_job_status( @@ -589,7 +589,7 @@ def retrieve_batch( ) # Try to use provider config first (for providers like bedrock) - model: str | None = kwargs.get("model", None) + model: Final[str | None] = kwargs.get("model", None) if model is not None: provider_config = ProviderConfigManager.get_provider_batches_config( model=model, @@ -599,7 +599,7 @@ def retrieve_batch( provider_config = None if provider_config is not None: - response = base_llm_http_handler.retrieve_batch( + response: Final = base_llm_http_handler.retrieve_batch( batch_id=batch_id, provider_config=provider_config, litellm_params=litellm_params, @@ -656,11 +656,11 @@ async def alist_batches( """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["alist_batches"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( list_batches, after, limit, @@ -671,9 +671,9 @@ async def alist_batches( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -700,8 +700,8 @@ def list_batches( """ try: # set API KEY - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params = get_litellm_params( + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params: Final = get_litellm_params( custom_llm_provider=custom_llm_provider, **kwargs, ) @@ -720,14 +720,14 @@ def list_batches( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("alist_batches", False) is True + _is_async: Final = kwargs.pop("alist_batches", False) is True 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 = ( @@ -737,7 +737,7 @@ def list_batches( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -783,13 +783,13 @@ def list_batches( ) elif custom_llm_provider == "vertex_ai": api_base = optional_params.api_base or "" - vertex_ai_project = ( + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_ai_batches_instance.list_batches( _is_async=_is_async, @@ -836,14 +836,14 @@ async def acancel_batch( LiteLLM Equivalent of POST https://api.openai.com/v1/batches/{batch_id}/cancel """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acancel_batch"] = True # Preserve model parameter - only pop from kwargs if it exists there # (to avoid passing it twice), otherwise keep the function parameter value model = kwargs.pop("model", None) or model # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( cancel_batch, batch_id, model, @@ -854,9 +854,9 @@ async def acancel_batch( **kwargs, ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -892,8 +892,8 @@ def cancel_batch( verbose_logger.exception( "litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - %s", e ) - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params = get_litellm_params( + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params: Final = get_litellm_params( custom_llm_provider=custom_llm_provider, **kwargs, ) @@ -906,20 +906,20 @@ def cancel_batch( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _cancel_batch_request = CancelBatchRequest( + _cancel_batch_request: Final = CancelBatchRequest( batch_id=batch_id, extra_headers=extra_headers, extra_body=extra_body, ) - _is_async = kwargs.pop("acancel_batch", False) is True + _is_async: Final = kwargs.pop("acancel_batch", False) is True api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: api_base = ( @@ -929,7 +929,7 @@ def cancel_batch( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None ) api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY") @@ -973,13 +973,13 @@ def cancel_batch( ) elif custom_llm_provider == "vertex_ai": api_base = optional_params.api_base or None - vertex_ai_project = ( + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_ai_batches_instance.cancel_batch( _is_async=_is_async, @@ -1025,10 +1025,10 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj async def _async_get_status(): # Create embedding handler instance - embedding_handler = BedrockEmbedding() + embedding_handler: Final = BedrockEmbedding() # Get the status of the async invoke job - status_response = await embedding_handler._get_async_invoke_status( + status_response: Final = await embedding_handler._get_async_invoke_status( invocation_arn=batch_id, aws_region_name=aws_region_name, logging_obj=logging_obj, @@ -1040,16 +1040,16 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj from litellm.types.utils import LiteLLMBatch # Normalize status to lowercase (AWS returns 'Completed', 'Failed', etc.) - aws_status_raw = status_response.get("status", "") - aws_status_lower = aws_status_raw.lower() + aws_status_raw: Final = status_response.get("status", "") + aws_status_lower: Final = aws_status_raw.lower() # Map AWS status values to LiteLLM expected values - status_mapping: dict[str, BatchJobStatus] = { + status_mapping: Final[dict[str, BatchJobStatus]] = { "completed": "completed", "failed": "failed", "inprogress": "in_progress", "in_progress": "in_progress", } - normalized_status: BatchJobStatus = status_mapping.get( + normalized_status: Final[BatchJobStatus] = status_mapping.get( aws_status_lower, "failed" ) # Default to "failed" if unknown status @@ -1073,7 +1073,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj _, _, ) = BedrockBatchesConfig()._parse_timestamps_and_status(status_response, aws_status_raw) - result = LiteLLMBatch( + result: Final = LiteLLMBatch( id=status_response["invocationArn"], object="batch", status=normalized_status, @@ -1105,7 +1105,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj import concurrent.futures def run_in_thread(): - new_loop = asyncio.new_event_loop() + new_loop: Final = asyncio.new_event_loop() asyncio.set_event_loop(new_loop) try: return new_loop.run_until_complete(_async_get_status()) @@ -1113,5 +1113,5 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj new_loop.close() with concurrent.futures.ThreadPoolExecutor() as executor: - future = executor.submit(run_in_thread) + future: Final = executor.submit(run_in_thread) return future.result() diff --git a/litellm/budget_manager.py b/litellm/budget_manager.py index cfe1775d7e8..dcb5a7cc183 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 +from typing import Final, Literal import litellm from litellm.constants import ( @@ -60,8 +60,8 @@ class BudgetManager: self.print_verbose(f"user dict from local: {self.user_dict}") elif self.client_type == "hosted": # Load the user_dict from hosted db - url = self.api_base + "/get_budget" - data = {"project_name": self.project_name} + url: Final = self.api_base + "/get_budget" + data: Final = {"project_name": self.project_name} response = litellm.module_level_client.post(url, headers=self.headers, json=data) response = response.json() if response["status"] == "error": @@ -100,11 +100,11 @@ class BudgetManager: return self.user_dict[user] def projected_cost(self, model: str, messages: list, user: str): - text = "".join(message["content"] for message in messages) - prompt_tokens = litellm.token_counter(model=model, text=text) + text: Final = "".join(message["content"] for message in messages) + prompt_tokens: Final = litellm.token_counter(model=model, text=text) prompt_cost, _ = litellm.cost_per_token(model=model, prompt_tokens=prompt_tokens, completion_tokens=0) - current_cost = self.user_dict[user].get("current_cost", 0) - projected_cost = prompt_cost + current_cost + current_cost: Final = self.user_dict[user].get("current_cost", 0) + projected_cost: Final = prompt_cost + current_cost return projected_cost def get_total_budget(self, user: str): @@ -178,11 +178,11 @@ class BudgetManager: def reset_on_duration(self, user: str): # Get current and creation time - last_updated_at = self.user_dict[user]["last_updated_at"] - current_time = time.time() + last_updated_at: Final = self.user_dict[user]["last_updated_at"] + current_time: Final = time.time() # Convert duration from days to seconds - duration_in_seconds = self.user_dict[user]["duration"] * HOURS_IN_A_DAY * 60 * 60 + duration_in_seconds: Final = self.user_dict[user]["duration"] * HOURS_IN_A_DAY * 60 * 60 # Check if duration has elapsed if current_time - last_updated_at >= duration_in_seconds: @@ -197,7 +197,7 @@ class BudgetManager: self.reset_on_duration(user) def _save_data_thread(self): - thread = threading.Thread(target=self.save_data) # [Non-Blocking]: saves data without blocking execution + thread: Final = threading.Thread(target=self.save_data) # [Non-Blocking]: saves data without blocking execution thread.start() def save_data(self): @@ -209,8 +209,8 @@ class BudgetManager: json.dump(self.user_dict, json_file, indent=4) # Indent for pretty formatting return {"status": "success"} elif self.client_type == "hosted": - url = self.api_base + "/set_budget" - data = {"project_name": self.project_name, "user_dict": self.user_dict} + url: Final = self.api_base + "/set_budget" + data: Final = {"project_name": self.project_name, "user_dict": self.user_dict} response = litellm.module_level_client.post(url, headers=self.headers, json=data) response = response.json() return response diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index ec886b14020..1073b34ef25 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -12,7 +12,7 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Final if TYPE_CHECKING: from litellm.router import Router @@ -26,7 +26,7 @@ def resolve_embedding_router( """Return ``llm_router`` iff it serves ``embedding_model`` as a deployment.""" if llm_router is None: return None - router_model_names: list[str] = ( + router_model_names: Final[list[str]] = ( [m["model_name"] for m in llm_model_list if "model_name" in m] if llm_model_list is not None else [] ) if embedding_model in router_model_names: @@ -38,6 +38,6 @@ def build_router_embedding_metadata( request_metadata: dict[str, Any] | None, ) -> dict[str, Any]: """Forward the caller's full metadata, flagged as a semantic-cache embedding.""" - metadata: dict[str, Any] = dict(request_metadata or {}) + metadata: Final[dict[str, Any]] = dict(request_metadata or {}) metadata["semantic-cache-embedding"] = True return metadata diff --git a/litellm/caching/_internal_lru_cache.py b/litellm/caching/_internal_lru_cache.py index df6e1fc0941..218ce2b9d79 100644 --- a/litellm/caching/_internal_lru_cache.py +++ b/litellm/caching/_internal_lru_cache.py @@ -1,6 +1,6 @@ from collections.abc import Callable from functools import lru_cache -from typing import TypeVar +from typing import Final, TypeVar T = TypeVar("T") @@ -21,7 +21,7 @@ def lru_cache_wrapper( return ("error", e) def wrapped(*args, **kwargs): - result = wrapper(*args, **kwargs) + result: Final = wrapper(*args, **kwargs) if result[0] == "error": raise result[1] return result[1] diff --git a/litellm/caching/azure_blob_cache.py b/litellm/caching/azure_blob_cache.py index 755b491f9a0..742932731b3 100644 --- a/litellm/caching/azure_blob_cache.py +++ b/litellm/caching/azure_blob_cache.py @@ -11,6 +11,7 @@ Has 4 methods: import asyncio import json from contextlib import suppress +from typing import Final from litellm._logging import print_verbose, verbose_logger @@ -41,7 +42,7 @@ class AzureBlobCache(BaseCache): def set_cache(self, key, value, **kwargs) -> None: print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}") - serialized_value = json.dumps(value) + serialized_value: Final = json.dumps(value) try: self.container_client.upload_blob(key, serialized_value) except Exception as e: @@ -50,7 +51,7 @@ class AzureBlobCache(BaseCache): async def async_set_cache(self, key, value, **kwargs) -> None: print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}") - serialized_value = json.dumps(value) + serialized_value: Final = json.dumps(value) try: await self.async_container_client.upload_blob(key, serialized_value, overwrite=True) except Exception as e: @@ -62,9 +63,9 @@ class AzureBlobCache(BaseCache): try: print_verbose(f"Get Azure Blob Cache: key: {key}") - as_bytes = self.container_client.download_blob(key).readall() - as_str = as_bytes.decode("utf-8") - cached_response = json.loads(as_str) + as_bytes: Final = self.container_client.download_blob(key).readall() + as_str: Final = as_bytes.decode("utf-8") + cached_response: Final = json.loads(as_str) verbose_logger.debug( "Got Azure Blob Cache: key: %s, cached_response %s. Type Response %s", @@ -82,10 +83,10 @@ class AzureBlobCache(BaseCache): try: print_verbose(f"Get Azure Blob Cache: key: {key}") - blob = await self.async_container_client.download_blob(key) - as_bytes = await blob.readall() - as_str = as_bytes.decode("utf-8") - cached_response = json.loads(as_str) + blob: Final = await self.async_container_client.download_blob(key) + as_bytes: Final = await blob.readall() + as_str: Final = as_bytes.decode("utf-8") + cached_response: Final = json.loads(as_str) verbose_logger.debug( "Got Azure Blob Cache: key: %s, cached_response %s. Type Response %s", key, @@ -105,7 +106,7 @@ class AzureBlobCache(BaseCache): await self.async_container_client.close() async def async_set_cache_pipeline(self, cache_list, **kwargs) -> None: - tasks = [] + tasks: Final = [] for val in cache_list: tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) diff --git a/litellm/caching/base_cache.py b/litellm/caching/base_cache.py index d1965772157..6fe0609445f 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, Union +from typing import TYPE_CHECKING, Any, Final, Union if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -24,7 +24,7 @@ class BaseCache(ABC): self.default_ttl = default_ttl def get_ttl(self, **kwargs) -> int | None: - kwargs_ttl: int | None = kwargs.get("ttl") + kwargs_ttl: Final[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 758a14afb17..446b7f8be13 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 +from typing import Any, Final from pydantic import BaseModel @@ -169,13 +169,13 @@ class Cache: if type == LiteLLMCacheType.REDIS: # Check REDIS_CLUSTER_NODES env var if no explicit startup nodes if not redis_startup_nodes: - _env_cluster_nodes = litellm.get_secret("REDIS_CLUSTER_NODES") + _env_cluster_nodes: Final = litellm.get_secret("REDIS_CLUSTER_NODES") if _env_cluster_nodes is not None and isinstance(_env_cluster_nodes, str): redis_startup_nodes = json.loads(_env_cluster_nodes) if redis_startup_nodes: # Only pass GCP parameters if they are provided - cluster_kwargs = { + cluster_kwargs: Final = { "host": host, "port": port, "password": password, @@ -312,9 +312,9 @@ class Cache: ) def _get_semantic_cache_tenant_scope(self, kwargs: dict) -> str: - metadata: dict = kwargs.get("metadata") or {} - litellm_params: dict = kwargs.get("litellm_params") or {} - metadata_in_litellm_params: dict = litellm_params.get("metadata") or {} + metadata: Final[dict] = kwargs.get("metadata") or {} + litellm_params: Final[dict] = kwargs.get("litellm_params") or {} + metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata") or {} scope = "" for field in self._SEMANTIC_CACHE_TENANT_SCOPE_FIELDS: @@ -338,15 +338,15 @@ class Cache: cache_key = "" # verbose_logger.debug("\nGetting Cache key. Kwargs: %s", kwargs) - preset_cache_key = self._get_preset_cache_key_from_kwargs(**kwargs) + preset_cache_key: Final = self._get_preset_cache_key_from_kwargs(**kwargs) if preset_cache_key is not None: verbose_logger.debug("\nReturning preset cache key: %s", preset_cache_key) return preset_cache_key - combined_kwargs = ModelParamHelper._get_all_llm_api_params() - litellm_param_kwargs = all_litellm_params - is_semantic_cache = self._is_semantic_cache() - scope_excluded_params = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset() + combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params() + litellm_param_kwargs: Final = all_litellm_params + is_semantic_cache: Final = self._is_semantic_cache() + scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset() for param in kwargs: if param in scope_excluded_params: continue @@ -373,7 +373,7 @@ class Cache: ) # Remove preset_cache_key from kwargs to avoid "got multiple values" TypeError # when kwargs already contains preset_cache_key from upstream callers - kwargs_for_preset = {k: v for k, v in kwargs.items() if k != "preset_cache_key"} + kwargs_for_preset: Final = {k: v for k, v in kwargs.items() if k != "preset_cache_key"} self._set_preset_cache_key_in_kwargs(preset_cache_key=hashed_cache_key, **kwargs_for_preset) return hashed_cache_key @@ -399,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: str | None = metadata.get("model_group") or metadata_in_litellm_params.get("model_group") - caching_group = self._get_caching_group(metadata, model_group) + metadata: Final[dict] = kwargs.get("metadata", {}) or {} + litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {} + metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata", {}) or {} + model_group: Final[str | None] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group") + caching_group: Final = 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: str | None) -> str | None: - caching_groups: list | None = metadata.get("caching_groups", []) + caching_groups: Final[list | None] = metadata.get("caching_groups", []) if caching_groups: for group in caching_groups: if model_group in group: @@ -418,9 +418,9 @@ class Cache: """ Handles getting the value for the 'file' param from kwargs. Used for `transcription` requests """ - file = kwargs.get("file") - metadata = kwargs.get("metadata", {}) - litellm_params = kwargs.get("litellm_params", {}) + file: Final = kwargs.get("file") + metadata: Final = kwargs.get("metadata", {}) + litellm_params: Final = kwargs.get("litellm_params", {}) return ( metadata.get("file_checksum") or getattr(file, "name", None) @@ -467,9 +467,9 @@ class Cache: Returns: str: The hashed cache key. """ - hash_object = hashlib.sha256(cache_key.encode()) + hash_object: Final = hashlib.sha256(cache_key.encode()) # Hexadecimal representation of the hash - hash_hex = hash_object.hexdigest() + hash_hex: Final = hash_object.hexdigest() verbose_logger.debug("Hashed cache key (SHA-256): %s", hash_hex) return hash_hex @@ -484,16 +484,16 @@ class Cache: Returns: str: The final hashed cache key with the redis namespace. """ - dynamic_cache_control: DynamicCacheControl = kwargs.get("cache", {}) - metadata = kwargs.get("metadata") or {} - namespace = dynamic_cache_control.get("namespace") or metadata.get("redis_namespace") or self.namespace + dynamic_cache_control: Final[DynamicCacheControl] = kwargs.get("cache", {}) + metadata: Final = kwargs.get("metadata") or {} + namespace: Final = dynamic_cache_control.get("namespace") or metadata.get("redis_namespace") or self.namespace if namespace: hash_hex = f"{namespace}:{hash_hex}" verbose_logger.debug("Final hashed key: %s", hash_hex) return hash_hex def generate_streaming_content(self, content): - chunk_size = 5 # Adjust the chunk size as needed + chunk_size: Final = 5 # Adjust the chunk size as needed for i in range(0, len(content), chunk_size): yield { "choices": [ @@ -517,11 +517,11 @@ class Cache: """ # Check if a timestamp was stored with the cached response if cached_result is not None and isinstance(cached_result, dict) and "timestamp" in cached_result: - timestamp = cached_result["timestamp"] - current_time = time.time() + timestamp: Final = cached_result["timestamp"] + current_time: Final = time.time() # Calculate age of the cached response - response_age = current_time - timestamp + response_age: Final = current_time - timestamp # Check if the cached response is older than the max-age if max_age is not None and response_age > max_age: @@ -544,12 +544,12 @@ class Cache: @staticmethod def _get_safe_cache_lookup_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]: - cache_lookup_kwargs: dict[str, Any] = {} + cache_lookup_kwargs: Final[dict[str, Any]] = {} for prompt_kwarg in ("messages", "input"): if prompt_kwarg in kwargs: cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg] - metadata = kwargs.get("metadata") + metadata: Final = kwargs.get("metadata") if isinstance(metadata, dict): cache_lookup_kwargs["metadata"] = dict(metadata) @@ -559,8 +559,8 @@ class Cache: def _update_metadata_from_cache_lookup_kwargs( 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") + original_metadata: Final = original_kwargs.get("metadata") + cache_lookup_metadata: Final = cache_lookup_kwargs.get("metadata") if not isinstance(original_metadata, dict) or not isinstance(cache_lookup_metadata, dict): return @@ -586,9 +586,9 @@ class Cache: else: cache_key = self.get_cache_key(**kwargs) if cache_key is not None: - cache_control_args: DynamicCacheControl = kwargs.get("cache", {}) + cache_control_args: Final[DynamicCacheControl] = kwargs.get("cache", {}) max_age = cache_control_args.get("s-maxage") or cache_control_args.get("s-max-age") or float("inf") - cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs) + cache_lookup_kwargs: Final = self._get_safe_cache_lookup_kwargs(kwargs) if dynamic_cache_object is not None: cached_result = dynamic_cache_object.get_cache(cache_key, **cache_lookup_kwargs) else: @@ -618,8 +618,8 @@ class Cache: else: cache_key = self.get_cache_key(**kwargs) if cache_key is not None: - cache_control_args = kwargs.get("cache", {}) - max_age = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf"))) + cache_control_args: Final = kwargs.get("cache", {}) + max_age: Final = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf"))) if dynamic_cache_object is not None: cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs) else: @@ -646,13 +646,13 @@ class Cache: if self.ttl is not None: kwargs["ttl"] = self.ttl ## Get Cache-Controls ## - _cache_kwargs = kwargs.get("cache", None) + _cache_kwargs: Final = kwargs.get("cache", None) if isinstance(_cache_kwargs, dict): for k, v in _cache_kwargs.items(): if k == "ttl": kwargs["ttl"] = v - cached_data = {"timestamp": time.time(), "response": result} + cached_data: Final = {"timestamp": time.time(), "response": result} return cache_key, cached_data, kwargs else: raise Exception("cache key is None") @@ -756,7 +756,7 @@ class Cache: if result.usage is None or result.usage.prompt_tokens_details is None: return None - details = result.usage.prompt_tokens_details + details: Final = result.usage.prompt_tokens_details if hasattr(details, "model_dump"): details_dict = details.model_dump(exclude_none=True) elif isinstance(details, dict): @@ -767,12 +767,12 @@ class Cache: if not details_dict: return None - num_items = len(result.data) + num_items: Final = len(result.data) if num_items <= 1: return details_dict # Distribute integer/float fields evenly across items - per_item: dict = {} + per_item: Final[dict] = {} for key, value in details_dict.items(): if isinstance(value, int): quotient, remainder = divmod(value, num_items) @@ -798,8 +798,8 @@ class Cache: if result.usage is None or result.usage.prompt_tokens is None: return None - total = result.usage.prompt_tokens - num_items = len(result.data) + total: Final = result.usage.prompt_tokens + num_items: Final = len(result.data) if num_items <= 1: return total @@ -813,23 +813,23 @@ class Cache: kwargs: dict, idx_in_result_data: int = 0, ) -> tuple[str, dict, dict]: - preset_cache_key = self.get_cache_key(**{**kwargs, "input": input}) + preset_cache_key: Final = self.get_cache_key(**{**kwargs, "input": input}) kwargs["cache_key"] = preset_cache_key - embedding_response = result.data[idx_in_result_data] + embedding_response: Final = result.data[idx_in_result_data] # Extract per-item prompt_tokens + details from response usage - prompt_tokens = self._get_per_item_prompt_tokens( + prompt_tokens: Final = self._get_per_item_prompt_tokens( result=result, idx_in_result_data=idx_in_result_data, ) - prompt_tokens_details = self._get_per_item_prompt_tokens_details( + prompt_tokens_details: Final = self._get_per_item_prompt_tokens_details( result=result, idx_in_result_data=idx_in_result_data, ) # Always convert to properly typed CachedEmbedding - model_name = result.model - embedding_dict: CachedEmbedding = self._convert_to_cached_embedding( + model_name: Final = result.model + embedding_dict: Final[CachedEmbedding] = self._convert_to_cached_embedding( embedding_response, model_name, prompt_tokens=prompt_tokens, @@ -856,7 +856,7 @@ class Cache: if self.ttl is not None: kwargs["ttl"] = self.ttl - cache_list = [] + cache_list: Final = [] if isinstance(kwargs["input"], list): for idx, i in enumerate(kwargs["input"]): ( @@ -887,7 +887,7 @@ class Cache: return True # when mode == default_off -> Cache is opt in only - _cache = kwargs.get("cache", None) + _cache: Final = kwargs.get("cache", None) verbose_logger.debug("should_use_cache: kwargs: %s; _cache: %s", kwargs, _cache) if _cache and isinstance(_cache, dict): if _cache.get("use-cache", False) is True: @@ -899,13 +899,13 @@ class Cache: await self.cache.batch_cache_write(cache_key, cached_data, **kwargs) async def ping(self): - cache_ping = getattr(self.cache, "ping") + cache_ping: Final = getattr(self.cache, "ping") if cache_ping: return await cache_ping() return None async def delete_cache_keys(self, keys): - cache_delete_cache_keys = getattr(self.cache, "delete_cache_keys") + cache_delete_cache_keys: Final = getattr(self.cache, "delete_cache_keys") if cache_delete_cache_keys: return await cache_delete_cache_keys(keys) return None diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 2655a5ad683..4747aac54c6 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -19,11 +19,7 @@ import datetime import inspect import time from collections.abc import AsyncGenerator, Callable, Generator -from typing import ( - TYPE_CHECKING, - Any, - Optional, -) +from typing import TYPE_CHECKING, Any, Final, Optional from pydantic import BaseModel @@ -76,7 +72,7 @@ class CachingHandlerResponse(BaseModel): 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 -in_memory_cache_obj = InMemoryCache() +in_memory_cache_obj: Final = InMemoryCache() def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str, object]: @@ -96,10 +92,10 @@ def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str def _is_chat_completion_cached_dict(cached_result: dict) -> bool: - cached_id = cached_result.get("id") + cached_id: Final = cached_result.get("id") if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"): return True - obj = cached_result.get("object") + obj: Final = cached_result.get("object") if isinstance(obj, str): return obj.startswith("chat.completion") return "choices" in cached_result @@ -184,10 +180,10 @@ class LLMCachingHandler: ######################################################### # Init cache timing metrics ######################################################### - cache_check_start_time = time.perf_counter() + cache_check_start_time: Final = time.perf_counter() cache_check_end_time: float | None = None ######################################################### - parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) kwargs["parent_otel_span"] = parent_otel_span if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function): @@ -201,15 +197,15 @@ class LLMCachingHandler: if cached_result is not None and not isinstance(cached_result, list): verbose_logger.debug("Cache Hit!") - cache_hit = True - end_time = datetime.datetime.now() + cache_hit: Final = True + end_time: Final = datetime.datetime.now() model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, custom_llm_provider=kwargs.get("custom_llm_provider", None), api_base=kwargs.get("api_base", None), api_key=kwargs.get("api_key", None), ) - cache_duration_ms = (cache_check_end_time - cache_check_start_time) * 1000 + cache_duration_ms: Final = (cache_check_end_time - cache_check_start_time) * 1000 self._update_litellm_logging_obj_environment( logging_obj=logging_obj, model=model, @@ -240,7 +236,7 @@ class LLMCachingHandler: end_time=end_time, cache_hit=cache_hit, ) - cache_key = ( + cache_key: Final = ( self.preset_cache_key or self.request_kwargs.get("cache_key") or litellm.cache.get_cache_key(**self.request_kwargs) @@ -295,7 +291,7 @@ class LLMCachingHandler: if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function): args = args or () # Now that we confirmed caching will happen, prepare kwargs - new_kwargs = kwargs.copy() + new_kwargs: Final = kwargs.copy() new_kwargs.update( convert_args_to_kwargs( self.original_function, @@ -326,8 +322,8 @@ class LLMCachingHandler: ) # LOG SUCCESS - cache_hit = True - end_time = datetime.datetime.now() + cache_hit: Final = True + end_time: Final = datetime.datetime.now() ( model, custom_llm_provider, @@ -354,7 +350,7 @@ class LLMCachingHandler: end_time=end_time, cache_hit=cache_hit, ) - cache_key = ( + cache_key: Final = ( self.preset_cache_key or self.request_kwargs.get("cache_key") or litellm.cache.get_cache_key(**self.request_kwargs) @@ -420,9 +416,9 @@ class LLMCachingHandler: """ embedding_all_elements_cache_hit: bool = False - remaining_list = [] - non_null_list = [] - kwargs_input_as_list = self.handle_kwargs_input_list_or_str(kwargs) + remaining_list: Final = [] + non_null_list: Final = [] + kwargs_input_as_list: Final = self.handle_kwargs_input_list_or_str(kwargs) for idx, cr in enumerate(cached_result): if cr is None: remaining_list.append(kwargs_input_as_list[idx]) @@ -479,7 +475,7 @@ class LLMCachingHandler: prompt_tokens_details = PromptTokensDetailsWrapper(**aggregated_details) except Exception: prompt_tokens_details = None - usage = Usage( + usage: Final = Usage( prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens, @@ -488,9 +484,9 @@ class LLMCachingHandler: final_embedding_cached_response.usage = usage if len(remaining_list) == 0: # LOG SUCCESS - cache_hit = True + cache_hit: Final = True embedding_all_elements_cache_hit = True - end_time = datetime.datetime.now() + end_time: Final = datetime.datetime.now() ( model, custom_llm_provider, @@ -546,10 +542,10 @@ class LLMCachingHandler: if details2 is None: return details1 - dict1 = details1.model_dump(exclude_none=True) if hasattr(details1, "model_dump") else {} - dict2 = details2.model_dump(exclude_none=True) if hasattr(details2, "model_dump") else {} + dict1: Final = details1.model_dump(exclude_none=True) if hasattr(details1, "model_dump") else {} + dict2: Final = details2.model_dump(exclude_none=True) if hasattr(details2, "model_dump") else {} - merged: dict = {} + merged: Final[dict] = {} for key in set(dict1.keys()) | set(dict2.keys()): v1 = dict1.get(key, 0) v2 = dict2.get(key, 0) @@ -607,7 +603,7 @@ class LLMCachingHandler: return embedding_response idx = 0 - final_data_list = [] + final_data_list: Final = [] for item in _caching_handler_response.final_embedding_cached_response.data: if item is None and embedding_response.data is not None: final_data_list.append(embedding_response.data[idx]) @@ -690,7 +686,7 @@ class LLMCachingHandler: if litellm.cache is None: return None - new_kwargs = kwargs.copy() + new_kwargs: Final = kwargs.copy() new_kwargs.update( convert_args_to_kwargs( self.original_function, @@ -708,7 +704,7 @@ class LLMCachingHandler: new_kwargs["input"] = [new_kwargs["input"]] elif not isinstance(new_kwargs["input"], list): raise ValueError("input must be a string or a list") - tasks = [] + tasks: Final = [] for idx, i in enumerate(new_kwargs["input"]): preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i}) tasks.append( @@ -724,8 +720,8 @@ class LLMCachingHandler: if all(result is None for result in cached_result): cached_result = None else: - request_kwargs = new_kwargs.copy() - request_cache_key = request_kwargs.pop("cache_key", None) + request_kwargs: Final = new_kwargs.copy() + request_cache_key: Final = request_kwargs.pop("cache_key", None) if litellm.cache._supports_async() is True: ## check if dual cache is supported ## self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs) @@ -828,7 +824,7 @@ class LLMCachingHandler: elif (call_type == CallTypes.atranscription.value or call_type == CallTypes.transcription.value) and isinstance( cached_result, dict ): - hidden_params = { + hidden_params: Final = { "model": "whisper-1", "custom_llm_provider": custom_llm_provider, "cache_hit": True, @@ -840,10 +836,10 @@ class LLMCachingHandler: hidden_params=hidden_params, ) elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict): - use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result) + use_chat_completion_cache: Final = _is_chat_completion_cached_dict(cached_result) if use_chat_completion_cache: if kwargs.get("stream", False) is True: - bridge_call_type = ( + bridge_call_type: Final = ( CallTypes.acompletion.value if call_type == "aresponses" else CallTypes.completion.value ) cached_result = self._convert_cached_stream_response( @@ -862,7 +858,7 @@ class LLMCachingHandler: CachedResponsesAPIStreamingIterator, ) - response_obj = ResponsesAPIResponse(**cached_result) + response_obj: Final = ResponsesAPIResponse(**cached_result) if ( hasattr(response_obj, "_hidden_params") and response_obj._hidden_params is not None @@ -957,14 +953,14 @@ class LLMCachingHandler: if litellm.cache is None: return - new_kwargs = kwargs.copy() + new_kwargs: Final = kwargs.copy() new_kwargs.update( convert_args_to_kwargs( original_function, args, ) ) - parent_otel_span = _get_parent_otel_span_from_kwargs(new_kwargs) + parent_otel_span: Final = _get_parent_otel_span_from_kwargs(new_kwargs) new_kwargs["parent_otel_span"] = parent_otel_span # [OPTIONAL] ADD TO CACHE if self._should_store_result_in_cache(original_function=original_function, kwargs=new_kwargs): @@ -1006,7 +1002,7 @@ class LLMCachingHandler: Sync internal method to add the result to the cache """ - new_kwargs = kwargs.copy() + new_kwargs: Final = kwargs.copy() new_kwargs.update( convert_args_to_kwargs( self.original_function, @@ -1067,7 +1063,7 @@ class LLMCachingHandler: """ - complete_streaming_response: ModelResponse | TextCompletionResponse | None = ( + complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = ( _assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, @@ -1089,7 +1085,7 @@ class LLMCachingHandler: """ Sync internal method to add the streaming response to the cache """ - complete_streaming_response: ModelResponse | TextCompletionResponse | None = ( + complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = ( _assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, @@ -1133,7 +1129,7 @@ class LLMCachingHandler: Returns: None """ - litellm_params = { + litellm_params: Final = { "logger_fn": kwargs.get("logger_fn", None), "acompletion": is_async, "api_base": kwargs.get("api_base", ""), @@ -1173,13 +1169,13 @@ def convert_args_to_kwargs( args: tuple[Any, ...] | None = None, ) -> dict[str, Any]: # Get the signature of the original function - signature = inspect.signature(original_function) + signature: Final = inspect.signature(original_function) # Get parameter names in the order they appear in the original function - param_names = list(signature.parameters.keys()) + param_names: Final = list(signature.parameters.keys()) # Create a mapping of positional arguments to parameter names - args_to_kwargs = {} + args_to_kwargs: Final = {} if args: for index, arg in enumerate(args): if index < len(param_names): diff --git a/litellm/caching/disk_cache.py b/litellm/caching/disk_cache.py index aec18d836b0..895f276eb20 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, Union +from typing import TYPE_CHECKING, Any, Final, Union from .base_cache import BaseCache @@ -41,7 +41,7 @@ class DiskCache(BaseCache): self.set_cache(key=cache_key, value=cache_value) def get_cache(self, key, **kwargs): - original_cached_response = self.disk_cache.get(key) + original_cached_response: Final = self.disk_cache.get(key) if original_cached_response: try: cached_response = json.loads(original_cached_response) # type: ignore @@ -51,7 +51,7 @@ class DiskCache(BaseCache): return None def batch_get_cache(self, keys: list, **kwargs): - return_val = [] + return_val: Final = [] for k in keys: val = self.get_cache(key=k, **kwargs) return_val.append(val) @@ -59,9 +59,9 @@ class DiskCache(BaseCache): def increment_cache(self, key, value: int, **kwargs) -> int: with self.disk_cache.transact(): - cached_value = self.get_cache(key=key) - init_value = cached_value if isinstance(cached_value, int) else 0 - new_value = init_value + value + cached_value: Final = self.get_cache(key=key) + init_value: Final = cached_value if isinstance(cached_value, int) else 0 + new_value: Final = init_value + value self.set_cache(key, new_value, **kwargs) return new_value @@ -69,7 +69,7 @@ class DiskCache(BaseCache): return self.get_cache(key=key, **kwargs) async def async_batch_get_cache(self, keys: list, **kwargs): - return_val = [] + return_val: Final = [] for k in keys: val = self.get_cache(key=k, **kwargs) return_val.append(val) diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index a242f4a818e..3b181ca23ff 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, Union +from typing import TYPE_CHECKING, Any, Final, Union if TYPE_CHECKING: from litellm.types.caching import RedisPipelineIncrementOperation @@ -161,14 +161,14 @@ class DualCache(BaseCache): try: result = None if self.in_memory_cache is not None: - in_memory_result = self.in_memory_cache.get_cache(key, **kwargs) + in_memory_result: Final = self.in_memory_cache.get_cache(key, **kwargs) if in_memory_result is not None: result = in_memory_result if result is None and self.redis_cache is not None and local_only is False: # If not found in in-memory cache, try fetching from Redis - redis_result = self.redis_cache.get_cache(key, parent_otel_span=parent_otel_span) + redis_result: Final = self.redis_cache.get_cache(key, parent_otel_span=parent_otel_span) if redis_result is not None: # Update in-memory cache with the value from Redis @@ -188,12 +188,12 @@ class DualCache(BaseCache): local_only: bool = False, **kwargs, ): - received_args = locals() + received_args: Final = locals() received_args.pop("self") def run_in_new_loop(): """Run the coroutine in a new event loop within this thread.""" - new_loop = asyncio.new_event_loop() + new_loop: Final = asyncio.new_event_loop() try: asyncio.set_event_loop(new_loop) return new_loop.run_until_complete(self.async_batch_get_cache(**received_args)) @@ -207,7 +207,7 @@ class DualCache(BaseCache): # If we're already in an event loop, run in a separate thread # to avoid nested event loop issues with ThreadPoolExecutor(max_workers=1) as executor: - future = executor.submit(run_in_new_loop) + future: Final = executor.submit(run_in_new_loop) return future.result() except RuntimeError: @@ -226,7 +226,7 @@ class DualCache(BaseCache): print_verbose(f"async get cache: cache key: {key}; local_only: {local_only}") result = None if self.in_memory_cache is not None: - in_memory_result = await self.in_memory_cache.async_get_cache(key, **kwargs) + in_memory_result: Final = await self.in_memory_cache.async_get_cache(key, **kwargs) print_verbose(f"in_memory_result: {in_memory_result}") if in_memory_result is not None: @@ -234,7 +234,7 @@ class DualCache(BaseCache): if result is None and self.redis_cache is not None and local_only is False: # If not found in in-memory cache, try fetching from Redis - redis_result = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span) + redis_result: Final = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span) if redis_result is not None: # Update in-memory cache with the value from Redis @@ -257,8 +257,8 @@ class DualCache(BaseCache): 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, float | None] = {} + sublist_keys: Final[list[str]] = [] + previous_access_times: Final[dict[str, float | None]] = {} with self._last_redis_batch_access_time_lock: for key, value in zip(keys, result): @@ -293,7 +293,7 @@ class DualCache(BaseCache): try: result = [None] * len(keys) if self.in_memory_cache is not None: - in_memory_result = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) + in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) if in_memory_result is not None: result = in_memory_result @@ -303,14 +303,14 @@ class DualCache(BaseCache): - for the none values in the result - check the redis cache """ - current_time = time.time() + current_time: Final = time.time() sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result) # Only hit Redis if enough time has passed since last access. if len(sublist_keys) > 0: try: # If not found in in-memory cache, try fetching from Redis - redis_result = await self.redis_cache.async_batch_get_cache( + redis_result: Final = await self.redis_cache.async_batch_get_cache( sublist_keys, parent_otel_span=parent_otel_span ) except Exception: @@ -323,7 +323,7 @@ class DualCache(BaseCache): return result # Pre-compute key-to-index mapping for O(1) lookup - key_to_index = {key: i for i, key in enumerate(keys)} + key_to_index: Final = {key: i for i, key in enumerate(keys)} # Update both result and in-memory cache in a single loop for key, value in redis_result.items(): diff --git a/litellm/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py index bca9656b252..c895669be2b 100644 --- a/litellm/caching/evicted_client_closer.py +++ b/litellm/caching/evicted_client_closer.py @@ -41,14 +41,15 @@ import weakref from collections import deque from collections.abc import Awaitable, Callable, Iterator from dataclasses import dataclass, replace +from typing import Final from litellm.constants import ( EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS, EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING, ) -_CLOSABLE_ANYWHERE = "closable-anywhere" -_CLOSABLE_ON_ANY_LOOP = "closable-on-any-loop" +_CLOSABLE_ANYWHERE: Final = "closable-anywhere" +_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop" _BucketKey = str | int @@ -88,7 +89,7 @@ def _running_loop_id() -> int | None: def _close_function(client: object) -> Callable[[], object] | None: - close_fn: Callable[[], object] | None = getattr(client, "aclose", None) or getattr(client, "close", None) + close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None) return close_fn @@ -103,7 +104,7 @@ def _transport_of(client: object) -> object: def _connection_is_idle(connection: object) -> bool: """A pooled connection is idle unless it is servicing a request.""" - is_idle: object = getattr(connection, "is_idle", None) + is_idle: Final[object] = getattr(connection, "is_idle", None) return bool(is_idle()) if callable(is_idle) else True @@ -112,7 +113,7 @@ def _pool_has_busy_connection(transport: object) -> bool | None: ``None`` when there is no such pool, so the caller can ask the other backend. """ - pooled: object = getattr(getattr(transport, "_pool", None), "connections", None) + pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None) if not isinstance(pooled, (list, tuple)): return None return any( @@ -134,11 +135,11 @@ def _has_connection_in_flight(client: object) -> bool: window as the only guard, exactly as it was before this check existed. """ try: - transport = _transport_of(client) - pooled_busy = _pool_has_busy_connection(transport) + transport: Final = _transport_of(client) + pooled_busy: Final = _pool_has_busy_connection(transport) if pooled_busy is not None: return pooled_busy - session: object = getattr(transport, "client", None) + session: Final[object] = getattr(transport, "client", None) return bool(getattr(getattr(session, "connector", None), "_acquired", None)) except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle return False @@ -192,7 +193,7 @@ class EvictedClientCloser: """ if client is None or not self._is_owned(client): return - close_fn = _close_function(client) + close_fn: Final = _close_function(client) if close_fn is None: return if self._pending_count >= self._max_pending: @@ -214,7 +215,7 @@ class EvictedClientCloser: """ if not self._pending_count: return - now = self._clock() + now: Final = self._clock() for pending in self._take_due(_running_loop_id(), now): client = pending.client_ref() if client is None: @@ -236,7 +237,7 @@ class EvictedClientCloser: the front rather than having to be searched for. """ with self._queue_lock: - bucket = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design + bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design while bucket and bucket[0].client_ref() is None: bucket.popleft() self._pending_count -= 1 @@ -249,7 +250,7 @@ class EvictedClientCloser: return tuple(pending for key in buckets for pending in self._drain_locked(key, now)) def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]: - bucket = self._buckets.get(key) + bucket: Final = self._buckets.get(key) if bucket is None: return while bucket and bucket[0].close_after <= now: @@ -259,18 +260,18 @@ class EvictedClientCloser: del self._buckets[key] def _close(self, client: object) -> None: - close_fn = _close_function(client) + close_fn: Final = _close_function(client) if close_fn is None: return try: - closing = close_fn() + closing: Final = close_fn() except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers return if not inspect.isawaitable(closing): return - task = asyncio.get_running_loop().create_task(_close_quietly(closing)) + task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing)) self._close_tasks.add(task) task.add_done_callback(self._close_tasks.discard) -default_evicted_client_closer = EvictedClientCloser() +default_evicted_client_closer: Final = EvictedClientCloser() diff --git a/litellm/caching/gcs_cache.py b/litellm/caching/gcs_cache.py index 1e1508669b3..e9922218828 100644 --- a/litellm/caching/gcs_cache.py +++ b/litellm/caching/gcs_cache.py @@ -4,6 +4,7 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests. import asyncio import json +from typing import Final from urllib.parse import quote from litellm._logging import print_verbose, verbose_logger @@ -33,7 +34,7 @@ class GCSCache(BaseCache): self.sync_client = _get_httpx_client() def _construct_headers(self) -> dict: - base = GCSBucketBase(bucket_name=self.bucket_name) + base: Final = GCSBucketBase(bucket_name=self.bucket_name) base.path_service_account_json = self.path_service_account base.BUCKET_NAME = self.bucket_name return base.sync_construct_request_headers() @@ -41,35 +42,35 @@ class GCSCache(BaseCache): def set_cache(self, key, value, **kwargs): try: print_verbose(f"LiteLLM SET Cache - GCS. Key={key}. Value={value}") - headers = self._construct_headers() - object_name = self.key_prefix + key - bucket_name = self.bucket_name + headers: Final = self._construct_headers() + object_name: Final = self.key_prefix + key + bucket_name: Final = self.bucket_name url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}" - data = json.dumps(value) + data: Final = json.dumps(value) self.sync_client.post(url=url, data=data, headers=headers) except Exception as e: print_verbose(f"GCS Caching: set_cache() - Got exception from GCS: {e}") async def async_set_cache(self, key, value, **kwargs): try: - headers = self._construct_headers() - object_name = self.key_prefix + key - bucket_name = self.bucket_name + headers: Final = self._construct_headers() + object_name: Final = self.key_prefix + key + bucket_name: Final = self.bucket_name url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}" - data = json.dumps(value) + data: Final = json.dumps(value) await self.async_client.post(url=url, data=data, headers=headers) except Exception as e: print_verbose(f"GCS Caching: async_set_cache() - Got exception from GCS: {e}") def get_cache(self, key, **kwargs): try: - headers = self._construct_headers() - object_name = self.key_prefix + key - bucket_name = self.bucket_name + headers: Final = self._construct_headers() + object_name: Final = self.key_prefix + key + bucket_name: Final = self.bucket_name url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media" - response = self.sync_client.get(url=url, headers=headers) + response: Final = self.sync_client.get(url=url, headers=headers) if response.status_code == 200: - cached_response = json.loads(response.text) + cached_response: Final = json.loads(response.text) verbose_logger.debug( "Got GCS Cache: key: %s, cached_response %s. Type Response %s", key, @@ -83,11 +84,11 @@ class GCSCache(BaseCache): async def async_get_cache(self, key, **kwargs): try: - headers = self._construct_headers() - object_name = self.key_prefix + key - bucket_name = self.bucket_name + headers: Final = self._construct_headers() + object_name: Final = self.key_prefix + key + bucket_name: Final = self.bucket_name url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media" - response = await self.async_client.get(url=url, headers=headers) + response: Final = await self.async_client.get(url=url, headers=headers) if response.status_code == 200: return json.loads(response.text) return None @@ -101,7 +102,7 @@ class GCSCache(BaseCache): pass async def async_set_cache_pipeline(self, cache_list, **kwargs): - tasks = [] + tasks: Final = [] for val in cache_list: tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index e8b071bd492..38a9966f9f9 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -13,7 +13,7 @@ import json import sys import threading import time -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Final if TYPE_CHECKING: from litellm.types.caching import RedisPipelineIncrementOperation @@ -67,7 +67,7 @@ class InMemoryCache(BaseCache): # Handle special types without full conversion when possible if hasattr(value, "__sizeof__"): # Use __sizeof__ if available - size = value.__sizeof__() / 1024 + size: Final = value.__sizeof__() / 1024 return size <= self.max_size_per_item # Fallback for complex types @@ -111,7 +111,7 @@ class InMemoryCache(BaseCache): - 3. the size of in-memory cache is bounded """ - current_time = time.time() + current_time: Final = time.time() # Step 1: Remove expired or outdated items while self.expiration_heap: @@ -144,7 +144,7 @@ class InMemoryCache(BaseCache): """ Check if ttl is set for a key """ - ttl_time = self.ttl_dict.get(key) + ttl_time: Final = self.ttl_dict.get(key) if ttl_time is None or float(ttl_time) < time.time(): # if ttl is not set, allow override return True else: @@ -186,7 +186,7 @@ class InMemoryCache(BaseCache): Add value to set """ # get the value - init_value = self.get_cache(key=key) or set() + init_value: Final = self.get_cache(key=key) or set() for val in value: init_value.add(val) self.set_cache(key, init_value, ttl=ttl) @@ -207,7 +207,7 @@ class InMemoryCache(BaseCache): if key in self.cache_dict: if self.evict_element_if_expired(key): return None - original_cached_response = self.cache_dict[key] + original_cached_response: Final = self.cache_dict[key] try: cached_response = json.loads(original_cached_response) except Exception: @@ -216,7 +216,7 @@ class InMemoryCache(BaseCache): return None def batch_get_cache(self, keys: list, **kwargs): - return_val = [] + return_val: Final = [] for k in keys: val = self.get_cache(key=k, **kwargs) return_val.append(val) @@ -225,7 +225,7 @@ class InMemoryCache(BaseCache): def increment_cache(self, key, value: float, **kwargs) -> float: with self._increment_lock: # keep read-modify-write atomic - init_value = self.get_cache(key=key) or 0 + init_value: Final = self.get_cache(key=key) or 0 value = init_value + value self.set_cache(key, value, **kwargs) return value @@ -234,7 +234,7 @@ class InMemoryCache(BaseCache): return self.get_cache(key=key, **kwargs) async def async_batch_get_cache(self, keys: list, **kwargs): - return_val = [] + return_val: Final = [] for k in keys: val = self.get_cache(key=k, **kwargs) return_val.append(val) @@ -246,7 +246,7 @@ class InMemoryCache(BaseCache): async def async_increment_pipeline( self, increment_list: list["RedisPipelineIncrementOperation"], **kwargs ) -> list[float] | None: - results = [] + results: Final = [] for increment in increment_list: result = await self.async_increment(increment["key"], increment["increment_value"], **kwargs) results.append(result) @@ -274,5 +274,5 @@ class InMemoryCache(BaseCache): Get the oldest n keys in the cache """ # sorted ttl dict by ttl - sorted_ttl_dict = sorted(self.ttl_dict.items(), key=lambda x: x[1]) + sorted_ttl_dict: Final = sorted(self.ttl_dict.items(), key=lambda x: x[1]) return [key for key, _ in sorted_ttl_dict[:n]] diff --git a/litellm/caching/llm_caching_handler.py b/litellm/caching/llm_caching_handler.py index 7eae8ee3749..6fa5963c99b 100644 --- a/litellm/caching/llm_caching_handler.py +++ b/litellm/caching/llm_caching_handler.py @@ -3,6 +3,7 @@ Add the event loop to the cache key, to prevent event loop closed errors. """ import asyncio +from typing import Final from .evicted_client_closer import EvictedClientCloser, default_evicted_client_closer from .in_memory_cache import InMemoryCache @@ -37,7 +38,7 @@ class LLMClientCache(InMemoryCache): self.evicted_client_closer = evicted_client_closer or default_evicted_client_closer def _remove_key(self, key: str) -> None: - evicted: object = self.cache_dict.get(key) + evicted: Final[object] = self.cache_dict.get(key) super()._remove_key(key) self.evicted_client_closer.schedule(evicted) self.evicted_client_closer.reap() @@ -48,8 +49,8 @@ class LLMClientCache(InMemoryCache): If none, use the key as is. """ try: - event_loop = asyncio.get_running_loop() - stringified_event_loop = str(id(event_loop)) + event_loop: Final = asyncio.get_running_loop() + stringified_event_loop: Final = str(id(event_loop)) return f"{key}-{stringified_event_loop}" except RuntimeError: # handle no current running event loop return key diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 98fd9cfd1d2..8f8323550f3 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, cast +from typing import Any, Final, cast import litellm from litellm._logging import print_verbose @@ -88,7 +88,7 @@ class QdrantSemanticCache(BaseCache): if quantization_config is None: print_verbose("Quantization config is not provided. Default binary quantization will be used.") - collection_exists = self.sync_client.get( + collection_exists: Final = self.sync_client.get( url=f"{self.qdrant_api_base}/collections/{self.collection_name}/exists", headers=self.headers, ) @@ -124,7 +124,7 @@ class QdrantSemanticCache(BaseCache): else: raise Exception("Quantization config must be one of 'scalar', 'binary' or 'product'") - new_collection_status = self.sync_client.put( + new_collection_status: Final = self.sync_client.put( url=f"{self.qdrant_api_base}/collections/{self.collection_name}", json={ "vectors": {"size": self.vector_size, "distance": "Cosine"}, @@ -167,7 +167,7 @@ class QdrantSemanticCache(BaseCache): def _ensure_cache_key_payload_index(self) -> None: try: - response = self.sync_client.put( + response: Final = self.sync_client.put( url=f"{self.qdrant_api_base}/collections/{self.collection_name}/index", headers=self.headers, json={ @@ -185,7 +185,7 @@ class QdrantSemanticCache(BaseCache): # payload field. Reassigning them to a caller's key would risk # cross-scope hits, so they're treated as misses and re-populated on # the next set_cache. - cached_key = payload.get(self.CACHE_KEY_FIELD_NAME) + cached_key: Final = 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: @@ -196,7 +196,7 @@ class QdrantSemanticCache(BaseCache): llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) if router is not None: return router.embedding( model=self.embedding_model, @@ -217,7 +217,7 @@ class QdrantSemanticCache(BaseCache): llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) if router is not None: return await router.aembedding( model=self.embedding_model, @@ -237,22 +237,22 @@ class QdrantSemanticCache(BaseCache): from litellm._uuid import uuid # get the prompt - messages = kwargs["messages"] - prompt = get_str_from_messages(messages) + messages: Final = kwargs["messages"] + prompt: Final = get_str_from_messages(messages) # create an embedding for prompt - embedding_response = cast( + embedding_response: Final = cast( EmbeddingResponse, self._get_embedding(prompt, metadata=kwargs.get("metadata")), ) # get the embedding - embedding = embedding_response["data"][0]["embedding"] + embedding: Final = embedding_response["data"][0]["embedding"] value = str(value) assert isinstance(value, str) - data = { + data: Final = { "points": [ { "id": str(uuid.uuid4()), @@ -275,19 +275,19 @@ class QdrantSemanticCache(BaseCache): print_verbose(f"sync qdrant semantic-cache get_cache, kwargs: {kwargs}") # get the messages - messages = kwargs["messages"] - prompt = get_str_from_messages(messages) + messages: Final = kwargs["messages"] + prompt: Final = get_str_from_messages(messages) # convert to embedding - embedding_response = cast( + embedding_response: Final = cast( EmbeddingResponse, self._get_embedding(prompt, metadata=kwargs.get("metadata")), ) # get the embedding - embedding = embedding_response["data"][0]["embedding"] + embedding: Final = embedding_response["data"][0]["embedding"] - data = { + data: Final = { "vector": embedding, "params": { "quantization": { @@ -301,12 +301,12 @@ class QdrantSemanticCache(BaseCache): } self._add_cache_key_filter_to_search_data(data=data, key=key) - search_response = self.sync_client.post( + search_response: Final = self.sync_client.post( url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search", headers=self.headers, json=data, ) - results = search_response.json()["result"] + results: Final = search_response.json()["result"] if results is None: kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 @@ -316,14 +316,14 @@ class QdrantSemanticCache(BaseCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - similarity = results[0]["score"] - payload = results[0]["payload"] + similarity: Final = results[0]["score"] + payload: Final = results[0]["payload"] if not self._payload_matches_cache_key(payload=payload, key=key): print_verbose("Qdrant semantic-cache hit did not match cache key scope") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - cached_prompt = payload["text"] + cached_prompt: Final = payload["text"] # check similarity, if more than self.similarity_threshold, return results print_verbose( @@ -335,7 +335,7 @@ class QdrantSemanticCache(BaseCache): if similarity >= self.similarity_threshold: # cache hit ! - cached_value = payload["response"] + cached_value: Final = payload["response"] print_verbose( f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}" ) @@ -350,17 +350,17 @@ class QdrantSemanticCache(BaseCache): print_verbose(f"async qdrant semantic-cache set_cache, kwargs: {kwargs}") # get the prompt - messages = kwargs["messages"] - prompt = get_str_from_messages(messages) - embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + messages: Final = kwargs["messages"] + prompt: Final = get_str_from_messages(messages) + embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) # get the embedding - embedding = embedding_response["data"][0]["embedding"] + embedding: Final = embedding_response["data"][0]["embedding"] value = str(value) assert isinstance(value, str) - data = { + data: Final = { "points": [ { "id": str(uuid.uuid4()), @@ -384,15 +384,15 @@ class QdrantSemanticCache(BaseCache): print_verbose(f"async qdrant semantic-cache get_cache, kwargs: {kwargs}") # get the messages - messages = kwargs["messages"] - prompt = get_str_from_messages(messages) + messages: Final = kwargs["messages"] + prompt: Final = get_str_from_messages(messages) - embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) # get the embedding - embedding = embedding_response["data"][0]["embedding"] + embedding: Final = embedding_response["data"][0]["embedding"] - data = { + data: Final = { "vector": embedding, "params": { "quantization": { @@ -406,13 +406,13 @@ class QdrantSemanticCache(BaseCache): } self._add_cache_key_filter_to_search_data(data=data, key=key) - search_response = await self.async_client.post( + search_response: Final = await self.async_client.post( url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search", headers=self.headers, json=data, ) - results = search_response.json()["result"] + results: Final = search_response.json()["result"] if results is None: kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 @@ -422,14 +422,14 @@ class QdrantSemanticCache(BaseCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - similarity = results[0]["score"] - payload = results[0]["payload"] + similarity: Final = results[0]["score"] + payload: Final = results[0]["payload"] if not self._payload_matches_cache_key(payload=payload, key=key): print_verbose("Qdrant semantic-cache hit did not match cache key scope") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - cached_prompt = payload["text"] + cached_prompt: Final = payload["text"] # check similarity, if more than self.similarity_threshold, return results print_verbose( @@ -441,7 +441,7 @@ class QdrantSemanticCache(BaseCache): if similarity >= self.similarity_threshold: # cache hit ! - cached_value = payload["response"] + cached_value: Final = payload["response"] print_verbose( f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}" ) @@ -454,7 +454,7 @@ class QdrantSemanticCache(BaseCache): return self.collection_info async def async_set_cache_pipeline(self, cache_list, **kwargs): - tasks = [] + tasks: Final = [] for val in cache_list: tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 2dfd123d46d..378260b954d 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -18,7 +18,7 @@ import time from collections.abc import Awaitable, Callable, Sequence from contextvars import ContextVar from datetime import timedelta -from typing import TYPE_CHECKING, Any, TypeVar, Union, cast +from typing import TYPE_CHECKING, Any, Final, TypeVar, Union, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -69,18 +69,18 @@ def _get_call_stack_info(num_frames: int = 2) -> str: A string with format "current_function <- caller_function [<- grandparent_function]" """ try: - current_frame = inspect.currentframe() + current_frame: Final = inspect.currentframe() if current_frame is None: return "unknown" # Skip this function and the immediate caller (which sets call_type) - f_back = current_frame.f_back + f_back: Final = current_frame.f_back if f_back is None: return "unknown" frame = f_back.f_back if frame is None: return "unknown" - function_names = [] + function_names: Final = [] for _ in range(num_frames): if frame is None: @@ -172,7 +172,7 @@ class RedisCircuitBreaker: _RedisCallResult = TypeVar("_RedisCallResult") -_swallowed_redis_failures: ContextVar[int] = ContextVar("litellm_swallowed_redis_failures", default=0) +_swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0) @functools.lru_cache(maxsize=1) @@ -230,9 +230,9 @@ async def _run_under_circuit_breaker( """ if breaker.is_open(): raise Exception(f"Redis circuit breaker is open — skipping {name}") - swallowed_before = _swallowed_redis_failures.get() + swallowed_before: Final = _swallowed_redis_failures.get() try: - result = await call() + result: Final = await call() except Exception as e: if _is_redis_health_failure(e): breaker.record_failure() @@ -282,7 +282,7 @@ class RedisCache(BaseCache): from .._redis import get_redis_client, get_redis_connection_pool - redis_kwargs = {} + redis_kwargs: Final = {} if host is not None: redis_kwargs["host"] = host if port is not None: @@ -363,9 +363,9 @@ class RedisCache(BaseCache): def _handle_async_ping_error(self, e: Exception): """Handle async ping error with service failure hook.""" try: - loop = asyncio.get_running_loop() - start_time = time.time() - end_time = start_time + loop: Final = asyncio.get_running_loop() + start_time: Final = time.time() + end_time: Final = start_time loop.create_task( self.service_logger_obj.async_service_failure_hook( service=ServiceTypes.REDIS, @@ -380,9 +380,9 @@ class RedisCache(BaseCache): def _handle_sync_ping_error(self, e: Exception): """Handle sync ping error with service failure hook.""" try: - loop = asyncio.get_running_loop() - start_time = time.time() - end_time = start_time + loop: Final = asyncio.get_running_loop() + start_time: Final = time.time() + end_time: Final = start_time loop.create_task( self.service_logger_obj.async_service_failure_hook( service=ServiceTypes.REDIS, @@ -401,9 +401,9 @@ class RedisCache(BaseCache): """ # Create a stable representation of redis_kwargs for hashing # Sort keys to ensure consistent hash regardless of parameter order - sorted_kwargs = sorted(self.redis_kwargs.items()) - kwargs_str = json.dumps(sorted_kwargs, sort_keys=True) - kwargs_hash = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16] + sorted_kwargs: Final = sorted(self.redis_kwargs.items()) + kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True) + kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16] return f"async-redis-client-{kwargs_hash}" def init_async_client( @@ -413,8 +413,8 @@ class RedisCache(BaseCache): from .._redis import get_redis_async_client, get_redis_connection_pool - cache_key = self._get_async_client_cache_key() - cached_client = in_memory_llm_clients_cache.get_cache(key=cache_key) + cache_key: Final = self._get_async_client_cache_key() + cached_client: Final = in_memory_llm_clients_cache.get_cache(key=cache_key) if cached_client is not None: redis_async_client = cast(async_redis_client | async_redis_cluster_client, cached_client) else: @@ -454,7 +454,7 @@ class RedisCache(BaseCache): return DEFAULT_REDIS_MAJOR_VERSION try: - version_str = str(self.redis_version).strip() + version_str: Final = str(self.redis_version).strip() # Handle cases where there's no dot (e.g., "7" or 7) if "." in version_str: major_version = int(version_str.split(".")[0]) @@ -467,14 +467,14 @@ class RedisCache(BaseCache): return DEFAULT_REDIS_MAJOR_VERSION def set_cache(self, key, value, **kwargs): - ttl = self.get_ttl(**kwargs) + ttl: Final = self.get_ttl(**kwargs) print_verbose(f"Set Redis Cache: key: {key}\nValue {value}\nttl={ttl}, redis_version={self.redis_version}") key = self.check_and_fix_namespace(key=key) try: - start_time = time.time() + start_time: Final = time.time() self.redis_client.set(name=key, value=str(value), ex=ttl) - end_time = time.time() - _duration = end_time - start_time + end_time: Final = time.time() + _duration: Final = end_time - start_time self.service_logger_obj.service_success_hook( service=ServiceTypes.REDIS, duration=_duration, @@ -487,13 +487,13 @@ class RedisCache(BaseCache): print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {e}") def increment_cache(self, key, value: int, ttl: float | None = None, **kwargs) -> int: - _redis_client = self.redis_client + _redis_client: Final = self.redis_client start_time = time.time() - set_ttl = self.get_ttl(ttl=ttl) + set_ttl: Final = self.get_ttl(ttl=ttl) key = self.check_and_fix_namespace(key=key) try: start_time = time.time() - result: int = _redis_client.incr(name=key, amount=value) # type: ignore + result: Final[int] = _redis_client.incr(name=key, amount=value) # type: ignore end_time = time.time() _duration = end_time - start_time self.service_logger_obj.service_success_hook( @@ -507,7 +507,7 @@ class RedisCache(BaseCache): if set_ttl is not None: # check if key already has ttl, if not -> set ttl start_time = time.time() - current_ttl = _redis_client.ttl(key) + current_ttl: Final = _redis_client.ttl(key) end_time = time.time() _duration = end_time - start_time self.service_logger_obj.service_success_hook( @@ -544,10 +544,10 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_scan_iter(self, pattern: str, count: int = 100) -> list: - start_time = time.time() + start_time: Final = time.time() try: - keys = [] - _redis_client = self.init_async_client() + keys: Final = [] + _redis_client: Final = self.init_async_client() if not hasattr(_redis_client, "scan_iter"): verbose_logger.debug( "Redis client does not support scan_iter, potentially using Redis Cluster. Returning empty list." @@ -620,7 +620,7 @@ class RedisCache(BaseCache): # different key prefixes never share an executor; in_memory_llm_clients_cache # then adds the running loop, completing the per-(client, namespace, loop) # scoping. - script_cache_key = ( + script_cache_key: Final = ( f"redis-registered-script-{self._get_async_client_cache_key()}-" f"{self.namespace}-{hashlib.sha256(script.encode()).hexdigest()[:16]}" ) @@ -646,21 +646,21 @@ class RedisCache(BaseCache): Kept separate from async_register_script so each loop caches its own executor; see that method for why the binding must be per loop. """ - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() if hasattr(_redis_client, "register_script"): - registered_script = _redis_client.register_script(script) + registered_script: Final = _redis_client.register_script(script) async def standalone_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any: - namespaced_keys = tuple(self.check_and_fix_namespace(key=key) for key in keys) + namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys) return await registered_script(keys=namespaced_keys, args=args, client=client) return standalone_executor if hasattr(_redis_client, "script_load"): - script_sha = _redis_client.script_load(script) + script_sha: Final = _redis_client.script_load(script) async def cluster_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any: - namespaced_keys = tuple(self.check_and_fix_namespace(key=key) for key in keys) + namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys) return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args) return cluster_executor @@ -678,9 +678,9 @@ class RedisCache(BaseCache): ) return None - start_time = time.time() + start_time: Final = time.time() try: - _redis_client: Redis = self.init_async_client() # type: ignore + _redis_client: Final[Redis] = self.init_async_client() # type: ignore except Exception as e: end_time = time.time() _duration = end_time - start_time @@ -704,14 +704,14 @@ class RedisCache(BaseCache): raise e key = self.check_and_fix_namespace(key=key) - ttl = self.get_ttl(**kwargs) - nx = kwargs.get("nx", False) + ttl: Final = self.get_ttl(**kwargs) + nx: Final = kwargs.get("nx", False) print_verbose(f"Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}") try: if not hasattr(_redis_client, "set"): raise Exception("Redis client cannot set cache. Attribute not found.") - result = await _redis_client.set( + result: Final = await _redis_client.set( name=key, value=json.dumps(value), nx=nx, @@ -779,7 +779,7 @@ class RedisCache(BaseCache): ex=_td, ) # Execute the pipeline and return the results. - results = await pipe.execute() + results: Final = await pipe.execute() return results @_redis_circuit_breaker_guard @@ -791,14 +791,14 @@ class RedisCache(BaseCache): if len(cache_list) == 0: return - _redis_client = self.init_async_client() - start_time = time.time() + _redis_client: Final = self.init_async_client() + start_time: Final = time.time() print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}") - cache_value: Any = None + cache_value: Final[Any] = None try: async with _redis_client.pipeline(transaction=False) as pipe: - results = await self._pipeline_helper(pipe, cache_list, ttl) + results: Final = await self._pipeline_helper(pipe, cache_list, ttl) print_verbose(f"pipeline results: {results}") # Optionally, you could process 'results' to make sure that all set operations were successful. @@ -851,7 +851,7 @@ class RedisCache(BaseCache): try: await redis_client.sadd(key, *value) # type: ignore if ttl is not None: - _td = timedelta(seconds=ttl) + _td: Final = timedelta(seconds=ttl) await redis_client.expire(key, _td) except Exception: raise @@ -860,9 +860,9 @@ class RedisCache(BaseCache): async def async_set_cache_sadd(self, key, value: list, ttl: float | None, **kwargs): from redis.asyncio import Redis - start_time = time.time() + start_time: Final = time.time() try: - _redis_client: Redis = self.init_async_client() # type: ignore + _redis_client: Final[Redis] = self.init_async_client() # type: ignore except Exception as e: end_time = time.time() _duration = end_time - start_time @@ -945,17 +945,17 @@ class RedisCache(BaseCache): ) -> float: from redis.asyncio import Redis - _redis_client: Redis = self.init_async_client() # type: ignore - start_time = time.time() - _used_ttl = self.get_ttl(ttl=ttl) + _redis_client: Final[Redis] = self.init_async_client() # type: ignore + start_time: Final = time.time() + _used_ttl: Final = self.get_ttl(ttl=ttl) key = self.check_and_fix_namespace(key=key) try: - result = await _redis_client.incrbyfloat(name=key, amount=value) + result: Final = await _redis_client.incrbyfloat(name=key, amount=value) if _used_ttl is not None: if refresh_ttl: await _redis_client.expire(key, _used_ttl) else: - current_ttl = await _redis_client.ttl(key) + current_ttl: Final = await _redis_client.ttl(key) if current_ttl == -1: await _redis_client.expire(key, _used_ttl) @@ -1012,10 +1012,10 @@ class RedisCache(BaseCache): GET/compare/SET runs in a single Lua call, so it is also atomic across racing callers and pods. Returns the resulting value. """ - _redis_client = self.init_async_client() - _used_ttl = self.get_ttl(ttl=ttl) + _redis_client: Final = self.init_async_client() + _used_ttl: Final = self.get_ttl(ttl=ttl) key = self.check_and_fix_namespace(key=key) - lua = ( + lua: Final = ( "local cur = redis.call('GET', KEYS[1]) " "if cur == false or tonumber(cur) < tonumber(ARGV[1]) then " "redis.call('SET', KEYS[1], ARGV[1]) " @@ -1056,10 +1056,10 @@ class RedisCache(BaseCache): try: key = self.check_and_fix_namespace(key=key) print_verbose(f"Get Redis Cache: key: {key}") - start_time = time.time() - cached_response = self.redis_client.get(key) - end_time = time.time() - _duration = end_time - start_time + start_time: Final = time.time() + cached_response: Final = self.redis_client.get(key) + end_time: Final = time.time() + _duration: Final = end_time - start_time self.service_logger_obj.service_success_hook( service=ServiceTypes.REDIS, duration=_duration, @@ -1088,7 +1088,7 @@ class RedisCache(BaseCache): We use a wrapper so RedisCluster can override this method """ - async_redis_client = self.init_async_client() + async_redis_client: Final = self.init_async_client() return await async_redis_client.mget(keys=keys) # type: ignore def batch_get_cache( @@ -1107,17 +1107,17 @@ class RedisCache(BaseCache): dict: A dictionary mapping keys to their cached values """ key_value_dict = {} - _key_list = [key for key in key_list if key is not None] + _key_list: Final = [key for key in key_list if key is not None] try: - _keys = [] + _keys: Final = [] for cache_key in _key_list: 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) - end_time = time.time() - _duration = end_time - start_time + start_time: Final = time.time() + results: Final[list] = self._run_redis_mget_operation(keys=_keys) + end_time: Final = time.time() + _duration: Final = end_time - start_time self.service_logger_obj.service_success_hook( service=ServiceTypes.REDIS, duration=_duration, @@ -1131,7 +1131,7 @@ class RedisCache(BaseCache): # 'results' is a list of values corresponding to the order of keys in '_key_list'. key_value_dict = dict(zip(_key_list, results)) - decoded_results = {} + decoded_results: Final = {} for k, v in key_value_dict.items(): if isinstance(k, bytes): k = k.decode("utf-8") @@ -1147,15 +1147,15 @@ class RedisCache(BaseCache): 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 + _redis_client: Final[Redis] = self.init_async_client() # type: ignore key = self.check_and_fix_namespace(key=key) - start_time = time.time() + start_time: Final = time.time() try: print_verbose(f"Get Async Redis Cache: key: {key}") - cached_response = await _redis_client.get(key) + cached_response: Final = await _redis_client.get(key) print_verbose(f"Got Async Redis Cache: key: {key}, cached_response {cached_response}") - response = self._get_cache_logic(cached_response=cached_response) + response: Final = self._get_cache_logic(cached_response=cached_response) end_time = time.time() _duration = end_time - start_time @@ -1209,14 +1209,14 @@ class RedisCache(BaseCache): """ # typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `mget` key_value_dict = {} - start_time = time.time() - _key_list = [key for key in key_list if key is not None] + start_time: Final = time.time() + _key_list: Final = [key for key in key_list if key is not None] try: - _keys = [] + _keys: Final = [] for cache_key in _key_list: cache_key = self.check_and_fix_namespace(key=cache_key) _keys.append(cache_key) - results = await self._async_run_redis_mget_operation(keys=_keys) + results: Final = await self._async_run_redis_mget_operation(keys=_keys) ## LOGGING ## end_time = time.time() _duration = end_time - start_time @@ -1235,7 +1235,7 @@ class RedisCache(BaseCache): # 'results' is a list of values corresponding to the order of keys in 'key_list'. key_value_dict = dict(zip(_key_list, results)) - decoded_results = {} + decoded_results: Final = {} for k, v in key_value_dict.items(): if isinstance(k, bytes): k = k.decode("utf-8") @@ -1267,9 +1267,9 @@ class RedisCache(BaseCache): Tests if the sync redis client is correctly setup. """ print_verbose("Pinging Sync Redis Cache") - start_time = time.time() + start_time: Final = time.time() try: - response: bool = self.redis_client.ping() # type: ignore + response: Final[bool] = self.redis_client.ping() # type: ignore print_verbose(f"Redis Cache PING: {response}") ## LOGGING ## end_time = time.time() @@ -1298,11 +1298,11 @@ class RedisCache(BaseCache): async def ping(self) -> bool: # typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ping` - _redis_client: Any = self.init_async_client() - start_time = time.time() + _redis_client: Final[Any] = self.init_async_client() + start_time: Final = time.time() print_verbose("Pinging Async Redis Cache") try: - response = await _redis_client.ping() + response: Final = await _redis_client.ping() ## LOGGING ## end_time = time.time() _duration = end_time - start_time @@ -1333,17 +1333,17 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def delete_cache_keys(self, keys): # typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete` - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() keys = [self.check_and_fix_namespace(key=key) for key in keys] # 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 + client_list: Final[list] = self.redis_client.client_list() # type: ignore return client_list def info(self): - info = self.redis_client.info() + info: Final = self.redis_client.info() return info def flush_cache(self): @@ -1373,10 +1373,10 @@ class RedisCache(BaseCache): import redis.asyncio as redis_async # Create a fresh Redis client with current settings - redis_client = redis_async.Redis(**self.redis_kwargs) + redis_client: Final = redis_async.Redis(**self.redis_kwargs) # Test the connection - ping_result = await redis_client.ping() # type: ignore[misc] + ping_result: Final = await redis_client.ping() # type: ignore[misc] # Close the connection await redis_client.aclose() # type: ignore[attr-defined] @@ -1399,7 +1399,7 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_delete_cache(self, key: str): # typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete` - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() key = self.check_and_fix_namespace(key=key) # keys is str return await _redis_client.delete(key) @@ -1425,7 +1425,7 @@ class RedisCache(BaseCache): _td = timedelta(seconds=increment_op["ttl"]) pipe.expire(cache_key, _td) # Execute the pipeline and return results - results = await pipe.execute() + results: Final = await pipe.execute() # only return float values verbose_logger.debug("Increment ASYNC Redis Cache PIPELINE: results: %s", results) return [r for r in results if isinstance(r, float)] @@ -1448,14 +1448,14 @@ class RedisCache(BaseCache): from redis.asyncio import Redis - _redis_client: Redis = self.init_async_client() # type: ignore - start_time = time.time() + _redis_client: Final[Redis] = self.init_async_client() # type: ignore + start_time: Final = time.time() print_verbose(f"Increment Async Redis Cache Pipeline: increment list: {increment_list}") try: async with _redis_client.pipeline(transaction=False) as pipe: - results = await self._pipeline_increment_helper(pipe, increment_list) + results: Final = await self._pipeline_increment_helper(pipe, increment_list) ## LOGGING ## end_time = time.time() @@ -1507,9 +1507,9 @@ class RedisCache(BaseCache): """ try: # typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ttl` - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() key = self.check_and_fix_namespace(key=key) - ttl = await _redis_client.ttl(key) + ttl: Final = await _redis_client.ttl(key) if ttl <= -1: # -1 means the key does not exist, -2 key does not exist return None return ttl @@ -1537,11 +1537,11 @@ class RedisCache(BaseCache): Returns: int: The length of the list after the push operation """ - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() key = self.check_and_fix_namespace(key=key) - start_time = time.time() + start_time: Final = time.time() try: - response = await _redis_client.rpush(key, *values) + response: Final = await _redis_client.rpush(key, *values) ## LOGGING ## end_time = time.time() _duration = end_time - start_time @@ -1578,7 +1578,7 @@ class RedisCache(BaseCache): for rpush_op in rpush_list: key = self.check_and_fix_namespace(key=rpush_op["key"]) pipe.rpush(key, *rpush_op["values"]) - results = await pipe.execute() + results: Final = await pipe.execute() # Preserve positional correspondence — raise on per-command errors for r in results: if isinstance(r, Exception): @@ -1604,12 +1604,12 @@ class RedisCache(BaseCache): if len(rpush_list) == 0: return [] - _redis_client: Any = self.init_async_client() - start_time = time.time() + _redis_client: Final[Any] = self.init_async_client() + start_time: Final = time.time() try: async with _redis_client.pipeline(transaction=False) as pipe: - results = await self._pipeline_rpush_helper(pipe, rpush_list) + results: Final = await self._pipeline_rpush_helper(pipe, rpush_list) ## LOGGING ## end_time = time.time() @@ -1641,7 +1641,7 @@ 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] = [] + result: Final[list[bytes]] = [] for _ in range(count): pipe.lpop(key) results = await pipe.execute() @@ -1661,12 +1661,12 @@ class RedisCache(BaseCache): parent_otel_span: Span | None = None, **kwargs, ) -> Any | list[Any]: - _redis_client: Any = self.init_async_client() + _redis_client: Final[Any] = self.init_async_client() key = self.check_and_fix_namespace(key=key) - start_time = time.time() + start_time: Final = time.time() print_verbose(f"LPOP from Redis list: key: {key}, count: {count}") try: - major_version = self._parse_redis_major_version() + major_version: Final = self._parse_redis_major_version() if count is not None and major_version < 7: # For Redis < 7.0, use pipeline to execute multiple LPOP commands @@ -1725,7 +1725,7 @@ class RedisCache(BaseCache): For Redis >= 7, queues one LPOP(key, count) per operation. For Redis < 7, queues `count` individual LPOP(key) commands per operation. """ - major_version = self._parse_redis_major_version() + major_version: Final = self._parse_redis_major_version() if major_version >= 7: for lpop_op in lpop_list: @@ -1735,14 +1735,14 @@ 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: Final[list[int]] = [] for lpop_op in lpop_list: key = self.check_and_fix_namespace(key=lpop_op["key"]) count = lpop_op["count"] or 1 counts.append(count) for _ in range(count): pipe.lpop(key) - flat_results = await pipe.execute() + flat_results: Final = await pipe.execute() # Re-group the flat results back into per-key lists raw_results = [] @@ -1758,7 +1758,7 @@ class RedisCache(BaseCache): raise r # Decode bytes -> str for each result set - decoded_results: list[list[str] | None] = [] + decoded_results: Final[list[list[str] | None]] = [] for r in raw_results: if r is None: decoded_results.append(None) @@ -1793,12 +1793,12 @@ class RedisCache(BaseCache): if len(lpop_list) == 0: return [] - _redis_client: Any = self.init_async_client() - start_time = time.time() + _redis_client: Final[Any] = self.init_async_client() + start_time: Final = time.time() try: async with _redis_client.pipeline(transaction=False) as pipe: - results = await self._pipeline_lpop_helper(pipe, lpop_list) + results: Final = await self._pipeline_lpop_helper(pipe, lpop_list) ## LOGGING ## end_time = time.time() diff --git a/litellm/caching/redis_cluster_cache.py b/litellm/caching/redis_cluster_cache.py index 926712e38ec..c275e3c1bf7 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, Union +from typing import TYPE_CHECKING, Any, Final, Union from litellm.caching.redis_cache import RedisCache @@ -37,7 +37,7 @@ class RedisClusterCache(RedisCache): if self.redis_async_redis_cluster_client: return self.redis_async_redis_cluster_client - _redis_client = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs) + _redis_client: Final = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs) if isinstance(_redis_client, RedisCluster): self.redis_async_redis_cluster_client = _redis_client @@ -53,7 +53,7 @@ class RedisClusterCache(RedisCache): """ Overrides `_async_run_redis_mget_operation` in redis_cache.py """ - async_redis_cluster_client = self.init_async_client() + async_redis_cluster_client: Final = self.init_async_client() return await async_redis_cluster_client.mget_nonatomic(keys=keys) # type: ignore async def test_connection(self) -> dict: @@ -68,21 +68,21 @@ class RedisClusterCache(RedisCache): from redis.cluster import ClusterNode # Create ClusterNode objects from startup_nodes - cluster_kwargs = self.redis_kwargs.copy() - startup_nodes = cluster_kwargs.pop("startup_nodes", []) + cluster_kwargs: Final = self.redis_kwargs.copy() + startup_nodes: Final = cluster_kwargs.pop("startup_nodes", []) - new_startup_nodes: list[ClusterNode] = [] + new_startup_nodes: Final[list[ClusterNode]] = [] for item in startup_nodes: new_startup_nodes.append(ClusterNode(**item)) # Create a fresh Redis Cluster client with current settings - redis_client = redis_async.RedisCluster( + redis_client: Final = redis_async.RedisCluster( startup_nodes=new_startup_nodes, **cluster_kwargs, # type: ignore ) # Test the connection - ping_result = await redis_client.ping() # type: ignore[attr-defined, misc] + ping_result: Final = await redis_client.ping() # type: ignore[attr-defined, misc] # Close the connection await redis_client.aclose() # type: ignore[attr-defined] diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 4fe42d1908e..b0c8fa963ee 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, cast +from typing import Any, Final, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -95,7 +95,7 @@ class RedisSemanticCache(BaseCache): password = password or os.environ["REDIS_PASSWORD"] except KeyError as e: # Raise a more informative exception if any of the required keys are missing - missing_var = e.args[0] + missing_var: Final = e.args[0] raise ValueError( f"Missing required Redis configuration: {missing_var}. Provide {missing_var} or redis_url." ) from e @@ -130,7 +130,7 @@ class RedisSemanticCache(BaseCache): from redisvl.utils.vectorize import CustomTextVectorizer # type: ignore[import-not-found, import-untyped] try: - cache_vectorizer = CustomTextVectorizer(self._get_embedding) + cache_vectorizer: Final = CustomTextVectorizer(self._get_embedding) return self._init_semantic_cache( semantic_cache_cls=SemanticCache, index_name=self._index_name, @@ -156,7 +156,7 @@ class RedisSemanticCache(BaseCache): cache_vectorizer: Any, ) -> Any: def _is_schema_mismatch(exc: ValueError) -> bool: - error_message = str(exc).lower() + error_message: Final = str(exc).lower() return any(phrase in error_message for phrase in ("schema does not match", "index schema")) try: @@ -172,7 +172,7 @@ class RedisSemanticCache(BaseCache): if not _is_schema_mismatch(exc): raise - isolated_index_name = f"{index_name}_isolated" + isolated_index_name: Final = f"{index_name}_isolated" print_verbose( "Redis semantic-cache existing index schema is not isolated; " f"using isolated index - {isolated_index_name}" @@ -239,16 +239,16 @@ class RedisSemanticCache(BaseCache): """ Extract a semantic-cache prompt from chat or Responses API request kwargs. """ - messages = kwargs.get("messages") + messages: Final = kwargs.get("messages") if messages: return get_str_from_messages(messages) if "input" not in kwargs: return None - prompt_parts: list[str] = [] + prompt_parts: Final[list[str]] = [] cls._collect_responses_input_text(kwargs.get("input"), prompt_parts) - prompt = "\n".join(prompt_parts).strip() + prompt: Final = "\n".join(prompt_parts).strip() return prompt or None @classmethod @@ -258,7 +258,7 @@ class RedisSemanticCache(BaseCache): return if isinstance(value, str): - stripped_value = value.strip() + stripped_value: Final = value.strip() if stripped_value: prompt_parts.append(stripped_value) return @@ -298,10 +298,10 @@ class RedisSemanticCache(BaseCache): @staticmethod def _coerce_response_input_value(value: Any) -> Any: - model_dump = getattr(value, "model_dump", None) + model_dump: Final = getattr(value, "model_dump", None) if callable(model_dump): return model_dump() - dict_method = getattr(value, "dict", None) + dict_method: Final = getattr(value, "dict", None) if callable(dict_method): return dict_method() return value @@ -318,7 +318,7 @@ class RedisSemanticCache(BaseCache): llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) if router is not None: embedding_response = cast( EmbeddingResponse, @@ -383,22 +383,22 @@ class RedisSemanticCache(BaseCache): value_str: str | None = None try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic caching") return value_str = str(value) - prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata")) + prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata")) - store_kwargs: dict[str, Any] = { + store_kwargs: Final[dict[str, Any]] = { "vector": prompt_embedding, "filters": self._get_cache_filters(key), } # Get TTL and store in Redis semantic cache - ttl = self._get_ttl(**kwargs) + ttl: Final = self._get_ttl(**kwargs) if ttl is not None: store_kwargs["ttl"] = int(ttl) self.llmcache.store(prompt, value_str, **store_kwargs) @@ -419,7 +419,7 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Redis semantic-cache get_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic cache lookup") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 @@ -427,13 +427,13 @@ class RedisSemanticCache(BaseCache): # Check the cache for semantically similar prompts in this exact # LiteLLM cache-key scope. - prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata")) - check_kwargs: dict[str, Any] = { + prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata")) + check_kwargs: Final[dict[str, Any]] = { "prompt": prompt, "vector": prompt_embedding, "filter_expression": self._get_cache_key_filter_expression(key), } - results = self.llmcache.check(**check_kwargs) + results: Final = self.llmcache.check(**check_kwargs) # Return None if no similar prompts found if not results: @@ -441,20 +441,20 @@ class RedisSemanticCache(BaseCache): return None # Process the best matching result - cache_hit = results[0] + cache_hit: Final = results[0] if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key): print_verbose("Redis semantic-cache hit did not match cache key scope") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - vector_distance = float(cache_hit["vector_distance"]) + vector_distance: Final = float(cache_hit["vector_distance"]) # Convert vector distance back to similarity score # For cosine distance: 0 = most similar, 2 = least similar # While similarity: 1 = most similar, 0 = least similar - similarity = 1 - vector_distance + similarity: Final = 1 - vector_distance - cached_prompt = cache_hit["prompt"] - cached_response = cache_hit["response"] + cached_prompt: Final = cache_hit["prompt"] + cached_response: Final = cache_hit["response"] # update kwargs["metadata"] with similarity, don't rewrite the original metadata kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity @@ -488,7 +488,7 @@ class RedisSemanticCache(BaseCache): llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) try: if router is not None: embedding_response = await router.aembedding( @@ -521,23 +521,23 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Async Redis semantic-cache set_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic caching") return - value_str = str(value) + value_str: Final = str(value) # Generate embedding for the value (response) to cache - prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + prompt_embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) - store_kwargs: dict[str, Any] = { + store_kwargs: Final[dict[str, Any]] = { "vector": prompt_embedding, "filters": self._get_cache_filters(key), } # Get TTL and store in Redis semantic cache - ttl = self._get_ttl(**kwargs) + ttl: Final = self._get_ttl(**kwargs) if ttl is not None: store_kwargs["ttl"] = ttl await self.llmcache.astore( @@ -562,43 +562,43 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Async Redis semantic-cache get_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic cache lookup") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None # Generate embedding for the prompt - prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + prompt_embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) # Check the cache for semantically similar prompts in this exact # LiteLLM cache-key scope. - check_kwargs: dict[str, Any] = { + check_kwargs: Final[dict[str, Any]] = { "prompt": prompt, "vector": prompt_embedding, "filter_expression": self._get_cache_key_filter_expression(key), } - results = await self.llmcache.acheck(**check_kwargs) + results: Final = await self.llmcache.acheck(**check_kwargs) # handle results / cache hit if not results: kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - cache_hit = results[0] + cache_hit: Final = results[0] if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key): print_verbose("Redis semantic-cache hit did not match cache key scope") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - vector_distance = float(cache_hit["vector_distance"]) + vector_distance: Final = float(cache_hit["vector_distance"]) # Convert vector distance back to similarity # For cosine distance: 0 = most similar, 2 = least similar # While similarity: 1 = most similar, 0 = least similar - similarity = 1 - vector_distance + similarity: Final = 1 - vector_distance - cached_prompt = cache_hit["prompt"] - cached_response = cache_hit["response"] + cached_prompt: Final = cache_hit["prompt"] + cached_response: Final = cache_hit["response"] # update kwargs["metadata"] with similarity, don't rewrite the original metadata kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity @@ -622,7 +622,7 @@ class RedisSemanticCache(BaseCache): Returns: Dict[str, Any]: Information about the Redis index """ - aindex = await self.llmcache._get_async_index() + aindex: Final = 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: @@ -634,7 +634,7 @@ class RedisSemanticCache(BaseCache): **kwargs: Additional arguments """ try: - tasks = [] + tasks: Final = [] for val in cache_list: tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) diff --git a/litellm/caching/s3_cache.py b/litellm/caching/s3_cache.py index baad0e29c5e..e953c9d67b0 100644 --- a/litellm/caching/s3_cache.py +++ b/litellm/caching/s3_cache.py @@ -13,6 +13,7 @@ import asyncio import json from datetime import datetime, timedelta, timezone from functools import partial +from typing import Final from litellm._logging import print_verbose, verbose_logger @@ -62,16 +63,16 @@ class S3Cache(BaseCache): def set_cache(self, key, value, **kwargs): try: print_verbose(f"LiteLLM SET Cache - S3. Key={key}. Value={value}") - ttl = kwargs.get("ttl", None) + ttl: Final = kwargs.get("ttl", None) # Convert value to JSON before storing in S3 - serialized_value = json.dumps(value) + serialized_value: Final = json.dumps(value) key = self._to_s3_key(key) if ttl is not None: cache_control = f"immutable, max-age={ttl}, s-maxage={ttl}" # Calculate expiration time - expiration_time = datetime.now(timezone.utc) + timedelta(seconds=ttl) + expiration_time: Final = datetime.now(timezone.utc) + timedelta(seconds=ttl) # Upload the data to S3 with the calculated expiration time self.s3_client.put_object( Bucket=self.bucket_name, @@ -105,8 +106,8 @@ class S3Cache(BaseCache): """ try: verbose_logger.debug("Set ASYNC S3 Cache: Key=%s. Value=%s", key, value) - loop = asyncio.get_event_loop() - func = partial(self.set_cache, key, value, **kwargs) + loop: Final = asyncio.get_event_loop() + func: Final = partial(self.set_cache, key, value, **kwargs) await loop.run_in_executor(None, func) except Exception as e: verbose_logger.error("S3 Caching: async_set_cache() - Got exception from S3: %s", e) @@ -123,8 +124,8 @@ class S3Cache(BaseCache): if cached_response is not None: if "Expires" in cached_response: - expires_time = cached_response["Expires"] - current_time = datetime.now(expires_time.tzinfo) + expires_time: Final = cached_response["Expires"] + current_time: Final = datetime.now(expires_time.tzinfo) if current_time > expires_time: return None @@ -160,9 +161,9 @@ class S3Cache(BaseCache): """ try: verbose_logger.debug("Get ASYNC S3 Cache: key: %s", key) - loop = asyncio.get_event_loop() - func = partial(self.get_cache, key, **kwargs) - result = await loop.run_in_executor(None, func) + loop: Final = asyncio.get_event_loop() + func: Final = partial(self.get_cache, key, **kwargs) + result: Final = await loop.run_in_executor(None, func) return result except Exception as e: verbose_logger.error("S3 Caching: async_get_cache() - Got exception from S3: %s", e) @@ -175,7 +176,7 @@ class S3Cache(BaseCache): pass async def async_set_cache_pipeline(self, cache_list, **kwargs): - tasks = [] + tasks: Final = [] for val in cache_list: tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) await asyncio.gather(*tasks) diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index e01bb430987..0fe8581df86 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -19,7 +19,7 @@ import hashlib import os import struct from dataclasses import dataclass -from typing import Any +from typing import Any, Final from redis import Redis from redis.asyncio import Redis as AsyncRedis @@ -106,8 +106,8 @@ class ValkeySemanticCache(RedisSemanticCache): "(or VALKEY_HOST/VALKEY_PORT), or pass redis_url." ) - credentials = f":{password}@" if password else "" - scheme = "rediss" if ssl else "redis" + credentials: Final = f":{password}@" if password else "" + scheme: Final = "rediss" if ssl else "redis" return f"{scheme}://{credentials}{host}:{port}" @classmethod @@ -154,7 +154,7 @@ class ValkeySemanticCache(RedisSemanticCache): return None def _assert_dim_matches(self, info: dict, dim: int) -> None: - existing_dim = self._extract_index_dim(info) + existing_dim: Final = self._extract_index_dim(info) if existing_dim is not None and existing_dim != dim: raise ValueError( f"Valkey semantic-cache index '{self.index_name}' already exists with " @@ -186,7 +186,7 @@ class ValkeySemanticCache(RedisSemanticCache): except Exception as exc: if not self._is_index_exists_error(exc): raise - info = await self.async_client.ft(self.index_name).info() + info: Final = await self.async_client.ft(self.index_name).info() self._assert_dim_matches(info, dim) self._index_dim = dim @@ -202,8 +202,8 @@ class ValkeySemanticCache(RedisSemanticCache): } def _knn_query(self, key: str) -> Query: - scope = self._scope_tag(key) - query_string = ( + scope: Final = self._scope_tag(key) + query_string: Final = ( f"(@{self.CACHE_KEY_FIELD_NAME}:{{{scope}}})" f"=>[KNN 1 @{self.EMBEDDING_FIELD_NAME} $vec AS {self.DISTANCE_FIELD_NAME}]" ) @@ -211,10 +211,10 @@ class ValkeySemanticCache(RedisSemanticCache): @classmethod def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None: - docs = getattr(search_result, "docs", []) + docs: Final = getattr(search_result, "docs", []) if not docs: return None - doc = docs[0] + doc: Final = docs[0] return _ValkeyCacheHit( response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)), distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)), @@ -225,7 +225,7 @@ class ValkeySemanticCache(RedisSemanticCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - similarity = 1 - hit.distance + similarity: Final = 1 - hit.distance kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity if similarity < self.similarity_threshold: @@ -235,17 +235,17 @@ class ValkeySemanticCache(RedisSemanticCache): def set_cache(self, key: str, value: Any, **kwargs: Any) -> None: print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic caching") return - embedding = self._get_embedding(prompt) + embedding: Final = self._get_embedding(prompt) self._ensure_index_sync(len(embedding)) - doc_key = self._doc_key(key) + doc_key: Final = self._doc_key(key) self.sync_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)) - ttl = self._get_ttl(**kwargs) + ttl: Final = self._get_ttl(**kwargs) if ttl is not None: self.sync_client.expire(doc_key, ttl) except Exception as e: @@ -254,15 +254,15 @@ class ValkeySemanticCache(RedisSemanticCache): def get_cache(self, key: str, **kwargs: Any) -> Any: print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - embedding = self._get_embedding(prompt) + embedding: Final = self._get_embedding(prompt) self._ensure_index_sync(len(embedding)) - search_result = self.sync_client.ft(self.index_name).search( + search_result: Final = self.sync_client.ft(self.index_name).search( self._knn_query(key), query_params={"vec": self._embedding_to_bytes(embedding)}, ) @@ -274,17 +274,17 @@ class ValkeySemanticCache(RedisSemanticCache): async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None: print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: print_verbose("No prompt provided for semantic caching") return - embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) await self._ensure_index_async(len(embedding)) - doc_key = self._doc_key(key) + doc_key: Final = self._doc_key(key) await self.async_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)) - ttl = self._get_ttl(**kwargs) + ttl: Final = self._get_ttl(**kwargs) if ttl is not None: await self.async_client.expire(doc_key, ttl) except Exception as e: @@ -293,15 +293,15 @@ class ValkeySemanticCache(RedisSemanticCache): async def async_get_cache(self, key: str, **kwargs: Any) -> Any: print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}") try: - prompt = self._get_prompt_from_kwargs(**kwargs) + prompt: Final = self._get_prompt_from_kwargs(**kwargs) if prompt is None: kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) + embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) await self._ensure_index_async(len(embedding)) - search_result = await self.async_client.ft(self.index_name).search( + search_result: Final = await self.async_client.ft(self.index_name).search( self._knn_query(key), query_params={"vec": self._embedding_to_bytes(embedding)}, ) diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 6083bc8e26b..1e5cccaf23f 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, Union +from typing import TYPE_CHECKING, Any, Final, Union from typing_extensions import TypedDict @@ -60,7 +60,7 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Unexpected responses stream payload") if hidden_params: - existing = getattr(response, "_hidden_params", None) + existing: Final = getattr(response, "_hidden_params", None) if not isinstance(existing, dict) or not existing: setattr(response, "_hidden_params", dict(hidden_params)) else: @@ -72,13 +72,13 @@ class ResponsesToCompletionBridgeHandler: for _ in stream_iter: pass - completed = getattr(stream_iter, "completed_response", None) - response_obj = getattr(completed, "response", None) if completed else None + completed: Final = getattr(stream_iter, "completed_response", None) + response_obj: Final = getattr(completed, "response", None) if completed else None if response_obj is None: raise ValueError("Stream ended without a completed response") - hidden_params = getattr(stream_iter, "_hidden_params", None) - response = self._coerce_response_object(response_obj, hidden_params) + hidden_params: Final = getattr(stream_iter, "_hidden_params", None) + response: Final = self._coerce_response_object(response_obj, hidden_params) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response @@ -87,13 +87,13 @@ class ResponsesToCompletionBridgeHandler: async for _ in stream_iter: pass - completed = getattr(stream_iter, "completed_response", None) - response_obj = getattr(completed, "response", None) if completed else None + completed: Final = getattr(stream_iter, "completed_response", None) + response_obj: Final = getattr(completed, "response", None) if completed else None if response_obj is None: raise ValueError("Stream ended without a completed response") - hidden_params = getattr(stream_iter, "_hidden_params", None) - response = self._coerce_response_object(response_obj, hidden_params) + hidden_params: Final = getattr(stream_iter, "_hidden_params", None) + response: Final = self._coerce_response_object(response_obj, hidden_params) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response @@ -102,35 +102,35 @@ class ResponsesToCompletionBridgeHandler: from litellm import LiteLLMLoggingObj from litellm.types.utils import ModelResponse - model = kwargs.get("model") + model: Final = kwargs.get("model") if model is None or not isinstance(model, str): raise ValueError("model is required") - custom_llm_provider = kwargs.get("custom_llm_provider") + custom_llm_provider: Final = kwargs.get("custom_llm_provider") if custom_llm_provider is None or not isinstance(custom_llm_provider, str): raise ValueError("custom_llm_provider is required") - messages = kwargs.get("messages") + messages: Final = kwargs.get("messages") if messages is None or not isinstance(messages, list): raise ValueError("messages is required") - optional_params = kwargs.get("optional_params") + optional_params: Final = kwargs.get("optional_params") if optional_params is None or not isinstance(optional_params, dict): raise ValueError("optional_params is required") - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if litellm_params is None or not isinstance(litellm_params, dict): raise ValueError("litellm_params is required") - headers = kwargs.get("headers") + headers: Final = kwargs.get("headers") if headers is None or not isinstance(headers, dict): raise ValueError("headers is required") - model_response = kwargs.get("model_response") + model_response: Final = kwargs.get("model_response") if model_response is None or not isinstance(model_response, ModelResponse): raise ValueError("model_response is required") - logging_obj = kwargs.get("logging_obj") + logging_obj: Final = kwargs.get("logging_obj") if logging_obj is None or not isinstance(logging_obj, LiteLLMLoggingObj): raise ValueError("logging_obj is required") @@ -158,19 +158,19 @@ class ResponsesToCompletionBridgeHandler: from litellm import responses from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - validated_kwargs = self.validate_input_kwargs(kwargs) - model = validated_kwargs["model"] - messages = validated_kwargs["messages"] + validated_kwargs: Final = self.validate_input_kwargs(kwargs) + model: Final = validated_kwargs["model"] + messages: Final = validated_kwargs["messages"] optional_params = validated_kwargs["optional_params"] - litellm_params = validated_kwargs["litellm_params"] - headers = validated_kwargs["headers"] - model_response = validated_kwargs["model_response"] - logging_obj = validated_kwargs["logging_obj"] - custom_llm_provider = validated_kwargs["custom_llm_provider"] + litellm_params: Final = validated_kwargs["litellm_params"] + headers: Final = validated_kwargs["headers"] + model_response: Final = validated_kwargs["model_response"] + logging_obj: Final = validated_kwargs["logging_obj"] + custom_llm_provider: Final = validated_kwargs["custom_llm_provider"] if kwargs.get("stream") is True and "stream" not in optional_params: optional_params = {**optional_params, "stream": True} - request_data = self.transformation_handler.transform_request( + request_data: Final = self.transformation_handler.transform_request( model=model, messages=messages, optional_params=optional_params, @@ -188,13 +188,13 @@ class ResponsesToCompletionBridgeHandler: # than adding an explicit kwarg) avoids the duplicate-keyword # TypeError that would otherwise fire on the real bridge path. request_data["custom_llm_provider"] = custom_llm_provider - result = responses( + result: Final = responses( **request_data, ) from litellm.types.utils import ModelResponse - stream = self._resolve_stream_flag(optional_params, litellm_params) + stream: Final = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( model=model, @@ -220,7 +220,7 @@ class ResponsesToCompletionBridgeHandler: json_mode=kwargs.get("json_mode"), ) elif not stream: - responses_api_response = self._collect_response_from_stream(result) + responses_api_response: Final = self._collect_response_from_stream(result) return self.transformation_handler.transform_response( model=model, raw_response=responses_api_response, @@ -237,12 +237,12 @@ class ResponsesToCompletionBridgeHandler: else: if self._is_preformatted_cached_chat_stream(result): return self._apply_post_stream_processing(result, model, custom_llm_provider) - completion_stream = self.transformation_handler.get_model_response_iterator( + completion_stream: Final = self.transformation_handler.get_model_response_iterator( streaming_response=result, # type: ignore sync_stream=True, json_mode=kwargs.get("json_mode"), ) - streamwrapper = CustomStreamWrapper( + streamwrapper: Final = CustomStreamWrapper( completion_stream=completion_stream, model=model, custom_llm_provider=custom_llm_provider, @@ -254,20 +254,20 @@ class ResponsesToCompletionBridgeHandler: from litellm import aresponses from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - validated_kwargs = self.validate_input_kwargs(kwargs) - model = validated_kwargs["model"] - messages = validated_kwargs["messages"] + validated_kwargs: Final = self.validate_input_kwargs(kwargs) + model: Final = validated_kwargs["model"] + messages: Final = validated_kwargs["messages"] optional_params = validated_kwargs["optional_params"] - litellm_params = validated_kwargs["litellm_params"] - headers = validated_kwargs["headers"] - model_response = validated_kwargs["model_response"] - logging_obj = validated_kwargs["logging_obj"] - custom_llm_provider = validated_kwargs["custom_llm_provider"] + litellm_params: Final = validated_kwargs["litellm_params"] + headers: Final = validated_kwargs["headers"] + model_response: Final = validated_kwargs["model_response"] + logging_obj: Final = validated_kwargs["logging_obj"] + custom_llm_provider: Final = validated_kwargs["custom_llm_provider"] if kwargs.get("stream") is True and "stream" not in optional_params: optional_params = {**optional_params, "stream": True} try: - request_data = self.transformation_handler.transform_request( + request_data: Final = self.transformation_handler.transform_request( model=model, messages=messages, optional_params=optional_params, @@ -285,14 +285,14 @@ class ResponsesToCompletionBridgeHandler: # keyword TypeError when `sanitized_litellm_params` already # carries `custom_llm_provider`. request_data["custom_llm_provider"] = custom_llm_provider - result = await aresponses( + result: Final = await aresponses( **request_data, aresponses=True, ) from litellm.types.utils import ModelResponse - stream = self._resolve_stream_flag(optional_params, litellm_params) + stream: Final = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( model=model, @@ -318,7 +318,7 @@ class ResponsesToCompletionBridgeHandler: json_mode=kwargs.get("json_mode"), ) elif not stream: - responses_api_response = await self._collect_response_from_stream_async(result) + responses_api_response: Final = await self._collect_response_from_stream_async(result) return self.transformation_handler.transform_response( model=model, raw_response=responses_api_response, @@ -335,12 +335,12 @@ class ResponsesToCompletionBridgeHandler: else: if self._is_preformatted_cached_chat_stream(result): return self._apply_post_stream_processing(result, model, custom_llm_provider) - completion_stream = self.transformation_handler.get_model_response_iterator( + completion_stream: Final = self.transformation_handler.get_model_response_iterator( streaming_response=result, # type: ignore sync_stream=False, json_mode=kwargs.get("json_mode"), ) - streamwrapper = CustomStreamWrapper( + streamwrapper: Final = CustomStreamWrapper( completion_stream=completion_stream, model=model, custom_llm_provider=custom_llm_provider, @@ -359,7 +359,7 @@ class ResponsesToCompletionBridgeHandler: from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - streamwrapper = CustomStreamWrapper( + streamwrapper: Final = CustomStreamWrapper( completion_stream=MockResponseIterator(model_response=response, json_mode=json_mode), model=model, custom_llm_provider=custom_llm_provider, @@ -378,7 +378,7 @@ class ResponsesToCompletionBridgeHandler: from litellm.utils import ProviderConfigManager try: - provider_config = ProviderConfigManager.get_provider_chat_config( + provider_config: Final = ProviderConfigManager.get_provider_chat_config( model=model, provider=LlmProviders(custom_llm_provider) ) except (ValueError, KeyError): @@ -389,4 +389,4 @@ class ResponsesToCompletionBridgeHandler: return stream -responses_api_bridge = ResponsesToCompletionBridgeHandler() +responses_api_bridge: Final = ResponsesToCompletionBridgeHandler() diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index ada67d1cfe5..4c6112952cc 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -5,13 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping -from typing import ( - TYPE_CHECKING, - Any, - Literal, - Union, - cast, -) +from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast from openai.types.responses.custom_tool_param import CustomToolParam from openai.types.responses.response_input_param import ( @@ -69,7 +63,7 @@ def _get_reasoning_items( msg: "AllMessageValues", ) -> list[ChatCompletionReasoningItem]: """Extract reasoning_items from a message dict with proper typing.""" - items = msg.get("reasoning_items") # type: ignore[union-attr] + items: Final = msg.get("reasoning_items") # type: ignore[union-attr] if items: return items # type: ignore[return-value] return [] @@ -84,7 +78,7 @@ def _build_reasoning_item( Handles both pydantic objects (attribute access) and plain dicts. """ - summary: list[dict[str, Any]] = [] + summary: Final[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", "")}) @@ -118,17 +112,17 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch LiteLLMCompletionResponsesConfig, ) - is_custom = item.get("type") == "custom_tool_call" - arguments = (item.get("input") if is_custom else item.get("arguments")) or "" - name = item.get("name") or ("custom_tool" if is_custom else "") - function_chunk = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments) - tool_call_dict = _ChatToolCallDict( + is_custom: Final = item.get("type") == "custom_tool_call" + arguments: Final = (item.get("input") if is_custom else item.get("arguments")) or "" + name: Final = item.get("name") or ("custom_tool" if is_custom else "") + function_chunk: Final = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments) + tool_call_dict: Final = _ChatToolCallDict( id=LiteLLMCompletionResponsesConfig._tool_call_id_from_responses_item(item.get("id"), item.get("call_id")), type="function", function=function_chunk, index=index, ) - raw_provider_fields = item.get("provider_specific_fields") + raw_provider_fields: Final = item.get("provider_specific_fields") if isinstance(raw_provider_fields, dict): provider_specific_fields = raw_provider_fields elif raw_provider_fields and hasattr(raw_provider_fields, "__dict__"): @@ -151,7 +145,7 @@ def _reasoning_item_to_response_input( 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: Final[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 @@ -174,15 +168,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): """Chat tool_choice nests the name under function/custom; Responses API expects top-level name.""" if not isinstance(tool_choice, dict): return tool_choice - choice_type = tool_choice.get("type") + choice_type: Final = tool_choice.get("type") if choice_type not in ("function", "custom"): return tool_choice if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"): # Return only Responses shape so stray chat ``function``/``custom`` keys are not sent upstream. return _flat_responses_tool_choice(choice_type, tool_choice["name"]) - nested = tool_choice.get(choice_type) + nested: Final = tool_choice.get(choice_type) if isinstance(nested, dict): - nested_name = nested.get("name") + nested_name: Final = nested.get("name") if isinstance(nested_name, str) and nested_name: return _flat_responses_tool_choice(choice_type, nested_name) return tool_choice @@ -200,7 +194,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): """ from litellm.types.utils import Choices, Message - item_type = item.get("type") + item_type: Final = item.get("type") # Ignore reasoning items for now if item_type == "reasoning": @@ -208,7 +202,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Handle message items with output_text content if item_type == "message": - content_list = item.get("content", []) + content_list: Final = item.get("content", []) for content_item in content_list: if isinstance(content_item, dict): content_type = content_item.get("type") @@ -235,9 +229,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def convert_chat_completion_messages_to_responses_api( self, messages: list["AllMessageValues"] ) -> tuple[list[Any], str | None]: - input_items: list[Any] = [] + input_items: Final[list[Any]] = [] instructions: str | None = None - custom_tool_call_ids = frozenset( + custom_tool_call_ids: Final = frozenset( tool_call["id"] for msg in messages if msg.get("role") == "assistant" and isinstance(msg.get("tool_calls"), list) @@ -386,13 +380,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): 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] = { + responses_optional_param_keys: Final = set(ResponsesAPIOptionalRequestParams.__annotations__.keys()) + sanitized: Final[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] = {} + legacy_metadata: Final = litellm_params.get("metadata") + existing_litellm_metadata: Final = litellm_params.get("litellm_metadata") + merged_litellm_metadata: Final[dict[str, Any]] = {} if isinstance(legacy_metadata, dict): merged_litellm_metadata.update(legacy_metadata) if isinstance(existing_litellm_metadata, dict): @@ -456,7 +450,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): optional_params = self._extract_extra_body_params(optional_params) # Build responses API request using the reverse transformation logic - responses_api_request = ResponsesAPIOptionalRequestParams() + responses_api_request: Final = ResponsesAPIOptionalRequestParams() # Set instructions if we found a system message if instructions: @@ -464,7 +458,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): self._map_optional_params_to_responses_api_request(optional_params, responses_api_request) - stream = optional_params.get("stream") or litellm_params.get("stream", False) + stream: Final = optional_params.get("stream") or litellm_params.get("stream", False) verbose_logger.debug("Chat provider: Stream parameter: %s", stream) # Ensure stream is properly set in the request @@ -472,22 +466,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): responses_api_request["stream"] = True # Handle session management if previous_response_id is provided - previous_response_id = optional_params.get("previous_response_id") + previous_response_id: Final = optional_params.get("previous_response_id") if previous_response_id: # Use the existing session handler for responses API verbose_logger.debug("Chat provider: Warning ignoring previous response ID: %s", previous_response_id) # Convert back to responses API format for the actual request - api_model = model + api_model: Final = model from litellm.types.utils import CallTypes setattr(litellm_logging_obj, "call_type", CallTypes.responses.value) - sanitized_litellm_params = self._build_sanitized_litellm_params(litellm_params) + sanitized_litellm_params: Final = self._build_sanitized_litellm_params(litellm_params) - request_data = { + request_data: Final = { "model": api_model, "input": input_items, "litellm_logging_obj": litellm_logging_obj, @@ -534,14 +528,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): from litellm.types.utils import Choices, Message - choices: list[Choices] = [] + choices: Final[list[Choices]] = [] index = 0 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: Final[list[dict[str, Any]]] = [] tool_call_index = 0 for item in output_items: @@ -649,10 +643,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): @classmethod def _extract_output_from_completed_event(cls, parsed_chunk: dict[str, Any]) -> list[dict[str, Any]] | None: - response_payload = parsed_chunk.get("response") + response_payload: Final = parsed_chunk.get("response") if not isinstance(response_payload, dict): return None - response_output = response_payload.get("output") + response_output: Final = 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) @@ -662,8 +656,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): 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: Final[dict[int, dict[str, Any]]] = {} + recovered_text_only_items: Final[dict[int, dict[str, Any]]] = {} for chunk in raw_sse.splitlines(): parsed_chunk = parse_sse_json_chunk(chunk) @@ -698,7 +692,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: Final[dict[int, dict[str, Any]]] = {**recovered_text_only_items} merged_items.update(recovered_output_items) if merged_items: @@ -708,8 +702,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): @classmethod 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") + model_call_details: Final = getattr(logging_obj, "model_call_details", {}) or {} + original_response: Final = model_call_details.get("original_response") return cls._recover_output_items_from_raw_sse(original_response) def transform_response( @@ -738,7 +732,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): output_items = raw_response.output if len(output_items) == 0: - recovered_output_items = self._recover_output_items_from_logging(logging_obj) + recovered_output_items: Final = self._recover_output_items_from_logging(logging_obj) if recovered_output_items: output_items = cast(Any, recovered_output_items) raw_response.output = cast(Any, recovered_output_items) @@ -748,7 +742,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) # Convert response output to choices using the static helper - choices = self._convert_response_output_to_choices( + choices: Final = self._convert_response_output_to_choices( output_items=output_items, handle_raw_dict_callback=self._handle_raw_dict_response_item, ) @@ -771,7 +765,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Preserve hidden params from the ResponsesAPIResponse, especially the headers # which contain important provider information like x-request-id - raw_response_hidden_params = getattr(raw_response, "_hidden_params", {}) + raw_response_hidden_params: Final = getattr(raw_response, "_hidden_params", {}) if raw_response_hidden_params: if not hasattr(model_response, "_hidden_params") or model_response._hidden_params is None: model_response._hidden_params = {} @@ -807,7 +801,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) -> "ResponseInputImageParam": from openai.types.responses import ResponseInputImageParam - content_image_url = content.get("image_url") + content_image_url: Final = content.get("image_url") actual_image_url: str | None = None detail: Literal["low", "high", "auto"] | None = None @@ -823,7 +817,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if actual_image_url is None: raise ValueError(f"Invalid image URL: {content_image_url}") - image_param = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image") + image_param: Final = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image") if detail: image_param["detail"] = detail @@ -921,7 +915,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): 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: Final[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": @@ -959,11 +953,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): unsupported params remain in extra_body. """ # Extract extra_body and separate supported params from unsupported ones - extra_body = optional_params.pop("extra_body", None) or {} + extra_body: Final = optional_params.pop("extra_body", None) or {} if not extra_body: return optional_params - supported_responses_api_params = set(ResponsesAPIOptionalRequestParams.__annotations__.keys()) + supported_responses_api_params: Final = set(ResponsesAPIOptionalRequestParams.__annotations__.keys()) # Also include params we handle specially supported_responses_api_params.update( { @@ -973,7 +967,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) # Extract supported params from extra_body and merge into optional_params - extra_body_copy = extra_body.copy() + extra_body_copy: Final = extra_body.copy() for key, value in extra_body_copy.items(): if key in supported_responses_api_params: # Prefer extra_body value if it exists (may have more complete info like summary in reasoning_effort) @@ -988,7 +982,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Check if auto-summary is enabled via flag or environment variable # Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var - auto_summary_enabled = ( + auto_summary_enabled: Final = ( litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" ) @@ -1032,7 +1026,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tools = [] responses_api_request["tools"] = tools - web_search_tool: dict[str, Any] = {"type": "web_search"} + web_search_tool: Final[dict[str, Any]] = {"type": "web_search"} if isinstance(web_search_options, dict): web_search_tool.update(web_search_options) @@ -1067,10 +1061,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return None if isinstance(response_format, dict): - format_type = response_format.get("type") + format_type: Final = response_format.get("type") if format_type == "json_schema": - json_schema = response_format.get("json_schema", {}) + json_schema: Final = response_format.get("json_schema", {}) return { "format": { "type": "json_schema", @@ -1099,7 +1093,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if not annotations: return None - result: list[ChatCompletionAnnotation] = [] + result: Final[list[ChatCompletionAnnotation]] = [] for annotation in annotations: try: # Convert Pydantic models to dicts (handles both v1 and v2) @@ -1127,7 +1121,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if not status: return "stop" - status_mapping = { + status_mapping: Final = { "completed": "stop", "incomplete": "length", "failed": "stop", @@ -1154,7 +1148,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): if not str_line or str_line.startswith("event:"): # ignore. return GenericStreamingChunk(text="", tool_use=None, is_finished=False, finish_reason="", usage=None) - index = str_line.find("data:") + index: Final = str_line.find("data:") if index != -1: str_line = str_line[index + 5 :] @@ -1240,10 +1234,10 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): # New output item added output_item = parsed_chunk.get("item", {}) if output_item.get("type") in ("function_call", "custom_tool_call"): - converted = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0)) - provider_specific_fields = converted.get("provider_specific_fields") + converted: Final = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0)) + provider_specific_fields: Final = converted.get("provider_specific_fields") - function_chunk = ChatCompletionToolCallFunctionChunk( + function_chunk: Final = ChatCompletionToolCallFunctionChunk( name=converted["function"]["name"] or None, arguments=converted["function"]["arguments"] or parsed_chunk.get("arguments") or "", ) @@ -1253,7 +1247,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): tool_call_index = OpenAiResponsesToChatCompletionStreamIterator._sequential_tool_call_index( tool_call_index_map, parsed_chunk.get("output_index", 0) ) - tool_call_chunk = ChatCompletionToolCallChunk( + tool_call_chunk: Final = ChatCompletionToolCallChunk( id=converted["id"], index=tool_call_index, type="function", @@ -1383,16 +1377,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): # Check if response contains function_call items in output # to determine correct finish_reason - response_data = parsed_chunk.get("response", {}) - output_items = response_data.get("output", []) if response_data else [] + response_data: Final = parsed_chunk.get("response", {}) + output_items: Final = response_data.get("output", []) if response_data else [] - has_function_calls = any( + has_function_calls: Final = any( item.get("type") in ("function_call", "custom_tool_call") for item in output_items if isinstance(item, dict) ) - finish_reason = "tool_calls" if has_function_calls else "stop" + finish_reason: Final = "tool_calls" if has_function_calls else "stop" # Extract reasoning items with encrypted_content for round-tripping completed_reasoning_items: list[dict[str, Any]] | None = None @@ -1408,7 +1402,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): summary_raw=item.get("summary"), ) ) - completed_reasoning_items_typed = cast( + completed_reasoning_items_typed: Final = cast( list[ChatCompletionReasoningItem] | None, completed_reasoning_items, ) diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index 042db79a71a..f844b3a3d7f 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, cast +from typing import Any, Final, cast from litellm.caching.dual_cache import DualCache from litellm.compression.message_stubbing import ( @@ -20,11 +20,11 @@ from litellm.types.utils import CallTypes # CallTypes that produce Anthropic-shaped messages (structured content blocks). # Everything else is treated as OpenAI chat-completions shape. -_ANTHROPIC_CALL_TYPES = frozenset({CallTypes.anthropic_messages.value}) +_ANTHROPIC_CALL_TYPES: Final = frozenset({CallTypes.anthropic_messages.value}) # CallTypes that are valid targets for compression. Compression operates on # message-shaped inputs, so we only accept call types whose payload is a list # of role/content messages. -_SUPPORTED_CALL_TYPES = frozenset( +_SUPPORTED_CALL_TYPES: Final = frozenset( { CallTypes.completion.value, CallTypes.acompletion.value, @@ -54,7 +54,7 @@ def _build_retrieval_tools(keys: list[str], call_type: str) -> list[dict]: if not keys: return [] - openai_tools = [build_retrieval_tool(keys)] + openai_tools: Final = [build_retrieval_tool(keys)] if not _is_anthropic_call_type(call_type): return openai_tools @@ -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: Final[list[str]] = [] + stack: Final[list[Any]] = [content] while stack: item = stack.pop() if isinstance(item, str): @@ -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: Final[list[dict[str, Any]]] = [dict(m) for m in messages] - normalized_messages: list[dict] = [] + normalized_messages: Final[list[dict]] = [] for msg in original_messages: normalized_messages.append( { @@ -135,7 +135,7 @@ def _extract_last_user_message(messages: list[dict]) -> str: def _extract_tool_use_ids(content: Any) -> list[str]: if not isinstance(content, list): return [] - tool_use_ids: list[str] = [] + tool_use_ids: Final[list[str]] = [] for part in content: if not isinstance(part, dict): continue @@ -150,7 +150,7 @@ def _extract_tool_use_ids(content: Any) -> list[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: Final[set[str]] = set() for part in content: if not isinstance(part, dict): continue @@ -171,7 +171,7 @@ def _extract_anthropic_tool_exchange_spans( 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: Final[list[set[int]]] = [] i = 0 while i < len(messages): current = messages[i] @@ -216,8 +216,8 @@ def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int so compressing it replaces the live instruction with a marker. Compression guardrails share this policy; see the Headroom guardrail. """ - system_indices = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system") - last_user = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:] + system_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system") + last_user: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:] last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:] return system_indices + last_user + last_assistant @@ -230,16 +230,16 @@ def _combine_scores( """Weighted average of BM25 and embedding scores, with min-max normalization.""" 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 + min_s: Final = min(scores) if scores else 0.0 + max_s: Final = max(scores) if scores else 0.0 + rng: Final = max_s - min_s if rng == 0: return [0.0] * len(scores) return [(s - min_s) / rng for s in scores] - norm_bm25 = _normalize(bm25_scores) - norm_emb = _normalize(emb_scores) - emb_weight = 1.0 - bm25_weight + norm_bm25: Final = _normalize(bm25_scores) + norm_emb: Final = _normalize(emb_scores) + emb_weight: Final = 1.0 - bm25_weight return [bm25_weight * b + emb_weight * e for b, e in zip(norm_bm25, norm_emb)] @@ -253,7 +253,7 @@ def _select_kept_indices_for_budget( initial_kept_indices: set[int], tool_exchange_spans: list[set[int]], ) -> tuple[set[int], dict[int, dict]]: - kept_indices = set(initial_kept_indices) + kept_indices: Final = set(initial_kept_indices) current_tokens = 0 for i in kept_indices: current_tokens += token_counter( @@ -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: Final[dict[int, dict]] = {} # idx -> truncated message dict + span_id_by_index: Final[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: Final[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 @@ -323,7 +323,7 @@ def _select_kept_indices_for_budget( 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() + dropped_tool_span_indices: Final[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) @@ -372,7 +372,7 @@ def compress( A ``CompressedResult`` dict containing compressed messages, token counts, a cache of original content, and the retrieval tool definition. """ - call_type_str = _normalize_call_type(call_type) + call_type_str: Final = _normalize_call_type(call_type) normalized_messages, original_messages = _normalize_messages_for_compression( messages=messages, call_type=call_type_str, @@ -381,7 +381,7 @@ def compress( if compression_target is None: compression_target = compression_trigger * 7 // 10 - original_tokens = token_counter( + original_tokens: Final = token_counter( model=model, messages=cast(list[Any], original_messages), ) @@ -399,17 +399,17 @@ def compress( ) # Extract query for relevance scoring - query = _extract_last_user_message(normalized_messages) + query: Final = _extract_last_user_message(normalized_messages) # Score each message - bm25_scores = bm25_score_messages(query, normalized_messages) + bm25_scores: Final = bm25_score_messages(query, normalized_messages) if embedding_model: from litellm.compression.scoring.embedding_scorer import ( embedding_score_messages, ) - emb_scores = embedding_score_messages( + emb_scores: Final = embedding_score_messages( query, normalized_messages, model=embedding_model, @@ -421,7 +421,7 @@ def compress( combined_scores = bm25_scores # Protected messages are never compressed - protected_indices = get_protected_indices(normalized_messages) + protected_indices: Final = get_protected_indices(normalized_messages) kept_indices: set[int] = set(protected_indices) tool_exchange_spans: list[set[int]] = [] @@ -454,10 +454,10 @@ def compress( ) # Build compressed messages and cache - compressed_messages: list[dict] = [] - cache: dict[str, str] = {} - used_keys: set[str] = set() - dropped_tool_span_indices = _get_dropped_tool_span_indices( + compressed_messages: Final[list[dict]] = [] + cache: Final[dict[str, str]] = {} + used_keys: Final[set[str]] = set() + dropped_tool_span_indices: Final = _get_dropped_tool_span_indices( kept_indices=kept_indices, tool_exchange_spans=tool_exchange_spans ) @@ -474,9 +474,9 @@ def compress( compressed_messages.append(stub_message(msg, key)) # Build retrieval tool in the target request schema - tools = _build_retrieval_tools(list(cache.keys()), call_type=call_type_str) + tools: Final = _build_retrieval_tools(list(cache.keys()), call_type=call_type_str) - compressed_tokens = token_counter( + compressed_tokens: Final = token_counter( model=model, messages=cast(list[Any], compressed_messages), ) diff --git a/litellm/compression/content_detection.py b/litellm/compression/content_detection.py index 4a072b63f2c..9cd475d662b 100644 --- a/litellm/compression/content_detection.py +++ b/litellm/compression/content_detection.py @@ -4,8 +4,9 @@ Auto-detect content type per message: code, JSON, or text. import json import re +from typing import Final -_CODE_KEYWORDS = re.compile( +_CODE_KEYWORDS: Final = re.compile( r"\b(?:def |function |class |import |from |require\(|#include|fn |func |const |let |var |public |private |static )\b" ) @@ -16,7 +17,7 @@ def detect_content_type(content: str) -> str: Returns one of: "code", "json", "text" """ - stripped = content.strip() + stripped: Final = content.strip() if not stripped: return "text" @@ -30,10 +31,10 @@ def detect_content_type(content: str) -> str: # Check code indicators # Sample first 5000 chars for performance - sample = stripped[:5000] - keyword_matches = len(_CODE_KEYWORDS.findall(sample)) - lines = sample.split("\n") - indented_lines = sum(1 for line in lines if line.startswith((" ", "\t")) and line.strip()) + sample: Final = stripped[:5000] + keyword_matches: Final = len(_CODE_KEYWORDS.findall(sample)) + lines: Final = sample.split("\n") + indented_lines: Final = sum(1 for line in lines if line.startswith((" ", "\t")) and line.strip()) # If we see multiple code keywords or significant indentation, it's likely code if keyword_matches >= 3 or (indented_lines > len(lines) * 0.3 and len(lines) > 5): diff --git a/litellm/compression/message_stubbing.py b/litellm/compression/message_stubbing.py index ebb2d19997a..a9a74646614 100644 --- a/litellm/compression/message_stubbing.py +++ b/litellm/compression/message_stubbing.py @@ -3,11 +3,12 @@ Replace messages with compact stubs and extract human-readable keys. """ import re +from typing import Final from litellm.compression.content_detection import detect_content_type # Patterns for extracting file paths from content -_FILE_PATH_PATTERNS = [ +_FILE_PATH_PATTERNS: Final = [ re.compile(r"^#\s*(\S+\.\w+)", re.MULTILINE), # # filename.py re.compile(r"^//\s*(\S+\.\w+)", re.MULTILINE), # // filename.js re.compile(r"^File:\s*(\S+)", re.MULTILINE), # File: path/to/file @@ -40,7 +41,7 @@ def extract_key(message: dict, fallback_index: int, used_keys: set[str]) -> str: key = f"message_{fallback_index}" # Handle duplicates - base_key = key + base_key: Final = key counter = 2 while key in used_keys: key = f"{base_key}_{counter}" @@ -61,10 +62,10 @@ def stub_message(message: dict, key: str) -> dict: if isinstance(content, list): content = " ".join(p.get("text", "") if isinstance(p, dict) else str(p) for p in content) - line_count = content.count("\n") + 1 - content_type = detect_content_type(content) + line_count: Final = content.count("\n") + 1 + content_type: Final = detect_content_type(content) - stub_content = ( + stub_content: Final = ( f"[Compressed: {key} — {line_count} lines, {content_type}. " f"Use litellm_content_retrieve tool to get full content.]" ) @@ -89,23 +90,23 @@ def truncate_message(message: dict, max_tokens: int) -> dict: content = " ".join(p.get("text", "") if isinstance(p, dict) else str(p) for p in content) # Rough conversion: 1 token ≈ 3 characters - target_chars = max(100, max_tokens * 3) + target_chars: Final = max(100, max_tokens * 3) if len(content) <= target_chars: return {**message, "content": content} - lines = content.split("\n") + lines: Final = content.split("\n") # Estimate target line count from character budget - avg_line_len = max(1, len(content) // max(1, len(lines))) - target_lines = max(2, target_chars // avg_line_len) + avg_line_len: Final = max(1, len(content) // max(1, len(lines))) + target_lines: Final = max(2, target_chars // avg_line_len) if len(lines) <= target_lines: return {**message, "content": content} - first_count = (target_lines * 7) // 10 - last_count = target_lines - first_count - truncated = ( + first_count: Final = (target_lines * 7) // 10 + last_count: Final = target_lines - first_count + truncated: Final = ( "\n".join(lines[:first_count]) + "\n...[truncated for context window]...\n" + "\n".join(lines[-last_count:]) ) return {**message, "content": truncated} diff --git a/litellm/compression/scoring/bm25.py b/litellm/compression/scoring/bm25.py index 1ff5962835c..a42ab7919f9 100644 --- a/litellm/compression/scoring/bm25.py +++ b/litellm/compression/scoring/bm25.py @@ -7,6 +7,7 @@ No external dependencies — uses only stdlib. import math import re from collections import Counter +from typing import Final def _tokenize(text: str) -> list[str]: @@ -16,11 +17,11 @@ def _tokenize(text: str) -> list[str]: def _extract_content(message: dict) -> str: """Extract text content from a message dict.""" - content = message.get("content", "") + content: Final = message.get("content", "") if isinstance(content, str): return content if isinstance(content, list): - parts = [] + parts: Final = [] for part in content: if isinstance(part, dict) and part.get("type") == "text": parts.append(part.get("text", "")) @@ -48,32 +49,32 @@ def bm25_score_messages( Returns: List of float scores, one per message. Higher = more relevant. """ - query_terms = _tokenize(query) + query_terms: Final = _tokenize(query) if not query_terms: return [0.0] * len(messages) # Tokenize all documents - doc_tokens: list[list[str]] = [] + doc_tokens: Final[list[list[str]]] = [] for msg in messages: doc_tokens.append(_tokenize(_extract_content(msg))) - n = len(doc_tokens) + n: Final = len(doc_tokens) if n == 0: return [] # Average document length - doc_lengths = [len(dt) for dt in doc_tokens] - avgdl = sum(doc_lengths) / n if n > 0 else 1.0 + doc_lengths: Final = [len(dt) for dt in doc_tokens] + avgdl: Final = sum(doc_lengths) / n if n > 0 else 1.0 # Document frequency for each term - df: dict[str, int] = {} + df: Final[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: Final[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) @@ -85,7 +86,7 @@ def bm25_score_messages( # stemmer dependency. def _expand_tf(query_term: str, tf_counts: Counter) -> int: # type: ignore[type-arg] """Sum TF across all doc tokens that are prefixed by query_term.""" - exact = tf_counts.get(query_term, 0) + exact: Final = tf_counts.get(query_term, 0) if exact: return exact if len(query_term) < 4: @@ -93,7 +94,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: Final[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 c65856a0414..aab1371e097 100644 --- a/litellm/compression/scoring/embedding_scorer.py +++ b/litellm/compression/scoring/embedding_scorer.py @@ -5,18 +5,18 @@ Computes cosine similarity between the query embedding and each message embeddin """ import math -from typing import Any +from typing import Any, Final from litellm.caching.dual_cache import DualCache def _extract_content(message: dict) -> str: """Extract text content from a message dict.""" - content = message.get("content", "") + content: Final = message.get("content", "") if isinstance(content, str): return content if isinstance(content, list): - parts = [] + parts: Final = [] for part in content: if isinstance(part, dict) and part.get("type") == "text": parts.append(part.get("text", "")) @@ -30,15 +30,15 @@ def _truncate_text(text: str, max_chars: int = 30000) -> str: """Truncate long text, keeping first and last portions.""" if len(text) <= max_chars: return text - half = max_chars // 2 + half: Final = max_chars // 2 return text[:half] + "\n...\n" + text[-half:] 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)) - norm_b = math.sqrt(sum(x * x for x in b)) + dot: Final = sum(x * y for x, y in zip(a, b)) + norm_a: Final = math.sqrt(sum(x * x for x in a)) + norm_b: Final = math.sqrt(sum(x * x for x in b)) if norm_a == 0 or norm_b == 0: return 0.0 return dot / (norm_a * norm_b) @@ -67,12 +67,12 @@ def embedding_score_messages( """ import litellm - texts = [_truncate_text(query)] + texts: Final = [_truncate_text(query)] for msg in messages: texts.append(_truncate_text(_extract_content(msg))) # Filter out empty texts — replace with a placeholder to maintain indexing - processed_texts = [t if t.strip() else "empty" for t in texts] + processed_texts: Final = [t if t.strip() else "empty" for t in texts] kwargs: dict[str, Any] = { "model": model, @@ -82,13 +82,13 @@ def embedding_score_messages( if embedding_model_params: kwargs = {**kwargs, **embedding_model_params} - response = litellm.embedding(**kwargs) + response: Final = litellm.embedding(**kwargs) # Extract embedding vectors - embeddings = [item["embedding"] for item in response.data] + embeddings: Final = [item["embedding"] for item in response.data] - query_embedding = embeddings[0] - scores: list[float] = [] + query_embedding: Final = embeddings[0] + scores: Final[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 164f5a77a76..663dbf3d6d1 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1,232 +1,232 @@ import os import sys -from typing import Literal +from typing import Final, Literal from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none -DEFAULT_HEALTH_CHECK_PROMPT = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) -AZURE_DEFAULT_RESPONSES_API_VERSION = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) -ROUTER_MAX_FALLBACKS = int(os.getenv("ROUTER_MAX_FALLBACKS", 5)) -DEFAULT_BATCH_SIZE = int(os.getenv("DEFAULT_BATCH_SIZE", 512)) -DEFAULT_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5)) -DEFAULT_S3_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) -DEFAULT_S3_BATCH_SIZE = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) -DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10)) -DEFAULT_NUM_WORKERS_LITELLM_PROXY = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1)) +DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) +AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) +ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5)) +DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512)) +DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5)) +DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) +DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) +DEFAULT_SQS_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10)) +DEFAULT_NUM_WORKERS_LITELLM_PROXY: Final = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1)) DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE = int(os.getenv("DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE", 1)) -DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512)) -SQS_SEND_MESSAGE_ACTION = "SendMessage" -SQS_API_VERSION = "2012-11-05" -DEFAULT_MAX_RETRIES = int(os.getenv("DEFAULT_MAX_RETRIES", 2)) +DEFAULT_SQS_BATCH_SIZE: Final = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512)) +SQS_SEND_MESSAGE_ACTION: Final = "SendMessage" +SQS_API_VERSION: Final = "2012-11-05" +DEFAULT_MAX_RETRIES: Final = int(os.getenv("DEFAULT_MAX_RETRIES", 2)) # Max records accepted in one POST /v1/callbacks/logs batch. Bounds the blast # radius: each record fans out to spend logs + every callback integration. -MAX_CALLBACK_LOG_RECORDS = 1000 -DEFAULT_MAX_RECURSE_DEPTH = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100)) +MAX_CALLBACK_LOG_RECORDS: Final = 1000 +DEFAULT_MAX_RECURSE_DEPTH: Final = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100)) DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10)) -DEFAULT_FAILURE_THRESHOLD_PERCENT = float( +DEFAULT_FAILURE_THRESHOLD_PERCENT: Final = float( os.getenv("DEFAULT_FAILURE_THRESHOLD_PERCENT", 0.5) ) # default cooldown a deployment if 50% of requests fail in a given minute -DEFAULT_MAX_TOKENS = int(os.getenv("DEFAULT_MAX_TOKENS", 4096)) -DEFAULT_ALLOWED_FAILS = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3)) -DEFAULT_REDIS_SYNC_INTERVAL = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1)) -DEFAULT_COOLDOWN_TIME_SECONDS = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5)) -DEFAULT_REPLICATE_POLLING_RETRIES = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5)) -DEFAULT_REPLICATE_POLLING_DELAY_SECONDS = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1)) -DEFAULT_IMAGE_TOKEN_COUNT = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250)) +DEFAULT_MAX_TOKENS: Final = int(os.getenv("DEFAULT_MAX_TOKENS", 4096)) +DEFAULT_ALLOWED_FAILS: Final = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3)) +DEFAULT_REDIS_SYNC_INTERVAL: Final = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1)) +DEFAULT_COOLDOWN_TIME_SECONDS: Final = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5)) +DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5)) +DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1)) +DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250)) # Maximum wall-clock seconds a streaming response is allowed to run. # Streams exceeding this duration are terminated with a Timeout error. # None (default) = no limit. Set env var to a number of seconds to enable globally. -_max_stream_duration_env = os.getenv("LITELLM_MAX_STREAMING_DURATION_SECONDS", None) -LITELLM_MAX_STREAMING_DURATION_SECONDS = ( +_max_stream_duration_env: Final = os.getenv("LITELLM_MAX_STREAMING_DURATION_SECONDS", None) +LITELLM_MAX_STREAMING_DURATION_SECONDS: Final = ( float(_max_stream_duration_env) if _max_stream_duration_env is not None else None ) # Maximum number of base64 characters to keep in logging payloads. # Data URIs exceeding this are replaced with a size placeholder. # Set to 0 to disable truncation. -MAX_BASE64_LENGTH_FOR_LOGGING = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64)) +MAX_BASE64_LENGTH_FOR_LOGGING: Final = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64)) # When true, adds detailed per-phase timing breakdown headers to responses. # Headers: x-litellm-timing-{pre-processing,llm-api,post-processing,message-copy}-ms -LITELLM_DETAILED_TIMING = os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true" +LITELLM_DETAILED_TIMING: Final = os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true" # Model cost map validation constants -MODEL_COST_MAP_MIN_MODEL_COUNT = int( +MODEL_COST_MAP_MIN_MODEL_COUNT: Final = int( os.getenv("MODEL_COST_MAP_MIN_MODEL_COUNT", 50) ) # Minimum number of models a fetched cost map must contain to be considered valid -MODEL_COST_MAP_MAX_SHRINK_RATIO = float( +MODEL_COST_MAP_MAX_SHRINK_RATIO: Final = float( os.getenv("MODEL_COST_MAP_MAX_SHRINK_RATIO", 0.5) ) # Maximum allowed shrinkage ratio vs local backup (0.5 = reject if fetched map is <50% of backup) -DEFAULT_IMAGE_WIDTH = int(os.getenv("DEFAULT_IMAGE_WIDTH", 300)) -DEFAULT_IMAGE_HEIGHT = int(os.getenv("DEFAULT_IMAGE_HEIGHT", 300)) +DEFAULT_IMAGE_WIDTH: Final = int(os.getenv("DEFAULT_IMAGE_WIDTH", 300)) +DEFAULT_IMAGE_HEIGHT: Final = int(os.getenv("DEFAULT_IMAGE_HEIGHT", 300)) # Maximum size for image URL downloads in MB (default 50MB, set to 0 to disable limit) # This prevents memory issues from downloading very large images # Maps to OpenAI's 50 MB payload limit - requests with images exceeding this size will be rejected # Set MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0 to disable image URL handling entirely -MAX_IMAGE_URL_DOWNLOAD_SIZE_MB = float(os.getenv("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB", 50)) -MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = int( +MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: Final = float(os.getenv("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB", 50)) +MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB: Final = int( os.getenv("MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB", 1024) ) # 1MB = 1024KB # Surrogate-repair fallback in _read_request_body runs two full-body re.sub passes # that block the event loop on multi-MB malformed bodies. Skip the repair above this # size and raise the existing 400 immediately. Set to 0 to disable the cap. -MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB = get_env_int("MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 1) -SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = int( +MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB: Final = get_env_int("MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 1) +SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD: Final = int( os.getenv("SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD", 1000) ) # Minimum number of requests to consider "reasonable traffic". Used for single-deployment cooldown logic. -DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS = int( +DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS: Final = int( os.getenv("DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS", 5) ) # Minimum number of requests before applying error rate cooldown. Prevents cooldown from triggering on first failure. DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET", 0)) # MCP Semantic Tool Filter Defaults -DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL = str( +DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL: Final = str( os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL", "text-embedding-3-small") ) -DEFAULT_MCP_SEMANTIC_FILTER_TOP_K = int(os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_TOP_K", 10)) -DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD = float( +DEFAULT_MCP_SEMANTIC_FILTER_TOP_K: Final = int(os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_TOP_K", 10)) +DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD: Final = float( os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3) ) -MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150)) +MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH: Final = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150)) # Semantic Guard Defaults -DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL = str( +DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL: Final = str( os.getenv("DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL", "text-embedding-3-small") ) DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD", 0.75)) # MCP OAuth2 Client Credentials Defaults -MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) -MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200")) -MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) +MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS: Final = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) +MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200")) +MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) # Default npm cache directory for STDIO MCP servers. # npm/npx needs a writable cache dir; in containers the default (~/.npm) # may not exist or be read-only. /tmp is always writable. -MCP_NPM_CACHE_DIR = os.getenv("MCP_NPM_CACHE_DIR", "/tmp/.npm_mcp_cache") -MCP_OAUTH2_TOKEN_CACHE_MIN_TTL = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MIN_TTL", "10")) +MCP_NPM_CACHE_DIR: Final = os.getenv("MCP_NPM_CACHE_DIR", "/tmp/.npm_mcp_cache") +MCP_OAUTH2_TOKEN_CACHE_MIN_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MIN_TTL", "10")) # Per-user OAuth token Redis cache (for server-side token storage) -MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX = "mcp:per_user_token" -MCP_PER_USER_TOKEN_DEFAULT_TTL = int( +MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX: Final = "mcp:per_user_token" +MCP_PER_USER_TOKEN_DEFAULT_TTL: Final = int( os.getenv("MCP_PER_USER_TOKEN_DEFAULT_TTL", "43200") # 12 hours ) -MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS = int(os.getenv("MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) +MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS: Final = int(os.getenv("MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) # MCP timeout defaults (seconds). Override via env vars for slow/custom MCP servers. -MCP_CLIENT_TIMEOUT = float(os.getenv("LITELLM_MCP_CLIENT_TIMEOUT", "60.0")) -MCP_TOOL_LISTING_TIMEOUT = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIMEOUT", "30.0")) -MCP_METADATA_TIMEOUT = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0")) -MCP_HEALTH_CHECK_TIMEOUT = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0")) +MCP_CLIENT_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_CLIENT_TIMEOUT", "60.0")) +MCP_TOOL_LISTING_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIMEOUT", "30.0")) +MCP_METADATA_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0")) +MCP_HEALTH_CHECK_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0")) # Allowlist of commands permitted for MCP stdio transport. # Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation. # Note: allowlisted runtimes can still execute code via args (e.g. python -c "..."). # This is an accepted residual risk since these endpoints require PROXY_ADMIN. # Extend via LITELLM_MCP_STDIO_EXTRA_COMMANDS env var (comma-separated). -_MCP_STDIO_EXTRA_COMMANDS = os.getenv("LITELLM_MCP_STDIO_EXTRA_COMMANDS", "") -MCP_STDIO_ALLOWED_COMMANDS: frozenset = frozenset( +_MCP_STDIO_EXTRA_COMMANDS: Final = os.getenv("LITELLM_MCP_STDIO_EXTRA_COMMANDS", "") +MCP_STDIO_ALLOWED_COMMANDS: Final[frozenset] = frozenset( {"npx", "uvx", "python", "python3", "node", "docker", "deno"} | (set(_MCP_STDIO_EXTRA_COMMANDS.split(",")) - {""}) ) # MCP OAuth2 Token Exchange (OBO) Defaults -MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE = int(os.getenv("MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE", "500")) +MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE", "500")) -LITELLM_UI_ALLOW_HEADERS = [ +LITELLM_UI_ALLOW_HEADERS: Final = [ "x-litellm-semantic-filter", "x-litellm-semantic-filter-tools", "x-litellm-adaptive-router-model", ] # Gemini model-specific minimal thinking budget constants -DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH = int( +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH: Final = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH", 1) ) -DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO = int( +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO: Final = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO", 128) ) -DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int( +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE: Final = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512) ) # Maximum number of callbacks that can be registered # This prevents callbacks from exponentially growing and consuming CPU resources # Override with LITELLM_MAX_CALLBACKS env var for large deployments (e.g., many teams with guardrails) -MAX_CALLBACKS = get_env_int("LITELLM_MAX_CALLBACKS", 100) +MAX_CALLBACKS: Final = get_env_int("LITELLM_MAX_CALLBACKS", 100) # Metadata key recording which pre_call guardrails the proxy loop already ran, # so the deployment-level hook does not re-run them for the same request -PRE_CALL_EXECUTED_GUARDRAILS_KEY = "_pre_call_executed_guardrails" +PRE_CALL_EXECUTED_GUARDRAILS_KEY: Final = "_pre_call_executed_guardrails" # Generic fallback for unknown models -DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET: Final = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) ) # Provider-specific API base URLs -XAI_API_BASE = "https://api.x.ai/v1" -OPEN_SANDBOX_API_BASE_ENV_VAR = "OPEN_SANDBOX_API_BASE" -OPEN_SANDBOX_API_KEY_ENV_VAR = "OPEN_SANDBOX_API_KEY" -OPEN_SANDBOX_DEFAULT_TEMPLATE = "opensandbox/code-interpreter:v1.1.0" -_OPEN_SANDBOX_FALLBACK_ENTRYPOINT = "/opt/code-interpreter/code-interpreter.sh" -OPEN_SANDBOX_DEFAULT_ENTRYPOINT = (_OPEN_SANDBOX_FALLBACK_ENTRYPOINT,) -OPEN_SANDBOX_DEFAULT_LANGUAGE = "python" -OPEN_SANDBOX_DEFAULT_CPU_LIMIT = "1" -OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT = "2Gi" -OPEN_SANDBOX_EXECD_PORT = 44772 -OPEN_SANDBOX_DEFAULT_TIMEOUT = 300 -OPEN_SANDBOX_READY_TIMEOUT = 30.0 -OPEN_SANDBOX_POLL_INTERVAL = 0.2 +XAI_API_BASE: Final = "https://api.x.ai/v1" +OPEN_SANDBOX_API_BASE_ENV_VAR: Final = "OPEN_SANDBOX_API_BASE" +OPEN_SANDBOX_API_KEY_ENV_VAR: Final = "OPEN_SANDBOX_API_KEY" +OPEN_SANDBOX_DEFAULT_TEMPLATE: Final = "opensandbox/code-interpreter:v1.1.0" +_OPEN_SANDBOX_FALLBACK_ENTRYPOINT: Final = "/opt/code-interpreter/code-interpreter.sh" +OPEN_SANDBOX_DEFAULT_ENTRYPOINT: Final = (_OPEN_SANDBOX_FALLBACK_ENTRYPOINT,) +OPEN_SANDBOX_DEFAULT_LANGUAGE: Final = "python" +OPEN_SANDBOX_DEFAULT_CPU_LIMIT: Final = "1" +OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT: Final = "2Gi" +OPEN_SANDBOX_EXECD_PORT: Final = 44772 +OPEN_SANDBOX_DEFAULT_TIMEOUT: Final = 300 +OPEN_SANDBOX_READY_TIMEOUT: Final = 30.0 +OPEN_SANDBOX_POLL_INTERVAL: Final = 0.2 DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024)) -DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET = int( +DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET: Final = int( os.getenv("DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET", 2048) ) DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET", 4096)) DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET", 8192)) DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET", 16384)) -MAX_TOKEN_TRIMMING_ATTEMPTS = int( +MAX_TOKEN_TRIMMING_ATTEMPTS: Final = int( os.getenv("MAX_TOKEN_TRIMMING_ATTEMPTS", 10) ) # Maximum number of attempts to trim the message -RUNWAYML_DEFAULT_API_VERSION = str(os.getenv("RUNWAYML_DEFAULT_API_VERSION", "2024-11-06")) +RUNWAYML_DEFAULT_API_VERSION: Final = str(os.getenv("RUNWAYML_DEFAULT_API_VERSION", "2024-11-06")) RUNWAYML_POLLING_TIMEOUT = int(os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)) # 10 minutes default for image generation ########## Networking constants ############################################################## -_DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client for 1 hour +_DEFAULT_TTL_FOR_HTTPX_CLIENTS: Final = 3600 # 1 hour, re-use the same httpx client for 1 hour # The earliest an evicted, litellm-created client may be closed. A request handed the # client just before eviction is still using it, so nothing is closed inside this window; # past it, the client is closed once it reports no connection in flight. -EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS = 900 +EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS: Final = 900 # How many evicted clients may be queued for closing at once. Past this, an evicted client # is left to the collector rather than letting a cache-churning workload grow the queue # without bound. Each queued entry is ~100 bytes and comes due within one grace window. -EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING = 10_000 +EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING: Final = 10_000 # Aiohttp connection pooling - prevents memory leaks from unbounded connection growth # Set to 0 for unlimited (not recommended for production) -AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000)) -AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 500)) -AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120)) -AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300)) +AIOHTTP_CONNECTOR_LIMIT: Final = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000)) +AIOHTTP_CONNECTOR_LIMIT_PER_HOST: Final = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 500)) +AIOHTTP_KEEPALIVE_TIMEOUT: Final = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120)) +AIOHTTP_TTL_DNS_CACHE: Final = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300)) # TCP keep-alive (SO_KEEPALIVE) — opt-in. Required when running behind NAT/LBs # whose idle timeout is shorter than provider response timeouts (e.g. AWS NAT # Gateway: 350s vs OpenAI/Azure: 600s). Without this, the kernel sends nothing # during a long provider call and the NAT reaps the flow before the response # arrives. Enabling SO_KEEPALIVE makes the kernel emit TCP probes that reset # the NAT idle timer. -AIOHTTP_SO_KEEPALIVE = os.getenv("AIOHTTP_SO_KEEPALIVE", "False").lower() == "true" -AIOHTTP_TCP_KEEPIDLE = int(os.getenv("AIOHTTP_TCP_KEEPIDLE", 60)) -AIOHTTP_TCP_KEEPINTVL = int(os.getenv("AIOHTTP_TCP_KEEPINTVL", 30)) -AIOHTTP_TCP_KEEPCNT = int(os.getenv("AIOHTTP_TCP_KEEPCNT", 5)) +AIOHTTP_SO_KEEPALIVE: Final = os.getenv("AIOHTTP_SO_KEEPALIVE", "False").lower() == "true" +AIOHTTP_TCP_KEEPIDLE: Final = int(os.getenv("AIOHTTP_TCP_KEEPIDLE", 60)) +AIOHTTP_TCP_KEEPINTVL: Final = int(os.getenv("AIOHTTP_TCP_KEEPINTVL", 30)) +AIOHTTP_TCP_KEEPCNT: Final = int(os.getenv("AIOHTTP_TCP_KEEPCNT", 5)) # enable_cleanup_closed is only needed for Python versions with the SSL leak bug # Fixed in Python 3.12.7+ and 3.13.1+ (see https://github.com/python/cpython/pull/118960) # Reference: https://github.com/aio-libs/aiohttp/blob/master/aiohttp/connector.py#L74-L78 -AIOHTTP_NEEDS_CLEANUP_CLOSED = (3, 13, 0) <= sys.version_info < ( +AIOHTTP_NEEDS_CLEANUP_CLOSED: Final = (3, 13, 0) <= sys.version_info < ( 3, 13, 1, @@ -235,13 +235,13 @@ AIOHTTP_NEEDS_CLEANUP_CLOSED = (3, 13, 0) <= sys.version_info < ( # WebSocket constants # Default to None (unlimited) to match OpenAI's official agents SDK behavior # https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235 -_max_size_env = os.getenv("REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES") -REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES = int(_max_size_env) if _max_size_env is not None else None +_max_size_env: Final = os.getenv("REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES") +REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES: Final = int(_max_size_env) if _max_size_env is not None else None # SSL/TLS cipher configuration for faster handshakes # Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones # This balances performance with broad compatibility -DEFAULT_SSL_CIPHERS = os.getenv( +DEFAULT_SSL_CIPHERS: Final = os.getenv( "LITELLM_SSL_CIPHERS", # Priority 1: TLS 1.3 ciphers (fastest, ~50ms handshake) "TLS_AES_256_GCM_SHA384:" # Fastest observed in testing @@ -263,141 +263,141 @@ DEFAULT_SSL_CIPHERS = os.getenv( ) ########### v2 Architecture constants for managing writing updates to the database ########### -REDIS_UPDATE_BUFFER_KEY = "litellm_spend_update_buffer" -REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_spend_update_buffer" -REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_team_spend_update_buffer" -REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_org_spend_update_buffer" -REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_end_user_spend_update_buffer" -REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_agent_spend_update_buffer" -REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer" -MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100)) +REDIS_UPDATE_BUFFER_KEY: Final = "litellm_spend_update_buffer" +REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_spend_update_buffer" +REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_team_spend_update_buffer" +REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_org_spend_update_buffer" +REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_end_user_spend_update_buffer" +REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_agent_spend_update_buffer" +REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_tag_spend_update_buffer" +MAX_REDIS_BUFFER_DEQUEUE_COUNT: Final = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100)) # Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth -LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000)) -TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60)) -GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS = int( +LITELLM_ASYNCIO_QUEUE_MAXSIZE: Final = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000)) +TOOL_POLICY_CACHE_TTL_SECONDS: Final = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60)) +GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS: Final = int( os.getenv("GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS", 24 * 60 * 60) ) # Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger. # Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire. -MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8))) -MAX_IN_MEMORY_QUEUE_FLUSH_COUNT = int(os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000)) +MAX_SIZE_IN_MEMORY_QUEUE: Final = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8))) +MAX_IN_MEMORY_QUEUE_FLUSH_COUNT: Final = int(os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000)) ############################################################################################### # Providers will not cache a prefix below a minimum size. That minimum is per-model, not global: # Anthropic's ranges from 512 to 4096 depending on the model, and can differ per platform for the # same model. The real minimum is resolved from `prompt_cache_min_tokens` in the model cost map; # this value is only the fallback for models the cost map has no entry for, and doubles as a global # escape hatch when `MINIMUM_PROMPT_CACHE_TOKEN_COUNT` is explicitly set. -MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE: int | None = get_env_int_or_none("MINIMUM_PROMPT_CACHE_TOKEN_COUNT") -DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT = 1024 -MINIMUM_PROMPT_CACHE_TOKEN_COUNT = ( +MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE: Final[int | None] = get_env_int_or_none("MINIMUM_PROMPT_CACHE_TOKEN_COUNT") +DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT: Final = 1024 +MINIMUM_PROMPT_CACHE_TOKEN_COUNT: Final = ( MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None else DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT ) -DEFAULT_TRIM_RATIO = float( +DEFAULT_TRIM_RATIO: Final = float( os.getenv("DEFAULT_TRIM_RATIO", 0.75) ) # default ratio of tokens to trim from the end of a prompt -HOURS_IN_A_DAY = int(os.getenv("HOURS_IN_A_DAY", 24)) -DAYS_IN_A_WEEK = int(os.getenv("DAYS_IN_A_WEEK", 7)) -DAYS_IN_A_MONTH = int(os.getenv("DAYS_IN_A_MONTH", 28)) -DAYS_IN_A_YEAR = int(os.getenv("DAYS_IN_A_YEAR", 365)) -REPLICATE_MODEL_NAME_WITH_ID_LENGTH = int(os.getenv("REPLICATE_MODEL_NAME_WITH_ID_LENGTH", 64)) +HOURS_IN_A_DAY: Final = int(os.getenv("HOURS_IN_A_DAY", 24)) +DAYS_IN_A_WEEK: Final = int(os.getenv("DAYS_IN_A_WEEK", 7)) +DAYS_IN_A_MONTH: Final = int(os.getenv("DAYS_IN_A_MONTH", 28)) +DAYS_IN_A_YEAR: Final = int(os.getenv("DAYS_IN_A_YEAR", 365)) +REPLICATE_MODEL_NAME_WITH_ID_LENGTH: Final = int(os.getenv("REPLICATE_MODEL_NAME_WITH_ID_LENGTH", 64)) #### TOKEN COUNTING #### -FUNCTION_DEFINITION_TOKEN_COUNT = int(os.getenv("FUNCTION_DEFINITION_TOKEN_COUNT", 9)) -SYSTEM_MESSAGE_TOKEN_COUNT = int(os.getenv("SYSTEM_MESSAGE_TOKEN_COUNT", 4)) -TOOL_CHOICE_OBJECT_TOKEN_COUNT = int(os.getenv("TOOL_CHOICE_OBJECT_TOKEN_COUNT", 4)) -DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT = int(os.getenv("DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT", 10)) -DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT = int(os.getenv("DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT", 20)) -MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES = int(os.getenv("MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES", 768)) -MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES = int(os.getenv("MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES", 2000)) -MAX_TILE_WIDTH = int(os.getenv("MAX_TILE_WIDTH", 512)) -MAX_TILE_HEIGHT = int(os.getenv("MAX_TILE_HEIGHT", 512)) -OPENAI_FILE_SEARCH_COST_PER_1K_CALLS = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000)) -GROQ_BROWSER_VISIT_WEBSITE_COST_PER_CALL = 1.0 / 1000 +FUNCTION_DEFINITION_TOKEN_COUNT: Final = int(os.getenv("FUNCTION_DEFINITION_TOKEN_COUNT", 9)) +SYSTEM_MESSAGE_TOKEN_COUNT: Final = int(os.getenv("SYSTEM_MESSAGE_TOKEN_COUNT", 4)) +TOOL_CHOICE_OBJECT_TOKEN_COUNT: Final = int(os.getenv("TOOL_CHOICE_OBJECT_TOKEN_COUNT", 4)) +DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT", 10)) +DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT", 20)) +MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES: Final = int(os.getenv("MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES", 768)) +MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES: Final = int(os.getenv("MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES", 2000)) +MAX_TILE_WIDTH: Final = int(os.getenv("MAX_TILE_WIDTH", 512)) +MAX_TILE_HEIGHT: Final = int(os.getenv("MAX_TILE_HEIGHT", 512)) +OPENAI_FILE_SEARCH_COST_PER_1K_CALLS: Final = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000)) +GROQ_BROWSER_VISIT_WEBSITE_COST_PER_CALL: Final = 1.0 / 1000 # Azure OpenAI Assistants feature costs # Source: https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/ -AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY = float( +AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY: Final = float( os.getenv("AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day ) -AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS = float( +AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS: Final = float( os.getenv("AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS", 3.0) # $0.003 USD per 1K Tokens ) -AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS = float( +AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS: Final = float( os.getenv("AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS", 12.0) # $0.012 USD per 1K Tokens ) -AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY = float( +AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY: Final = float( os.getenv("AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day (same as file search) ) -MIN_NON_ZERO_TEMPERATURE = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001)) +MIN_NON_ZERO_TEMPERATURE: Final = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001)) #### RELIABILITY #### -REPEATED_STREAMING_CHUNK_LIMIT = int( +REPEATED_STREAMING_CHUNK_LIMIT: Final = int( os.getenv("REPEATED_STREAMING_CHUNK_LIMIT", 100) ) # catch if model starts looping the same chunk while streaming. Uses high default to prevent false positives. # Shared maxsize for functools.lru_cache usage across hot paths. # Defaulted to 64 to avoid cache thrash in multi-model production workloads. -DEFAULT_MAX_LRU_CACHE_SIZE = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 64)) +DEFAULT_MAX_LRU_CACHE_SIZE: Final = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 64)) _REALTIME_BODY_CACHE_SIZE = 1000 # Keep realtime helper caches bounded; workloads rarely exceed 1k models/intents -INITIAL_RETRY_DELAY = float(os.getenv("INITIAL_RETRY_DELAY", 0.5)) -MAX_RETRY_DELAY = float(os.getenv("MAX_RETRY_DELAY", 8.0)) -JITTER = float(os.getenv("JITTER", 0.75)) +INITIAL_RETRY_DELAY: Final = float(os.getenv("INITIAL_RETRY_DELAY", 0.5)) +MAX_RETRY_DELAY: Final = float(os.getenv("MAX_RETRY_DELAY", 8.0)) +JITTER: Final = float(os.getenv("JITTER", 0.75)) DEFAULT_IN_MEMORY_TTL = int(os.getenv("DEFAULT_IN_MEMORY_TTL", 5)) # default time to live for the in-memory cache -DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE = int( +DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE: Final = int( os.getenv("DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE", 1000) ) # default max size for redis batch cache -DEFAULT_POLLING_INTERVAL = float( +DEFAULT_POLLING_INTERVAL: Final = float( os.getenv("DEFAULT_POLLING_INTERVAL", 0.03) ) # default polling interval for the scheduler -AZURE_OPERATION_POLLING_TIMEOUT = int(os.getenv("AZURE_OPERATION_POLLING_TIMEOUT", 120)) -AZURE_DOCUMENT_INTELLIGENCE_API_VERSION = str(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_API_VERSION", "2024-11-30")) -AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI = int(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI", 96)) -REDIS_SOCKET_TIMEOUT = float(os.getenv("REDIS_SOCKET_TIMEOUT", 0.1)) -REDIS_CONNECTION_POOL_TIMEOUT = int(os.getenv("REDIS_CONNECTION_POOL_TIMEOUT", 5)) -REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5)) -REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60)) -REDIS_CIRCUIT_BREAKER_ENABLED = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true" +AZURE_OPERATION_POLLING_TIMEOUT: Final = int(os.getenv("AZURE_OPERATION_POLLING_TIMEOUT", 120)) +AZURE_DOCUMENT_INTELLIGENCE_API_VERSION: Final = str(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_API_VERSION", "2024-11-30")) +AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI: Final = int(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI", 96)) +REDIS_SOCKET_TIMEOUT: Final = float(os.getenv("REDIS_SOCKET_TIMEOUT", 0.1)) +REDIS_CONNECTION_POOL_TIMEOUT: Final = int(os.getenv("REDIS_CONNECTION_POOL_TIMEOUT", 5)) +REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD: Final = int(os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5)) +REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT: Final = int(os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60)) +REDIS_CIRCUIT_BREAKER_ENABLED: Final = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true" # Seconds of idle before a Redis cluster connection is validated with a PING and # reconnected if dead, so a connection silently dropped by a cluster restart # (e.g. ElastiCache Serverless maintenance) is not reused while broken -REDIS_CLUSTER_HEALTH_CHECK_INTERVAL = 25 +REDIS_CLUSTER_HEALTH_CHECK_INTERVAL: Final = 25 # Default Redis major version to assume when version cannot be determined # Using 7 as it's the modern version that supports LPOP with count parameter -DEFAULT_REDIS_MAJOR_VERSION = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7)) -NON_LLM_CONNECTION_TIMEOUT = int( +DEFAULT_REDIS_MAJOR_VERSION: Final = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7)) +NON_LLM_CONNECTION_TIMEOUT: Final = int( os.getenv("NON_LLM_CONNECTION_TIMEOUT", 15) ) # timeout for adjacent services (e.g. jwt auth) -MAX_EXCEPTION_MESSAGE_LENGTH = int(os.getenv("MAX_EXCEPTION_MESSAGE_LENGTH", 2000)) -MAX_STRING_LENGTH_PROMPT_IN_DB = int(os.getenv("MAX_STRING_LENGTH_PROMPT_IN_DB", 2048)) -BEDROCK_MAX_POLICY_SIZE = int(os.getenv("BEDROCK_MAX_POLICY_SIZE", 75)) +MAX_EXCEPTION_MESSAGE_LENGTH: Final = int(os.getenv("MAX_EXCEPTION_MESSAGE_LENGTH", 2000)) +MAX_STRING_LENGTH_PROMPT_IN_DB: Final = int(os.getenv("MAX_STRING_LENGTH_PROMPT_IN_DB", 2048)) +BEDROCK_MAX_POLICY_SIZE: Final = int(os.getenv("BEDROCK_MAX_POLICY_SIZE", 75)) # One entry per distinct AWS credential-argument set. Per-user cost attribution passes the attributed # identity as aws_session_name, so this bounds how many attributed identities keep a cached STS session. -BEDROCK_IAM_CACHE_MAX_ENTRIES = 1000 +BEDROCK_IAM_CACHE_MAX_ENTRIES: Final = 1000 # Single-flight lock stripes over that cache. Only keys landing on the same stripe wait for each # other, so a burst of distinct identities still resolves its credentials in parallel. -BEDROCK_IAM_CACHE_FETCH_LOCK_STRIPES = 64 +BEDROCK_IAM_CACHE_FETCH_LOCK_STRIPES: Final = 64 # Retire a cached STS credential this many seconds before AWS expires it, so a request that reads it # still has a usable credential for the whole call. -STS_CREDENTIAL_EXPIRY_SAFETY_MARGIN_SECONDS = 60 -BEDROCK_MIN_THINKING_BUDGET_TOKENS = int(os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024)) +STS_CREDENTIAL_EXPIRY_SAFETY_MARGIN_SECONDS: Final = 60 +BEDROCK_MIN_THINKING_BUDGET_TOKENS: Final = int(os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024)) # Anthropic's Messages API rejects thinking.budget_tokens < 1024. -ANTHROPIC_MIN_THINKING_BUDGET_TOKENS = 1024 -REPLICATE_POLLING_DELAY_SECONDS = float(os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5)) -DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int(os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096)) -DEFAULT_OCI_CHAT_MAX_TOKENS = 4096 -TOGETHER_AI_4_B = int(os.getenv("TOGETHER_AI_4_B", 4)) -TOGETHER_AI_8_B = int(os.getenv("TOGETHER_AI_8_B", 8)) -TOGETHER_AI_21_B = int(os.getenv("TOGETHER_AI_21_B", 21)) -TOGETHER_AI_41_B = int(os.getenv("TOGETHER_AI_41_B", 41)) -TOGETHER_AI_80_B = int(os.getenv("TOGETHER_AI_80_B", 80)) -TOGETHER_AI_110_B = int(os.getenv("TOGETHER_AI_110_B", 110)) -TOGETHER_AI_EMBEDDING_150_M = int(os.getenv("TOGETHER_AI_EMBEDDING_150_M", 150)) -TOGETHER_AI_EMBEDDING_350_M = int(os.getenv("TOGETHER_AI_EMBEDDING_350_M", 350)) -QDRANT_SCALAR_QUANTILE = float(os.getenv("QDRANT_SCALAR_QUANTILE", 0.99)) -QDRANT_VECTOR_SIZE = int(os.getenv("QDRANT_VECTOR_SIZE", 1536)) -CACHED_STREAMING_CHUNK_DELAY = float(os.getenv("CACHED_STREAMING_CHUNK_DELAY", 0.02)) -AUDIO_SPEECH_CHUNK_SIZE = int( +ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: Final = 1024 +REPLICATE_POLLING_DELAY_SECONDS: Final = float(os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5)) +DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS: Final = int(os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096)) +DEFAULT_OCI_CHAT_MAX_TOKENS: Final = 4096 +TOGETHER_AI_4_B: Final = int(os.getenv("TOGETHER_AI_4_B", 4)) +TOGETHER_AI_8_B: Final = int(os.getenv("TOGETHER_AI_8_B", 8)) +TOGETHER_AI_21_B: Final = int(os.getenv("TOGETHER_AI_21_B", 21)) +TOGETHER_AI_41_B: Final = int(os.getenv("TOGETHER_AI_41_B", 41)) +TOGETHER_AI_80_B: Final = int(os.getenv("TOGETHER_AI_80_B", 80)) +TOGETHER_AI_110_B: Final = int(os.getenv("TOGETHER_AI_110_B", 110)) +TOGETHER_AI_EMBEDDING_150_M: Final = int(os.getenv("TOGETHER_AI_EMBEDDING_150_M", 150)) +TOGETHER_AI_EMBEDDING_350_M: Final = int(os.getenv("TOGETHER_AI_EMBEDDING_350_M", 350)) +QDRANT_SCALAR_QUANTILE: Final = float(os.getenv("QDRANT_SCALAR_QUANTILE", 0.99)) +QDRANT_VECTOR_SIZE: Final = int(os.getenv("QDRANT_VECTOR_SIZE", 1536)) +CACHED_STREAMING_CHUNK_DELAY: Final = float(os.getenv("CACHED_STREAMING_CHUNK_DELAY", 0.02)) +AUDIO_SPEECH_CHUNK_SIZE: Final = int( os.getenv("AUDIO_SPEECH_CHUNK_SIZE", 8192) ) # chunk_size for audio speech streaming. Balance between latency and memory usage -DEFAULT_MAX_TOKENS_FOR_TRITON = int(os.getenv("DEFAULT_MAX_TOKENS_FOR_TRITON", 2000)) +DEFAULT_MAX_TOKENS_FOR_TRITON: Final = int(os.getenv("DEFAULT_MAX_TOKENS_FOR_TRITON", 2000)) #### Networking settings #### # Sentinel used when `REQUEST_TIMEOUT` is unset: `litellm.request_timeout` keeps this # value so longer-running surfaces (Router `timeout or litellm.request_timeout`, @@ -405,70 +405,70 @@ DEFAULT_MAX_TOKENS_FOR_TRITON = int(os.getenv("DEFAULT_MAX_TOKENS_FOR_TRITON", 2 # `completion()` maps this sentinel down to 600s when the caller did not set a # per-request/model timeout—see ``CompletionTimeout.resolve`` in completion_timeout.py. MCP uses # dedicated timeouts (e.g. `MCP_CLIENT_TIMEOUT`), not `request_timeout`. -DEFAULT_REQUEST_TIMEOUT_SECONDS: float = 6000.0 +DEFAULT_REQUEST_TIMEOUT_SECONDS: Final[float] = 6000.0 # Pair used for default httpx clients when no custom timeout is passed: read/write # deadline and connect handshake (see ``http_handler`` cached handler paths). -COMPLETION_HTTP_FALLBACK_SECONDS: float = 600.0 -HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS: float = 5.0 +COMPLETION_HTTP_FALLBACK_SECONDS: Final[float] = 600.0 +HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS: Final[float] = 5.0 request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS)))) request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ -DEFAULT_A2A_AGENT_TIMEOUT: float = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes +DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes # 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: Final[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: Final[list[str]] = [ "connect", "connection", "network", "refused", ] -STREAM_SSE_DONE_STRING: str = "[DONE]" -STREAM_SSE_DATA_PREFIX: str = "data: " +STREAM_SSE_DONE_STRING: Final[str] = "[DONE]" +STREAM_SSE_DATA_PREFIX: Final[str] = "data: " ### SPEND TRACKING ### -DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND = float( +DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND: Final = float( os.getenv("DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND", 0.001400) ) # price per second for a100 80GB -FIREWORKS_AI_56_B_MOE = int(os.getenv("FIREWORKS_AI_56_B_MOE", 56)) -FIREWORKS_AI_176_B_MOE = int(os.getenv("FIREWORKS_AI_176_B_MOE", 176)) -FIREWORKS_AI_4_B = int(os.getenv("FIREWORKS_AI_4_B", 4)) -FIREWORKS_AI_16_B = int(os.getenv("FIREWORKS_AI_16_B", 16)) -FIREWORKS_AI_80_B = int(os.getenv("FIREWORKS_AI_80_B", 80)) +FIREWORKS_AI_56_B_MOE: Final = int(os.getenv("FIREWORKS_AI_56_B_MOE", 56)) +FIREWORKS_AI_176_B_MOE: Final = int(os.getenv("FIREWORKS_AI_176_B_MOE", 176)) +FIREWORKS_AI_4_B: Final = int(os.getenv("FIREWORKS_AI_4_B", 4)) +FIREWORKS_AI_16_B: Final = int(os.getenv("FIREWORKS_AI_16_B", 16)) +FIREWORKS_AI_80_B: Final = int(os.getenv("FIREWORKS_AI_80_B", 80)) #### Logging callback constants #### -REDACTED_BY_LITELM_STRING = "REDACTED_BY_LITELM" -MAX_LANGFUSE_INITIALIZED_CLIENTS = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50)) -LOGGING_WORKER_CONCURRENCY = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0 -LOGGING_WORKER_MAX_QUEUE_SIZE = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000)) -LOGGING_WORKER_MAX_TIME_PER_COROUTINE = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0)) -LOGGING_WORKER_CLEAR_PERCENTAGE = int( +REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM" +MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50)) +LOGGING_WORKER_CONCURRENCY: Final = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0 +LOGGING_WORKER_MAX_QUEUE_SIZE: Final = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000)) +LOGGING_WORKER_MAX_TIME_PER_COROUTINE: Final = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0)) +LOGGING_WORKER_CLEAR_PERCENTAGE: Final = int( os.getenv("LOGGING_WORKER_CLEAR_PERCENTAGE", 50) ) # Percentage of queue to clear (default: 50%) -MAX_ITERATIONS_TO_CLEAR_QUEUE = int(os.getenv("MAX_ITERATIONS_TO_CLEAR_QUEUE", 200)) -MAX_TIME_TO_CLEAR_QUEUE = float(os.getenv("MAX_TIME_TO_CLEAR_QUEUE", 5.0)) -LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS = float( +MAX_ITERATIONS_TO_CLEAR_QUEUE: Final = int(os.getenv("MAX_ITERATIONS_TO_CLEAR_QUEUE", 200)) +MAX_TIME_TO_CLEAR_QUEUE: Final = float(os.getenv("MAX_TIME_TO_CLEAR_QUEUE", 5.0)) +LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS: Final = float( os.getenv("LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS", 0.5) ) # Cooldown time in seconds before allowing another aggressive clear (default: 0.5s) -DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE = os.getenv( +DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE: Final = os.getenv( "DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield" ) -LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED = 499 +LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED: Final = 499 -EMAIL_BUDGET_ALERT_TTL = int(os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)) # 24 hours in seconds -EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE = float( +EMAIL_BUDGET_ALERT_TTL: Final = int(os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)) # 24 hours in seconds +EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE: Final = float( os.getenv("EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE", 0.8) ) # 80% of max budget ############### LLM Provider Constants ############### ### ANTHROPIC CONSTANTS ### ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01") -ANTHROPIC_SKILLS_API_BETA_VERSION = "skills-2025-10-02" -ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = { +ANTHROPIC_SKILLS_API_BETA_VERSION: Final = "skills-2025-10-02" +ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES: Final = { "low": 1, "medium": 5, "high": 10, @@ -476,19 +476,19 @@ ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = { # LiteLLM standard web search tool name # Used for web search interception across providers -LITELLM_WEB_SEARCH_TOOL_NAME = "litellm_web_search" +LITELLM_WEB_SEARCH_TOOL_NAME: Final = "litellm_web_search" -DEFAULT_IMAGE_ENDPOINT_MODEL = "dall-e-2" -DEFAULT_VIDEO_ENDPOINT_MODEL = "sora-2" +DEFAULT_IMAGE_ENDPOINT_MODEL: Final = "dall-e-2" +DEFAULT_VIDEO_ENDPOINT_MODEL: Final = "sora-2" -DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS = int(os.getenv("DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS", 8)) +DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS: Final = int(os.getenv("DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS", 8)) ### DATAFORSEO CONSTANTS ### -DEFAULT_DATAFORSEO_LOCATION_CODE = int( +DEFAULT_DATAFORSEO_LOCATION_CODE: Final = int( os.getenv("DEFAULT_DATAFORSEO_LOCATION_CODE", 2250) ) # Default to France (2250) - lower number, commonly used location -LITELLM_CHAT_PROVIDERS = [ +LITELLM_CHAT_PROVIDERS: Final = [ "openai", "openai_like", "bytez", @@ -585,7 +585,7 @@ LITELLM_CHAT_PROVIDERS = [ "amazon_nova", ] -LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [ +LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS: Final = [ "openai", "azure", "hosted_vllm", @@ -593,7 +593,7 @@ LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [ ] -OPENAI_CHAT_COMPLETION_PARAMS = [ +OPENAI_CHAT_COMPLETION_PARAMS: Final = [ "functions", "function_call", "temperature", @@ -642,22 +642,22 @@ OPENAI_CHAT_COMPLETION_PARAMS = [ "store", ] -OPENAI_TRANSCRIPTION_PARAMS = [ +OPENAI_TRANSCRIPTION_PARAMS: Final = [ "language", "response_format", "timestamp_granularities", ] -OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user"] +OPENAI_EMBEDDING_PARAMS: Final = ["dimensions", "encoding_format", "user"] -DEFAULT_EMBEDDING_PARAM_VALUES = { +DEFAULT_EMBEDDING_PARAM_VALUES: Final = { **{k: None for k in OPENAI_EMBEDDING_PARAMS}, "model": None, "custom_llm_provider": "", "input": None, } -DEFAULT_CHAT_COMPLETION_PARAM_VALUES = { +DEFAULT_CHAT_COMPLETION_PARAM_VALUES: Final = { "functions": None, "function_call": None, "temperature": None, @@ -705,7 +705,7 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = { "context_management": None, } -openai_compatible_endpoints: list = [ +openai_compatible_endpoints: Final[list] = [ "api.perplexity.ai", "api.endpoints.anyscale.com/v1", "api.deepinfra.com/v1/openai", @@ -751,7 +751,7 @@ openai_compatible_endpoints: list = [ ] -openai_compatible_providers: list = [ +openai_compatible_providers: Final[list] = [ "anyscale", "groq", "nvidia_nim", @@ -816,7 +816,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: Final[list] = [ # providers that support `/v1/completions` "together_ai", "fireworks_ai", "hosted_vllm", @@ -839,14 +839,14 @@ openai_text_completion_compatible_providers: list = [ # providers that support "hyperbolic", "wandb", ] -_openai_like_providers: list = [ +_openai_like_providers: Final[list] = [ "predibase", "databricks", "lemonade", "watsonx", ] # private helper. similar to openai but require some custom auth / endpoint handling, so can't use the openai sdk # well supported replicate llms -replicate_models: set = set( +replicate_models: Final[set] = set( [ # llama replicate supported LLMs "replicate/llama-2-70b-chat:2796ee9483c3fd7aa2e171d38f4ca12251a30609463dcfd4cd76703f22e96cdf", @@ -863,7 +863,7 @@ replicate_models: set = set( ] ) -clarifai_models: set = set( +clarifai_models: Final[set] = set( [ "clarifai/openai.chat-completion.gpt-oss-20b", "clarifai/qwen.qwenLM.Qwen3-30B-A3B-Instruct-2507", @@ -899,7 +899,7 @@ clarifai_models: set = set( ) -huggingface_models: set = set( +huggingface_models: Final[set] = set( [ "meta-llama/Llama-2-7b-hf", "meta-llama/Llama-2-7b-chat-hf", @@ -915,14 +915,14 @@ huggingface_models: set = set( "meta-llama/Llama-2-70b-chat", ] ) # these have been tested on extensively. But by default all text2text-generation and text-generation models are supported by liteLLM. - https://docs.litellm.ai/docs/providers -empower_models = set( +empower_models: Final = set( [ "empower/empower-functions", "empower/empower-functions-small", ] ) -together_ai_models: set = set( +together_ai_models: Final[set] = set( [ # llama llms - chat "togethercomputer/llama-2-70b-chat", @@ -956,7 +956,7 @@ together_ai_models: set = set( # supports all together ai models, just pass in the model id e.g. completion(model="together_computer/replit_code_3b",...) -baseten_models: set = set( +baseten_models: Final[set] = set( [ "qvv0xeq", "q841o8w", @@ -964,7 +964,7 @@ baseten_models: set = set( ] ) # FALCON 7B # WizardLM # Mosaic ML -featherless_ai_models: set = set( +featherless_ai_models: Final[set] = set( [ "featherless-ai/Qwerky-72B", "featherless-ai/Qwerky-QwQ-32B", @@ -978,7 +978,7 @@ featherless_ai_models: set = set( ] ) -nebius_models: set = set( +nebius_models: Final[set] = set( [ # deepseek models "deepseek-ai/DeepSeek-R1-0528", @@ -1032,7 +1032,7 @@ nebius_models: set = set( ] ) -dashscope_models: set = set( +dashscope_models: Final[set] = set( [ "qwen-turbo", "qwen-plus", @@ -1047,7 +1047,7 @@ dashscope_models: set = set( ] ) -nebius_embedding_models: set = set( +nebius_embedding_models: Final[set] = set( [ "BAAI/bge-en-icl", "BAAI/bge-multilingual-gemma2", @@ -1055,7 +1055,7 @@ nebius_embedding_models: set = set( ] ) -WANDB_MODELS: set = set( +WANDB_MODELS: Final[set] = set( [ # openai models "openai/gpt-oss-120b", @@ -1084,7 +1084,7 @@ WANDB_MODELS: set = set( ] ) -modelscope_models: set = set( +modelscope_models: Final[set] = set( [ # Qwen series models "Qwen/Qwen3-0.6B", @@ -1151,7 +1151,7 @@ BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[ "nova", ] -BEDROCK_CONVERSE_MODELS = [ +BEDROCK_CONVERSE_MODELS: Final = [ "qwen.qwen3-coder-480b-a35b-v1:0", "qwen.qwen3-coder-next", "qwen.qwen3-235b-a22b-2507-v1:0", @@ -1212,8 +1212,8 @@ BEDROCK_CONVERSE_MODELS = [ ] -open_ai_embedding_models: set = set(["text-embedding-ada-002"]) -cohere_embedding_models: set = set( +open_ai_embedding_models: Final[set] = set(["text-embedding-ada-002"]) +cohere_embedding_models: Final[set] = set( [ "embed-v4.0", "embed-english-v3.0", @@ -1224,7 +1224,7 @@ cohere_embedding_models: set = set( "embed-multilingual-v2.0", ] ) -bedrock_embedding_models: set = set( +bedrock_embedding_models: Final[set] = set( [ "amazon.titan-embed-text-v1", "amazon.nova-2-multimodal-embeddings-v1:0", @@ -1235,7 +1235,7 @@ bedrock_embedding_models: set = set( ] ) -known_tokenizer_config = { +known_tokenizer_config: Final = { "mistralai/Mistral-7B-Instruct-v0.1": { "tokenizer": { "chat_template": "{{ bos_token }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if message['role'] == 'user' %}{{ '[INST] ' + message['content'] + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ message['content'] + eos_token + ' ' }}{% else %}{{ raise_exception('Only user and assistant roles are supported!') }}{% endif %}{% endfor %}", @@ -1293,33 +1293,33 @@ known_tokenizer_config = { } -OPENAI_FINISH_REASONS = [ +OPENAI_FINISH_REASONS: Final = [ "stop", "length", "function_call", "tool_calls", "content_filter", ] -HUMANLOOP_PROMPT_CACHE_TTL_SECONDS = int(os.getenv("HUMANLOOP_PROMPT_CACHE_TTL_SECONDS", 60)) # 1 minute +HUMANLOOP_PROMPT_CACHE_TTL_SECONDS: Final = int(os.getenv("HUMANLOOP_PROMPT_CACHE_TTL_SECONDS", 60)) # 1 minute RESPONSE_FORMAT_TOOL_NAME = "json_tool_call" # default tool name used when converting response format to tool call ########################### Logging Callback Constants ########################### -AZURE_STORAGE_MSFT_VERSION = "2019-07-07" -PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES = int( +AZURE_STORAGE_MSFT_VERSION: Final = "2019-07-07" +PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int( os.getenv("PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES", 5) ) -CLOUDZERO_EXPORT_INTERVAL_MINUTES = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60)) -MCP_TOOL_NAME_PREFIX = "mcp_tool" -MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100)) +CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60)) +MCP_TOOL_NAME_PREFIX: Final = "mcp_tool" +MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100)) # Headers to control callbacks -X_LITELLM_DISABLE_CALLBACKS = "x-litellm-disable-callbacks" -LITELLM_METADATA_FIELD = "litellm_metadata" -OLD_LITELLM_METADATA_FIELD = "metadata" -RETURN_RAW_MODEL_NAME_METADATA_KEY = "_complexity_router_return_raw_model_name" -INTERNAL_CALL_ORIGIN_METADATA_KEY = "internal_call_origin" -LITELLM_TRUNCATED_PAYLOAD_FIELD = "litellm_truncated" -LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE = ( +X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks" +LITELLM_METADATA_FIELD: Final = "litellm_metadata" +OLD_LITELLM_METADATA_FIELD: Final = "metadata" +RETURN_RAW_MODEL_NAME_METADATA_KEY: Final = "_complexity_router_return_raw_model_name" +INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin" +LITELLM_TRUNCATED_PAYLOAD_FIELD: Final = "litellm_truncated" +LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE: Final = ( "Truncation is a DB storage safeguard. " "Full, untruncated data is logged to logging callbacks (OTEL, Datadog, etc.). " "To increase the truncation limit, set `MAX_STRING_LENGTH_PROMPT_IN_DB` in your env." @@ -1330,25 +1330,25 @@ LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE = ( # Standard headers that are always checked for customer/end-user ID (no configuration required) # These headers work out-of-the-box for tools like Claude Code that support custom headers -STANDARD_CUSTOMER_ID_HEADERS = [ +STANDARD_CUSTOMER_ID_HEADERS: Final = [ "x-litellm-customer-id", "x-litellm-end-user-id", ] -MAX_SPENDLOG_ROWS_TO_QUERY = int( +MAX_SPENDLOG_ROWS_TO_QUERY: Final = int( os.getenv("MAX_SPENDLOG_ROWS_TO_QUERY", 1_000_000) ) # if spendLogs has more than 1M rows, do not query the DB -DEFAULT_SOFT_BUDGET = float( +DEFAULT_SOFT_BUDGET: Final = float( os.getenv("DEFAULT_SOFT_BUDGET", 50.0) ) # by default all litellm proxy keys have a soft budget of 50.0 # makes it clear this is a rate limit error for a litellm virtual key -RATE_LIMIT_ERROR_MESSAGE_FOR_VIRTUAL_KEY = "LiteLLM Virtual Key user_api_key_hash" +RATE_LIMIT_ERROR_MESSAGE_FOR_VIRTUAL_KEY: Final = "LiteLLM Virtual Key user_api_key_hash" # Python garbage collection threshold configuration # Format: "gen0,gen1,gen2" e.g., "1000,50,50" -PYTHON_GC_THRESHOLD = os.getenv("PYTHON_GC_THRESHOLD") +PYTHON_GC_THRESHOLD: Final = os.getenv("PYTHON_GC_THRESHOLD") # pass through route constansts -BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES = [ +BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES: Final = [ "agents/", "knowledgebases/", "flows/", @@ -1361,7 +1361,7 @@ BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES = [ # Headers that are safe to forward from incoming requests to Vertex AI # Using an allowlist approach for security - only forward headers we explicitly trust -ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS = { +ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = { "anthropic-beta", # Required for Anthropic features like extended context windows "content-type", # Required for request body parsing } @@ -1369,17 +1369,17 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS = { # Prefix for headers that should be forwarded to the provider with the prefix stripped # e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value' # Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.) -PASS_THROUGH_HEADER_PREFIX = "x-pass-" +PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-" -BASE_MCP_ROUTE = "/mcp" +BASE_MCP_ROUTE: Final = "/mcp" -BATCH_STATUS_POLL_INTERVAL_SECONDS = int(os.getenv("BATCH_STATUS_POLL_INTERVAL_SECONDS", 3600)) # 1 hour -BATCH_STATUS_POLL_MAX_ATTEMPTS = int(os.getenv("BATCH_STATUS_POLL_MAX_ATTEMPTS", 24)) # for 24 hours +BATCH_STATUS_POLL_INTERVAL_SECONDS: Final = int(os.getenv("BATCH_STATUS_POLL_INTERVAL_SECONDS", 3600)) # 1 hour +BATCH_STATUS_POLL_MAX_ATTEMPTS: Final = int(os.getenv("BATCH_STATUS_POLL_MAX_ATTEMPTS", 24)) # for 24 hours -HEALTH_CHECK_TIMEOUT_SECONDS = int(os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)) # 60 seconds -_background_health_check_max_tokens_env = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS") +HEALTH_CHECK_TIMEOUT_SECONDS: Final = int(os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)) # 60 seconds +_background_health_check_max_tokens_env: Final = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS") try: - _raw_background_health_check_max_tokens = ( + _raw_background_health_check_max_tokens: Final = ( _background_health_check_max_tokens_env.strip() if _background_health_check_max_tokens_env is not None else "" ) BACKGROUND_HEALTH_CHECK_MAX_TOKENS: int | None = ( @@ -1389,9 +1389,9 @@ except (ValueError, TypeError): BACKGROUND_HEALTH_CHECK_MAX_TOKENS = None -_background_health_check_max_tokens_reasoning_env = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING") +_background_health_check_max_tokens_reasoning_env: Final = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING") try: - _raw_background_health_check_max_tokens_reasoning = ( + _raw_background_health_check_max_tokens_reasoning: Final = ( _background_health_check_max_tokens_reasoning_env.strip() if _background_health_check_max_tokens_reasoning_env is not None else "" @@ -1404,141 +1404,141 @@ try: except (ValueError, TypeError): BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING = None -LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check" -LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli" -LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME = "litellm_internal_jobs" +LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME: Final = "litellm-internal-health-check" +LITTELM_CLI_SERVICE_ACCOUNT_NAME: Final = "litellm-cli" +LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME: Final = "litellm_internal_jobs" # Stable identifier substituted in place of the master key on UserAPIKeyAuth # objects so the master key (or its hash) never propagates to spend logs, # Prometheus metrics, audit trails, or any other downstream consumer. -LITELLM_PROXY_MASTER_KEY_ALIAS = "litellm_proxy_master_key" +LITELLM_PROXY_MASTER_KEY_ALIAS: Final = "litellm_proxy_master_key" # Marker placed in ``model_call_details`` on a synthetic ``Logging`` object that # records a proxy-gate error (auth/rate-limit rejection) for a request that never # reached an upstream provider. Tracing callbacks key off it to avoid fabricating # an LLM-call span for a call that did not happen. See # ``ProxyLogging._handle_logging_proxy_only_error``. -LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL = "litellm_no_upstream_llm_call" +LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL: Final = "litellm_no_upstream_llm_call" # Key Rotation Constants -LITELLM_KEY_ROTATION_ENABLED = os.getenv("LITELLM_KEY_ROTATION_ENABLED", "false") -LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS = int( +LITELLM_KEY_ROTATION_ENABLED: Final = os.getenv("LITELLM_KEY_ROTATION_ENABLED", "false") +LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS: Final = int( os.getenv("LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS", 86400) ) # 24 hours default -LITELLM_KEY_ROTATION_GRACE_PERIOD: str = os.getenv( +LITELLM_KEY_ROTATION_GRACE_PERIOD: Final[str] = os.getenv( "LITELLM_KEY_ROTATION_GRACE_PERIOD", "" ) # Duration to keep old key valid after rotation (e.g. "24h", "2d"); empty = immediate revoke (default) -LITELLM_KEY_ROTATION_LOCK_TTL_SECONDS = int( +LITELLM_KEY_ROTATION_LOCK_TTL_SECONDS: Final = int( os.getenv("LITELLM_KEY_ROTATION_LOCK_TTL_SECONDS", 600) ) # 10 minutes default — caps the deadlock window if a pod crashes mid-rotation -UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard" +UI_SESSION_TOKEN_TEAM_ID: Final = "litellm-dashboard" LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED = os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED", "false") -LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS = int( +LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS: Final = int( os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS", 86400) ) # 24 hours default -LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE = int( +LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE: Final = int( os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE", 1000) ) -LITELLM_PROXY_ADMIN_NAME = "default_user_id" -LITELLM_PROXY_BUDGET_NAME = "litellm-proxy-budget" -GLOBAL_PROXY_SPEND_CACHE_KEY = f"{LITELLM_PROXY_ADMIN_NAME}:spend" +LITELLM_PROXY_ADMIN_NAME: Final = "default_user_id" +LITELLM_PROXY_BUDGET_NAME: Final = "litellm-proxy-budget" +GLOBAL_PROXY_SPEND_CACHE_KEY: Final = f"{LITELLM_PROXY_ADMIN_NAME}:spend" ########################### CLI SSO AUTHENTICATION CONSTANTS ########################### -LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli" -LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token" -CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session" -CLI_SSO_SESSION_TTL_SECONDS = 600 -CLI_SESSION_KEY_PREFIX = "cli-session" +LITELLM_CLI_SOURCE_IDENTIFIER: Final = "litellm-cli" +LITELLM_CLI_SESSION_TOKEN_PREFIX: Final = "litellm-session-token" +CLI_SSO_SESSION_CACHE_KEY_PREFIX: Final = "cli_sso_session" +CLI_SSO_SESSION_TTL_SECONDS: Final = 600 +CLI_SESSION_KEY_PREFIX: Final = "cli-session" # Support both CLI_JWT_EXPIRATION_HOURS and LITELLM_CLI_JWT_EXPIRATION_HOURS for backwards compatibility -CLI_JWT_EXPIRATION_HOURS = int( +CLI_JWT_EXPIRATION_HOURS: Final = int( os.getenv("CLI_JWT_EXPIRATION_HOURS") or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS") or 24 ) # Comma-separated allowlisted OIDC claim map for CLI SSO polling, e.g. # "employment_type->acme_employment_type,org_info.department->department" -CLI_SSO_CLAIM_MAP = os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or "" -CLI_SSO_CLAIM_MAX_SCALAR_LENGTH = 1024 +CLI_SSO_CLAIM_MAP: Final = os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or "" +CLI_SSO_CLAIM_MAX_SCALAR_LENGTH: Final = 1024 ########################### UI SESSION DURATION ########################### # Duration for UI login session (username/password, SSO, invitation links). Format: "30s", "30m", "24h", "7d" # Does NOT apply to EXPERIMENTAL_UI_LOGIN flow, which intentionally uses a fixed 10-minute expiry for security. -LITELLM_UI_SESSION_DURATION = os.getenv("LITELLM_UI_SESSION_DURATION", "24h") +LITELLM_UI_SESSION_DURATION: Final = os.getenv("LITELLM_UI_SESSION_DURATION", "24h") ########################### DB CRON JOB NAMES ########################### -DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" -DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME = "db_daily_tag_spend_update_job" -PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics" -CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data" -MAVVRIK_FOCUS_EXPORT_JOB_NAME = "mavvrik_focus_export_usage_data" -CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)) -SPEND_LOG_CLEANUP_JOB_NAME = "spend_log_cleanup" -KEY_ROTATION_JOB_NAME = "litellm_key_rotation_job" -EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME = "litellm_expired_ui_session_key_cleanup_job" -SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) -SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) +DB_SPEND_UPDATE_JOB_NAME: Final = "db_spend_update_job" +DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME: Final = "db_daily_tag_spend_update_job" +PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME: Final = "prometheus_emit_budget_metrics" +CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME: Final = "cloudzero_export_usage_data" +MAVVRIK_FOCUS_EXPORT_JOB_NAME: Final = "mavvrik_focus_export_usage_data" +CLOUDZERO_MAX_FETCHED_DATA_RECORDS: Final = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)) +SPEND_LOG_CLEANUP_JOB_NAME: Final = "spend_log_cleanup" +KEY_ROTATION_JOB_NAME: Final = "litellm_key_rotation_job" +EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME: Final = "litellm_expired_ui_session_key_cleanup_job" +SPEND_LOG_RUN_LOOPS: Final = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) +SPEND_LOG_CLEANUP_BATCH_SIZE: Final = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)) -SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float( +SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS: Final = float( os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5) ) -TOOL_SPEND_TOP_TOOLS = 100 -SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day") -SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) -SPEND_LOG_WRITE_BATCH_MAX_BYTES = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) -SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) -SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) -SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) -DEFAULT_CRON_JOB_LOCK_TTL_SECONDS = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute -PROXY_BUDGET_RESCHEDULER_MIN_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)) -PROXY_BATCH_POLLING_INTERVAL = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) -MAX_OBJECTS_PER_POLL_CYCLE = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) -MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))) -STALE_OBJECT_CLEANUP_BATCH_SIZE = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))) +TOOL_SPEND_TOP_TOOLS: Final = 100 +SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day") +SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) +SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) +SPEND_LOG_QUEUE_SIZE_THRESHOLD: Final = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) +SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) +SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) +DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute +PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)) +PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) +MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) +MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))) +STALE_OBJECT_CLEANUP_BATCH_SIZE: Final = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))) # Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and # CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on # installations with large numbers of stale managed objects). -_batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower() -PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true" -PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605)) -PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10 -PROXY_CONFIG_RELOAD_INTERVAL_SECONDS = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30) +_batch_polling_env: Final = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower() +PROXY_BATCH_POLLING_ENABLED: Final = _batch_polling_env == "true" +PROXY_BUDGET_RESCHEDULER_MAX_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605)) +PROXY_BATCH_WRITE_AT: Final = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10 +PROXY_CONFIG_RELOAD_INTERVAL_SECONDS: Final = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30) # APScheduler Configuration - MEMORY LEAK FIX # These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions -APSCHEDULER_COALESCE = os.getenv("APSCHEDULER_COALESCE", "True").lower() in [ +APSCHEDULER_COALESCE: Final = os.getenv("APSCHEDULER_COALESCE", "True").lower() in [ "true", "1", ] # collapse many missed runs into one -APSCHEDULER_MISFIRE_GRACE_TIME = int( +APSCHEDULER_MISFIRE_GRACE_TIME: Final = int( os.getenv("APSCHEDULER_MISFIRE_GRACE_TIME", 3600) ) # ignore runs older than 1 hour (was 120) -APSCHEDULER_MAX_INSTANCES = int(os.getenv("APSCHEDULER_MAX_INSTANCES", 1)) # prevent concurrent job instances -APSCHEDULER_REPLACE_EXISTING = os.getenv("APSCHEDULER_REPLACE_EXISTING", "True").lower() in [ +APSCHEDULER_MAX_INSTANCES: Final = int(os.getenv("APSCHEDULER_MAX_INSTANCES", 1)) # prevent concurrent job instances +APSCHEDULER_REPLACE_EXISTING: Final = os.getenv("APSCHEDULER_REPLACE_EXISTING", "True").lower() in [ "true", "1", ] # always replace existing jobs # The number of tag entries are higher than number of user, team entries. This leads to a higher QPS. # This will run tag spcific tasks at a later time to smooth QPS -DAILY_TAG_SPEND_BATCH_MULTIPLIER = 2.3 +DAILY_TAG_SPEND_BATCH_MULTIPLIER: Final = 2.3 -DEFAULT_HEALTH_CHECK_INTERVAL = int(os.getenv("DEFAULT_HEALTH_CHECK_INTERVAL", 300)) # 5 minutes -DEFAULT_SHARED_HEALTH_CHECK_TTL = int( +DEFAULT_HEALTH_CHECK_INTERVAL: Final = int(os.getenv("DEFAULT_HEALTH_CHECK_INTERVAL", 300)) # 5 minutes +DEFAULT_SHARED_HEALTH_CHECK_TTL: Final = int( os.getenv("DEFAULT_SHARED_HEALTH_CHECK_TTL", 300) ) # 5 minutes - TTL for cached health check results -DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL = int( +DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL: Final = int( os.getenv("DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL", 60) ) # 1 minute - TTL for health check lock -DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER = 2 # health state is stale after interval * this -PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS = int(os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9)) -DEFAULT_MODEL_CREATED_AT_TIME = int( +DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER: Final = 2 # health state is stale after interval * this +PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS: Final = int(os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9)) +DEFAULT_MODEL_CREATED_AT_TIME: Final = int( os.getenv("DEFAULT_MODEL_CREATED_AT_TIME", 1677610602) ) # returns on `/models` endpoint -DEFAULT_SLACK_ALERTING_THRESHOLD = int(os.getenv("DEFAULT_SLACK_ALERTING_THRESHOLD", 300)) -MAX_TEAM_LIST_LIMIT = int(os.getenv("MAX_TEAM_LIST_LIMIT", 20)) -MAX_POLICY_ESTIMATE_IMPACT_ROWS = int(os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000)) +DEFAULT_SLACK_ALERTING_THRESHOLD: Final = int(os.getenv("DEFAULT_SLACK_ALERTING_THRESHOLD", 300)) +MAX_TEAM_LIST_LIMIT: Final = int(os.getenv("MAX_TEAM_LIST_LIMIT", 20)) +MAX_POLICY_ESTIMATE_IMPACT_ROWS: Final = int(os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000)) DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD", 0.7)) -LENGTH_OF_LITELLM_GENERATED_KEY = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY", 16)) -MINIMUM_CUSTOM_KEY_LENGTH = int(os.getenv("MINIMUM_CUSTOM_KEY_LENGTH", 16)) -SECRET_MANAGER_REFRESH_INTERVAL = int(os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400)) -LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [ +LENGTH_OF_LITELLM_GENERATED_KEY: Final = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY", 16)) +MINIMUM_CUSTOM_KEY_LENGTH: Final = int(os.getenv("MINIMUM_CUSTOM_KEY_LENGTH", 16)) +SECRET_MANAGER_REFRESH_INTERVAL: Final = int(os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400)) +LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [ "default_internal_user_params", "default_team_params", "public_mcp_servers", @@ -1556,21 +1556,21 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [ "anthropic_prompt_caching_ttl", "max_ui_session_budget", ] -SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"] +SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"] DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)) -DEFAULT_ACCESS_GROUP_CACHE_TTL = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600)) +DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600)) # Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated # callers from forcing a DB query per request for unknown names, while bounding # staleness so a transient DB error (which surfaces as an empty list) cannot # hide a real group for long. -DEFAULT_MCP_ACCESS_GROUP_NEGATIVE_CACHE_TTL = 10 +DEFAULT_MCP_ACCESS_GROUP_NEGATIVE_CACHE_TTL: Final = 10 # Maximum number of comma-separated MCP server / access-group tokens accepted # in a single ``/{name1,name2,...}/mcp`` URL. Bounds the per-request DB / cache # fan-out an authenticated caller can trigger by stuffing the path with tokens. -DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS = 16 +DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS: Final = 16 # Sentry Scrubbing Configuration -SENTRY_DENYLIST = [ +SENTRY_DENYLIST: Final = [ # API Keys and Tokens "api_key", "token", @@ -1626,7 +1626,7 @@ SENTRY_DENYLIST = [ "proxy_key", "environment_variables", ] -SENTRY_PII_DENYLIST = [ +SENTRY_PII_DENYLIST: Final = [ "user_id", "email", "phone", @@ -1637,41 +1637,41 @@ SENTRY_PII_DENYLIST = [ ] # CoroutineChecker cache configuration -COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY = int(os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000)) +COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY: Final = int(os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000)) ########################### RAG Text Splitter Constants ########################### -DEFAULT_CHUNK_SIZE = int(os.getenv("DEFAULT_CHUNK_SIZE", 1000)) -DEFAULT_CHUNK_OVERLAP = int(os.getenv("DEFAULT_CHUNK_OVERLAP", 200)) +DEFAULT_CHUNK_SIZE: Final = int(os.getenv("DEFAULT_CHUNK_SIZE", 1000)) +DEFAULT_CHUNK_OVERLAP: Final = int(os.getenv("DEFAULT_CHUNK_OVERLAP", 200)) ########################### S3 Vectors RAG Constants ########################### -S3_VECTORS_DEFAULT_DIMENSION = int(os.getenv("S3_VECTORS_DEFAULT_DIMENSION", 1024)) -S3_VECTORS_DEFAULT_DISTANCE_METRIC = str(os.getenv("S3_VECTORS_DEFAULT_DISTANCE_METRIC", "cosine")) -S3_VECTORS_DEFAULT_NON_FILTERABLE_METADATA_KEYS = ["source_text"] +S3_VECTORS_DEFAULT_DIMENSION: Final = int(os.getenv("S3_VECTORS_DEFAULT_DIMENSION", 1024)) +S3_VECTORS_DEFAULT_DISTANCE_METRIC: Final = str(os.getenv("S3_VECTORS_DEFAULT_DISTANCE_METRIC", "cosine")) +S3_VECTORS_DEFAULT_NON_FILTERABLE_METADATA_KEYS: Final = ["source_text"] ########################### Microsoft SSO Constants ########################### -MICROSOFT_USER_EMAIL_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_EMAIL_ATTRIBUTE", "userPrincipalName")) -MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "displayName")) -MICROSOFT_USER_ID_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_ID_ATTRIBUTE", "id")) -MICROSOFT_USER_FIRST_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "givenName")) -MICROSOFT_USER_LAST_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname")) +MICROSOFT_USER_EMAIL_ATTRIBUTE: Final = str(os.getenv("MICROSOFT_USER_EMAIL_ATTRIBUTE", "userPrincipalName")) +MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE: Final = str(os.getenv("MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "displayName")) +MICROSOFT_USER_ID_ATTRIBUTE: Final = str(os.getenv("MICROSOFT_USER_ID_ATTRIBUTE", "id")) +MICROSOFT_USER_FIRST_NAME_ATTRIBUTE: Final = str(os.getenv("MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "givenName")) +MICROSOFT_USER_LAST_NAME_ATTRIBUTE: Final = str(os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname")) # Maximum payload size (in bytes) to fully serialize for DEBUG logging. # Payloads larger than this are truncated to avoid multi-second json.dumps blocking the response. -MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG = int(os.getenv("MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG", 102400)) # 100 KB +MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG: Final = int(os.getenv("MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG", 102400)) # 100 KB # Policy template enrichment -MAX_COMPETITOR_NAMES = int(os.getenv("MAX_COMPETITOR_NAMES", 100)) -COMPETITOR_LLM_TEMPERATURE = float(os.getenv("COMPETITOR_LLM_TEMPERATURE", 0.3)) -DEFAULT_COMPETITOR_DISCOVERY_MODEL = "gpt-4o-mini" +MAX_COMPETITOR_NAMES: Final = int(os.getenv("MAX_COMPETITOR_NAMES", 100)) +COMPETITOR_LLM_TEMPERATURE: Final = float(os.getenv("COMPETITOR_LLM_TEMPERATURE", 0.3)) +DEFAULT_COMPETITOR_DISCOVERY_MODEL: Final = "gpt-4o-mini" # Advisor tool orchestration # Providers that support advisor_20260301 natively (no LiteLLM orchestration needed). # Add vertex_ai here once verified. -ADVISOR_NATIVE_PROVIDERS: frozenset = frozenset({"anthropic"}) +ADVISOR_NATIVE_PROVIDERS: Final[frozenset] = frozenset({"anthropic"}) # Hard cap on advisor iterations per request to prevent runaway loops. -ADVISOR_MAX_USES: int = 5 +ADVISOR_MAX_USES: Final[int] = 5 # Description injected into the synthetic advisor tool definition sent to non-native providers. -ADVISOR_TOOL_DESCRIPTION: str = ( +ADVISOR_TOOL_DESCRIPTION: Final[str] = ( "Consult a highly intelligent advisor model when you need expert guidance, " "want to verify your reasoning, or face a complex decision. " "Describe your question or challenge clearly in the 'question' field." diff --git a/litellm/containers/endpoint_factory.py b/litellm/containers/endpoint_factory.py index 5a1b8da351f..09bc7eda41f 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, Literal +from typing import Any, Final, Literal import litellm from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT @@ -28,7 +28,7 @@ from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client # Response type mapping -RESPONSE_TYPES: dict[str, type] = { +RESPONSE_TYPES: Final[dict[str, type]] = { "ContainerFileListResponse": ContainerFileListResponse, "ContainerFileObject": ContainerFileObject, "DeleteContainerFileResponse": DeleteContainerFileResponse, @@ -37,7 +37,7 @@ RESPONSE_TYPES: dict[str, type] = { def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" - config_path = Path(__file__).parent / "endpoints.json" + config_path: Final = Path(__file__).parent / "endpoints.json" with open(config_path) as f: return json.load(f) @@ -48,9 +48,9 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable: Uses the generic container handler instead of individual handler methods. """ - endpoint_name = endpoint_config["name"] - response_type = RESPONSE_TYPES.get(endpoint_config["response_type"]) - path_params = endpoint_config.get("path_params", []) + endpoint_name: Final = endpoint_config["name"] + response_type: Final = RESPONSE_TYPES.get(endpoint_config["response_type"]) + path_params: Final = endpoint_config.get("path_params", []) @client def endpoint_func( @@ -61,12 +61,12 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable: extra_body: dict[str, Any] | None = None, **kwargs, ): - local_vars = locals() + local_vars: Final = locals() try: resolved_custom_llm_provider: str = custom_llm_provider - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response mock_response = kwargs.get("mock_response") @@ -99,7 +99,7 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable: raise ValueError(f"Container provider config not found for: {resolved_custom_llm_provider}") # Build optional params for logging - optional_params = {k: kwargs.get(k) for k in path_params if k in kwargs} + optional_params: Final = {k: kwargs.get(k) for k in path_params if k in kwargs} # Pre-call logging litellm_logging_obj.update_from_kwargs( @@ -150,12 +150,12 @@ def create_async_endpoint_function( extra_body: dict[str, Any] | None = None, **kwargs, ): - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( sync_func, timeout=timeout, custom_llm_provider=custom_llm_provider, @@ -165,9 +165,9 @@ def create_async_endpoint_function( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -193,8 +193,8 @@ def generate_container_endpoints() -> dict[str, Callable]: Returns a dict mapping function names to their implementations. """ - config = _load_endpoints_config() - endpoints = {} + config: Final = _load_endpoints_config() + endpoints: Final = {} for endpoint_config in config["endpoints"]: # Create sync function @@ -210,8 +210,8 @@ def generate_container_endpoints() -> dict[str, Callable]: def get_all_endpoint_names() -> list[str]: """Get all endpoint names (sync and async) from config.""" - config = _load_endpoints_config() - names = [] + config: Final = _load_endpoints_config() + names: Final = [] for endpoint in config["endpoints"]: names.append(endpoint["name"]) names.append(endpoint["async_name"]) @@ -220,21 +220,21 @@ def get_all_endpoint_names() -> list[str]: def get_async_endpoint_names() -> list[str]: """Get all async endpoint names for router registration.""" - config = _load_endpoints_config() + config: Final = _load_endpoints_config() return [endpoint["async_name"] for endpoint in config["endpoints"]] # Generate endpoints on module load -_generated_endpoints = generate_container_endpoints() +_generated_endpoints: Final = generate_container_endpoints() # Export generated functions dynamically -list_container_files = _generated_endpoints.get("list_container_files") -alist_container_files = _generated_endpoints.get("alist_container_files") -upload_container_file = _generated_endpoints.get("upload_container_file") -aupload_container_file = _generated_endpoints.get("aupload_container_file") -retrieve_container_file = _generated_endpoints.get("retrieve_container_file") -aretrieve_container_file = _generated_endpoints.get("aretrieve_container_file") -delete_container_file = _generated_endpoints.get("delete_container_file") -adelete_container_file = _generated_endpoints.get("adelete_container_file") -retrieve_container_file_content = _generated_endpoints.get("retrieve_container_file_content") -aretrieve_container_file_content = _generated_endpoints.get("aretrieve_container_file_content") +list_container_files: Final = _generated_endpoints.get("list_container_files") +alist_container_files: Final = _generated_endpoints.get("alist_container_files") +upload_container_file: Final = _generated_endpoints.get("upload_container_file") +aupload_container_file: Final = _generated_endpoints.get("aupload_container_file") +retrieve_container_file: Final = _generated_endpoints.get("retrieve_container_file") +aretrieve_container_file: Final = _generated_endpoints.get("aretrieve_container_file") +delete_container_file: Final = _generated_endpoints.get("delete_container_file") +adelete_container_file: Final = _generated_endpoints.get("adelete_container_file") +retrieve_container_file_content: Final = _generated_endpoints.get("retrieve_container_file_content") +aretrieve_container_file_content: Final = _generated_endpoints.get("aretrieve_container_file_content") diff --git a/litellm/containers/main.py b/litellm/containers/main.py index 466237ebce3..c13f8bc75a6 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, Literal, overload +from typing import Any, Final, Literal, overload import litellm from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT @@ -76,12 +76,12 @@ async def acreate_container( Returns: - `response` (ContainerObject): The created container object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( create_container, name=name, expires_after=expires_after, @@ -94,9 +94,9 @@ async def acreate_container( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -185,11 +185,11 @@ def create_container( print(response) ``` """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -197,12 +197,12 @@ def create_container( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = ContainerObject(**mock_response) + response: Final = ContainerObject(**mock_response) return response # get llm provider logic # Pass credential params explicitly since they're named args, not in kwargs - litellm_params = GenericLiteLLMParams( + litellm_params: Final = GenericLiteLLMParams( api_key=api_key, api_base=api_base, api_version=api_version, @@ -218,12 +218,12 @@ def create_container( local_vars.update(kwargs) # Get ContainerCreateOptionalRequestParams with only valid parameters - container_create_optional_params: ContainerCreateOptionalRequestParams = ( + container_create_optional_params: Final[ContainerCreateOptionalRequestParams] = ( ContainerRequestUtils.get_requested_container_create_optional_param(local_vars) ) # Get optional parameters for the container API - container_create_request_params: dict = ContainerRequestUtils.get_optional_params_container_create( + container_create_request_params: Final[dict] = ContainerRequestUtils.get_optional_params_container_create( container_provider_config=container_provider_config, container_create_optional_params=container_create_optional_params, ) @@ -306,12 +306,12 @@ async def alist_containers( Returns: - `response` (ContainerListResponse): The list of containers """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( list_containers, after=after, limit=limit, @@ -324,9 +324,9 @@ async def alist_containers( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -403,11 +403,11 @@ def list_containers( Currently supports OpenAI """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -415,12 +415,12 @@ def list_containers( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = ContainerListResponse(**mock_response) + response: Final = ContainerListResponse(**mock_response) return response # get llm provider logic # Pass credential params explicitly since they're named args, not in kwargs - litellm_params = GenericLiteLLMParams( + litellm_params: Final = GenericLiteLLMParams( api_key=api_key, api_base=api_base, api_version=api_version, @@ -435,7 +435,7 @@ def list_containers( raise ValueError(f"Container provider config not found for provider: {custom_llm_provider}") # Get container list request parameters - container_list_optional_params: ContainerListOptionalRequestParams = ( + container_list_optional_params: Final[ContainerListOptionalRequestParams] = ( ContainerRequestUtils.get_requested_container_list_optional_param(local_vars) ) @@ -504,12 +504,12 @@ async def aretrieve_container( Returns: - `response` (ContainerObject): The container object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( retrieve_container, container_id=container_id, timeout=timeout, @@ -520,9 +520,9 @@ async def aretrieve_container( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -593,12 +593,12 @@ def retrieve_container( Currently supports OpenAI """ - local_vars = locals() + local_vars: Final = locals() try: resolved_custom_llm_provider: str = custom_llm_provider - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -606,7 +606,7 @@ def retrieve_container( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = ContainerObject(**mock_response) + response: Final = ContainerObject(**mock_response) return response # get llm provider logic @@ -625,7 +625,7 @@ def retrieve_container( litellm_params=litellm_params, ) # True when input was a LiteLLM-managed ID (any length); needed to re-encode output for routing affinity - was_encoded = original_container_id != container_id + was_encoded: Final = original_container_id != container_id # get provider config container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( @@ -719,12 +719,12 @@ async def adelete_container( Returns: - `response` (DeleteContainerResult): The deletion result """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( delete_container, container_id=container_id, timeout=timeout, @@ -735,9 +735,9 @@ async def adelete_container( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -808,12 +808,12 @@ def delete_container( Currently supports OpenAI """ - local_vars = locals() + local_vars: Final = locals() try: resolved_custom_llm_provider: str = custom_llm_provider - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -821,7 +821,7 @@ def delete_container( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = DeleteContainerResult(**mock_response) + response: Final = DeleteContainerResult(**mock_response) return response # get llm provider logic @@ -840,7 +840,7 @@ def delete_container( litellm_params=litellm_params, ) # True when input was a LiteLLM-managed ID (any length); needed to re-encode output for routing affinity - was_encoded = original_container_id != container_id + was_encoded: Final = original_container_id != container_id # get provider config container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config( @@ -938,12 +938,12 @@ async def alist_container_files( Returns: - `response` (ContainerFileListResponse): The list of container files """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( list_container_files, container_id=container_id, after=after, @@ -957,9 +957,9 @@ async def alist_container_files( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1037,12 +1037,12 @@ def list_container_files( Currently supports OpenAI """ - local_vars = locals() + local_vars: Final = locals() try: resolved_custom_llm_provider: str = custom_llm_provider - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -1050,7 +1050,7 @@ def list_container_files( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = ContainerFileListResponse(**mock_response) + response: Final = ContainerFileListResponse(**mock_response) return response # get llm provider logic @@ -1168,12 +1168,12 @@ async def aupload_container_file( print(response) ``` """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - func = partial( + func: Final = partial( upload_container_file, container_id=container_id, file=file, @@ -1185,9 +1185,9 @@ async def aupload_container_file( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1288,12 +1288,12 @@ def upload_container_file( """ from litellm.llms.custom_httpx.container_handler import generic_container_handler - local_vars = locals() + local_vars: Final = locals() try: resolved_custom_llm_provider: str = custom_llm_provider - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id") - _is_async = kwargs.pop("async_call", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id") + _is_async: Final = kwargs.pop("async_call", False) is True # Check for mock response first mock_response = kwargs.get("mock_response") @@ -1301,7 +1301,7 @@ def upload_container_file( if isinstance(mock_response, str): mock_response = json.loads(mock_response) - response = ContainerFileObject(**mock_response) + response: Final = ContainerFileObject(**mock_response) return response # get llm provider logic diff --git a/litellm/containers/utils.py b/litellm/containers/utils.py index 9b48a4f0f98..f07820602bf 100644 --- a/litellm/containers/utils.py +++ b/litellm/containers/utils.py @@ -1,4 +1,4 @@ -from typing import Any, TypeVar +from typing import Any, Final, TypeVar from litellm.llms.base_llm.containers.transformation import BaseContainerConfig from litellm.responses.utils import ResponsesAPIRequestUtils @@ -19,14 +19,14 @@ def decode_managed_container_id_for_request( Returns: (original_container_id, resolved_provider, updated_litellm_params) """ - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) + decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + original_container_id: Final = decoded.get("response_id", container_id) - decoded_provider = decoded.get("custom_llm_provider") + decoded_provider: Final = decoded.get("custom_llm_provider") if decoded_provider and custom_llm_provider == "openai": custom_llm_provider = decoded_provider - decoded_model_id = decoded.get("model_id") + decoded_model_id: Final = decoded.get("model_id") if decoded_model_id and not litellm_params.get("model_id"): litellm_params["model_id"] = decoded_model_id @@ -42,9 +42,9 @@ class ContainerRequestUtils: passed_params: dict, ) -> ContainerCreateOptionalRequestParams: """Extract only valid container creation parameters from the passed parameters.""" - container_create_optional_params = ContainerCreateOptionalRequestParams() + container_create_optional_params: Final = ContainerCreateOptionalRequestParams() - valid_params = [ + valid_params: Final = [ "expires_after", "file_ids", "extra_headers", @@ -63,10 +63,10 @@ class ContainerRequestUtils: container_create_optional_params: ContainerCreateOptionalRequestParams, ) -> dict: """Get the optional parameters for container creation.""" - supported_params = container_provider_config.get_supported_openai_params() + supported_params: Final = container_provider_config.get_supported_openai_params() # Filter out unsupported parameters - filtered_params = {k: v for k, v in container_create_optional_params.items() if k in supported_params} + filtered_params: Final = {k: v for k, v in container_create_optional_params.items() if k in supported_params} return container_provider_config.map_openai_params( container_create_optional_params=filtered_params, # type: ignore @@ -78,9 +78,9 @@ class ContainerRequestUtils: passed_params: dict, ) -> ContainerListOptionalRequestParams: """Extract only valid container list parameters from the passed parameters.""" - container_list_optional_params = ContainerListOptionalRequestParams() + container_list_optional_params: Final = ContainerListOptionalRequestParams() - valid_params = [ + valid_params: Final = [ "after", "limit", "order", @@ -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: Final[dict[str, Any]] = litellm_metadata.get("model_info", {}) or {} model_id = model_info.get("id") # Check if we should encode based on routing metadata @@ -139,7 +139,7 @@ class ContainerRequestUtils: should_encode = True # Extract model_id from target_model_names if not already set if model_id is None: - target_models = extra_body["target_model_names"] + target_models: Final = extra_body["target_model_names"] # Use first model as model_id for encoding if isinstance(target_models, str): model_id = target_models.split(",")[0].strip() @@ -148,7 +148,7 @@ class ContainerRequestUtils: # Only encode if we have routing metadata if should_encode and response_obj and hasattr(response_obj, "id"): - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id: Final = ResponsesAPIRequestUtils._build_container_id( custom_llm_provider=custom_llm_provider, model_id=model_id, container_id=response_obj.id, diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index bd03feb03c2..b894bd48c7e 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, Literal, cast +from typing import TYPE_CHECKING, Any, Final, Literal, cast from httpx import Response from pydantic import BaseModel @@ -131,14 +131,14 @@ else: LitellmLoggingObject = Any # Pre-resolved CallTypes enum values for fast membership checks -_A2A_CALL_TYPES = frozenset( +_A2A_CALL_TYPES: Final = frozenset( { CallTypes.asend_message.value, CallTypes.send_message.value, } ) -_VIDEO_CALL_TYPES = frozenset( +_VIDEO_CALL_TYPES: Final = frozenset( { CallTypes.create_video.value, CallTypes.acreate_video.value, @@ -149,36 +149,36 @@ _VIDEO_CALL_TYPES = frozenset( } ) -_SPEECH_CALL_TYPES = frozenset( +_SPEECH_CALL_TYPES: Final = frozenset( { CallTypes.speech.value, CallTypes.aspeech.value, } ) -_TRANSCRIPTION_CALL_TYPES = frozenset( +_TRANSCRIPTION_CALL_TYPES: Final = frozenset( { CallTypes.atranscription.value, CallTypes.transcription.value, } ) -_RERANK_CALL_TYPES = frozenset( +_RERANK_CALL_TYPES: Final = frozenset( { CallTypes.rerank.value, CallTypes.arerank.value, } ) -_SEARCH_CALL_TYPES = frozenset( +_SEARCH_CALL_TYPES: Final = frozenset( { CallTypes.search.value, CallTypes.asearch.value, } ) -_AREALTIME_CALL_TYPE = CallTypes.arealtime.value -_MCP_CALL_TYPE = CallTypes.call_mcp_tool.value +_AREALTIME_CALL_TYPE: Final = CallTypes.arealtime.value +_MCP_CALL_TYPE: Final = CallTypes.call_mcp_tool.value def _cost_per_token_custom_pricing_helper( @@ -201,24 +201,24 @@ def _cost_per_token_custom_pricing_helper( return None if custom_cost_per_token is not None: - input_cost_per_token = custom_cost_per_token["input_cost_per_token"] - output_cost_per_token = custom_cost_per_token["output_cost_per_token"] + input_cost_per_token: Final = custom_cost_per_token["input_cost_per_token"] + output_cost_per_token: Final = custom_cost_per_token["output_cost_per_token"] - cache_read_input_token_cost = custom_cost_per_token.get( + cache_read_input_token_cost: Final = custom_cost_per_token.get( "cache_read_input_token_cost", input_cost_per_token, ) - cache_creation_input_token_cost = custom_cost_per_token.get( + cache_creation_input_token_cost: Final = custom_cost_per_token.get( "cache_creation_input_token_cost", input_cost_per_token, ) - regular_prompt_tokens = max( + regular_prompt_tokens: Final = max( prompt_tokens - cached_tokens - cache_creation_tokens, 0, ) - input_cost = ( + input_cost: Final = ( regular_prompt_tokens * input_cost_per_token + cached_tokens * cache_read_input_token_cost + cache_creation_tokens * cache_creation_input_token_cost @@ -284,13 +284,13 @@ def _transcription_usage_has_token_details( if usage_block is None: return False - prompt_tokens_val = getattr(usage_block, "prompt_tokens", 0) or 0 - completion_tokens_val = getattr(usage_block, "completion_tokens", 0) or 0 - prompt_details = getattr(usage_block, "prompt_tokens_details", None) + prompt_tokens_val: Final = getattr(usage_block, "prompt_tokens", 0) or 0 + completion_tokens_val: Final = getattr(usage_block, "completion_tokens", 0) or 0 + prompt_details: Final = getattr(usage_block, "prompt_tokens_details", None) if prompt_details is not None: - audio_token_count = getattr(prompt_details, "audio_tokens", 0) or 0 - text_token_count = getattr(prompt_details, "text_tokens", 0) or 0 + audio_token_count: Final = getattr(prompt_details, "audio_tokens", 0) or 0 + text_token_count: Final = getattr(prompt_details, "text_tokens", 0) or 0 if audio_token_count > 0 or text_token_count > 0: return True @@ -375,7 +375,7 @@ def cost_per_token( _is_anthropic_style = False if usage_object is not None: - _pt_details = getattr(usage_object, "prompt_tokens_details", None) + _pt_details: Final = getattr(usage_object, "prompt_tokens_details", None) if _pt_details is not None: _cache_read_tokens = float(getattr(_pt_details, "cached_tokens", 0) or 0) # OpenAI-compatible providers report cache-write tokens under @@ -385,8 +385,8 @@ def cost_per_token( getattr(_pt_details, "cache_write_tokens", 0) or getattr(_pt_details, "cache_creation_tokens", 0) or 0 ) - _anthropic_read = getattr(usage_object, "cache_read_input_tokens", None) - _anthropic_create = getattr(usage_object, "cache_creation_input_tokens", None) + _anthropic_read: Final = getattr(usage_object, "cache_read_input_tokens", None) + _anthropic_create: Final = getattr(usage_object, "cache_creation_input_tokens", None) if _anthropic_read is not None or _anthropic_create is not None: _is_anthropic_style = True if _anthropic_read is not None: @@ -407,7 +407,7 @@ def cost_per_token( if _is_anthropic_style: _normalized_prompt_tokens += _cache_read_tokens + _cache_creation_tokens - response_cost = _cost_per_token_custom_pricing_helper( + response_cost: Final = _cost_per_token_custom_pricing_helper( prompt_tokens=_normalized_prompt_tokens, completion_tokens=completion_tokens, response_time_ms=response_time_ms, @@ -423,23 +423,23 @@ def cost_per_token( # given prompt_tokens_cost_usd_dollar: float = 0 completion_tokens_cost_usd_dollar: float = 0 - model_cost_ref = litellm.model_cost + model_cost_ref: Final = litellm.model_cost # Only callers that explicitly pass `custom_llm_provider` get the # dedup/prefix-join treatment. When provider is omitted, preserve legacy # behavior: `model_with_provider` stays equal to the raw `model` string # (provider is detected below for downstream use only). - caller_supplied_provider = custom_llm_provider is not None + caller_supplied_provider: Final = custom_llm_provider is not None # `model` is normally a string, but callers that mock the transport can pass # non-string objects. Only run the string-based dedup/prefix-join when it is # actually a string — e.g. a MagicMock's `.startswith()` is always truthy and # its slices return new mocks, which would spin the dedup loop forever. - model_is_str = isinstance(model, str) + model_is_str: Final = isinstance(model, str) # Router/proxy deployments may repeat the provider segment (e.g. model_name # "openai/openai/gpt-5.5"). Strip duplicated `{provider}/` chains before joining. if caller_supplied_provider and model_is_str: - _dup_prefix = f"{custom_llm_provider}/" + _dup_prefix: Final = f"{custom_llm_provider}/" while model.startswith(_dup_prefix): _remainder = model[len(_dup_prefix) :] if _remainder.startswith(_dup_prefix): @@ -449,13 +449,13 @@ def cost_per_token( model_with_provider = model if caller_supplied_provider: - _prov_prefix = f"{custom_llm_provider}/" + _prov_prefix: Final = f"{custom_llm_provider}/" if model_is_str and model.startswith(_prov_prefix): model_with_provider = model else: model_with_provider = f"{custom_llm_provider}/{model}" if region_name is not None: - model_with_provider_and_region = f"{custom_llm_provider}/{region_name}/{model}" + model_with_provider_and_region: Final = f"{custom_llm_provider}/{region_name}/{model}" if model_with_provider_and_region in model_cost_ref: # use region based pricing, if it's available model_with_provider = model_with_provider_and_region else: @@ -464,7 +464,7 @@ def cost_per_token( assert custom_llm_provider is not None # caller-supplied or get_llm_provider model_without_prefix = model - model_parts = model.split("/", 1) + model_parts: Final = model.split("/", 1) if len(model_parts) > 1: model_without_prefix = model_parts[1] else: @@ -487,7 +487,7 @@ def cost_per_token( # see this https://learn.microsoft.com/en-us/azure/ai-services/openai/concepts/models if call_type == "speech" or call_type == "aspeech": speech_model_info = litellm.get_model_info(model=model_without_prefix, custom_llm_provider=custom_llm_provider) - cost_metric = select_cost_metric_for_model(speech_model_info) + cost_metric: Final = select_cost_metric_for_model(speech_model_info) prompt_cost: float = 0.0 completion_cost: float = 0.0 if cost_metric == "cost_per_character": @@ -574,7 +574,7 @@ def cost_per_token( optional_params=(response._hidden_params if response and hasattr(response, "_hidden_params") else None), ) elif custom_llm_provider == "vertex_ai": - cost_router = google_cost_router( + cost_router: Final = google_cost_router( model=model_without_prefix, custom_llm_provider=custom_llm_provider, call_type=call_type, @@ -643,7 +643,7 @@ def cost_per_token( service_tier=service_tier, ) else: - model_info = _cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) + model_info: Final = _cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) if (model_info.get("input_cost_per_token") or 0.0) > 0 or (model_info.get("output_cost_per_token") or 0.0) > 0: return generic_cost_per_token( @@ -654,7 +654,7 @@ def cost_per_token( data_residency=data_residency, ) - input_cost_per_second = model_info.get("input_cost_per_second") + input_cost_per_second: Final = model_info.get("input_cost_per_second") if input_cost_per_second is not None and response_time_ms is not None: verbose_logger.debug( "For model=%s - input_cost_per_second: %s; response time: %s", @@ -665,7 +665,7 @@ def cost_per_token( ## COST PER SECOND ## prompt_tokens_cost_usd_dollar = input_cost_per_second * response_time_ms / 1000 - output_cost_per_second = model_info.get("output_cost_per_second") + output_cost_per_second: Final = model_info.get("output_cost_per_second") if output_cost_per_second is not None and response_time_ms is not None: verbose_logger.debug( "For model=%s - output_cost_per_second: %s; response time: %s", @@ -688,12 +688,12 @@ def cost_per_token( def get_replicate_completion_pricing(completion_response: dict, total_time=0.0): # see https://replicate.com/pricing # for all litellm currently supported LLMs, almost all requests go to a100_80gb - a100_80gb_price_per_second_public = ( + a100_80gb_price_per_second_public: Final = ( DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND # assume all calls sent to A100 80GB for now ) if total_time == 0.0: # total time is in ms - start_time = completion_response.get("created", time.time()) - end_time = getattr(completion_response, "ended", time.time()) + start_time: Final = completion_response.get("created", time.time()) + end_time: Final = getattr(completion_response, "ended", time.time()) total_time = end_time - start_time return a100_80gb_price_per_second_public * total_time / 1000 @@ -747,11 +747,11 @@ def _select_model_name_for_cost_calc( completion_response_model = getattr(completion_response, "model", None) elif isinstance(completion_response, dict): completion_response_model = completion_response.get("model", None) - hidden_params: dict | None = getattr(completion_response, "_hidden_params", None) + hidden_params: Final[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: - entry = litellm.model_cost[router_model_id] + entry: Final = litellm.model_cost[router_model_id] if ( entry.get("input_cost_per_token") is not None or entry.get("input_cost_per_second") is not None @@ -796,7 +796,7 @@ def _model_contains_known_llm_provider(model: str) -> bool: """ Check if the model contains a known llm provider """ - _provider_prefix = model.split("/")[0] + _provider_prefix: Final = model.split("/")[0] return _provider_prefix in LlmProvidersSet @@ -818,7 +818,7 @@ def _get_response_model(completion_response: Any) -> str | None: return None -_GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: dict = { +_GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: Final[dict] = { # ON_DEMAND_PRIORITY maps to "priority" — selects input_cost_per_token_priority, etc. "ON_DEMAND_PRIORITY": "priority", # FLEX / BATCH maps to "flex" — selects input_cost_per_token_flex, etc. @@ -844,7 +844,7 @@ def _map_traffic_type_to_service_tier(traffic_type: str | None) -> str | None: """ if traffic_type is None: return None - service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper()) + service_tier: Final = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper()) return service_tier @@ -865,7 +865,7 @@ def _normalize_service_tier(service_tier: object) -> str | None: def _get_usage_object( completion_response: Any, ) -> Usage | None: - usage_obj = cast( + usage_obj: Final = cast( Usage | ResponseAPIUsage | dict | BaseModel, ( completion_response.get("usage") @@ -950,14 +950,14 @@ def _apply_cost_discount( Returns: Tuple of (final_cost, discount_percent, discount_amount) """ - original_cost = base_cost + original_cost: Final = base_cost discount_percent = 0.0 discount_amount = 0.0 if custom_llm_provider and custom_llm_provider in litellm.cost_discount_config: discount_percent = litellm.cost_discount_config[custom_llm_provider] discount_amount = original_cost * discount_percent - final_cost = original_cost - discount_amount + final_cost: Final = original_cost - discount_amount if verbose_logger.isEnabledFor(logging.DEBUG): verbose_logger.debug( @@ -984,7 +984,7 @@ def _apply_cost_margin( Returns: Tuple of (final_cost, margin_percent, margin_fixed_amount, margin_total_amount) """ - original_cost = base_cost + original_cost: Final = base_cost margin_percent = 0.0 margin_fixed_amount = 0.0 margin_total_amount = 0.0 @@ -1022,7 +1022,7 @@ def _apply_cost_margin( margin_fixed_amount = float(margin_config["fixed_amount"]) margin_total_amount += margin_fixed_amount - final_cost = original_cost + margin_total_amount + final_cost: Final = original_cost + margin_total_amount if verbose_logger.isEnabledFor(logging.DEBUG): verbose_logger.debug( @@ -1180,7 +1180,7 @@ def completion_cost( 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: Usage | None = _get_usage_object(completion_response=completion_response) + cost_per_token_usage_object: Final[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 @@ -1207,7 +1207,7 @@ def completion_cost( service_tier = _normalize_service_tier(service_tier) - selected_model = _select_model_name_for_cost_calc( + selected_model: Final = _select_model_name_for_cost_calc( model=model, completion_response=completion_response, custom_llm_provider=custom_llm_provider, @@ -1216,7 +1216,7 @@ def completion_cost( router_model_id=router_model_id, ) - potential_model_names = [ + potential_model_names: Final = [ selected_model, _get_response_model(completion_response), ] @@ -1691,9 +1691,9 @@ def get_response_cost_from_hidden_params( else: _hidden_params_dict = hidden_params - additional_headers = _hidden_params_dict.get("additional_headers", {}) + additional_headers: Final = _hidden_params_dict.get("additional_headers", {}) if additional_headers and "llm_provider-x-litellm-response-cost" in additional_headers: - response_cost = additional_headers["llm_provider-x-litellm-response-cost"] + response_cost: Final = additional_headers["llm_provider-x-litellm-response-cost"] if response_cost is None: return None return float(additional_headers["llm_provider-x-litellm-response-cost"]) @@ -1761,7 +1761,7 @@ def response_cost_calculator( if isinstance(response_object, BaseModel): if hasattr(response_object, "_hidden_params"): response_object._hidden_params["optional_params"] = optional_params - provider_response_cost = get_response_cost_from_hidden_params(response_object._hidden_params) + provider_response_cost: Final = get_response_cost_from_hidden_params(response_object._hidden_params) if provider_response_cost is not None: return provider_response_cost @@ -1818,7 +1818,7 @@ def ocr_cost( except Exception: model_info = None - credits = getattr(response.usage_info, "credits", None) + credits: Final = getattr(response.usage_info, "credits", None) cost_per_credit = None if model_info is not None: cost_per_credit = model_info.get("ocr_cost_per_credit") @@ -1829,7 +1829,7 @@ def ocr_cost( if model_info is not None: ocr_cost_per_page = model_info.get("ocr_cost_per_page") - pages_processed = response.usage_info.pages_processed + pages_processed: Final = response.usage_info.pages_processed if pages_processed is None: if cost_per_credit is not None or ocr_cost_per_page is None: # Surface missing usage data instead of silently under-reporting @@ -1862,7 +1862,7 @@ def ocr_cost( ) return 0.0, 0.0 - total_ocr_processing_cost: float = ocr_cost_per_page * pages_processed + total_ocr_processing_cost: Final[float] = ocr_cost_per_page * pages_processed return total_ocr_processing_cost, 0.0 @@ -1884,7 +1884,7 @@ def vector_store_search_cost( model=model, ) - config = ProviderConfigManager.get_provider_vector_stores_config( + config: Final = ProviderConfigManager.get_provider_vector_stores_config( provider=LlmProviders(custom_llm_provider), api_type=api_type, ) @@ -1910,7 +1910,7 @@ def rerank_cost( _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) try: - config = ProviderConfigManager.get_provider_rerank_config( + config: Final = ProviderConfigManager.get_provider_rerank_config( model=model, api_base=None, present_version_params=[], @@ -1973,19 +1973,19 @@ def default_image_cost_calculator( 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}" - model_name_with_quality = f"{quality}/{base_model_name}" if quality else base_model_name + model_name_with_quality: Final = f"{quality}/{base_model_name}" if quality else base_model_name # gpt-image-1 models use low, medium, high quality. If user did not specify quality, use medium fot gpt-image-1 model family - model_name_with_v2_quality = f"{ImageGenerationRequestQuality.HIGH.value}/{base_model_name}" + model_name_with_v2_quality: Final = f"{ImageGenerationRequestQuality.HIGH.value}/{base_model_name}" verbose_logger.debug("Looking up cost for models: %s, %s", model_name_with_quality, base_model_name) - model_without_provider = f"{size_str}/{model.split('/')[-1]}" + model_without_provider: Final = f"{size_str}/{model.split('/')[-1]}" 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: dict | None = None - models_to_check: list[str | None] = [ + models_to_check: Final[list[str | None]] = [ model_name_with_quality, base_model_name, model_name_with_v2_quality, @@ -2050,10 +2050,10 @@ def default_video_cost_calculator( verbose_logger.debug("Looking up cost for video model: %s", base_model_name) - model_without_provider = model.split("/")[-1] + model_without_provider: Final = model.split("/")[-1] # Try model with provider first, fall back to base model name - models_to_check: list[str | None] = [ + models_to_check: Final[list[str | None]] = [ base_model_name, model, model_without_provider, @@ -2066,7 +2066,7 @@ def default_video_cost_calculator( # If still not found, try with custom_llm_provider prefix if cost_info is None and custom_llm_provider: - prefixed_model = f"{custom_llm_provider}/{model}" + prefixed_model: Final = f"{custom_llm_provider}/{model}" if prefixed_model in litellm.model_cost: cost_info = litellm.model_cost[prefixed_model] @@ -2074,11 +2074,11 @@ def default_video_cost_calculator( raise Exception(f"Model not found in cost map for model={model}") # Check for video-specific cost per second first - video_cost_per_second = cost_info.get("output_cost_per_video_per_second") + video_cost_per_second: Final = cost_info.get("output_cost_per_video_per_second") if video_cost_per_second is not None: return video_cost_per_second * duration_seconds - output_cost_per_second = _video_output_cost_per_second(cost_info, video_resolution) + output_cost_per_second: Final = _video_output_cost_per_second(cost_info, video_resolution) if output_cost_per_second is not None: return output_cost_per_second * duration_seconds @@ -2133,7 +2133,7 @@ def batch_cost_calculator( # but carries no pricing fields. Fall back to the global pricing table so # that standard model pricing is used instead of silently returning $0. try: - global_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + global_info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) if global_info: model_info = global_info except Exception: @@ -2142,31 +2142,31 @@ def batch_cost_calculator( if not model_info: return 0.0, 0.0 - input_cost_per_token_batches = model_info.get("input_cost_per_token_batches") - input_cost_per_token = model_info.get("input_cost_per_token") - output_cost_per_token_batches = model_info.get("output_cost_per_token_batches") - output_cost_per_token = model_info.get("output_cost_per_token") + input_cost_per_token_batches: Final = model_info.get("input_cost_per_token_batches") + input_cost_per_token: Final = model_info.get("input_cost_per_token") + output_cost_per_token_batches: Final = model_info.get("output_cost_per_token_batches") + output_cost_per_token: Final = model_info.get("output_cost_per_token") total_prompt_cost = 0.0 total_completion_cost = 0.0 if input_cost_per_token_batches: total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches elif input_cost_per_token: - details = _parse_prompt_tokens_details(usage) - cache_read_tokens = details["cache_hit_tokens"] - cache_creation_tokens = details["cache_creation_tokens"] + details: Final = _parse_prompt_tokens_details(usage) + cache_read_tokens: Final = details["cache_hit_tokens"] + cache_creation_tokens: Final = details["cache_creation_tokens"] # Subtract cached tokens from prompt_tokens before calculating cost # Fixes issue where cached tokens are being charged again - base_input_tokens = get_billable_input_tokens(usage) - cache_creation_tokens + base_input_tokens: Final = get_billable_input_tokens(usage) - cache_creation_tokens total_prompt_cost = ( base_input_tokens * (input_cost_per_token) / 2 ) # batch cost is usually half of the regular token cost # Add cache read cost if applicable - cache_read_cost_key = _get_service_tier_cost_key("cache_read_input_token_cost", None) + cache_read_cost_key: Final = _get_service_tier_cost_key("cache_read_input_token_cost", None) total_prompt_cost += calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens) / 2 - cache_creation_cost = model_info.get("cache_creation_input_token_cost") or input_cost_per_token + cache_creation_cost: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token total_prompt_cost += cache_creation_tokens * cache_creation_cost / 2 if output_cost_per_token_batches: total_completion_cost = usage.completion_tokens * output_cost_per_token_batches @@ -2175,7 +2175,7 @@ def batch_cost_calculator( usage.completion_tokens * (output_cost_per_token) / 2 ) # batch cost is usually half of the regular token cost - uplift = _get_regional_uplift_multiplier(model_info, data_residency) + uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency) if uplift != 1.0: total_prompt_cost *= uplift total_completion_cost *= uplift @@ -2184,7 +2184,7 @@ def batch_cost_calculator( def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]: - field_names = list(type(prompt_tokens_details).model_fields) + field_names: Final = list(type(prompt_tokens_details).model_fields) if getattr(prompt_tokens_details, "cache_write_tokens", None) is None: return field_names return [attr for attr in field_names if attr != "cache_creation_tokens"] @@ -2202,7 +2202,7 @@ class BaseTokenUsageProcessor: Usage, ) - combined = Usage() + combined: Final = Usage() # Sum basic token counts for usage in usage_objects: @@ -2268,11 +2268,11 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): """ Collect usage from realtime stream results """ - response_done_events: list[OpenAIRealtimeStreamResponseBaseObject] = cast( + response_done_events: Final[list[OpenAIRealtimeStreamResponseBaseObject]] = cast( list[OpenAIRealtimeStreamResponseBaseObject], [result for result in results if result["type"] == "response.done"], ) - usage_objects: list[Usage] = [] + usage_objects: Final[list[Usage]] = [] for result in response_done_events: usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result["response"].get("usage", {}) @@ -2288,7 +2288,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): Collect and combine usage from realtime stream results """ collected_usage_objects = RealtimeAPITokenUsageProcessor.collect_usage_from_realtime_stream_results(results) - combined_usage_object = RealtimeAPITokenUsageProcessor.combine_usage_objects(collected_usage_objects) + combined_usage_object: Final = RealtimeAPITokenUsageProcessor.combine_usage_objects(collected_usage_objects) return combined_usage_object @staticmethod @@ -2301,7 +2301,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): ) -_TRANSCRIPTION_COMPLETED_EVENT_TYPE = "conversation.item.input_audio_transcription.completed" +_TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed" def handle_realtime_stream_cost_calculation( @@ -2321,7 +2321,7 @@ def handle_realtime_stream_cost_calculation( results: A list of OpenAIRealtimeStreamBaseObject objects """ received_model = None - potential_model_names = [] + potential_model_names: Final = [] for result in results: if result["type"] == "session.created": received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None) @@ -2346,7 +2346,7 @@ def handle_realtime_stream_cost_calculation( input_cost_per_token += _input_cost_per_token output_cost_per_token += _output_cost_per_token break # exit if we find a valid model - transcription_cost = ( + transcription_cost: Final = ( handle_realtime_transcription_cost_calculation( results=results, custom_llm_provider=custom_llm_provider, @@ -2355,7 +2355,7 @@ def handle_realtime_stream_cost_calculation( if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 ) - total_cost = input_cost_per_token + output_cost_per_token + transcription_cost + total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost _store_cost_breakdown_in_logging_obj( litellm_logging_obj=litellm_logging_obj, @@ -2384,13 +2384,13 @@ def handle_realtime_transcription_cost_calculation( - {"type": "duration", "seconds": } → priced via input_cost_per_second - {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost """ - completed_events = [ + completed_events: Final = [ cast(dict, result) for result in results if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE ] if not completed_events: return 0.0 - model_name = _get_transcription_model_name_from_results(results) or litellm_model_name + model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name try: model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider) except Exception: @@ -2427,20 +2427,20 @@ def _get_transcription_model_name_from_results( def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float: if model_info is None: return 0.0 - usage_type = usage.get("type") + usage_type: Final = usage.get("type") if usage_type == "duration": - seconds = usage.get("seconds") or 0.0 - per_second = model_info.get("input_cost_per_second") or 0.0 + seconds: Final = usage.get("seconds") or 0.0 + per_second: Final = model_info.get("input_cost_per_second") or 0.0 return float(seconds) * float(per_second) if usage_type == "tokens": - input_token_details = usage.get("input_token_details") or {} - audio_tokens = input_token_details.get("audio_tokens") or 0 - text_tokens = input_token_details.get("text_tokens") or 0 - output_tokens = usage.get("output_tokens") or 0 - audio_cost = float(audio_tokens) * float( + input_token_details: Final = usage.get("input_token_details") or {} + audio_tokens: Final = input_token_details.get("audio_tokens") or 0 + text_tokens: Final = input_token_details.get("text_tokens") or 0 + output_tokens: Final = usage.get("output_tokens") or 0 + audio_cost: Final = float(audio_tokens) * float( model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0 ) - text_cost = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0) - output_cost = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0) + text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0) + output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0) return audio_cost + text_cost + output_cost return 0.0 diff --git a/litellm/endpoints/speech/speech_to_completion_bridge/handler.py b/litellm/endpoints/speech/speech_to_completion_bridge/handler.py index babb1811d1b..9e949db625a 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 +from typing import TYPE_CHECKING, Final from typing_extensions import TypedDict @@ -32,23 +32,23 @@ class SpeechToCompletionBridgeHandler: def validate_input_kwargs(self, kwargs: dict) -> SpeechToCompletionBridgeHandlerInputKwargs: from litellm import LiteLLMLoggingObj - model = kwargs.get("model") + model: Final = kwargs.get("model") if model is None or not isinstance(model, str): raise ValueError("model is required") - custom_llm_provider = kwargs.get("custom_llm_provider") + custom_llm_provider: Final = kwargs.get("custom_llm_provider") if custom_llm_provider is None or not isinstance(custom_llm_provider, str): raise ValueError("custom_llm_provider is required") - input = kwargs.get("input") + input: Final = kwargs.get("input") if input is None or not isinstance(input, str): raise ValueError("input is required") - optional_params = kwargs.get("optional_params") + optional_params: Final = kwargs.get("optional_params") if optional_params is None or not isinstance(optional_params, dict): raise ValueError("optional_params is required") - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") if litellm_params is None or not isinstance(litellm_params, dict): raise ValueError("litellm_params is required") @@ -60,7 +60,7 @@ class SpeechToCompletionBridgeHandler: if headers is None or not isinstance(headers, dict): raise ValueError("headers is required") - logging_obj = kwargs.get("logging_obj") + logging_obj: Final = kwargs.get("logging_obj") if logging_obj is None or not isinstance(logging_obj, LiteLLMLoggingObj): raise ValueError("logging_obj is required") @@ -86,11 +86,11 @@ class SpeechToCompletionBridgeHandler: logging_obj: "LiteLLMLoggingObj", custom_llm_provider: str, ) -> "HttpxBinaryResponseContent": - received_args = locals() + received_args: Final = locals() from litellm import completion from litellm.types.utils import ModelResponse - validated_kwargs = self.validate_input_kwargs(received_args) + validated_kwargs: Final = self.validate_input_kwargs(received_args) model = validated_kwargs["model"] input = validated_kwargs["input"] optional_params = validated_kwargs["optional_params"] @@ -100,7 +100,7 @@ class SpeechToCompletionBridgeHandler: custom_llm_provider = validated_kwargs["custom_llm_provider"] voice = validated_kwargs["voice"] - request_data = self.transformation_handler.transform_request( + request_data: Final = self.transformation_handler.transform_request( model=model, input=input, optional_params=optional_params, @@ -111,7 +111,7 @@ class SpeechToCompletionBridgeHandler: voice=voice, ) - result = completion( + result: Final = completion( **request_data, ) diff --git a/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py b/litellm/endpoints/speech/speech_to_completion_bridge/transformation.py index 2f2861dfc26..a9429b673e4 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, cast +from typing import TYPE_CHECKING, Final, cast from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS @@ -20,7 +20,7 @@ class SpeechToCompletionBridgeTransformationHandler: litellm_logging_obj: "LiteLLMLoggingObj", custom_llm_provider: str, ) -> dict: - passed_optional_params = {} + passed_optional_params: Final = {} for op in optional_params: if op in OPENAI_CHAT_COMPLETION_PARAMS: passed_optional_params[op] = optional_params[op] @@ -66,13 +66,13 @@ class SpeechToCompletionBridgeTransformationHandler: import struct # WAV header parameters - byte_rate = sample_rate * channels * 2 # 2 bytes per sample (16-bit) - block_align = channels * 2 - data_size = len(pcm_data) - file_size = 36 + data_size + byte_rate: Final = sample_rate * channels * 2 # 2 bytes per sample (16-bit) + block_align: Final = channels * 2 + data_size: Final = len(pcm_data) + file_size: Final = 36 + data_size # Create WAV header - wav_header = struct.pack( + wav_header: Final = struct.pack( "<4sI4s4sIHHIIHH4sI", b"RIFF", # Chunk ID file_size, # Chunk Size @@ -103,17 +103,17 @@ class SpeechToCompletionBridgeTransformationHandler: from litellm.types.llms.openai import HttpxBinaryResponseContent from litellm.types.utils import Choices - audio_part = cast(Choices, model_response.choices[0]).message.audio + audio_part: Final = cast(Choices, model_response.choices[0]).message.audio if audio_part is None: raise ValueError("No audio part found in the response") - audio_content = audio_part.data + audio_content: Final = audio_part.data # Decode base64 to get binary content binary_data = base64.b64decode(audio_content) # Check if this is a Gemini TTS model that returns raw PCM16 data - model = getattr(model_response, "model", "") - headers = {} + model: Final = getattr(model_response, "model", "") + headers: Final = {} if self._is_gemini_tts_model(model): # Convert PCM16 to WAV format for proper audio file playback binary_data = self._convert_pcm16_to_wav(binary_data) @@ -122,5 +122,5 @@ class SpeechToCompletionBridgeTransformationHandler: headers["Content-Type"] = "audio/mpeg" # Create an httpx.Response object - response = httpx.Response(status_code=200, content=binary_data, headers=headers) + response: Final = httpx.Response(status_code=200, content=binary_data, headers=headers) return HttpxBinaryResponseContent(response) diff --git a/litellm/evals/main.py b/litellm/evals/main.py index 0bcccb73cb1..bf6337bd234 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 +from typing import Any, Final import httpx @@ -36,7 +36,7 @@ from litellm.utils import ProviderConfigManager, client # Initialize HTTP handler base_llm_http_handler = BaseLLMHTTPHandler() -DEFAULT_OPENAI_API_BASE = "https://api.openai.com" +DEFAULT_OPENAI_API_BASE: Final = "https://api.openai.com" @client @@ -70,12 +70,12 @@ async def acreate_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acreate_eval"] = True - func = partial( + func: Final = partial( create_eval, data_source_config=data_source_config, testing_criteria=testing_criteria, @@ -89,9 +89,9 @@ async def acreate_eval( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -139,14 +139,14 @@ def create_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("acreate_eval", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("acreate_eval", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -161,7 +161,7 @@ def create_eval( raise ValueError(f"CREATE eval is not supported for {custom_llm_provider}") # Build create request - create_request: CreateEvalRequest = { + create_request: Final[CreateEvalRequest] = { "data_source_config": data_source_config, # type: ignore "testing_criteria": testing_criteria, # type: ignore } @@ -177,15 +177,15 @@ def create_eval( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - request_body = evals_api_provider_config.transform_create_eval_request( + request_body: Final = evals_api_provider_config.transform_create_eval_request( create_request=create_request, litellm_params=litellm_params, headers=headers, ) # Get API base and URL - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE - url = evals_api_provider_config.get_complete_url(api_base=api_base, endpoint="evals") + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + url: Final = evals_api_provider_config.get_complete_url(api_base=api_base, endpoint="evals") # Pre-call logging litellm_logging_obj.update_from_kwargs( @@ -199,7 +199,7 @@ def create_eval( ) # Make HTTP request - response = base_llm_http_handler.create_eval_handler( # type: ignore + response: Final = base_llm_http_handler.create_eval_handler( # type: ignore url=url, request_body=request_body, evals_api_provider_config=evals_api_provider_config, @@ -255,12 +255,12 @@ async def alist_evals( Returns: ListEvalsResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["alist_evals"] = True - func = partial( + func: Final = partial( list_evals, limit=limit, after=after, @@ -274,9 +274,9 @@ async def alist_evals( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -324,14 +324,14 @@ def list_evals( Returns: ListEvalsResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("alist_evals", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("alist_evals", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -346,7 +346,7 @@ def list_evals( raise ValueError(f"LIST evals is not supported for {custom_llm_provider}") # Build list parameters - list_params: ListEvalsParams = {} + list_params: Final[ListEvalsParams] = {} if limit is not None: list_params["limit"] = limit if after is not None: @@ -385,7 +385,7 @@ def list_evals( ) # Make HTTP request - response = base_llm_http_handler.list_evals_handler( # type: ignore + response: Final = base_llm_http_handler.list_evals_handler( # type: ignore url=url, query_params=query_params, evals_api_provider_config=evals_api_provider_config, @@ -433,12 +433,12 @@ async def aget_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["aget_eval"] = True - func = partial( + func: Final = partial( get_eval, eval_id=eval_id, extra_headers=extra_headers, @@ -448,9 +448,9 @@ async def aget_eval( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -490,14 +490,14 @@ def get_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("aget_eval", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("aget_eval", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -516,7 +516,7 @@ def get_eval( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE url, headers = evals_api_provider_config.transform_get_eval_request( eval_id=eval_id, api_base=api_base, @@ -536,7 +536,7 @@ def get_eval( ) # Make HTTP request - response = base_llm_http_handler.get_eval_handler( # type: ignore + response: Final = base_llm_http_handler.get_eval_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, @@ -589,12 +589,12 @@ async def aupdate_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["aupdate_eval"] = True - func = partial( + func: Final = partial( update_eval, eval_id=eval_id, name=name, @@ -607,9 +607,9 @@ async def aupdate_eval( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -655,14 +655,14 @@ def update_eval( Returns: Eval object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("aupdate_eval", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("aupdate_eval", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -677,14 +677,14 @@ def update_eval( raise ValueError(f"UPDATE eval is not supported for {custom_llm_provider}") # Build update request - update_request: UpdateEvalRequest = {} + update_request: Final[UpdateEvalRequest] = {} if name is not None: update_request["name"] = name # Filter metadata to exclude internal LiteLLM fields if metadata is not None: # List of internal LiteLLM metadata keys that should NOT be sent to OpenAI - internal_keys = { + internal_keys: Final = { "headers", "requester_metadata", "user_api_key_hash", @@ -717,7 +717,7 @@ def update_eval( "user_agent", } # Only include user-provided metadata keys - filtered_metadata = {k: v for k, v in metadata.items() if k not in internal_keys} + filtered_metadata: Final = {k: v for k, v in metadata.items() if k not in internal_keys} if filtered_metadata: # Only add if there's user metadata update_request["metadata"] = filtered_metadata @@ -730,7 +730,7 @@ def update_eval( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE ( url, headers, @@ -755,7 +755,7 @@ def update_eval( ) # Make HTTP request - response = base_llm_http_handler.update_eval_handler( # type: ignore + response: Final = base_llm_http_handler.update_eval_handler( # type: ignore url=url, request_body=request_body, evals_api_provider_config=evals_api_provider_config, @@ -803,12 +803,12 @@ async def adelete_eval( Returns: DeleteEvalResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["adelete_eval"] = True - func = partial( + func: Final = partial( delete_eval, eval_id=eval_id, extra_headers=extra_headers, @@ -818,9 +818,9 @@ async def adelete_eval( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -860,14 +860,14 @@ def delete_eval( Returns: DeleteEvalResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("adelete_eval", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("adelete_eval", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -886,7 +886,7 @@ def delete_eval( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE url, headers = evals_api_provider_config.transform_delete_eval_request( eval_id=eval_id, api_base=api_base, @@ -906,7 +906,7 @@ def delete_eval( ) # Make HTTP request - response = base_llm_http_handler.delete_eval_handler( # type: ignore + response: Final = base_llm_http_handler.delete_eval_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, @@ -953,12 +953,12 @@ async def acancel_eval( Returns: CancelEvalResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acancel_eval"] = True - func = partial( + func: Final = partial( cancel_eval, eval_id=eval_id, extra_headers=extra_headers, @@ -968,9 +968,9 @@ async def acancel_eval( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1010,14 +1010,14 @@ def cancel_eval( Returns: CancelEvalResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("acancel_eval", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("acancel_eval", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1036,7 +1036,7 @@ def cancel_eval( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE ( url, headers, @@ -1060,7 +1060,7 @@ def cancel_eval( ) # Make HTTP request - response = base_llm_http_handler.cancel_eval_handler( # type: ignore + response: Final = base_llm_http_handler.cancel_eval_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, @@ -1120,12 +1120,12 @@ async def acreate_run( Returns: Run object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acreate_run"] = True - func = partial( + func: Final = partial( create_run, eval_id=eval_id, data_source=data_source, @@ -1139,9 +1139,9 @@ async def acreate_run( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1189,14 +1189,14 @@ def create_run( Returns: Run object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("acreate_run", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("acreate_run", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1211,7 +1211,7 @@ def create_run( raise ValueError(f"CREATE run is not supported for {custom_llm_provider}") # Build create request - create_request: CreateRunRequest = { + create_request: Final[CreateRunRequest] = { "data_source": data_source, # type: ignore } if name is not None: @@ -1228,7 +1228,7 @@ def create_run( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE url, request_body = evals_api_provider_config.transform_create_run_request( eval_id=eval_id, create_request=create_request, @@ -1248,7 +1248,7 @@ def create_run( ) # Make HTTP request (default 600s timeout for long-running operations) - response = base_llm_http_handler.create_run_handler( # type: ignore + response: Final = base_llm_http_handler.create_run_handler( # type: ignore url=url, request_body=request_body, evals_api_provider_config=evals_api_provider_config, @@ -1304,12 +1304,12 @@ async def alist_runs( Returns: ListRunsResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["alist_runs"] = True - func = partial( + func: Final = partial( list_runs, eval_id=eval_id, limit=limit, @@ -1323,9 +1323,9 @@ async def alist_runs( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1373,14 +1373,14 @@ def list_runs( Returns: ListRunsResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("alist_runs", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("alist_runs", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1395,7 +1395,7 @@ def list_runs( raise ValueError(f"LIST runs is not supported for {custom_llm_provider}") # Build list parameters - list_params: ListRunsParams = {} + list_params: Final[ListRunsParams] = {} if limit is not None: list_params["limit"] = limit if after is not None: @@ -1433,7 +1433,7 @@ def list_runs( ) # Make HTTP request - response = base_llm_http_handler.list_runs_handler( # type: ignore + response: Final = base_llm_http_handler.list_runs_handler( # type: ignore url=url, query_params=query_params, evals_api_provider_config=evals_api_provider_config, @@ -1483,12 +1483,12 @@ async def aget_run( Returns: Run object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["aget_run"] = True - func = partial( + func: Final = partial( get_run, eval_id=eval_id, run_id=run_id, @@ -1499,9 +1499,9 @@ async def aget_run( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1543,14 +1543,14 @@ def get_run( Returns: Run object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("aget_run", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("aget_run", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1569,7 +1569,7 @@ def get_run( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE url, headers = evals_api_provider_config.transform_get_run_request( eval_id=eval_id, run_id=run_id, @@ -1590,7 +1590,7 @@ def get_run( ) # Make HTTP request - response = base_llm_http_handler.get_run_handler( # type: ignore + response: Final = base_llm_http_handler.get_run_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, @@ -1639,12 +1639,12 @@ async def acancel_run( Returns: CancelRunResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acancel_run"] = True - func = partial( + func: Final = partial( cancel_run, eval_id=eval_id, run_id=run_id, @@ -1655,9 +1655,9 @@ async def acancel_run( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1699,14 +1699,14 @@ def cancel_run( Returns: CancelRunResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("acancel_run", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("acancel_run", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1725,7 +1725,7 @@ def cancel_run( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE ( url, headers, @@ -1750,7 +1750,7 @@ def cancel_run( ) # Make HTTP request - response = base_llm_http_handler.cancel_run_handler( # type: ignore + response: Final = base_llm_http_handler.cancel_run_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, @@ -1804,12 +1804,12 @@ async def adelete_run( Returns: RunDeleteResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["adelete_run"] = True - func = partial( + func: Final = partial( delete_run, eval_id=eval_id, run_id=run_id, @@ -1820,9 +1820,9 @@ async def adelete_run( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1864,14 +1864,14 @@ def delete_run( Returns: RunDeleteResponse object """ - local_vars = locals() + local_vars: Final = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("adelete_run", False) is True + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + _is_async: Final = kwargs.pop("adelete_run", False) is True # Get LiteLLM parameters - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) # Determine provider if custom_llm_provider is None: @@ -1890,7 +1890,7 @@ def delete_run( headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params) # Transform request - api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE + api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE ( url, headers, @@ -1915,7 +1915,7 @@ def delete_run( ) # Make HTTP request - response = base_llm_http_handler.delete_run_handler( # type: ignore + response: Final = base_llm_http_handler.delete_run_handler( # type: ignore url=url, evals_api_provider_config=evals_api_provider_config, custom_llm_provider=custom_llm_provider, diff --git a/litellm/exceptions.py b/litellm/exceptions.py index c4a64e0ad9b..dfb0fc32f5f 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 +from typing import Any, Final import httpx import openai @@ -81,8 +81,8 @@ class RateLimitType(str, enum.Enum): """Per-session max-iterations cap reached (agent-style flows).""" -_RATE_LIMIT_CATEGORY_VALUES = frozenset(c.value for c in RateLimitErrorCategory) -_RATE_LIMIT_TYPE_VALUES = frozenset(t.value for t in RateLimitType) +_RATE_LIMIT_CATEGORY_VALUES: Final = frozenset(c.value for c in RateLimitErrorCategory) +_RATE_LIMIT_TYPE_VALUES: Final = frozenset(t.value for t in RateLimitType) def validate_rate_limit_category(value: Any) -> str | None: @@ -339,7 +339,7 @@ class Timeout(openai.APITimeoutError): # type: ignore headers: dict | None = None, exception_status_code: int | None = None, ): - request = httpx.Request( + request: Final = httpx.Request( method="POST", url="https://api.openai.com/v1", ) @@ -464,7 +464,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore # headers stay reachable on `e.response.headers` for callers that # 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 + _response_headers: Final = getattr(response, "headers", None) if response is not None 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. @@ -558,8 +558,8 @@ class RejectedRequestError(BadRequestError): # type: ignore self.llm_provider = llm_provider self.litellm_debug_info = litellm_debug_info self.request_data = request_data - request = httpx.Request(method="POST", url="https://api.openai.com/v1") - response = httpx.Response(status_code=400, request=request) + request: Final = httpx.Request(method="POST", url="https://api.openai.com/v1") + response: Final = httpx.Response(status_code=400, request=request) super().__init__( message=self.message, model=self.model, # type: ignore @@ -648,7 +648,7 @@ class ServiceUnavailableError(openai.APIStatusError): # type: ignore self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries - _response_headers = getattr(response, "headers", None) if response is not None else None + _response_headers: Final = getattr(response, "headers", None) if response is not None else None self.response = httpx.Response( status_code=self.status_code, headers=_response_headers, @@ -696,7 +696,7 @@ class BadGatewayError(openai.APIStatusError): # type: ignore self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries - _response_headers = getattr(response, "headers", None) if response is not None else None + _response_headers: Final = getattr(response, "headers", None) if response is not None else None self.response = httpx.Response( status_code=self.status_code, headers=_response_headers, @@ -744,7 +744,7 @@ class InternalServerError(openai.InternalServerError): # type: ignore self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries - _response_headers = getattr(response, "headers", None) if response is not None else None + _response_headers: Final = getattr(response, "headers", None) if response is not None else None self.response = httpx.Response( status_code=self.status_code, headers=_response_headers, @@ -868,8 +868,8 @@ class APIResponseValidationError(openai.APIResponseValidationError): # type: ig 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") - response = httpx.Response(status_code=500, request=request) + request: Final = httpx.Request(method="POST", url="https://api.openai.com/v1") + response: Final = httpx.Response(status_code=500, request=request) self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries @@ -933,7 +933,7 @@ class UnsupportedParamsError(BadRequestError): self.num_retries = num_retries -LITELLM_EXCEPTION_TYPES = [ +LITELLM_EXCEPTION_TYPES: Final = [ AuthenticationError, NotFoundError, BadRequestError, @@ -1046,7 +1046,7 @@ class GuardrailRaisedException(Exception): should_wrap_with_default_message: bool = True, status_code: int = 400, ): - default_message = f"Guardrail raised an exception, Guardrail: {guardrail_name}, Message: {message}" + default_message: Final = f"Guardrail raised an exception, Guardrail: {guardrail_name}, Message: {message}" self.guardrail_name = guardrail_name self.status_code = status_code self.message = default_message if should_wrap_with_default_message else message @@ -1084,7 +1084,7 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore generated_content: str = "", is_pre_first_chunk: bool = False, ): - original_status = getattr(original_exception, "status_code", None) + original_status: Final = getattr(original_exception, "status_code", None) self.status_code = int(original_status) if original_status is not None else 503 self.message = f"litellm.MidStreamFallbackError: {message}" self.model = model @@ -1109,11 +1109,11 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore self.response = response # Save the original attributes before they are overridden by ServiceUnavailableError - _saved_response = self.response - _saved_request = getattr(self.response, "request", None) or httpx.Request( + _saved_response: Final = self.response + _saved_request: Final = getattr(self.response, "request", None) or httpx.Request( method="POST", url=f"https://{llm_provider}.com/v1/" ) - _saved_message = self.message + _saved_message: Final = self.message # Call the parent constructor (which hardcodes status_code=503 and modifies the response object) super().__init__( diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index c9d73363242..64f4a773901 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -6,10 +6,7 @@ import asyncio import base64 import os from collections.abc import Awaitable, Callable, Generator -from typing import ( - Any, - TypeVar, -) +from typing import Any, Final, TypeVar import httpx from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters @@ -61,7 +58,7 @@ def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]: def _first_non_cancelled_cause(exc: BaseException) -> BaseException | None: - queue: list[BaseException] = [exc] + queue: Final[list[BaseException]] = [exc] while queue: current = queue.pop(0) nested = getattr(current, "exceptions", None) @@ -123,7 +120,7 @@ class MCPSigV4Auth(httpx.Auth): # Fall back to default boto3 credential chain import botocore.session - session = botocore.session.get_session() + session: Final = botocore.session.get_session() self.credentials = session.get_credentials() if self.credentials is None: raise ValueError( @@ -145,19 +142,19 @@ class MCPSigV4Auth(httpx.Auth): import boto3 from botocore.credentials import Credentials - session_name = aws_session_name or f"litellm-mcp-{int(__import__('time').time())}" - sts_kwargs: dict = {"region_name": aws_region_name} + session_name: Final = aws_session_name or f"litellm-mcp-{int(__import__('time').time())}" + sts_kwargs: Final[dict] = {"region_name": aws_region_name} if aws_access_key_id and aws_secret_access_key: sts_kwargs["aws_access_key_id"] = aws_access_key_id sts_kwargs["aws_secret_access_key"] = aws_secret_access_key if aws_session_token: sts_kwargs["aws_session_token"] = aws_session_token - sts_client = boto3.client("sts", **sts_kwargs) - sts_response = sts_client.assume_role( + sts_client: Final = boto3.client("sts", **sts_kwargs) + sts_response: Final = sts_client.assume_role( RoleArn=aws_role_name, RoleSessionName=session_name, ) - sts_creds = sts_response["Credentials"] + sts_creds: Final = sts_response["Credentials"] return Credentials( access_key=sts_creds["AccessKeyId"], secret_key=sts_creds["SecretAccessKey"], @@ -170,7 +167,7 @@ class MCPSigV4Auth(httpx.Auth): # Build AWSRequest from the httpx Request. # Pass all request headers so the canonical SigV4 signature covers them. - aws_request = AWSRequest( + aws_request: Final = AWSRequest( method=request.method, url=str(request.url), data=request.content, @@ -179,7 +176,7 @@ class MCPSigV4Auth(httpx.Auth): # Sign the request — SigV4Auth.add_auth() adds Authorization, # X-Amz-Date, and X-Amz-Security-Token (if session token present). # Host header is derived automatically from the URL. - sigv4 = SigV4Auth(self.credentials, self.service_name, self.region_name) + sigv4: Final = SigV4Auth(self.credentials, self.service_name, self.region_name) sigv4.add_auth(aws_request) # Copy SigV4 headers back to the httpx request for header_name, header_value in aws_request.headers.items(): @@ -246,7 +243,7 @@ class MCPClient: if self.transport_type == MCPTransport.stdio: if not self.stdio_config: raise ValueError("stdio_config is required for stdio transport") - server_params = StdioServerParameters( + server_params: Final = StdioServerParameters( command=self.stdio_config.get("command", ""), args=self.stdio_config.get("args", []), env=self._get_safe_stdio_env(self.stdio_config.get("env")), @@ -274,7 +271,7 @@ class MCPClient: headers=headers, timeout=httpx.Timeout(self.timeout), ) - transport_ctx = streamable_http_client( + transport_ctx: Final = streamable_http_client( url=self.server_url, http_client=http_client, ) @@ -292,7 +289,7 @@ class MCPClient: return provided_env # Minimal allowlist of safe/standard environment variables - safe_keys = { + safe_keys: Final = { "PATH", "HOME", "USER", @@ -316,7 +313,7 @@ class MCPClient: "WINDIR", } - safe_env = {} + safe_env: Final = {} for key in safe_keys: if key in os.environ: safe_env[key] = os.environ[key] @@ -338,25 +335,25 @@ class MCPClient: so that upstream MCP servers can request LLM inference (sampling), user input (elicitation), or send log messages. """ - transport = await transport_ctx.__aenter__() + transport: Final = await transport_ctx.__aenter__() 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: Final[dict[str, Any]] = {} if self._sampling_callback is not None: session_kwargs["sampling_callback"] = self._sampling_callback if self._elicitation_callback is not None: session_kwargs["elicitation_callback"] = self._elicitation_callback if self._logging_callback is not None: session_kwargs["logging_callback"] = self._logging_callback - session_ctx = ClientSession(read_stream, write_stream, **session_kwargs) - session = await session_ctx.__aenter__() + session_ctx: Final = ClientSession(read_stream, write_stream, **session_kwargs) + session: Final = await session_ctx.__aenter__() try: - init_result = await session.initialize() + init_result: Final = await session.initialize() self._last_initialize_instructions = None if init_result is not None: - ins = getattr(init_result, "instructions", None) + ins: Final = getattr(init_result, "instructions", None) if isinstance(ins, str) and ins.strip(): self._last_initialize_instructions = ins.strip() return await operation(session) @@ -373,7 +370,7 @@ class MCPClient: await transport_ctx.__aexit__(None, None, None) except BaseException as exit_error: verbose_logger.debug("Error during transport context exit: %s", exit_error) - root_cause = _first_non_cancelled_cause(exit_error) + root_cause: Final = _first_non_cancelled_cause(exit_error) if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError): raise root_cause from in_flight_error @@ -394,7 +391,7 @@ class MCPClient: transport_ctx, http_client = self._create_transport_context() return await self._execute_session_operation(transport_ctx, operation) except Exception: - _log = verbose_logger.debug if quiet_on_error else verbose_logger.warning + _log: Final = verbose_logger.debug if quiet_on_error else verbose_logger.warning _log("MCP client run_with_session failed for %s", self.server_url or "stdio") raise finally: @@ -418,7 +415,7 @@ class MCPClient: def _get_auth_headers(self) -> dict: """Generate authentication headers based on auth type.""" - headers = {} + headers: Final = {} if self._mcp_auth_value: if isinstance(self._mcp_auth_value, str): if self.auth_type == MCPAuth.bearer_token: @@ -463,13 +460,13 @@ class MCPClient: ) -> httpx.AsyncClient: """Create an httpx.AsyncClient with LiteLLM's SSL configuration.""" # Get unified SSL configuration using the same logic as http_handler.py - ssl_config = get_ssl_configuration(self.ssl_verify) + ssl_config: Final = get_ssl_configuration(self.ssl_verify) verbose_logger.debug("MCP client using SSL configuration: %s", type(ssl_config).__name__) # The MCP SDK's sse_client and streamable_http_client call this factory without # passing auth=, so the fallback is used: a v2-resolved auth if present, else the # SigV4 aws_auth. Both are None for the common case — no behavior change. - fallback_auth = self._resolved_auth if self._resolved_auth is not None else self._aws_auth - effective_auth = auth if auth is not None else fallback_auth + fallback_auth: Final = self._resolved_auth if self._resolved_auth is not None else self._aws_auth + effective_auth: Final = auth if auth is not None else fallback_auth return httpx.AsyncClient( headers=headers, timeout=timeout, @@ -496,9 +493,9 @@ class MCPClient: return await session.list_tools() try: - result = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error) - tool_count = len(result.tools) - tool_names = [tool.name for tool in result.tools] + result: Final = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error) + tool_count: Final = len(result.tools) + tool_names: Final = [tool.name for tool in result.tools] verbose_logger.info( "MCP client listed %s tools from %s: %s", tool_count, self.server_url or "stdio", tool_names ) @@ -507,13 +504,13 @@ class MCPClient: verbose_logger.warning("MCP client list_tools was cancelled") raise except Exception as e: - error_type = type(e).__name__ + error_type: Final = type(e).__name__ # Mirror call_tool: when the caller opted into raise_on_error it owns the exception and # logs it at the fitting level (an expected pass-through re-auth 401 is info, not an # error), so log at debug here to avoid an error-level line + traceback that would trip # error-rate alerts on that expected signal. The swallow path still logs the full # exception because nothing downstream will surface the failure. - _log = verbose_logger.debug if raise_on_error else verbose_logger.exception + _log: Final = verbose_logger.debug if raise_on_error else verbose_logger.exception _log( f"MCP client list_tools failed - " f"Error Type: {error_type}, " @@ -523,7 +520,7 @@ class MCPClient: ) # Check if it's a stream/connection error if "BrokenResourceError" in error_type or "Broken" in error_type: - _log_broken = verbose_logger.debug if raise_on_error else verbose_logger.error + _log_broken: Final = verbose_logger.debug if raise_on_error else verbose_logger.error _log_broken( "MCP client detected broken connection/stream during list_tools - " "the MCP server may have crashed, disconnected, or timed out" @@ -560,7 +557,7 @@ class MCPClient: verbose_logger.info("MCP client calling tool '%s'", call_tool_request_params.name) async def on_progress(progress: float, total: float | None, message: str | None): - percentage = (progress / total * 100) if total else 0 + percentage: Final = (progress / total * 100) if total else 0 verbose_logger.info( f"MCP Tool '{call_tool_request_params.name}' progress: " f"{progress}/{total} ({percentage:.0f}%) - {message or ''}" @@ -581,7 +578,7 @@ class MCPClient: ) try: - tool_result = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error) + tool_result: Final = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error) verbose_logger.info("MCP client tool call '%s' completed successfully", call_tool_request_params.name) return tool_result except asyncio.CancelledError: @@ -590,16 +587,16 @@ class MCPClient: except Exception as e: import traceback - error_trace = traceback.format_exc() + error_trace: Final = traceback.format_exc() verbose_logger.debug("MCP client tool call traceback:\n%s", error_trace) # Log detailed error information - error_type = type(e).__name__ + error_type: Final = type(e).__name__ # When the caller opted into raise_on_error it owns the exception and logs it at the # level that fits (an expected pass-through re-auth 401 is info, not an operator-actionable # error), so log at debug here to avoid an error-level line that would trip error-rate # alerts on that expected signal. The swallow path (raise_on_error=False) still logs at # error because nothing downstream will surface the failure. - _log = verbose_logger.debug if raise_on_error else verbose_logger.error + _log: Final = verbose_logger.debug if raise_on_error else verbose_logger.error _log( f"MCP client call_tool failed - " f"Error Type: {error_type}, " @@ -627,9 +624,9 @@ class MCPClient: return await session.list_prompts() try: - result = await self.run_with_session(_list_prompts_operation) - prompt_count = len(result.prompts) - prompt_names = [prompt.name for prompt in result.prompts] + result: Final = await self.run_with_session(_list_prompts_operation) + prompt_count: Final = len(result.prompts) + prompt_names: Final = [prompt.name for prompt in result.prompts] verbose_logger.info( "MCP client listed %s tools from %s: %s", prompt_count, self.server_url or "stdio", prompt_names ) @@ -638,7 +635,7 @@ class MCPClient: verbose_logger.warning("MCP client list_prompts was cancelled") raise except Exception as e: - error_type = type(e).__name__ + error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_prompts failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", error_type, @@ -667,7 +664,7 @@ class MCPClient: ) try: - get_prompt_result = await self.run_with_session(_get_prompt_operation) + get_prompt_result: Final = await self.run_with_session(_get_prompt_operation) verbose_logger.info("MCP client get_prompt '%s' completed successfully", get_prompt_request_params.name) return get_prompt_result except asyncio.CancelledError: @@ -676,10 +673,10 @@ class MCPClient: except Exception as e: import traceback - error_trace = traceback.format_exc() + error_trace: Final = traceback.format_exc() verbose_logger.debug("MCP client get_prompt traceback:\n%s", error_trace) # Log detailed error information - error_type = type(e).__name__ + error_type: Final = type(e).__name__ verbose_logger.error( "MCP client get_prompt failed - Error Type: %s, Error: %s, Prompt: %s, Server: %s, Transport: %s", error_type, @@ -704,9 +701,9 @@ class MCPClient: return await session.list_resources() try: - result = await self.run_with_session(_list_resources_operation) - resource_count = len(result.resources) - resource_names = [resource.name for resource in result.resources] + result: Final = await self.run_with_session(_list_resources_operation) + resource_count: Final = len(result.resources) + resource_names: Final = [resource.name for resource in result.resources] verbose_logger.info( "MCP client listed %s resources from %s: %s", resource_count, self.server_url or "stdio", resource_names ) @@ -715,7 +712,7 @@ class MCPClient: verbose_logger.warning("MCP client list_resources was cancelled") raise except Exception as e: - error_type = type(e).__name__ + error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_resources failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", error_type, @@ -740,9 +737,9 @@ class MCPClient: return await session.list_resource_templates() try: - result = await self.run_with_session(_list_resource_templates_operation) - resource_template_count = len(result.resourceTemplates) - resource_template_names = [resourceTemplate.name for resourceTemplate in result.resourceTemplates] + result: Final = await self.run_with_session(_list_resource_templates_operation) + resource_template_count: Final = len(result.resourceTemplates) + resource_template_names: Final = [resourceTemplate.name for resourceTemplate in result.resourceTemplates] verbose_logger.info( "MCP client listed %s resource templates from %s: %s", resource_template_count, @@ -754,7 +751,7 @@ class MCPClient: verbose_logger.warning("MCP client list_resource_templates was cancelled") raise except Exception as e: - error_type = type(e).__name__ + error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_resource_templates failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", error_type, @@ -780,7 +777,7 @@ class MCPClient: return await session.read_resource(url) try: - read_resource_result = await self.run_with_session(_read_resource_operation) + read_resource_result: Final = await self.run_with_session(_read_resource_operation) verbose_logger.info("MCP client read_resource '%s' completed successfully", url) return read_resource_result except asyncio.CancelledError: @@ -789,10 +786,10 @@ class MCPClient: except Exception as e: import traceback - error_trace = traceback.format_exc() + error_trace: Final = traceback.format_exc() verbose_logger.debug("MCP client read_resource traceback:\n%s", error_trace) # Log detailed error information - error_type = type(e).__name__ + error_type: Final = type(e).__name__ verbose_logger.error( "MCP client read_resource failed - Error Type: %s, Error: %s, Url: %s, Server: %s, Transport: %s", error_type, diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py index 23ab77f0037..30d50e2a74b 100644 --- a/litellm/experimental_mcp_client/tools.py +++ b/litellm/experimental_mcp_client/tools.py @@ -1,5 +1,5 @@ import json -from typing import Literal +from typing import Final, Literal from mcp import ClientSession from mcp.types import CallToolRequestParams as MCPCallToolRequestParams @@ -18,7 +18,7 @@ from litellm.types.utils import ChatCompletionMessageToolCall ######################################################## def transform_mcp_tool_to_openai_tool(mcp_tool: MCPTool) -> ChatCompletionToolParam: """Convert an MCP tool to an OpenAI tool.""" - normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema) + normalized_parameters: Final = _normalize_mcp_input_schema(mcp_tool.inputSchema) return ChatCompletionToolParam( type="function", @@ -44,7 +44,7 @@ def _normalize_mcp_input_schema(input_schema: dict) -> dict: return {"type": "object", "properties": {}, "additionalProperties": False} # Make a copy to avoid modifying the original - normalized_schema = dict(input_schema) + normalized_schema: Final = dict(input_schema) # Ensure type is 'object' if "type" not in normalized_schema: @@ -65,7 +65,7 @@ def transform_mcp_tool_to_openai_responses_api_tool( mcp_tool: MCPTool, ) -> FunctionToolParam: """Convert an MCP tool to an OpenAI Responses API tool.""" - normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema) + normalized_parameters: Final = _normalize_mcp_input_schema(mcp_tool.inputSchema) return FunctionToolParam( name=mcp_tool.name, @@ -103,7 +103,7 @@ async def load_mcp_tools( If format is set to "openai", the tools are converted to OpenAI API compatible tools. """ - tools = await session.list_tools() + tools: Final = await session.list_tools() if format == "openai": return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools] return tools.tools @@ -119,7 +119,7 @@ async def call_mcp_tool( call_tool_request_params: MCPCallToolRequestParams, ) -> MCPCallToolResult: """Call an MCP tool.""" - tool_result = await session.call_tool( + tool_result: Final = await session.call_tool( name=call_tool_request_params.name, arguments=call_tool_request_params.arguments, ) @@ -141,7 +141,7 @@ def transform_openai_tool_call_request_to_mcp_tool_call_request( openai_tool: ChatCompletionMessageToolCall | dict, ) -> MCPCallToolRequestParams: """Convert an OpenAI ChatCompletionMessageToolCall to an MCP CallToolRequestParams.""" - function = openai_tool["function"] + function: Final = openai_tool["function"] return MCPCallToolRequestParams( name=function["name"], arguments=_get_function_arguments(function), @@ -161,7 +161,7 @@ async def call_openai_tool( Returns: The result of the MCP tool call. """ - mcp_tool_call_request_params = transform_openai_tool_call_request_to_mcp_tool_call_request( + mcp_tool_call_request_params: Final = transform_openai_tool_call_request_to_mcp_tool_call_request( openai_tool=openai_tool, ) return await call_mcp_tool( diff --git a/litellm/files/main.py b/litellm/files/main.py index e692cdc7c76..e137c7587c0 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, Literal, cast +from typing import Any, Final, Literal, cast import httpx @@ -78,17 +78,17 @@ def _should_sdk_support_streaming( return custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS -openai_files_instance = OpenAIFilesAPI() -azure_files_instance = AzureOpenAIFilesAPI() -vertex_ai_files_instance = VertexAIFilesHandler() -bedrock_files_instance = BedrockFilesHandler() +openai_files_instance: Final = OpenAIFilesAPI() +azure_files_instance: Final = AzureOpenAIFilesAPI() +vertex_ai_files_instance: Final = VertexAIFilesHandler() +bedrock_files_instance: Final = BedrockFilesHandler() ################################################# def _add_trusted_model_credentials_to_litellm_params( litellm_params_dict: dict[str, Any], kwargs: dict[str, Any] ) -> None: - trusted_model_credentials = kwargs.get("_litellm_internal_model_credentials") + trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials") if isinstance(trusted_model_credentials, type(MappingProxyType({}))): litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials @@ -109,10 +109,10 @@ async def acreate_file( LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acreate_file"] = True - call_args = { + call_args: Final = { "file": file, "purpose": purpose, "expires_after": expires_after, @@ -123,11 +123,11 @@ async def acreate_file( } # Use a partial function to pass your keyword arguments - func = partial(create_file, **call_args) + func: Final = partial(create_file, **call_args) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -156,13 +156,13 @@ def create_file( Specify either provider_list or custom_llm_provider. """ try: - _is_async = kwargs.pop("acreate_file", False) is True - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = dict(**kwargs) - logging_obj = cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj")) + _is_async: Final = kwargs.pop("acreate_file", False) is True + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = dict(**kwargs) + logging_obj: Final = cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj")) if logging_obj is None: raise ValueError("logging_obj is required") - client = kwargs.get("client") + client: Final = kwargs.get("client") ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -173,7 +173,7 @@ def create_file( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(cast(str, custom_llm_provider)) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore @@ -196,7 +196,7 @@ def create_file( extra_body=extra_body, ) - provider_config = ProviderConfigManager.get_provider_files_config( + provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", provider=LlmProviders(custom_llm_provider), ) @@ -214,7 +214,7 @@ def create_file( timeout=timeout, ) elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, @@ -229,7 +229,7 @@ def create_file( create_file_data=_create_file_request, ) elif custom_llm_provider == "azure": - azure_creds = get_azure_credentials( + azure_creds: Final = get_azure_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, api_version=optional_params.api_version, @@ -274,11 +274,11 @@ async def afile_retrieve( LiteLLM Equivalent of GET https://api.openai.com/v1/files """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["is_async"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( file_retrieve, file_id, custom_llm_provider, @@ -288,9 +288,9 @@ async def afile_retrieve( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -315,7 +315,7 @@ def file_retrieve( LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files """ try: - optional_params = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -325,17 +325,17 @@ def file_retrieve( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("is_async", False) is True + _is_async: Final = kwargs.pop("is_async", False) is True if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, @@ -350,7 +350,7 @@ def file_retrieve( organization=openai_creds.organization, ) elif custom_llm_provider == "azure": - azure_creds = get_azure_credentials( + azure_creds: Final = get_azure_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, api_version=optional_params.api_version, @@ -366,12 +366,12 @@ def file_retrieve( ) else: # Try using provider config pattern (for Manus, Bedrock, etc.) - provider_config = ProviderConfigManager.get_provider_files_config( + provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", provider=LlmProviders(custom_llm_provider), ) if provider_config is not None: - litellm_params_dict = get_litellm_params(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) _add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, @@ -395,7 +395,7 @@ def file_retrieve( function_id=str(kwargs.get("id") or ""), ) - client = kwargs.get("client") + client: Final = kwargs.get("client") response = base_llm_http_handler.retrieve_file( file_id=file_id, provider_config=provider_config, @@ -443,12 +443,12 @@ async def afile_delete( LiteLLM Equivalent of DELETE https://api.openai.com/v1/files """ try: - loop = asyncio.get_event_loop() - model = kwargs.pop("model", None) + loop: Final = asyncio.get_event_loop() + model: Final = kwargs.pop("model", None) kwargs["is_async"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( file_delete, file_id, model, @@ -459,9 +459,9 @@ async def afile_delete( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -492,8 +492,8 @@ def file_delete( _, custom_llm_provider, _, _ = get_llm_provider(model, custom_llm_provider) except Exception: pass - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) _add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, @@ -501,22 +501,22 @@ def file_delete( ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default - client = kwargs.get("client") + client: Final = kwargs.get("client") if ( timeout is not None and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("is_async", False) is True + _is_async: Final = kwargs.pop("is_async", False) is True if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, @@ -531,7 +531,7 @@ def file_delete( organization=openai_creds.organization, ) elif custom_llm_provider == "azure": - azure_creds = get_azure_credentials( + azure_creds: Final = get_azure_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, api_version=optional_params.api_version, @@ -549,7 +549,7 @@ def file_delete( ) else: # Try using provider config pattern (for Manus, Bedrock, etc.) - provider_config = ProviderConfigManager.get_provider_files_config( + provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", provider=LlmProviders(custom_llm_provider), ) @@ -619,11 +619,11 @@ async def afile_list( LiteLLM Equivalent of GET https://api.openai.com/v1/files """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["is_async"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( file_list, custom_llm_provider, purpose, @@ -633,9 +633,9 @@ async def afile_list( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -660,7 +660,7 @@ def file_list( LiteLLM Equivalent of GET https://api.openai.com/v1/files """ try: - optional_params = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -670,22 +670,22 @@ def file_list( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("is_async", False) is True + _is_async: Final = kwargs.pop("is_async", False) is True # Check if provider has a custom files config (e.g., Manus, Bedrock, Vertex AI) - provider_config = ProviderConfigManager.get_provider_files_config( + provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", provider=LlmProviders(custom_llm_provider), ) if provider_config is not None: - litellm_params_dict = get_litellm_params(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) litellm_params_dict["api_key"] = optional_params.api_key litellm_params_dict["api_base"] = optional_params.api_base @@ -705,7 +705,7 @@ def file_list( function_id=str(kwargs.get("id", "")), ) - client = kwargs.get("client") + client: Final = kwargs.get("client") response = base_llm_http_handler.list_files( purpose=purpose, provider_config=provider_config, @@ -718,7 +718,7 @@ def file_list( ) return response elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, @@ -733,7 +733,7 @@ def file_list( organization=openai_creds.organization, ) elif custom_llm_provider == "azure": - azure_creds = get_azure_credentials( + azure_creds: Final = get_azure_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, api_version=optional_params.api_version, @@ -779,12 +779,12 @@ async def afile_content( LiteLLM Equivalent of GET https://api.openai.com/v1/files """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["afile_content"] = True - model = kwargs.pop("model", None) + model: Final = kwargs.pop("model", None) # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( file_content, file_id=file_id, model=model, @@ -797,9 +797,9 @@ async def afile_content( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -832,15 +832,15 @@ def file_content( LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files """ try: - optional_params = GenericLiteLLMParams(**kwargs) - litellm_params_dict = get_litellm_params(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) _add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, ) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 - client = kwargs.get("client") + client: Final = kwargs.get("client") # set timeout for 10 minutes by default try: @@ -854,20 +854,20 @@ def file_content( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(cast(str, custom_llm_provider)) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _file_content_request = FileContentRequest( + _file_content_request: Final = FileContentRequest( file_id=file_id, extra_headers=extra_headers, extra_body=extra_body, ) - _is_async = kwargs.pop("afile_content", False) is True + _is_async: Final = kwargs.pop("afile_content", False) is True if stream and _should_sdk_support_streaming(custom_llm_provider): return file_content_streaming( @@ -885,7 +885,7 @@ def file_content( ) # Check if provider has a custom files config (e.g., Anthropic, Manus) - provider_config = ProviderConfigManager.get_provider_files_config( + provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", provider=LlmProviders(custom_llm_provider), ) @@ -918,7 +918,7 @@ def file_content( return response if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, @@ -933,7 +933,7 @@ def file_content( organization=openai_creds.organization, ) elif custom_llm_provider == "azure": - azure_creds = get_azure_credentials( + azure_creds: Final = get_azure_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, api_version=optional_params.api_version, @@ -950,14 +950,14 @@ def file_content( litellm_params=litellm_params_dict, ) elif custom_llm_provider == "vertex_ai": - api_base = optional_params.api_base or "" - vertex_ai_project = ( + api_base: Final = optional_params.api_base or "" + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_ai_files_instance.file_content( _is_async=_is_async, @@ -1014,7 +1014,7 @@ def file_content_streaming( logging_obj.model_call_details["model"] = model or "" logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider - litellm_params = logging_obj.model_call_details.get("litellm_params", {}) or {} + litellm_params: Final = logging_obj.model_call_details.get("litellm_params", {}) or {} if optional_params.api_base is not None: litellm_params["api_base"] = optional_params.api_base logging_obj.model_call_details["litellm_params"] = litellm_params @@ -1037,7 +1037,7 @@ def file_content_streaming( stream_iterator=iter(()), headers={} ) if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: - openai_creds = get_openai_credentials( + openai_creds: Final = get_openai_credentials( api_base=optional_params.api_base, api_key=optional_params.api_key, organization=optional_params.organization, diff --git a/litellm/files/streaming.py b/litellm/files/streaming.py index 7c2b53b395e..d9df05e7135 100644 --- a/litellm/files/streaming.py +++ b/litellm/files/streaming.py @@ -1,12 +1,7 @@ import datetime import traceback from collections.abc import AsyncIterator, Iterator -from typing import ( - TYPE_CHECKING, - Any, - Optional, - cast, -) +from typing import TYPE_CHECKING, Any, Final, Optional, cast import anyio @@ -91,7 +86,7 @@ class FileContentStreamingResponse: self._close_completed = True self._logging_completed = True - stream_to_close = self.stream_iterator + stream_to_close: Final = self.stream_iterator self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(())) # Shield cleanup from request cancellation so upstream HTTP connections @@ -100,7 +95,7 @@ class FileContentStreamingResponse: if hasattr(stream_to_close, "aclose"): await cast(AsyncIterator[bytes], stream_to_close).aclose() # type: ignore[attr-defined] elif hasattr(stream_to_close, "close"): - result = cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined] + result: Final = cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined] if result is not None: await result @@ -110,14 +105,14 @@ class FileContentStreamingResponse: self._close_completed = True self._logging_completed = True - stream_to_close = self.stream_iterator + stream_to_close: Final = self.stream_iterator 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]: - response = { + response: Final = { "id": self.file_id, "object": "file.content", } @@ -154,7 +149,7 @@ class FileContentStreamingResponse: ) self._sync_hidden_params() - payload = get_standard_logging_object_payload( + payload: Final = get_standard_logging_object_payload( kwargs=self.logging_obj.model_call_details, init_response_obj=self._build_logging_response(), start_time=self._start_time, @@ -165,7 +160,7 @@ class FileContentStreamingResponse: if payload is None: return None - merged_hidden_params = cast( + merged_hidden_params: Final = cast( "StandardLoggingHiddenParams", { **cast(dict[str, Any], payload.get("hidden_params") or {}), @@ -189,8 +184,8 @@ class FileContentStreamingResponse: return self._logging_completed = True - end_time = datetime.datetime.now() - standard_logging_object = self._build_standard_logging_object(end_time=end_time) + end_time: Final = datetime.datetime.now() + standard_logging_object: Final = self._build_standard_logging_object(end_time=end_time) await self.logging_obj.async_success_handler( result=self._build_logging_response(), start_time=self._start_time, @@ -208,8 +203,8 @@ class FileContentStreamingResponse: return self._logging_completed = True - end_time = datetime.datetime.now() - standard_logging_object = self._build_standard_logging_object(end_time=end_time) + end_time: Final = datetime.datetime.now() + standard_logging_object: Final = self._build_standard_logging_object(end_time=end_time) self.logging_obj.success_handler( result=self._build_logging_response(), start_time=self._start_time, @@ -222,8 +217,8 @@ class FileContentStreamingResponse: return self._logging_completed = True - end_time = datetime.datetime.now() - traceback_str = traceback.format_exc() + end_time: Final = datetime.datetime.now() + traceback_str: Final = traceback.format_exc() self.logging_obj.failure_handler(error, traceback_str, self._start_time, end_time) await self.logging_obj.async_failure_handler(error, traceback_str, self._start_time, end_time) @@ -232,5 +227,5 @@ class FileContentStreamingResponse: return self._logging_completed = True - end_time = datetime.datetime.now() + end_time: Final = datetime.datetime.now() self.logging_obj.failure_handler(error, traceback.format_exc(), self._start_time, end_time) diff --git a/litellm/files/utils.py b/litellm/files/utils.py index 3c58533f66a..f470931115f 100644 --- a/litellm/files/utils.py +++ b/litellm/files/utils.py @@ -1,3 +1,5 @@ +from typing import Final + from litellm.types.llms.openai import CreateFileRequest from litellm.types.utils import ExtractedFileData @@ -6,7 +8,7 @@ from litellm.types.utils import ExtractedFileData # batch file must not silently bypass the streaming path just because of its # declared type. ``purpose == "batch"`` is the authoritative signal; non-JSONL # content still fails loudly when the rows are parsed. -_BATCH_JSONL_CONTENT_TYPES = frozenset( +_BATCH_JSONL_CONTENT_TYPES: Final = frozenset( { "application/jsonl", "application/json", diff --git a/litellm/fine_tuning/main.py b/litellm/fine_tuning/main.py index 1987bf6a284..e89defedabe 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, Literal +from typing import Any, Final, Literal import httpx @@ -29,9 +29,9 @@ from litellm.types.utils import LiteLLMFineTuningJob from litellm.utils import client, supports_httpx_timeout ####### ENVIRONMENT VARIABLES ################### -openai_fine_tuning_apis_instance = OpenAIFineTuningAPI() -azure_fine_tuning_apis_instance = AzureOpenAIFineTuningAPI() -vertex_fine_tuning_apis_instance = VertexFineTuningAPI() +openai_fine_tuning_apis_instance: Final = OpenAIFineTuningAPI() +azure_fine_tuning_apis_instance: Final = AzureOpenAIFineTuningAPI() +vertex_fine_tuning_apis_instance: Final = VertexFineTuningAPI() ################################################# @@ -61,7 +61,7 @@ def _prepare_azure_extra_body( extra_body = {} # Azure-specific root-level parameters - azure_specific_params = ["trainingType"] + azure_specific_params: Final = ["trainingType"] for param in azure_specific_params: if param in kwargs: extra_body[param] = kwargs[param] @@ -93,11 +93,11 @@ async def acreate_fine_tuning_job( """ verbose_logger.debug("inside acreate_fine_tuning_job model=%s and kwargs=%s", model, kwargs) try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acreate_fine_tuning_job"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( create_fine_tuning_job, model, training_file, @@ -113,9 +113,9 @@ async def acreate_fine_tuning_job( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -171,24 +171,24 @@ def create_fine_tuning_job( """ try: - _is_async = kwargs.pop("acreate_fine_tuning_job", False) is True - optional_params = GenericLiteLLMParams(**kwargs) + _is_async: Final = kwargs.pop("acreate_fine_tuning_job", False) is True + optional_params: Final = GenericLiteLLMParams(**kwargs) # handle hyperparameters hyperparameters = hyperparameters or {} # original hyperparameters # For Azure, extract Azure-specific hyperparameters before creating OpenAI-spec hyperparameters - azure_specific_hyperparams = {} + azure_specific_hyperparams: Final = {} if custom_llm_provider == "azure": - azure_hyperparameter_keys = ["prompt_loss_weight"] + azure_hyperparameter_keys: Final = ["prompt_loss_weight"] for key in azure_hyperparameter_keys: if key in hyperparameters: azure_specific_hyperparams[key] = hyperparameters.pop(key) - _oai_hyperparameters: Hyperparameters = Hyperparameters( + _oai_hyperparameters: Final[Hyperparameters] = Hyperparameters( **hyperparameters ) # Typed Hyperparameters for OpenAI Spec - timeout = _resolve_fine_tuning_timeout( + timeout: Final = _resolve_fine_tuning_timeout( optional_params.timeout or kwargs.get("request_timeout", 600), custom_llm_provider, ) @@ -203,7 +203,7 @@ def create_fine_tuning_job( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -287,13 +287,13 @@ def create_fine_tuning_job( ) elif custom_llm_provider == "vertex_ai": api_base = optional_params.api_base or "" - vertex_ai_project = ( + vertex_ai_project: Final = ( optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT") ) - vertex_ai_location = ( + vertex_ai_location: Final = ( optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION") ) - vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") + vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS") response = vertex_fine_tuning_apis_instance.create_fine_tuning_job( _is_async=_is_async, create_fine_tuning_job_data=_build_fine_tuning_job_data( @@ -342,11 +342,11 @@ async def acancel_fine_tuning_job( Async: Immediately cancel a fine-tune job. """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["acancel_fine_tuning_job"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( cancel_fine_tuning_job, fine_tuning_job_id, custom_llm_provider, @@ -356,9 +356,9 @@ async def acancel_fine_tuning_job( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -383,7 +383,7 @@ def cancel_fine_tuning_job( """ try: - optional_params = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -393,14 +393,14 @@ def cancel_fine_tuning_job( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("acancel_fine_tuning_job", False) is True + _is_async: Final = kwargs.pop("acancel_fine_tuning_job", False) is True # OpenAI if custom_llm_provider == "openai": @@ -412,7 +412,7 @@ def cancel_fine_tuning_job( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -493,11 +493,11 @@ async def alist_fine_tuning_jobs( Async: List your organization's fine-tuning jobs """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["alist_fine_tuning_jobs"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( list_fine_tuning_jobs, after, limit, @@ -508,9 +508,9 @@ async def alist_fine_tuning_jobs( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -537,7 +537,7 @@ def list_fine_tuning_jobs( - limit: Optional[int] = None, Number of fine-tuning jobs to retrieve. Defaults to 20 """ try: - optional_params = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -547,14 +547,14 @@ def list_fine_tuning_jobs( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("alist_fine_tuning_jobs", False) is True + _is_async: Final = kwargs.pop("alist_fine_tuning_jobs", False) is True # OpenAI if custom_llm_provider == "openai": @@ -566,7 +566,7 @@ def list_fine_tuning_jobs( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) @@ -649,11 +649,11 @@ async def aretrieve_fine_tuning_job( Async: Get info about a fine-tuning job. """ try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["aretrieve_fine_tuning_job"] = True # Use a partial function to pass your keyword arguments - func = partial( + func: Final = partial( retrieve_fine_tuning_job, fine_tuning_job_id, custom_llm_provider, @@ -663,9 +663,9 @@ async def aretrieve_fine_tuning_job( ) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response else: @@ -687,7 +687,7 @@ def retrieve_fine_tuning_job( Get info about a fine-tuning job. """ try: - optional_params = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams(**kwargs) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 # set timeout for 10 minutes by default @@ -697,14 +697,14 @@ def retrieve_fine_tuning_job( and isinstance(timeout, httpx.Timeout) and supports_httpx_timeout(custom_llm_provider) is False ): - read_timeout = timeout.read or 600 + read_timeout: Final = timeout.read or 600 timeout = read_timeout # default 10 min timeout elif timeout is not None and not isinstance(timeout, httpx.Timeout): timeout = float(timeout) # type: ignore elif timeout is None: timeout = 600.0 - _is_async = kwargs.pop("aretrieve_fine_tuning_job", False) is True + _is_async: Final = kwargs.pop("aretrieve_fine_tuning_job", False) is True # OpenAI if custom_llm_provider == "openai": @@ -715,7 +715,7 @@ def retrieve_fine_tuning_job( or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1" ) - organization = ( + organization: Final = ( optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None ) api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY") diff --git a/litellm/google_genai/adapters/handler.py b/litellm/google_genai/adapters/handler.py index 5236e207cc5..5dafe2befee 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, cast +from typing import Any, Final, cast import litellm from litellm.types.router import GenericLiteLLMParams @@ -8,7 +8,7 @@ from litellm.types.utils import ModelResponse from .transformation import GoogleGenAIAdapter # Initialize adapter -GOOGLE_GENAI_ADAPTER = GoogleGenAIAdapter() +GOOGLE_GENAI_ADAPTER: Final = GoogleGenAIAdapter() class GenerateContentToCompletionHandler: @@ -26,7 +26,7 @@ class GenerateContentToCompletionHandler: """Prepare kwargs for litellm.completion/acompletion""" # Transform generate_content request to completion format - completion_request = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion( + completion_request: Final = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion( model=model, contents=contents, config=config, @@ -34,7 +34,7 @@ class GenerateContentToCompletionHandler: **(extra_kwargs or {}), ) - completion_kwargs: dict[str, Any] = dict(completion_request) + completion_kwargs: Final[dict[str, Any]] = dict(completion_request) # Forward extra_kwargs that should be passed to completion call if extra_kwargs is not None: @@ -61,7 +61,7 @@ class GenerateContentToCompletionHandler: ) -> dict[str, Any] | AsyncIterator[bytes]: """Handle generate_content call asynchronously using completion adapter""" - completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs( + completion_kwargs: Final = GenerateContentToCompletionHandler._prepare_completion_kwargs( model=model, contents=contents, config=config, @@ -71,7 +71,7 @@ class GenerateContentToCompletionHandler: ) try: - completion_response = await litellm.acompletion(**completion_kwargs) + completion_response: Final = await litellm.acompletion(**completion_kwargs) if stream: # Check if completion_response is actually a stream or a ModelResponse @@ -84,7 +84,7 @@ class GenerateContentToCompletionHandler: return generate_content_response else: # Transform streaming completion response to generate_content format - transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming( + transformed_stream: Final = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming( completion_response ) if transformed_stream is not None: @@ -122,7 +122,7 @@ class GenerateContentToCompletionHandler: **kwargs, ) - completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs( + completion_kwargs: Final = GenerateContentToCompletionHandler._prepare_completion_kwargs( model=model, contents=contents, config=config, @@ -132,7 +132,7 @@ class GenerateContentToCompletionHandler: ) try: - completion_response = litellm.completion(**completion_kwargs) + completion_response: Final = litellm.completion(**completion_kwargs) if stream: # Check if completion_response is actually a stream or a ModelResponse @@ -145,7 +145,7 @@ class GenerateContentToCompletionHandler: return generate_content_response else: # Transform streaming completion response to generate_content format - transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming( + transformed_stream: Final = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming( completion_response ) if transformed_stream is not None: diff --git a/litellm/google_genai/adapters/transformation.py b/litellm/google_genai/adapters/transformation.py index b13fb71690b..4f127f476c3 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, cast +from typing import Any, Final, cast from litellm import verbose_logger from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema @@ -85,7 +85,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper): # After the stream is exhausted, check for any remaining accumulated tool calls if self.accumulated_tool_calls: try: - parts = [] + parts: Final = [] for ( tool_call_index, tool_call_data, @@ -110,7 +110,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper): tool_call_data["arguments"], ) if parts: - final_chunk = { + final_chunk: Final = { "candidates": [ { "content": {"parts": parts, "role": "model"}, @@ -197,9 +197,9 @@ class GoogleGenAIAdapter: """ # Extract top-level fields from kwargs - system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction") - tools = kwargs.get("tools") - tool_config = kwargs.get("toolConfig") or kwargs.get("tool_config") + system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction") + tools: Final = kwargs.get("tools") + tool_config: Final = kwargs.get("toolConfig") or kwargs.get("tool_config") # Normalize contents to list format if isinstance(contents, dict): @@ -208,10 +208,10 @@ class GoogleGenAIAdapter: contents_list = contents # Transform contents to OpenAI messages format - messages = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction) + messages: Final = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction) # Create base request as dict (which is compatible with ChatCompletionRequest) - completion_request: ChatCompletionRequest = { + completion_request: Final[ChatCompletionRequest] = { "model": model, "messages": messages, } @@ -248,14 +248,14 @@ class GoogleGenAIAdapter: # Check if tools are already in OpenAI format or Google GenAI format if isinstance(tools, list) and len(tools) > 0: # Tools are in Google GenAI format, transform them - openai_tools = self._transform_google_genai_tools_to_openai(tools) + openai_tools: Final = self._transform_google_genai_tools_to_openai(tools) if openai_tools: completion_request["tools"] = openai_tools # Handle tool_config (tool choice) if tool_config: - tool_choice = self._transform_google_genai_tool_config_to_openai(tool_config) + tool_choice: Final = self._transform_google_genai_tool_config_to_openai(tool_config) if tool_choice: completion_request["tool_choice"] = tool_choice @@ -285,9 +285,9 @@ class GoogleGenAIAdapter: Returns: Dict[str, Any] """ - allowed_fields = GenericLiteLLMParams.model_fields.keys() + allowed_fields: Final = GenericLiteLLMParams.model_fields.keys() if litellm_params: - litellm_dict = litellm_params.model_dump(exclude_none=True) + litellm_dict: Final = litellm_params.model_dump(exclude_none=True) for key, value in litellm_dict.items(): if key in allowed_fields: completion_request_dict[key] = value @@ -298,7 +298,7 @@ class GoogleGenAIAdapter: completion_stream: Any, ) -> AsyncIterator[bytes] | None: """Transform streaming completion output to Google GenAI format""" - google_genai_wrapper = GoogleGenAIStreamWrapper(completion_stream=completion_stream) + google_genai_wrapper: Final = GoogleGenAIStreamWrapper(completion_stream=completion_stream) # Return the SSE-wrapped version for proper event formatting return google_genai_wrapper.async_google_genai_sse_wrapper() @@ -307,7 +307,7 @@ class GoogleGenAIAdapter: tools: list[dict[str, Any]], ) -> list[ChatCompletionToolParam]: """Transform Google GenAI tools to OpenAI tools format""" - openai_tools: list[dict[str, Any]] = [] + openai_tools: Final[list[dict[str, Any]]] = [] for tool in tools: if "functionDeclarations" in tool: @@ -325,7 +325,7 @@ class GoogleGenAIAdapter: openai_tools.append(openai_tool) # normalize the tool schemas - normalized_tools = [normalize_tool_schema(tool) for tool in openai_tools] + normalized_tools: Final = [normalize_tool_schema(tool) for tool in openai_tools] return cast(list[ChatCompletionToolParam], normalized_tools) @@ -334,12 +334,12 @@ class GoogleGenAIAdapter: 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") + function_calling_config: Final = tool_config.get("functionCallingConfig", {}) + mode: Final = function_calling_config.get("mode", "AUTO") - mode_mapping = {"AUTO": "auto", "ANY": "required", "NONE": "none"} + mode_mapping: Final = {"AUTO": "auto", "ANY": "required", "NONE": "none"} - tool_choice = mode_mapping.get(mode, "auto") + tool_choice: Final = mode_mapping.get(mode, "auto") return cast(ChatCompletionToolChoiceValues, tool_choice) def _transform_contents_to_messages( @@ -348,11 +348,11 @@ class GoogleGenAIAdapter: system_instruction: dict[str, Any] | None = None, ) -> list[AllMessageValues]: """Transform Google GenAI contents to OpenAI messages format""" - messages: list[AllMessageValues] = [] + messages: Final[list[AllMessageValues]] = [] # Handle system instruction if system_instruction: - system_parts = system_instruction.get("parts", []) + system_parts: Final = system_instruction.get("parts", []) if system_parts and "text" in system_parts[0]: messages.append(ChatCompletionSystemMessage(role="system", content=system_parts[0]["text"])) @@ -473,7 +473,7 @@ class GoogleGenAIAdapter: """ # Extract the main response content - choice = response.choices[0] if response.choices else None + choice: Final = response.choices[0] if response.choices else None if not choice: raise ValueError("Invalid completion response: no choices found") @@ -490,7 +490,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: Final[dict[str, Any]] = { "candidates": [ { "content": {"parts": parts, "role": "model"}, @@ -537,7 +537,7 @@ class GoogleGenAIAdapter: """ # Extract the main response content from streaming chunk - choice = response.choices[0] if response.choices else None + choice: Final = response.choices[0] if response.choices else None if not choice: # Return empty chunk if no choices return None @@ -551,7 +551,7 @@ class GoogleGenAIAdapter: finish_reason = getattr(choice, "finish_reason", None) else: # Fallback for generic choice objects - message_content = getattr(choice, "delta", {}).get("content", "") + message_content: Final = getattr(choice, "delta", {}).get("content", "") parts = [{"text": message_content}] if message_content else [] finish_reason = getattr(choice, "finish_reason", None) @@ -560,7 +560,7 @@ class GoogleGenAIAdapter: return None # Create Google GenAI streaming format response - streaming_chunk: dict[str, Any] = { + streaming_chunk: Final[dict[str, Any]] = { "candidates": [ { "content": {"parts": parts, "role": "model"}, @@ -573,7 +573,7 @@ class GoogleGenAIAdapter: # Add usage metadata only in the final chunk (when finish_reason is present) if finish_reason: - usage_metadata = ( + usage_metadata: Final = ( self._map_usage(getattr(response, "usage", None)) if hasattr(response, "usage") and getattr(response, "usage", None) else { @@ -599,7 +599,7 @@ class GoogleGenAIAdapter: message: Any, ) -> list[dict[str, Any]]: """Transform OpenAI message to Google GenAI parts format""" - parts: list[dict[str, Any]] = [] + parts: Final[list[dict[str, Any]]] = [] # Add text content if present if hasattr(message, "content") and message.content: @@ -633,13 +633,13 @@ class GoogleGenAIAdapter: if not hasattr(wrapper, "accumulated_tool_calls"): wrapper.accumulated_tool_calls = {} - parts: list[dict[str, Any]] = [] + parts: Final[list[dict[str, Any]]] = [] if hasattr(delta, "content") and delta.content: parts.append({"text": delta.content}) # 2. Ensure tool_calls is iterable - tool_calls = delta.tool_calls or [] + tool_calls: Final = delta.tool_calls or [] for tool_call in tool_calls: if not hasattr(tool_call, "function"): @@ -704,7 +704,7 @@ class GoogleGenAIAdapter: if not finish_reason: return "STOP" - mapping = { + mapping: Final = { "stop": "STOP", "length": "MAX_TOKENS", "content_filter": "SAFETY", diff --git a/litellm/google_genai/main.py b/litellm/google_genai/main.py index dbb124a3106..634739d86f8 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 +from typing import TYPE_CHECKING, Any, ClassVar, Final import httpx from pydantic import BaseModel, ConfigDict @@ -110,11 +110,11 @@ class GenerateContentHelper: Returns: GenerateContentSetupResult containing all setup information """ - litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj") - litellm_call_id: str | None = kwargs.get("litellm_call_id", None) + litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj") + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) # get llm provider logic - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) ## MOCK RESPONSE LOGIC (only for non-streaming) if ( @@ -140,7 +140,7 @@ class GenerateContentHelper: litellm_params.custom_llm_provider = custom_llm_provider # get provider config - generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None = ( + generate_content_provider_config: Final[BaseGoogleGenAIGenerateContentConfig | None] = ( ProviderConfigManager.get_provider_google_genai_generate_content_config( model=model, provider=litellm.LlmProviders(custom_llm_provider), @@ -168,19 +168,19 @@ class GenerateContentHelper: # Construct request body ######################################################################################### # Create Google Optional Params Config - generate_content_config_dict = generate_content_provider_config.map_generate_content_optional_params( + generate_content_config_dict: Final = generate_content_provider_config.map_generate_content_optional_params( generate_content_config_dict=config or {}, model=model, ) # Extract systemInstruction from kwargs to pass to transform - system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction") + system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction") # Native top-level REST fields arrive as loose kwargs and are otherwise dropped. - native_request_fields: dict[str, object] = { + native_request_fields: Final[dict[str, object]] = { field: kwargs[field] for field in generate_content_provider_config.get_generate_content_request_top_level_fields() if field in kwargs } - request_body = generate_content_provider_config.transform_generate_content_request( + request_body: Final = generate_content_provider_config.transform_generate_content_request( model=model, contents=contents, tools=tools, @@ -250,9 +250,9 @@ async def agenerate_content( """ Async: Generate content using Google GenAI """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["agenerate_content"] = True # Handle generationConfig parameter from kwargs for backward compatibility @@ -265,7 +265,7 @@ async def agenerate_content( custom_llm_provider=custom_llm_provider, ) - func = partial( + func: Final = partial( generate_content, model=model, contents=contents, @@ -279,9 +279,9 @@ async def agenerate_content( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -318,9 +318,9 @@ def generate_content( """ Generate content using Google GenAI """ - local_vars = locals() + local_vars: Final = locals() try: - _is_async = kwargs.pop("agenerate_content", False) + _is_async: Final = kwargs.pop("agenerate_content", False) _mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content.value, _is_async) @@ -328,12 +328,12 @@ def generate_content( if "generationConfig" in kwargs and config is None: config = kwargs.pop("generationConfig") # Check for mock response first - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) if litellm_params.mock_response and isinstance(litellm_params.mock_response, str): return GenerateContentHelper.mock_generate_content_response(mock_response=litellm_params.mock_response) # Setup the call - setup_result = GenerateContentHelper.setup_generate_content_call( + setup_result: Final = GenerateContentHelper.setup_generate_content_call( model=model, contents=contents, config=config, @@ -343,7 +343,7 @@ def generate_content( ) # Extract systemInstruction from kwargs to pass to handler - system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction") + system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction") # Check if we should use the adapter (when provider config is None) if setup_result.generate_content_provider_config is None: @@ -360,7 +360,7 @@ def generate_content( ) # Call the standard handler - response = base_llm_http_handler.generate_content_handler( + response: Final = base_llm_http_handler.generate_content_handler( model=setup_result.model, contents=contents, tools=tools, @@ -408,7 +408,7 @@ async def agenerate_content_stream( """ Async: Generate content using Google GenAI with streaming response """ - local_vars = locals() + local_vars: Final = locals() try: kwargs["agenerate_content_stream"] = True @@ -424,7 +424,7 @@ async def agenerate_content_stream( ) # Setup the call - setup_result = GenerateContentHelper.setup_generate_content_call( + setup_result: Final = GenerateContentHelper.setup_generate_content_call( model=model, contents=contents, config=config, @@ -434,7 +434,7 @@ async def agenerate_content_stream( ) # Extract systemInstruction from kwargs to pass to handler - system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction") + system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction") # Check if we should use the adapter (when provider config is None) if setup_result.generate_content_provider_config is None: @@ -503,10 +503,10 @@ def generate_content_stream( """ Generate content using Google GenAI with streaming response """ - local_vars = locals() + local_vars: Final = locals() try: # Remove any async-related flags since this is the sync function - _is_async = kwargs.pop("agenerate_content_stream", False) + _is_async: Final = kwargs.pop("agenerate_content_stream", False) _mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, _is_async) @@ -514,7 +514,7 @@ def generate_content_stream( if "generationConfig" in kwargs and config is None: config = kwargs.pop("generationConfig") # Setup the call - setup_result = GenerateContentHelper.setup_generate_content_call( + setup_result: Final = GenerateContentHelper.setup_generate_content_call( model=model, contents=contents, config=config, @@ -524,7 +524,7 @@ def generate_content_stream( ) # Extract systemInstruction from kwargs to pass to handler - system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction") + system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction") # Check if we should use the adapter (when provider config is None) if setup_result.generate_content_provider_config is None: diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index 2829699492d..e03f7ee745f 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 +from typing import TYPE_CHECKING, Any, Final from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.success_handler import ( @@ -15,7 +15,7 @@ if TYPE_CHECKING: else: BaseGoogleGenAIGenerateContentConfig = Any -GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ = PassThroughEndpointLogging() +GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging() def _encode_google_genai_sse_event(event_lines: list[str]) -> bytes: @@ -23,7 +23,7 @@ def _encode_google_genai_sse_event(event_lines: list[str]) -> bytes: def _next_google_genai_sse_chunk(line_iter) -> bytes: - event_lines: list[str] = [] + event_lines: Final[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: Final[list[str]] = [] while True: try: line = await line_iter.__anext__() @@ -82,7 +82,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: PassThroughStreamingHandler, ) - end_time = datetime.now() + end_time: Final = datetime.now() asyncio.create_task( PassThroughStreamingHandler._route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, @@ -134,7 +134,7 @@ class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContent def __next__(self): try: - chunk = _next_google_genai_sse_chunk(self.stream_iterator) + chunk: Final = _next_google_genai_sse_chunk(self.stream_iterator) self.collected_chunks.append(chunk) return chunk except StopIteration: @@ -185,7 +185,7 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateCo async def __anext__(self): try: - chunk = await _anext_google_genai_sse_chunk(self.stream_iterator) + chunk: Final = await _anext_google_genai_sse_chunk(self.stream_iterator) self.collected_chunks.append(chunk) return chunk except StopAsyncIteration: diff --git a/litellm/images/main.py b/litellm/images/main.py index d26c9d54f83..4430bb5beb4 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -3,14 +3,7 @@ import contextvars import importlib from collections.abc import Coroutine from functools import partial -from typing import ( - TYPE_CHECKING, - Any, - Literal, - Optional, - cast, - overload, -) +from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload if TYPE_CHECKING: from litellm.images.utils import ImageEditRequestUtils @@ -76,7 +69,7 @@ def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils": global _ImageEditRequestUtils_cache if _ImageEditRequestUtils_cache is None: # Access via module to trigger __getattr__ if not cached - module = importlib.import_module(__name__) + module: Final = importlib.import_module(__name__) _ImageEditRequestUtils_cache = module.ImageEditRequestUtils assert _ImageEditRequestUtils_cache is not None # Type narrowing for type checker return _ImageEditRequestUtils_cache @@ -95,23 +88,23 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: Returns: - `response` (Any): The response returned by the `image_generation` function. """ - loop = asyncio.get_event_loop() - model = args[0] if len(args) > 0 else kwargs["model"] + loop: Final = asyncio.get_event_loop() + model: Final = args[0] if len(args) > 0 else kwargs["model"] ### PASS ARGS TO Image Generation ### kwargs["aimg_generation"] = True custom_llm_provider = None try: # Use a partial function to pass your keyword arguments - func = partial(image_generation, *args, **kwargs) + func: Final = partial(image_generation, *args, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) # Await normally - init_response = await loop.run_in_executor(None, func_with_context) + init_response: Final = await loop.run_in_executor(None, func_with_context) response: ImageResponse | None = None if isinstance(init_response, dict): @@ -210,20 +203,20 @@ def image_generation( Currently supports just Azure + OpenAI. """ try: - args = locals() - 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: str | None = kwargs.get("mock_response", None) # type: ignore - proxy_server_request = kwargs.get("proxy_server_request", None) + args: Final = locals() + aimg_generation: Final = kwargs.get("aimg_generation", False) + litellm_call_id: Final = kwargs.get("litellm_call_id", None) + logger_fn: Final = kwargs.get("logger_fn", None) + mock_response: Final[str | None] = kwargs.get("mock_response", None) # type: ignore + proxy_server_request: Final = kwargs.get("proxy_server_request", None) azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) - model_info = kwargs.get("model_info", None) - metadata = kwargs.get("metadata", {}) - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - client = kwargs.get("client", None) - extra_headers = kwargs.get("extra_headers", None) - headers: dict = kwargs.get("headers", None) or {} - base_model = kwargs.get("base_model", None) + model_info: Final = kwargs.get("model_info", None) + metadata: Final = kwargs.get("metadata", {}) + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + client: Final = kwargs.get("client", None) + extra_headers: Final = kwargs.get("extra_headers", None) + headers: Final[dict] = kwargs.get("headers", None) or {} + base_model: Final = kwargs.get("base_model", None) if extra_headers is not None: headers.update(extra_headers) model_response: ImageResponse = litellm.utils.ImageResponse() @@ -238,7 +231,7 @@ def image_generation( model = "dall-e-2" custom_llm_provider = "openai" # default to dall-e-2 on openai model_response._hidden_params["model"] = model - openai_params = [ + openai_params: Final = [ "user", "request_timeout", "api_base", @@ -255,9 +248,9 @@ def image_generation( "size", "style", ] - litellm_params = all_litellm_params - default_params = openai_params + litellm_params - non_default_params = { + litellm_params: Final = all_litellm_params + default_params: Final = openai_params + litellm_params + non_default_params: Final = { k: v for k, v in kwargs.items() if k not in default_params } # model-specific params - pass them straight to the model/provider @@ -268,7 +261,7 @@ def image_generation( provider=LlmProviders(custom_llm_provider), ) - optional_params = get_optional_params_image_gen( + optional_params: Final = get_optional_params_image_gen( model=base_model or model, n=n, quality=quality, @@ -281,9 +274,9 @@ def image_generation( **non_default_params, ) - litellm_params_dict = get_litellm_params(**kwargs) + litellm_params_dict: Final = get_litellm_params(**kwargs) - logging: Logging = litellm_logging_obj + logging: Final[Logging] = litellm_logging_obj logging.update_from_kwargs( kwargs=kwargs, model=model, @@ -308,7 +301,7 @@ def image_generation( if custom_llm_provider == "azure": # azure configs - api_type = get_secret_str("AZURE_API_TYPE") or "azure" + api_type: Final = get_secret_str("AZURE_API_TYPE") or "azure" api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") @@ -322,7 +315,7 @@ def image_generation( or get_secret_str("AZURE_API_KEY") ) - azure_ad_token = optional_params.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN") + azure_ad_token: Final = optional_params.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN") # Create azure_ad_token_provider from tenant_id, client_id, client_secret if not already provided if azure_ad_token_provider is None: @@ -331,9 +324,9 @@ def image_generation( ) # Extract Azure AD credentials from litellm_params - tenant_id = litellm_params_dict.get("tenant_id") - client_id = litellm_params_dict.get("client_id") - client_secret = litellm_params_dict.get("client_secret") + tenant_id: Final = litellm_params_dict.get("tenant_id") + client_id: Final = litellm_params_dict.get("client_id") + client_secret: Final = litellm_params_dict.get("client_secret") azure_scope = litellm_params_dict.get("azure_scope") or "https://cognitiveservices.azure.com/.default" # Create token provider if credentials are available @@ -392,7 +385,7 @@ def image_generation( raise ValueError(f"image generation config is not supported for {custom_llm_provider}") # Resolve api_base from litellm.api_base if not explicitly provided - _api_base = api_base or litellm.api_base + _api_base: Final = api_base or litellm.api_base litellm_params_dict["api_base"] = _api_base return llm_http_handler.image_generation_handler( @@ -468,7 +461,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: str | None = kwargs.get("organization", None) + organization: Final[str | None] = kwargs.get("organization", None) model_response = openai_chat_completions.image_generation( model=model, prompt=prompt, @@ -568,18 +561,18 @@ async def aimage_variation(*args, **kwargs) -> ImageResponse: Returns: - `response` (Any): The response returned by the `image_variation` function. """ - loop = asyncio.get_event_loop() - model = kwargs.get("model", None) + loop: Final = asyncio.get_event_loop() + model: Final = kwargs.get("model", None) custom_llm_provider = kwargs.get("custom_llm_provider", None) ### PASS ARGS TO Image Generation ### kwargs["async_call"] = True try: # Use a partial function to pass your keyword arguments - func = partial(image_variation, *args, **kwargs) + func: Final = partial(image_variation, *args, **kwargs) # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) if custom_llm_provider is None and model is not None: _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) @@ -618,12 +611,12 @@ def image_variation( **kwargs, ) -> ImageResponse: # get non-default params - client = kwargs.get("client", None) + client: Final = kwargs.get("client", None) # get logging object - litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")) + litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")) # get the litellm params - litellm_params = get_litellm_params(**kwargs) + litellm_params: Final = get_litellm_params(**kwargs) # get the custom llm provider model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( model=model, @@ -634,17 +627,17 @@ def image_variation( # route to the correct provider w/ the params try: - llm_provider = LlmProviders(custom_llm_provider) - image_variation_provider = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider) + llm_provider: Final = LlmProviders(custom_llm_provider) + image_variation_provider: Final = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider) except ValueError: raise ValueError( f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" ) - model_response = ImageResponse() + model_response: Final = ImageResponse() response: ImageResponse | None = None - provider_config = ProviderConfigManager.get_provider_model_info( + provider_config: Final = ProviderConfigManager.get_provider_model_info( model=model or "", # openai defaults to dall-e-2 provider=llm_provider, ) @@ -654,7 +647,7 @@ def image_variation( f"image variation provider has no known model info config - required for getting api keys, etc.: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" ) - api_key = provider_config.get_api_key(litellm_params.get("api_key", None)) + api_key: Final = provider_config.get_api_key(litellm_params.get("api_key", None)) api_base = provider_config.get_api_base(litellm_params.get("api_base", None)) if image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.OPENAI: @@ -727,9 +720,9 @@ def image_edit( """ Maps the image edit functionality, similar to OpenAI's images/edits endpoint. """ - local_vars = locals() + local_vars: Final = locals() try: - openai_params = [ + openai_params: Final = [ "user", "request_timeout", "api_base", @@ -747,22 +740,22 @@ def image_edit( "style", "async_call", ] - litellm_params_list = all_litellm_params - default_params = openai_params + litellm_params_list - non_default_params = { + litellm_params_list: Final = all_litellm_params + default_params: Final = openai_params + litellm_params_list + non_default_params: Final = { 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: 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 + litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) + model_info: Final = kwargs.get("model_info", None) + metadata: Final = kwargs.get("metadata", {}) + _is_async: Final = kwargs.pop("async_call", False) is True # add images / or return a single image - images = image if isinstance(image, list) else ([image] if image is not None else []) + images: Final = 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] = {} + headers_from_kwargs: Final = kwargs.get("headers") + merged_extra_headers: Final[dict[str, Any]] = {} if isinstance(headers_from_kwargs, dict): merged_extra_headers.update(headers_from_kwargs) if isinstance(extra_headers, dict): @@ -772,7 +765,7 @@ def image_edit( extra_headers = dict(merged_extra_headers) # get llm provider logic - litellm_params = GenericLiteLLMParams(**kwargs) + litellm_params: Final = GenericLiteLLMParams(**kwargs) model, custom_llm_provider, _, _ = get_llm_provider( model=model or DEFAULT_IMAGE_ENDPOINT_MODEL, custom_llm_provider=custom_llm_provider, @@ -788,7 +781,7 @@ def image_edit( if custom_handler is None: raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider) - model_response = ImageResponse() + model_response: Final = ImageResponse() if _is_async: async_custom_client: AsyncHTTPHandler | None = None @@ -836,11 +829,11 @@ def image_edit( local_vars.update(kwargs) # Get ImageEditOptionalRequestParams with only valid parameters - image_edit_optional_params: ImageEditOptionalRequestParams = ( + image_edit_optional_params: Final[ImageEditOptionalRequestParams] = ( _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: Final[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, @@ -973,9 +966,9 @@ async def aimage_edit( Returns: - `response` (Any): The response returned by the `image_edit` function. """ - local_vars = locals() + local_vars: Final = locals() try: - loop = asyncio.get_event_loop() + loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True # get custom llm provider so we can use this for mapping exceptions @@ -984,9 +977,9 @@ async def aimage_edit( model=model, api_base=local_vars.get("base_url", None) ) - images = image if isinstance(image, list) else [image] + images: Final = image if isinstance(image, list) else [image] - func = partial( + func: Final = partial( image_edit, image=images, prompt=prompt, @@ -1002,9 +995,9 @@ async def aimage_edit( **kwargs, ) - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) + ctx: Final = contextvars.copy_context() + func_with_context: Final = partial(ctx.run, func) + init_response: Final = await loop.run_in_executor(None, func_with_context) if asyncio.iscoroutine(init_response): response = await init_response @@ -1029,7 +1022,7 @@ def __getattr__(name: str) -> Any: from .utils import ImageEditRequestUtils as _ImageEditRequestUtils # Cache it in the module's __dict__ for subsequent accesses - module = importlib.import_module(__name__) + module: Final = importlib.import_module(__name__) module.__dict__["ImageEditRequestUtils"] = _ImageEditRequestUtils return _ImageEditRequestUtils raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/litellm/images/utils.py b/litellm/images/utils.py index e906cb5094b..2f080d88de4 100644 --- a/litellm/images/utils.py +++ b/litellm/images/utils.py @@ -1,5 +1,5 @@ from io import BufferedReader, BytesIO -from typing import Any, cast, get_type_hints +from typing import Any, Final, cast, get_type_hints import litellm from litellm.litellm_core_utils.token_counter import get_image_type @@ -30,16 +30,16 @@ class ImageEditRequestUtils: Returns: A dictionary of supported parameters for the image edit API """ - supported_params = image_edit_provider_config.get_supported_openai_params(model) + supported_params: Final = image_edit_provider_config.get_supported_openai_params(model) - should_drop = litellm.drop_params is True or drop_params is True + should_drop: Final = litellm.drop_params is True or drop_params is True - filtered_optional_params = dict(image_edit_optional_params) + filtered_optional_params: Final = dict(image_edit_optional_params) if additional_drop_params: for param in additional_drop_params: filtered_optional_params.pop(param, None) - unsupported_params = [param for param in filtered_optional_params if param not in supported_params] + unsupported_params: Final = [param for param in filtered_optional_params if param not in supported_params] if unsupported_params: if should_drop: @@ -51,7 +51,7 @@ class ImageEditRequestUtils: message=f"The following parameters are not supported for model {model}: {', '.join(unsupported_params)}", ) - mapped_params = image_edit_provider_config.map_openai_params( + mapped_params: Final = image_edit_provider_config.map_openai_params( image_edit_optional_params=cast(ImageEditOptionalRequestParams, filtered_optional_params), model=model, drop_params=should_drop, @@ -72,8 +72,8 @@ class ImageEditRequestUtils: Returns: ImageEditOptionalRequestParams instance with only the valid parameters """ - valid_keys = get_type_hints(ImageEditOptionalRequestParams).keys() - filtered_params = {k: v for k, v in params.items() if k in valid_keys and v is not None} + valid_keys: Final = get_type_hints(ImageEditOptionalRequestParams).keys() + filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None} return cast(ImageEditOptionalRequestParams, filtered_params) @staticmethod @@ -118,13 +118,13 @@ class ImageEditRequestUtils: return FILE_MIME_TYPES[FileType.PNG] # Default fallback # Use the existing get_image_type function to detect image type - image_type_str = get_image_type(bytes_data) + image_type_str: Final = get_image_type(bytes_data) if image_type_str is None: return FILE_MIME_TYPES[FileType.PNG] # Default if detection fails # Map detected type string to FileType enum and get MIME type - type_mapping = { + type_mapping: Final = { "png": FileType.PNG, "jpeg": FileType.JPEG, "gif": FileType.GIF, @@ -132,7 +132,7 @@ class ImageEditRequestUtils: "heic": FileType.HEIC, } - file_type = type_mapping.get(image_type_str) + file_type: Final = type_mapping.get(image_type_str) if file_type is None: return FILE_MIME_TYPES[FileType.PNG] # Default to PNG if unknown diff --git a/litellm/integrations/SlackAlerting/batching_handler.py b/litellm/integrations/SlackAlerting/batching_handler.py index 12b1f772616..a7febdadacd 100644 --- a/litellm/integrations/SlackAlerting/batching_handler.py +++ b/litellm/integrations/SlackAlerting/batching_handler.py @@ -6,7 +6,7 @@ Slack alerts are sent every 10s or when events are greater than X events see custom_batch_logger.py for more details / defaults """ -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_proxy_logger @@ -19,7 +19,7 @@ else: def squash_payloads(queue): - squashed = {} + squashed: Final = {} if len(queue) == 0: return squashed if len(queue) == 1: @@ -57,12 +57,12 @@ async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item, count) """ import json - payload = item.get("payload", {}) + payload: Final = item.get("payload", {}) try: if count > 1: payload["text"] = f"[Num Alerts: {count}]\n\n{payload['text']}" - response = await slackAlertingInstance.async_http_handler.post( + response: Final = await slackAlertingInstance.async_http_handler.post( url=item["url"], headers=item["headers"], data=json.dumps(payload), diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index 50700774ea6..f35ff7b5f82 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import Literal +from typing import Final, Literal from litellm.proxy._types import CallInfo, Litellm_EntityType @@ -97,7 +97,7 @@ def get_budget_alert_type( ) -> BaseBudgetAlertType: """Factory function to get the appropriate budget alert type class""" - alert_types = { + alert_types: Final = { "proxy_budget": ProxyBudgetAlert(), "soft_budget": SoftBudgetAlert(), "user_budget": UserBudgetAlert(), diff --git a/litellm/integrations/SlackAlerting/hanging_request_check.py b/litellm/integrations/SlackAlerting/hanging_request_check.py index 55dff2fde1f..4d7cbfe8fd1 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 +from typing import TYPE_CHECKING, Any, Final import litellm from litellm._logging import verbose_proxy_logger @@ -57,8 +57,8 @@ class AlertingHangingRequestCheck: if request_data is None: return - request_metadata = get_litellm_metadata_from_kwargs(kwargs=request_data) - model = request_data.get("model", "") + request_metadata: Final = get_litellm_metadata_from_kwargs(kwargs=request_data) + model: Final = request_data.get("model", "") api_base: str | None = None if request_data.get("deployment", None) is not None and isinstance(request_data["deployment"], dict): @@ -67,7 +67,7 @@ class AlertingHangingRequestCheck: optional_params=request_data["deployment"].get("litellm_params", {}), ) - hanging_request_data = HangingRequestData( + hanging_request_data: Final = HangingRequestData( request_id=request_data.get("litellm_call_id", ""), model=model, api_base=api_base, @@ -96,7 +96,7 @@ class AlertingHangingRequestCheck: if proxy_logging_obj.internal_usage_cache is None: return - hanging_requests = await self.hanging_request_cache.async_get_oldest_n_keys( + hanging_requests: Final = await self.hanging_request_cache.async_get_oldest_n_keys( n=MAX_OLDEST_HANGING_REQUESTS_TO_CHECK, ) @@ -166,7 +166,7 @@ class AlertingHangingRequestCheck: ################ # Send the Alert on Slack ################ - request_info = f"""Request Model: `{hanging_request_data.model}` + request_info: Final = f"""Request Model: `{hanging_request_data.model}` API Base: `{hanging_request_data.api_base}` Key Alias: `{hanging_request_data.key_alias}` Team Alias: `{hanging_request_data.team_alias}`""" diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 0d842a5889a..3e81e7fa92b 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, Literal +from typing import TYPE_CHECKING, Any, Final, Literal from openai import APIError @@ -128,7 +128,7 @@ class SlackAlerting(CustomBatchLogger): if self.alert_to_webhook_url is None: self.alert_to_webhook_url = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url) else: - _new_values = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url) or {} + _new_values: Final = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url) or {} self.alert_to_webhook_url.update(_new_values) if llm_router is not None: self.llm_router = llm_router @@ -139,7 +139,7 @@ class SlackAlerting(CustomBatchLogger): Converts set objects to lists for JSON serialization. """ # Convert to dict for processing - cache_value = dict(outage_value) + cache_value: Final = dict(outage_value) if "deployment_ids" in cache_value and isinstance(cache_value["deployment_ids"], set): cache_value["deployment_ids"] = list(cache_value["deployment_ids"]) @@ -173,19 +173,19 @@ class SlackAlerting(CustomBatchLogger): end_time, # start/end time ): try: - time_difference = end_time - start_time + time_difference: Final = end_time - start_time # Convert the timedelta to float (in seconds) - time_difference_float = time_difference.total_seconds() - litellm_params = kwargs.get("litellm_params", {}) - model = kwargs.get("model", "") - api_base = litellm.get_api_base(model=model, optional_params=litellm_params) + time_difference_float: Final = time_difference.total_seconds() + litellm_params: Final = kwargs.get("litellm_params", {}) + model: Final = kwargs.get("model", "") + api_base: Final = litellm.get_api_base(model=model, optional_params=litellm_params) messages = kwargs.get("messages", None) # if messages does not exist fallback to "input" if messages is None: messages = kwargs.get("input", None) # only use first 100 chars for alerting - _messages = str(messages)[:100] + _messages: Final = str(messages)[:100] return time_difference_float, model, api_base, _messages except Exception as e: @@ -251,10 +251,10 @@ class SlackAlerting(CustomBatchLogger): if time_difference_float > self.alerting_threshold: # add deployment latencies to alert if kwargs is not None and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"]: - _metadata: dict = kwargs["litellm_params"]["metadata"] + _metadata: Final[dict] = kwargs["litellm_params"]["metadata"] request_info = _add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata) - _deployment_latency_map = self._get_deployment_latencies_to_alert(metadata=_metadata) + _deployment_latency_map: Final = self._get_deployment_latencies_to_alert(metadata=_metadata) if _deployment_latency_map is not None: request_info += f"\nAvailable Deployment Latencies\n{_deployment_latency_map}" @@ -324,15 +324,15 @@ class SlackAlerting(CustomBatchLogger): False -> if not sent """ - ids = router.get_model_ids() + ids: Final = router.get_model_ids() # get keys - 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] + failed_request_keys: Final = [f"{id}:{SlackAlertingCacheKeys.failed_requests_key.value}" for id in ids] + latency_keys: Final = [f"{id}:{SlackAlertingCacheKeys.latency_key.value}" for id in ids] - combined_metrics_keys = failed_request_keys + latency_keys # reduce cache calls + combined_metrics_keys: Final = failed_request_keys + latency_keys # reduce cache calls - combined_metrics_values = await self.internal_usage_cache.async_batch_get_cache( + combined_metrics_values: Final = await self.internal_usage_cache.async_batch_get_cache( keys=combined_metrics_keys ) # [1, 2, None, ..] @@ -348,8 +348,8 @@ class SlackAlerting(CustomBatchLogger): if all_none: return False - failed_request_values = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..] - latency_values = combined_metrics_values[len(failed_request_keys) :] + failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..] + latency_values: Final = combined_metrics_values[len(failed_request_keys) :] # find top 5 failed ## Replace None values with a placeholder value (-1 in this case) @@ -367,7 +367,7 @@ class SlackAlerting(CustomBatchLogger): # find top 5 slowest # Replace None values with a placeholder value (-1 in this case) placeholder_value = 0 - replaced_slowest_values = [value if value is not None else placeholder_value for value in latency_values] + replaced_slowest_values: Final = [value if value is not None else placeholder_value for value in latency_values] # Get the indices of top 5 values with the highest numerical values (ignoring None and 0 values) top_5_slowest = sorted( @@ -420,9 +420,9 @@ class SlackAlerting(CustomBatchLogger): message += f"\t{i + 1}. Deployment: `{deployment_name}`, Latency per output token: `{value}s/token`, API Base: `{api_base}`\n\n" # cache cleanup -> reset values to 0 - latency_cache_keys = [(key, 0) for key in latency_keys] - failed_request_cache_keys = [(key, 0) for key in failed_request_keys] - combined_metrics_cache_keys = latency_cache_keys + failed_request_cache_keys + latency_cache_keys: Final = [(key, 0) for key in latency_keys] + failed_request_cache_keys: Final = [(key, 0) for key in failed_request_keys] + combined_metrics_cache_keys: Final = latency_cache_keys + failed_request_cache_keys await self.internal_usage_cache.async_set_cache_pipeline(cache_list=combined_metrics_cache_keys) message += f"\n\nNext Run is at: `{time.time() + self.alerting_args.daily_report_frequency}`s" @@ -463,10 +463,10 @@ class SlackAlerting(CustomBatchLogger): if "failed_tracking_spend" not in self.alert_types: return - _cache: DualCache = self.internal_usage_cache - message = "Failed Tracking Cost for " + error_message - _cache_key = f"budget_alerts:failed_tracking:{failing_model}" - result = await _cache.async_get_cache(key=_cache_key) + _cache: Final[DualCache] = self.internal_usage_cache + message: Final = "Failed Tracking Cost for " + error_message + _cache_key: Final = f"budget_alerts:failed_tracking:{failing_model}" + result: Final = await _cache.async_get_cache(key=_cache_key) if result is None: await self.send_alert( message=message, @@ -506,7 +506,7 @@ class SlackAlerting(CustomBatchLogger): # - Alert once within 24hr period # - Cache this information # - Don't re-alert, if alert already sent - _cache: DualCache = self.internal_usage_cache + _cache: Final[DualCache] = self.internal_usage_cache if self.alerting is None or self.alert_types is None: # do nothing if alerting is not switched on @@ -515,10 +515,10 @@ class SlackAlerting(CustomBatchLogger): return # Get the appropriate budget alert type handler - budget_alert_class = get_budget_alert_type(type) - _id = budget_alert_class.get_id(user_info) - user_info_json = user_info.model_dump(exclude_none=True) - user_info_str = self._get_user_info_str(user_info) + budget_alert_class: Final = get_budget_alert_type(type) + _id: Final = budget_alert_class.get_id(user_info) + user_info_json: Final = user_info.model_dump(exclude_none=True) + user_info_str: Final = self._get_user_info_str(user_info) event_message = budget_alert_class.get_event_message() # Set default event unless we're in projected_limit_exceeded @@ -541,8 +541,8 @@ class SlackAlerting(CustomBatchLogger): # send alert if event is not None and user_info.event_group is not None: - _cache_key = f"budget_alerts:{event}:{_id}" - result = await _cache.async_get_cache(key=_cache_key) + _cache_key: Final = f"budget_alerts:{event}:{_id}" + result: Final = await _cache.async_get_cache(key=_cache_key) if result is None: webhook_event = WebhookEvent( event=event, @@ -581,7 +581,7 @@ class SlackAlerting(CustomBatchLogger): Handles Max Budget and Soft Budget Alerts """ - percent_left: float = self._get_percent_of_max_budget_left(user_info=user_info) + percent_left: Final[float] = self._get_percent_of_max_budget_left(user_info=user_info) ##################################################################### # SOFT BUDGET CHECK @@ -616,8 +616,8 @@ class SlackAlerting(CustomBatchLogger): Get the percent of the max budget that is left """ percent_left: float = 0.0 - current_spend: float = user_info.spend - max_budget: float | None = user_info.max_budget + current_spend: Final[float] = user_info.spend + max_budget: Final[float | None] = user_info.max_budget if max_budget is None: return percent_left if max_budget <= 0: @@ -629,7 +629,7 @@ class SlackAlerting(CustomBatchLogger): """ Create a standard message for a budget alert """ - _all_fields_as_dict = user_info.model_dump(exclude_none=True) + _all_fields_as_dict: Final = user_info.model_dump(exclude_none=True) _all_fields_as_dict.pop("token") msg = "" for k, v in _all_fields_as_dict.items(): @@ -655,7 +655,7 @@ class SlackAlerting(CustomBatchLogger): and response_cost is not None ): # log customer spend - event = WebhookEvent( + event: Final = WebhookEvent( spend=response_cost, max_budget=max_budget, token=token, @@ -681,7 +681,7 @@ class SlackAlerting(CustomBatchLogger): Returns: - str -> formatted string. This is an alert message, giving a human-friendly description of the errors. """ - error_breakdown = {"Timeout Errors": 0, "API Errors": 0, "Unknown Errors": 0} + error_breakdown: Final = {"Timeout Errors": 0, "API Errors": 0, "Unknown Errors": 0} for alert in alerts: if alert == 408: error_breakdown["Timeout Errors"] += 1 @@ -707,7 +707,7 @@ class SlackAlerting(CustomBatchLogger): outage_value: BaseOutageModel, ) -> str: """Format an alert message for slack""" - headers = {f"{key} Name": key_val, "Provider": provider} + headers: Final = {f"{key} Name": key_val, "Provider": provider} if api_base is not None: headers["API Base"] = api_base # type: ignore @@ -739,7 +739,7 @@ class SlackAlerting(CustomBatchLogger): if self.llm_router is None: return - deployment = self.llm_router.get_deployment(model_id=deployment_id) + deployment: Final = self.llm_router.get_deployment(model_id=deployment_id) if deployment is None: return @@ -761,7 +761,7 @@ class SlackAlerting(CustomBatchLogger): return ### UNIQUE CACHE KEY ### - cache_key = provider + region_name + cache_key: Final = provider + region_name outage_value: ProviderRegionOutageModel | None = await self.internal_usage_cache.async_get_cache(key=cache_key) @@ -896,7 +896,7 @@ class SlackAlerting(CustomBatchLogger): return ### EXTRACT MODEL DETAILS ### - deployment = self.llm_router.get_deployment(model_id=deployment_id) + deployment: Final = self.llm_router.get_deployment(model_id=deployment_id) if deployment is None: return @@ -907,7 +907,7 @@ class SlackAlerting(CustomBatchLogger): model, provider, _, _ = litellm.get_llm_provider(model=model) except Exception: provider = "" - api_base = litellm.get_api_base(model=model, optional_params=deployment.litellm_params) + api_base: Final = litellm.get_api_base(model=model, optional_params=deployment.litellm_params) if outage_value is None: outage_value = OutageModel( @@ -979,13 +979,13 @@ class SlackAlerting(CustomBatchLogger): ## update cache ## # Convert set to list for JSON serialization - cache_value = self._prepare_outage_value_for_cache(outage_value) + cache_value: Final = self._prepare_outage_value_for_cache(outage_value) await self.internal_usage_cache.async_set_cache(key=deployment_id, value=cache_value) except Exception: pass async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: Any): - base_model_from_user = getattr(passed_model_info, "base_model", None) + base_model_from_user: Final = getattr(passed_model_info, "base_model", None) model_info = {} base_model = "" if base_model_from_user is not None: @@ -1001,7 +1001,7 @@ class SlackAlerting(CustomBatchLogger): model_info_str += f"{k}: {v}\n" - message = f""" + message: Final = f""" *🚅 New Model Added* Model Name: `{model_name}` {base_model} @@ -1031,7 +1031,7 @@ Model Info: ``` """ - alert_val = self.send_alert( + alert_val: Final = self.send_alert( message=message, level="Low", alert_type=AlertType.new_model_added, @@ -1056,14 +1056,14 @@ Model Info: - if WEBHOOK_URL is not set """ - webhook_url = os.getenv("WEBHOOK_URL", None) + webhook_url: Final = os.getenv("WEBHOOK_URL", None) if webhook_url is None: raise Exception("Missing webhook_url from environment") - payload = webhook_event.model_dump_json() - headers = {"Content-type": "application/json"} + payload: Final = webhook_event.model_dump_json() + headers: Final = {"Content-type": "application/json"} - response = await self.async_http_handler.post( + response: Final = await self.async_http_handler.post( url=webhook_url, headers=headers, data=payload, @@ -1108,18 +1108,18 @@ Model Info: if email_support_contact is None: email_support_contact = LITELLM_SUPPORT_CONTACT - event_name = webhook_event.event_message + event_name: Final = webhook_event.event_message recipient_email = webhook_event.user_email - recipient_user_id = webhook_event.user_id + recipient_user_id: Final = webhook_event.user_id if recipient_email is None and recipient_user_id is not None and prisma_client is not None: user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": recipient_user_id}) if user_row is not None: recipient_email = user_row.user_email - key_token = webhook_event.token - key_budget = webhook_event.max_budget - base_url = os.getenv("PROXY_BASE_URL", "http://0.0.0.0:4000") + key_token: Final = webhook_event.token + key_budget: Final = webhook_event.max_budget + base_url: Final = os.getenv("PROXY_BASE_URL", "http://0.0.0.0:4000") email_html_content = "Alert from LiteLLM Server" if recipient_email is None: @@ -1139,10 +1139,10 @@ Model Info: ) elif webhook_event.event == "internal_user_created": # GET TEAM NAME - team_id = webhook_event.team_id + team_id: Final = webhook_event.team_id team_name = "Default Team" if team_id is not None and prisma_client is not None: - team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) + team_row: Final = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if team_row is not None: team_name = team_row.team_alias or "-" email_html_content = USER_INVITED_EMAIL_TEMPLATE.format( @@ -1159,7 +1159,7 @@ Model Info: ) webhook_event.model_dump_json() - email_event = { + email_event: Final = { "to": recipient_email, "subject": f"LiteLLM: {event_name}", "html": email_html_content, @@ -1197,10 +1197,10 @@ Model Info: if email_support_contact is None: email_support_contact = LITELLM_SUPPORT_CONTACT - event_name = webhook_event.event_message - recipient_email = webhook_event.user_email - user_name = webhook_event.user_id - max_budget = webhook_event.max_budget + event_name: Final = webhook_event.event_message + recipient_email: Final = webhook_event.user_email + user_name: Final = webhook_event.user_id + max_budget: Final = webhook_event.max_budget email_html_content = "Alert from LiteLLM Server" if recipient_email is None: verbose_proxy_logger.error("Trying to send email alert to no recipient", extra=webhook_event.dict()) @@ -1222,7 +1222,7 @@ Model Info: """ webhook_event.model_dump_json() - email_event = { + email_event: Final = { "to": recipient_email, "subject": f"LiteLLM: {event_name}", "html": email_html_content, @@ -1290,8 +1290,8 @@ Model Info: from datetime import datetime # Check if digest mode is enabled for this alert type - alert_type_name_str = getattr(alert_type, "value", str(alert_type)) - _atc = self.alert_type_config.get(alert_type_name_str) + alert_type_name_str: Final = getattr(alert_type, "value", str(alert_type)) + _atc: Final = self.alert_type_config.get(alert_type_name_str) 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: @@ -1303,10 +1303,10 @@ Model Info: if _digest_webhook is None: raise ValueError("Missing SLACK_WEBHOOK_URL from environment") - digest_key = f"{alert_type_name_str}:{request_model or ''}:{api_base or ''}" + digest_key: Final = f"{alert_type_name_str}:{request_model or ''}:{api_base or ''}" async with self.digest_lock: - now = datetime.now() + now: Final = datetime.now() if digest_key in self.digest_buckets: self.digest_buckets[digest_key]["count"] += 1 self.digest_buckets[digest_key]["last_time"] = now @@ -1325,11 +1325,11 @@ Model Info: return # Suppress immediate alert; will be emitted by _flush_digest_buckets # Get the current timestamp - current_time = datetime.now().strftime("%H:%M:%S") - _proxy_base_url = os.getenv("PROXY_BASE_URL", None) + current_time: Final = datetime.now().strftime("%H:%M:%S") + _proxy_base_url: Final = os.getenv("PROXY_BASE_URL", None) # Use .name if it's an enum, otherwise use as is - alert_type_name = getattr(alert_type, "name", alert_type) - alert_type_formatted = f"Alert type: `{alert_type_name}`" + alert_type_name: Final = getattr(alert_type, "name", alert_type) + alert_type_formatted: Final = f"Alert type: `{alert_type_name}`" if alert_type == "daily_reports" or alert_type == "new_model_added": formatted_message = alert_type_formatted + message else: @@ -1356,8 +1356,8 @@ Model Info: if slack_webhook_url is None: raise ValueError("Missing SLACK_WEBHOOK_URL from environment") - payload = {"text": formatted_message} - headers = {"Content-type": "application/json"} + payload: Final = {"text": formatted_message} + headers: Final = {"Content-type": "application/json"} if isinstance(slack_webhook_url, list): for url in slack_webhook_url: @@ -1386,8 +1386,8 @@ Model Info: if not self.log_queue: return - squashed_queue = squash_payloads(self.log_queue) - tasks = [ + squashed_queue: Final = squash_payloads(self.log_queue) + tasks: Final = [ send_to_webhook(slackAlertingInstance=self, item=item["item"], count=item["count"]) for item in squashed_queue.values() ] @@ -1402,8 +1402,8 @@ Model Info: """ from datetime import datetime - now = datetime.now() - flushed_keys: list[str] = [] + now: Final = datetime.now() + flushed_keys: Final[list[str]] = [] async with self.digest_lock: for key, entry in self.digest_buckets.items(): @@ -1474,10 +1474,10 @@ Model Info: """Log deployment latency""" try: if "daily_reports" in self.alert_types: - litellm_params = kwargs.get("litellm_params", {}) or {} - model_info = litellm_params.get("model_info", {}) or {} - model_id = model_info.get("id", "") or "" - response_s: timedelta = end_time - start_time + litellm_params: Final = kwargs.get("litellm_params", {}) or {} + model_info: Final = litellm_params.get("model_info", {}) or {} + model_id: Final = model_info.get("id", "") or "" + response_s: Final[timedelta] = end_time - start_time final_value = response_s @@ -1486,7 +1486,7 @@ Model Info: and response_obj.usage is not None # type: ignore and hasattr(response_obj.usage, "completion_tokens") # type: ignore ): - completion_tokens = response_obj.usage.completion_tokens # type: ignore + completion_tokens: Final = response_obj.usage.completion_tokens # type: ignore if completion_tokens is not None and completion_tokens > 0: final_value = float(response_s.total_seconds() / completion_tokens) if isinstance(final_value, timedelta): @@ -1507,9 +1507,9 @@ Model Info: async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """Log failure + deployment latency""" - _litellm_params = kwargs.get("litellm_params", {}) - _model_info = _litellm_params.get("model_info", {}) or {} - model_id = _model_info.get("id", "") + _litellm_params: Final = kwargs.get("litellm_params", {}) + _model_info: Final = _litellm_params.get("model_info", {}) or {} + model_id: Final = _model_info.get("id", "") try: if "daily_reports" in self.alert_types: try: @@ -1544,12 +1544,12 @@ Model Info: """ report_sent_bool = False - report_sent = await self.internal_usage_cache.async_get_cache( + report_sent: Final = await self.internal_usage_cache.async_get_cache( key=SlackAlertingCacheKeys.report_sent_key.value, parent_otel_span=None, ) # None | float - current_time = time.time() + current_time: Final = time.time() if report_sent is None: await self.internal_usage_cache.async_set_cache( @@ -1558,7 +1558,7 @@ Model Info: ) elif isinstance(report_sent, float): # Check if current time - interval >= time last sent - interval_seconds = self.alerting_args.daily_report_frequency + interval_seconds: Final = self.alerting_args.daily_report_frequency if current_time - report_sent >= interval_seconds: # Sneak in the reporting logic here @@ -1612,20 +1612,20 @@ Model Info: ) # Parse the time range - days = int(time_range[:-1]) + days: Final = int(time_range[:-1]) if time_range[-1].lower() != "d": raise ValueError("Time range must be specified in days, e.g., '7d'") - todays_date = datetime.datetime.now().date() - start_date = todays_date - datetime.timedelta(days=days) + todays_date: Final = datetime.datetime.now().date() + start_date: Final = todays_date - datetime.timedelta(days=days) - _event_cache_key = ( + _event_cache_key: Final = ( f"weekly_spend_report_sent_{start_date.strftime('%Y-%m-%d')}_{todays_date.strftime('%Y-%m-%d')}" ) if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp = await _get_spend_report_for_time_range( + _resp: Final = await _get_spend_report_for_time_range( start_date=start_date.strftime("%Y-%m-%d"), end_date=todays_date.strftime("%Y-%m-%d"), ) @@ -1675,8 +1675,8 @@ Model Info: _get_spend_report_for_time_range, ) - todays_date = datetime.datetime.now().date() - first_day_of_month = todays_date.replace(day=1) + todays_date: Final = datetime.datetime.now().date() + first_day_of_month: Final = todays_date.replace(day=1) _, last_day_of_month = monthrange(todays_date.year, todays_date.month) last_day_of_month = first_day_of_month + datetime.timedelta(days=last_day_of_month - 1) @@ -1684,7 +1684,7 @@ Model Info: if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp = await _get_spend_report_for_time_range( + _resp: Final = await _get_spend_report_for_time_range( start_date=first_day_of_month.strftime("%Y-%m-%d"), end_date=last_day_of_month.strftime("%Y-%m-%d"), ) @@ -1742,9 +1742,9 @@ Model Info: ) # call prometheuslogger. - falllback_success_info_prometheus = await get_fallback_metric_from_prometheus() + falllback_success_info_prometheus: Final = await get_fallback_metric_from_prometheus() - fallback_message = f"*Fallback Statistics:*\n{falllback_success_info_prometheus}" + fallback_message: Final = f"*Fallback Statistics:*\n{falllback_success_info_prometheus}" await self.send_alert( message=fallback_message, @@ -1773,7 +1773,7 @@ Model Info: try: message = f"`{event_name}`\n" - key_event_dict = key_event.model_dump() + key_event_dict: Final = key_event.model_dump() # Add Created by information first message += "*Action Done by:*\n" @@ -1783,7 +1783,7 @@ Model Info: # Add args sent to function in the alert message += "\n*Arguments passed:*\n" - request_kwargs = key_event.request_kwargs + request_kwargs: Final = key_event.request_kwargs for key, value in request_kwargs.items(): if key == "user_api_key_dict": continue @@ -1808,8 +1808,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: str | None = await self.internal_usage_cache.async_get_cache( + litellm_call_id: Final = request_data.get("litellm_call_id", "") + status: Final[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"): diff --git a/litellm/integrations/SlackAlerting/utils.py b/litellm/integrations/SlackAlerting/utils.py index 9587e0ae78b..297d069a868 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 +from typing import TYPE_CHECKING, Any, Final import litellm from litellm.proxy._types import AlertType @@ -74,7 +74,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: str | None = None - litellm_logging_obj: Logging = request_data["litellm_logging_obj"] + litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"] for _ in range(3): trace_id = litellm_logging_obj._get_trace_id(service_name="langfuse") @@ -82,9 +82,9 @@ async def _add_langfuse_trace_id_to_alert( break await asyncio.sleep(3) # wait 3s before retrying for trace id ######################################################### - langfuse_object = litellm_logging_obj._get_callback_object(service_name="langfuse") + langfuse_object: Final = litellm_logging_obj._get_callback_object(service_name="langfuse") if langfuse_object is not None: - base_url = langfuse_object.Langfuse.base_url + base_url: Final = langfuse_object.Langfuse.base_url return f"{base_url}/trace/{trace_id}" return None diff --git a/litellm/integrations/agentops/agentops.py b/litellm/integrations/agentops/agentops.py index 5295d8bf2be..399ee49238c 100644 --- a/litellm/integrations/agentops/agentops.py +++ b/litellm/integrations/agentops/agentops.py @@ -4,7 +4,7 @@ AgentOps integration for LiteLLM - Provides OpenTelemetry tracing for LLM calls import os from dataclasses import dataclass -from typing import Any +from typing import Any, Final from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig from litellm.llms.custom_httpx.http_handler import _get_httpx_client @@ -58,21 +58,21 @@ class AgentOps(OpenTelemetry): project_id = None if config.api_key: try: - response = self._fetch_auth_token(config.api_key, config.auth_endpoint) + response: Final = self._fetch_auth_token(config.api_key, config.auth_endpoint) jwt_token = response.get("token") project_id = response.get("project_id") except Exception: pass - headers = f"Authorization=Bearer {jwt_token}" if jwt_token else None + headers: Final = f"Authorization=Bearer {jwt_token}" if jwt_token else None - otel_config = OpenTelemetryConfig(exporter="otlp_http", endpoint=config.endpoint, headers=headers) + otel_config: Final = OpenTelemetryConfig(exporter="otlp_http", endpoint=config.endpoint, headers=headers) # Initialize OpenTelemetry with our config super().__init__(config=otel_config, callback_name="agentops") # Set AgentOps-specific resource attributes - resource_attrs = { + resource_attrs: Final = { "service.name": config.service_name or "litellm", "deployment.environment": config.deployment_environment or "production", "telemetry.sdk.name": "agentops", @@ -94,14 +94,14 @@ class AgentOps(OpenTelemetry): Returns: Dict containing JWT token and project ID """ - headers = { + headers: Final = { "Content-Type": "application/json", "Connection": "keep-alive", } - client = _get_httpx_client() + client: Final = _get_httpx_client() try: - response = client.post( + response: Final = client.post( url=auth_endpoint, headers=headers, json={"api_key": api_key}, diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index f52b6bd8415..34b3c4dacde 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, cast +from typing import TYPE_CHECKING, Any, Final, cast from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -32,7 +32,7 @@ else: # Anthropic (and Bedrock Claude) reject requests with more than 4 cache_control # breakpoints: "A maximum of 4 blocks with cache_control may be provided." -MAX_CACHE_CONTROL_BLOCKS = 4 +MAX_CACHE_CONTROL_BLOCKS: Final = 4 class AnthropicCacheControlHook(CustomPromptManagement): @@ -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: Final[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: Final[list[CacheControlMessageInjectionPoint]] = [] + remaining_points: Final[list[CacheControlInjectionPoint]] = [] for point in injection_points: if point.get("location") == "message": message_points.append(cast(CacheControlMessageInjectionPoint, point)) @@ -81,7 +81,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): # provider transform, where each tool_config point appends at most one # cachePoint to the tools. That block also counts toward Anthropic's # limit, so reserve a slot for it here to leave room. - reserved_blocks = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0 + reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0 processed_messages = self._apply_message_injections( points=message_points, @@ -154,7 +154,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues] ) -> list[int]: """Resolve which message indices an injection point targets.""" - _targetted_index: int | str | None = point.get("index", None) + _targetted_index: Final[int | str | None] = point.get("index", None) targetted_index: int | None = None if isinstance(_targetted_index, str): try: @@ -166,7 +166,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): # Case 1: Target by specific index if targetted_index is not None: - original_index = targetted_index + original_index: Final = targetted_index if targetted_index < 0: targetted_index += len(messages) @@ -182,7 +182,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): return [] # Case 2: Target by role - targetted_role = point.get("role", None) + targetted_role: Final = point.get("role", None) if targetted_role is not None: return [idx for idx, msg in enumerate(messages) if msg.get("role") == targetted_role] @@ -194,7 +194,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): count = 0 if message.get("cache_control") is not None: count += 1 - content = message.get("content") + content: Final = message.get("content") if isinstance(content, list): for block in content: if isinstance(block, dict) and block.get("cache_control") is not None: @@ -221,7 +221,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): Per Anthropic's API specification, when using multiple content blocks, only the last content block can have cache_control. """ - message_content = message.get("content", None) + message_content: Final = message.get("content", None) # 1. if string, insert cache control in the message if isinstance(message_content, str): @@ -248,9 +248,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): 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: Final[list[CacheControlMessageInjectionPoint]] = [] + system_points: Final[list[CacheControlMessageInjectionPoint]] = [] + remaining_points: Final[list[CacheControlInjectionPoint]] = [] for point in injection_points: if point.get("location") == "message": @@ -262,8 +262,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): else: remaining_points.append(point) - reserved_blocks = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0 - max_blocks = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks + reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0 + max_blocks: Final = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks used_blocks = sum( AnthropicCacheControlHook._count_cache_control_blocks(cast(AllMessageValues, msg)) @@ -275,11 +275,11 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) if system_points and processed_system is not None and used_blocks < max_blocks: - system_already_has_cc = isinstance(processed_system, list) and any( + system_already_has_cc: Final = isinstance(processed_system, list) and any( isinstance(b, dict) and b.get("cache_control") is not None for b in processed_system ) if not system_already_has_cc: - control = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral") + control: Final = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral") if isinstance(processed_system, str): processed_system = [{"type": "text", "text": processed_system, "cache_control": control}] used_blocks += 1 @@ -309,7 +309,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ import litellm - ttl = litellm.anthropic_prompt_caching_ttl + ttl: Final = litellm.anthropic_prompt_caching_ttl if ttl == "5m" or ttl == "1h": return ChatCompletionCachedContent(type="ephemeral", ttl=ttl) return ChatCompletionCachedContent(type="ephemeral") @@ -419,8 +419,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools): return [] - control = AnthropicCacheControlHook._default_control() - points: list[CacheControlInjectionPoint] = [ + control: Final = AnthropicCacheControlHook._default_control() + points: Final[list[CacheControlInjectionPoint]] = [ CacheControlMessageInjectionPoint(location="message", role="system", index=None, control=control), CacheControlMessageInjectionPoint(location="message", role=None, index=-1, control=control), ] @@ -452,7 +452,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): ): non_default_params.pop("cache_control_injection_points") return - points = AnthropicCacheControlHook.get_default_injection_points( + points: Final = AnthropicCacheControlHook.get_default_injection_points( messages=messages, system=None, model=model, @@ -484,7 +484,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): downstream transforms can handle them. """ typed_messages = cast(list[AllMessageValues], messages) # cast-ok: Anthropic-shaped dicts from v1/messages - configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list + configured: Final = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None) ) if configured and AnthropicCacheControlHook._should_stand_down(configured, typed_messages, system, tools): diff --git a/litellm/integrations/argilla.py b/litellm/integrations/argilla.py index b1cda6a5593..76a63f75897 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 +from typing import Any, Final import httpx from pydantic import BaseModel # type: ignore @@ -29,7 +29,7 @@ from litellm.types.utils import StandardLoggingPayload def is_serializable(value): - non_serializable_types = ( + non_serializable_types: Final = ( types.CoroutineType, types.FunctionType, types.GeneratorType, @@ -62,7 +62,7 @@ class ArgillaLogger(CustomBatchLogger): ) self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - _batch_size = os.getenv("ARGILLA_BATCH_SIZE", None) or litellm.argilla_batch_size + _batch_size: Final = os.getenv("ARGILLA_BATCH_SIZE", None) or litellm.argilla_batch_size if _batch_size: self.batch_size = int(_batch_size) asyncio.create_task(self.periodic_flush()) @@ -85,11 +85,11 @@ class ArgillaLogger(CustomBatchLogger): argilla_dataset_name: str | None, argilla_base_url: str | None, ) -> ArgillaCredentialsObject: - _credentials_api_key = argilla_api_key or os.getenv("ARGILLA_API_KEY") + _credentials_api_key: Final = argilla_api_key or os.getenv("ARGILLA_API_KEY") if _credentials_api_key is None: raise Exception("Invalid Argilla API Key given. _credentials_api_key=None.") - _credentials_base_url = argilla_base_url or os.getenv("ARGILLA_BASE_URL") or "http://localhost:6900/" + _credentials_base_url: Final = argilla_base_url or os.getenv("ARGILLA_BASE_URL") or "http://localhost:6900/" if _credentials_base_url is None: raise Exception("Invalid Argilla Base URL given. _credentials_base_url=None.") @@ -97,11 +97,11 @@ class ArgillaLogger(CustomBatchLogger): if _credentials_dataset_name is None: raise Exception("Invalid Argilla Dataset give. Value=None.") else: - dataset_response = litellm.module_level_client.get( + dataset_response: Final = litellm.module_level_client.get( url=f"{_credentials_base_url}/api/v1/me/datasets?name={_credentials_dataset_name}", headers={"X-Argilla-Api-Key": _credentials_api_key}, ) - json_response = dataset_response.json() + json_response: Final = dataset_response.json() if ( "items" in json_response and isinstance(json_response["items"], list) @@ -116,7 +116,7 @@ class ArgillaLogger(CustomBatchLogger): ) def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, Any]]: - payload_messages = payload.get("messages", None) + payload_messages: Final = payload.get("messages", None) if payload_messages is None: raise Exception("No chat messages found in payload.") @@ -129,7 +129,7 @@ class ArgillaLogger(CustomBatchLogger): raise Exception(f"Invalid chat messages format: {payload_messages}") def get_str_response(self, payload: StandardLoggingPayload) -> str: - response = payload["response"] + response: Final = payload["response"] if response is None: raise Exception("No response found in payload.") @@ -144,14 +144,14 @@ class ArgillaLogger(CustomBatchLogger): 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: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if payload is None: raise Exception("Error logging request payload. Payload=none.") - argilla_message = self.get_chat_messages(payload) - argilla_response = self.get_str_response(payload) - argilla_item: ArgillaItem = {"fields": {}} + argilla_message: Final = self.get_chat_messages(payload) + argilla_response: Final = self.get_str_response(payload) + argilla_item: Final[ArgillaItem] = {"fields": {}} for k, v in self.argilla_transformation_object.items(): if v == "messages": argilla_item["fields"][k] = argilla_message @@ -168,17 +168,17 @@ class ArgillaLogger(CustomBatchLogger): if not self.log_queue: return - argilla_api_base = self.default_credentials["ARGILLA_BASE_URL"] - argilla_dataset_name = self.default_credentials["ARGILLA_DATASET_NAME"] + argilla_api_base: Final = self.default_credentials["ARGILLA_BASE_URL"] + argilla_dataset_name: Final = self.default_credentials["ARGILLA_DATASET_NAME"] - url = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk" + url: Final = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk" - argilla_api_key = self.default_credentials["ARGILLA_API_KEY"] + argilla_api_key: Final = self.default_credentials["ARGILLA_API_KEY"] - headers = {"X-Argilla-Api-Key": argilla_api_key} + headers: Final = {"X-Argilla-Api-Key": argilla_api_key} try: - response = litellm.module_level_client.post( + response: Final = litellm.module_level_client.post( url=url, json=self.log_queue, headers=headers, @@ -195,13 +195,13 @@ class ArgillaLogger(CustomBatchLogger): def log_success_event(self, kwargs, response_obj, start_time, end_time): try: - sampling_rate = ( + sampling_rate: Final = ( float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore if os.getenv("LANGSMITH_SAMPLING_RATE") is not None and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore else 1.0 ) - random_sample = random.random() + random_sample: Final = random.random() if random_sample > sampling_rate: verbose_logger.info( "Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample @@ -212,7 +212,7 @@ class ArgillaLogger(CustomBatchLogger): kwargs, response_obj, ) - data = self._prepare_log_data(kwargs, response_obj, start_time, end_time) + data: Final = self._prepare_log_data(kwargs, response_obj, start_time, end_time) if data is None: return @@ -227,8 +227,8 @@ class ArgillaLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: - sampling_rate = self.sampling_rate - random_sample = random.random() + sampling_rate: Final = self.sampling_rate + random_sample: Final = random.random() if random_sample > sampling_rate: verbose_logger.info( "Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample @@ -239,7 +239,7 @@ class ArgillaLogger(CustomBatchLogger): kwargs, response_obj, ) - payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) data = self._prepare_log_data(kwargs, response_obj, start_time, end_time) @@ -268,8 +268,8 @@ class ArgillaLogger(CustomBatchLogger): verbose_logger.exception("Argilla Layer Error - error logging async success event.") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): - sampling_rate = self.sampling_rate - random_sample = random.random() + sampling_rate: Final = self.sampling_rate + random_sample: Final = random.random() if random_sample > sampling_rate: verbose_logger.info( "Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample @@ -277,7 +277,7 @@ class ArgillaLogger(CustomBatchLogger): return # Skip logging verbose_logger.info("Langsmith Failure Event Logging!") try: - data = self._prepare_log_data(kwargs, response_obj, start_time, end_time) + data: Final = self._prepare_log_data(kwargs, response_obj, start_time, end_time) self.log_queue.append(data) verbose_logger.debug( "Langsmith logging: queue length %s, batch size %s", @@ -302,17 +302,17 @@ class ArgillaLogger(CustomBatchLogger): if not self.log_queue: return - argilla_api_base = self.default_credentials["ARGILLA_BASE_URL"] - argilla_dataset_name = self.default_credentials["ARGILLA_DATASET_NAME"] + argilla_api_base: Final = self.default_credentials["ARGILLA_BASE_URL"] + argilla_dataset_name: Final = self.default_credentials["ARGILLA_DATASET_NAME"] - url = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk" + url: Final = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk" - argilla_api_key = self.default_credentials["ARGILLA_API_KEY"] + argilla_api_key: Final = self.default_credentials["ARGILLA_API_KEY"] - headers = {"X-Argilla-Api-Key": argilla_api_key} + headers: Final = {"X-Argilla-Api-Key": argilla_api_key} try: - response = await self.async_httpx_client.put( + response: Final = await self.async_httpx_client.put( url=url, data=json.dumps( { diff --git a/litellm/integrations/arize/__init__.py b/litellm/integrations/arize/__init__.py index 24271a9b926..21da835ee16 100644 --- a/litellm/integrations/arize/__init__.py +++ b/litellm/integrations/arize/__init__.py @@ -1,5 +1,5 @@ import os -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Final if TYPE_CHECKING: from litellm.integrations.custom_prompt_management import CustomPromptManagement @@ -10,22 +10,22 @@ from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .arize_phoenix_prompt_manager import ArizePhoenixPromptManager # Global instances -global_arize_config: dict | None = None +global_arize_config: Final[dict | None] = None def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement": """ Initialize a prompt from Arize Phoenix. """ - api_key = getattr(litellm_params, "api_key", None) or os.environ.get("PHOENIX_API_KEY") - api_base = getattr(litellm_params, "api_base", None) - prompt_id = getattr(litellm_params, "prompt_id", None) + api_key: Final = getattr(litellm_params, "api_key", None) or os.environ.get("PHOENIX_API_KEY") + api_base: Final = getattr(litellm_params, "api_base", None) + prompt_id: Final = getattr(litellm_params, "prompt_id", None) if not api_key or not api_base: raise ValueError("api_key and api_base are required for Arize Phoenix prompt integration") try: - arize_prompt_manager = ArizePhoenixPromptManager( + arize_prompt_manager: Final = ArizePhoenixPromptManager( **{ "api_key": api_key, "api_base": api_base, @@ -39,6 +39,6 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom raise e -prompt_initializer_registry = { +prompt_initializer_registry: Final = { SupportedPromptIntegrations.ARIZE_PHOENIX.value: prompt_initializer, } diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index fe5d235e51b..8c494794858 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 +from typing import TYPE_CHECKING, Any, Final from typing_extensions import override @@ -32,12 +32,12 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @override def set_messages(span: "Span", kwargs: dict[str, Any]): - messages = kwargs.get("messages") + messages: Final = kwargs.get("messages") # for /chat/completions # https://docs.arize.com/arize/large-language-models/tracing/semantic-conventions if messages: - last_message = messages[-1] + last_message: Final = messages[-1] safe_set_attribute( span, SpanAttributes.INPUT_VALUE, @@ -129,7 +129,7 @@ def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs): def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs): - images = response_obj.get("data", []) + images: Final = response_obj.get("data", []) for i, image in enumerate(images): img_url = image.get("url") if img_url is None and image.get("b64_json"): @@ -145,7 +145,7 @@ def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs): def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs): - audio = response_obj.get("audio", []) + audio: Final = response_obj.get("audio", []) for i, audio_item in enumerate(audio): audio_url = audio_item.get("url") if audio_url is None and audio_item.get("b64_json"): @@ -166,7 +166,7 @@ def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs): def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_attrs): - embeddings = response_obj.get("data", []) + embeddings: Final = response_obj.get("data", []) for i, embedding_item in enumerate(embeddings): embedding_vector = embedding_item.get("embedding") if embedding_vector: @@ -193,7 +193,7 @@ def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_att def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs): - output_items = response_obj.get("output", []) + output_items: Final = response_obj.get("output", []) for i, item in enumerate(output_items): prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{i}" if not hasattr(item, "type"): @@ -232,7 +232,7 @@ def _safe_get(obj, key, default=None): """ if obj is None: return default - getter = getattr(obj, "get", None) + getter: Final = getattr(obj, "get", None) if callable(getter): try: return getter(key, default) @@ -243,15 +243,15 @@ def _safe_get(obj, key, default=None): def _set_usage_outputs(span: "Span", response_obj, span_attrs): - usage = response_obj and response_obj.get("usage") + usage: Final = response_obj and response_obj.get("usage") if not usage: return safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_TOTAL, _safe_get(usage, "total_tokens")) - completion_tokens = _safe_get(usage, "completion_tokens") or _safe_get(usage, "output_tokens") + completion_tokens: Final = _safe_get(usage, "completion_tokens") or _safe_get(usage, "output_tokens") if completion_tokens: safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_COMPLETION, completion_tokens) - prompt_tokens = _safe_get(usage, "prompt_tokens") or _safe_get(usage, "input_tokens") + prompt_tokens: Final = _safe_get(usage, "prompt_tokens") or _safe_get(usage, "input_tokens") if prompt_tokens: safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens) @@ -259,8 +259,8 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): # API (Usage) and in `output_tokens_details` for Responses API # (ResponseAPIUsage). Both nested objects may be plain Pydantic models # without `.get`. - token_details = _safe_get(usage, "completion_tokens_details") or _safe_get(usage, "output_tokens_details") - reasoning_tokens = _safe_get(token_details, "reasoning_tokens") + token_details: Final = _safe_get(usage, "completion_tokens_details") or _safe_get(usage, "output_tokens_details") + reasoning_tokens: Final = _safe_get(token_details, "reasoning_tokens") if reasoning_tokens: safe_set_attribute( span, @@ -275,8 +275,8 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): # `cache_creation_input_tokens` # All emits are conditional, so when none of these fields exist (the # situation in the existing test fixtures) no extra attributes are set. - prompt_token_details = _safe_get(usage, "prompt_tokens_details") or _safe_get(usage, "input_tokens_details") - cache_read = _safe_get(prompt_token_details, "cached_tokens") or _safe_get(usage, "cache_read_input_tokens") + prompt_token_details: Final = _safe_get(usage, "prompt_tokens_details") or _safe_get(usage, "input_tokens_details") + cache_read: Final = _safe_get(prompt_token_details, "cached_tokens") or _safe_get(usage, "cache_read_input_tokens") if cache_read: safe_set_attribute( span, @@ -285,7 +285,7 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): ) # Anthropic / Bedrock-Anthropic only — OpenAI's `prompt_tokens_details` # does not expose a cache-write count, so we read straight off `usage`. - cache_write = _safe_get(usage, "cache_creation_input_tokens") + cache_write: Final = _safe_get(usage, "cache_creation_input_tokens") if cache_write: safe_set_attribute( span, @@ -293,7 +293,7 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): cache_write, ) - audio_prompt_tokens = _safe_get(prompt_token_details, "audio_tokens") + audio_prompt_tokens: Final = _safe_get(prompt_token_details, "audio_tokens") if audio_prompt_tokens: safe_set_attribute( span, @@ -310,7 +310,7 @@ def _infer_open_inference_span_kind(call_type: str | None) -> str: if not call_type: return OpenInferenceSpanKindValues.UNKNOWN.value - lowered = str(call_type).lower() + lowered: Final = str(call_type).lower() if "embed" in lowered: return OpenInferenceSpanKindValues.EMBEDDING.value @@ -416,7 +416,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO # routes) into a dict so downstream `.get()` calls don't crash. Existing # dict / `.get()`-bearing objects (incl. Pydantic OpenAI Responses API # models) are returned unchanged, preserving the existing test behavior. - response_obj_for_attrs = _coerce_response_obj_for_attrs(response_obj) + response_obj_for_attrs: Final = _coerce_response_obj_for_attrs(response_obj) # Set span.kind defensively before anything else. If a downstream step # throws, the span still has a kind so Arize can render it correctly @@ -425,17 +425,17 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO _safe_emit("early span kind", _set_early_span_kind, span, kwargs) try: - optional_params = _sanitize_optional_params(kwargs.get("optional_params")) - litellm_params = kwargs.get("litellm_params", {}) or {} - standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") + optional_params: Final = _sanitize_optional_params(kwargs.get("optional_params")) + litellm_params: Final = kwargs.get("litellm_params", {}) or {} + standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") - metadata = standard_logging_payload.get("metadata") if standard_logging_payload else None + metadata: Final = standard_logging_payload.get("metadata") if standard_logging_payload else None _set_metadata_attributes(span, metadata, SpanAttributes) - metadata_tools = _extract_metadata_tools(metadata) - optional_tools = _extract_optional_tools(optional_params) + metadata_tools: Final = _extract_metadata_tools(metadata) + optional_tools: Final = _extract_optional_tools(optional_params) _set_request_attributes( span=span, @@ -455,7 +455,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO _set_tool_attributes(span, optional_tools, metadata_tools) attributes.set_messages(span, kwargs) - model_params = standard_logging_payload.get("model_parameters") if standard_logging_payload else None + model_params: Final = standard_logging_payload.get("model_parameters") if standard_logging_payload else None _set_model_params(span, model_params, SpanAttributes) _set_response_attributes(span=span, response_obj=response_obj_for_attrs) @@ -468,7 +468,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO # Additive emitters. Each is independently guarded so a failure can never # blank the attributes set by the main try-block above. New attributes are # written under new keys; existing attributes are not overwritten. - slp = kwargs.get("standard_logging_object") + slp: Final = kwargs.get("standard_logging_object") _safe_emit("session/user attrs", _set_session_and_user_attrs, span, kwargs, slp) _safe_emit("response cost", _set_response_cost_attr, span, slp) _safe_emit( @@ -497,7 +497,7 @@ def _set_metadata_attributes(span: "Span", metadata: Any | None, span_attrs) -> def _extract_metadata_tools(metadata: Any | None) -> list | None: if not isinstance(metadata, dict): return None - llm_obj = metadata.get("llm") + llm_obj: Final = metadata.get("llm") if isinstance(llm_obj, dict): return llm_obj.get("tools") return None @@ -550,7 +550,7 @@ def _set_model_params(span: "Span", model_params: dict | None, span_attrs) -> No safe_set_attribute(span, span_attrs.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params)) if model_params.get("user"): - user_id = model_params.get("user") + user_id: Final = model_params.get("user") if user_id is not None: safe_set_attribute(span, span_attrs.USER_ID, user_id) @@ -573,8 +573,8 @@ def _safe_emit(label: str, fn, *args, **kwargs) -> None: def _set_early_span_kind(span: "Span", kwargs: dict) -> None: """Defensively set OPENINFERENCE_SPAN_KIND before any other logic runs.""" - slp = kwargs.get("standard_logging_object") - call_type = slp.get("call_type") if isinstance(slp, dict) else None + slp: Final = kwargs.get("standard_logging_object") + call_type: Final = slp.get("call_type") if isinstance(slp, dict) else None safe_set_attribute( span, SpanAttributes.OPENINFERENCE_SPAN_KIND, @@ -595,10 +595,10 @@ def _coerce_response_obj_for_attrs(response_obj): """ if response_obj is None or hasattr(response_obj, "get"): return response_obj - text = getattr(response_obj, "text", None) + text: Final = getattr(response_obj, "text", None) if isinstance(text, str) and text: try: - parsed = json.loads(text) + parsed: Final = json.loads(text) if isinstance(parsed, dict): return parsed except Exception: @@ -620,7 +620,7 @@ def _coerce_text(value) -> str | None: if isinstance(value, str): return value if isinstance(value, list): - parts = [] + parts: Final = [] for part in value: if isinstance(part, str): parts.append(part) @@ -641,7 +641,7 @@ def _to_plain_dict(value): """ if value is None or isinstance(value, dict): return value - model_dump = getattr(value, "model_dump", None) + model_dump: Final = getattr(value, "model_dump", None) if callable(model_dump): try: return model_dump() @@ -655,7 +655,7 @@ def _get_tool_calls(message) -> list | None: Works for dicts and Pydantic message objects via ``_safe_get``. """ - tool_calls = _safe_get(message, "tool_calls") + tool_calls: Final = _safe_get(message, "tool_calls") return tool_calls if isinstance(tool_calls, list) and tool_calls else None @@ -667,11 +667,11 @@ def _normalize_tool_call(raw_tc) -> dict[str, Any] | None: Arguments are coerced to a JSON string per OpenInference convention. Returns ``None`` when ``raw_tc`` cannot be coerced to a dict. """ - tc = _to_plain_dict(raw_tc) + tc: Final = _to_plain_dict(raw_tc) if not isinstance(tc, dict): return None - function = _to_plain_dict(tc.get("function")) - name = function.get("name") if isinstance(function, dict) else None + function: Final = _to_plain_dict(tc.get("function")) + name: Final = function.get("name") if isinstance(function, dict) else None args = function.get("arguments") if isinstance(function, dict) else None if args is not None and not isinstance(args, str): try: @@ -692,7 +692,7 @@ def _summarize_tool_calls_for_output(tool_calls) -> str: so OUTPUT_VALUE is never blanked on a malformed payload. """ try: - normalized = [n for n in (_normalize_tool_call(tc) for tc in tool_calls) if n] + normalized: Final = [n for n in (_normalize_tool_call(tc) for tc in tool_calls) if n] return json.dumps({"tool_calls": normalized}) except Exception: return str(tool_calls) @@ -705,7 +705,7 @@ def _emit_message_tool_calls(span: "Span", prefix: str, message) -> None: Accepts dicts or Pydantic message objects (e.g. ``litellm.Message``); the same applies to each tool_call entry. """ - tool_calls = _get_tool_calls(message) + tool_calls: Final = _get_tool_calls(message) if not tool_calls: return for tc_idx, raw_tc in enumerate(tool_calls): @@ -744,11 +744,11 @@ def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None if not isinstance(message, dict): return - name = message.get("name") + name: Final = message.get("name") if name: safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_NAME}", name) - tool_call_id = message.get("tool_call_id") + tool_call_id: Final = message.get("tool_call_id") if tool_call_id: safe_set_attribute( span, @@ -758,9 +758,9 @@ def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None _emit_message_tool_calls(span, prefix, message) - content = message.get("content") + content: Final = message.get("content") if isinstance(content, list): - contents_prefix = f"{prefix}.{MessageAttributes.MESSAGE_CONTENTS}" + contents_prefix: Final = f"{prefix}.{MessageAttributes.MESSAGE_CONTENTS}" for part_idx, part in enumerate(content): if not isinstance(part, dict): continue @@ -823,36 +823,36 @@ def _set_session_and_user_attrs(span: "Span", kwargs: dict, standard_logging_pay """ if not isinstance(standard_logging_payload, dict): return - metadata = standard_logging_payload.get("metadata") or {} + metadata: Final = standard_logging_payload.get("metadata") or {} if not isinstance(metadata, dict): return - session_id = metadata.get("user_api_key_end_user_id") + session_id: Final = metadata.get("user_api_key_end_user_id") if session_id: safe_set_attribute(span, SpanAttributes.SESSION_ID, str(session_id)) - trace_id = standard_logging_payload.get("trace_id") + trace_id: Final = standard_logging_payload.get("trace_id") if trace_id: safe_set_attribute(span, "litellm.trace_id", str(trace_id)) - optional_params = kwargs.get("optional_params") or {} - model_params = standard_logging_payload.get("model_parameters") or {} - has_user_already = bool( + optional_params: Final = kwargs.get("optional_params") or {} + model_params: Final = standard_logging_payload.get("model_parameters") or {} + has_user_already: Final = bool( (isinstance(optional_params, dict) and optional_params.get("user")) or (isinstance(model_params, dict) and model_params.get("user")) ) if not has_user_already: - user_id = metadata.get("user_api_key_user_id") + user_id: Final = metadata.get("user_api_key_user_id") if user_id: safe_set_attribute(span, SpanAttributes.USER_ID, str(user_id)) - team_id = metadata.get("user_api_key_team_id") + team_id: Final = metadata.get("user_api_key_team_id") if team_id: safe_set_attribute(span, "litellm.team_id", str(team_id)) - team_alias = metadata.get("user_api_key_team_alias") + team_alias: Final = metadata.get("user_api_key_team_alias") if team_alias: safe_set_attribute(span, "litellm.team_alias", str(team_alias)) - key_alias = metadata.get("user_api_key_alias") + key_alias: Final = metadata.get("user_api_key_alias") if key_alias: safe_set_attribute(span, "litellm.key_alias", str(key_alias)) @@ -868,11 +868,11 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None: """ if not isinstance(standard_logging_payload, dict): return - cost = standard_logging_payload.get("response_cost") + cost: Final = standard_logging_payload.get("response_cost") if cost is None: return try: - cost_value = float(cost) + cost_value: Final = float(cost) except (TypeError, ValueError): return safe_set_attribute(span, "llm.cost.total", cost_value) @@ -882,7 +882,7 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None: def _is_passthrough_call_type(call_type: str | None) -> bool: if not call_type: return False - lowered = str(call_type).lower() + lowered: Final = str(call_type).lower() return "passthrough" in lowered or "pass_through" in lowered @@ -913,7 +913,7 @@ def _maybe_normalize_passthrough( passthrough I/O (with central redaction) for free and this helper's `complete_input_dict` fallback can be deleted. See follow-up issue. """ - call_type = standard_logging_payload.get("call_type") if isinstance(standard_logging_payload, dict) else None + call_type: Final = standard_logging_payload.get("call_type") if isinstance(standard_logging_payload, dict) else None if not _is_passthrough_call_type(call_type): return @@ -927,13 +927,13 @@ def _maybe_normalize_passthrough( return # --- INPUT -------------------------------------------------------------- - additional_args = kwargs.get("additional_args") or {} + additional_args: Final = kwargs.get("additional_args") or {} complete_input_dict = additional_args.get("complete_input_dict") if isinstance(additional_args, dict) else None if isinstance(complete_input_dict, dict): _set_passthrough_input_attributes(span, complete_input_dict.get("messages")) # --- OUTPUT ------------------------------------------------------------- - parsed_response = _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs) + parsed_response: Final = _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs) if not isinstance(parsed_response, dict): return @@ -977,13 +977,13 @@ def _set_passthrough_input_attributes(span: "Span", messages) -> None: def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> None: """Render passthrough response into OUTPUT_VALUE + LLM_OUTPUT_MESSAGES.""" # Anthropic / Bedrock-Anthropic: `content` is a list of typed parts. - content_list = parsed_response.get("content") + content_list: Final = parsed_response.get("content") if isinstance(content_list, list) and content_list: - texts = [] + texts: Final = [] for part in content_list: if isinstance(part, dict) and isinstance(part.get("text"), str): texts.append(part["text"]) - joined = "\n\n".join(t for t in texts if t) + joined: Final = "\n\n".join(t for t in texts if t) if joined: safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, joined) prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" @@ -999,13 +999,13 @@ def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> N ) # OpenAI-style passthrough: `choices[0].message.content` - choices = parsed_response.get("choices") + choices: Final = parsed_response.get("choices") if isinstance(choices, list) and choices: - first = choices[0] + first: Final = choices[0] if isinstance(first, dict): - msg = first.get("message") + msg: Final = first.get("message") if isinstance(msg, dict): - text = _coerce_text(msg.get("content")) + text: Final = _coerce_text(msg.get("content")) if text: safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text) prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" @@ -1024,7 +1024,7 @@ def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> N def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): """Return a dict view of the provider response for passthrough routes.""" # Prefer the coerced view (already JSON-parsed for httpx.Response). - candidates = [] + candidates: Final = [] if isinstance(coerced_response_obj, dict): candidates.append(coerced_response_obj) if isinstance(raw_response_obj, dict) and raw_response_obj is not coerced_response_obj: @@ -1047,7 +1047,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): return candidate # Fallback: kwargs["original_response"] from the OTel base path. - original = kwargs.get("original_response") if isinstance(kwargs, dict) else None + original: Final = kwargs.get("original_response") if isinstance(kwargs, dict) else None if isinstance(original, dict): return original if isinstance(original, str): diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py index 86e861afb8a..bcab610835c 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, Union +from typing import TYPE_CHECKING, Any, Final, Union from litellm.integrations.arize import _utils from litellm.integrations.arize._utils import ArizeOTELAttributes @@ -50,7 +50,7 @@ class ArizeLogger(OpenTelemetry): self.span_kind = SpanKind return - provider = TracerProvider(resource=self._get_litellm_resource(self.config)) + provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config)) provider.add_span_processor(self._get_span_processor()) self.tracer = provider.get_tracer("litellm") self.span_kind = SpanKind @@ -80,13 +80,13 @@ class ArizeLogger(OpenTelemetry): Raises: ValueError: If required environment variables are not set. """ - space_id = os.environ.get("ARIZE_SPACE_ID") - space_key = os.environ.get("ARIZE_SPACE_KEY") - api_key = os.environ.get("ARIZE_API_KEY") - project_name = os.environ.get("ARIZE_PROJECT_NAME") + space_id: Final = os.environ.get("ARIZE_SPACE_ID") + space_key: Final = os.environ.get("ARIZE_SPACE_KEY") + api_key: Final = os.environ.get("ARIZE_API_KEY") + project_name: Final = os.environ.get("ARIZE_PROJECT_NAME") - grpc_endpoint = os.environ.get("ARIZE_ENDPOINT") - http_endpoint = os.environ.get("ARIZE_HTTP_ENDPOINT") + grpc_endpoint: Final = os.environ.get("ARIZE_ENDPOINT") + http_endpoint: Final = os.environ.get("ARIZE_HTTP_ENDPOINT") endpoint = None protocol: Protocol = "otlp_grpc" @@ -147,7 +147,7 @@ class ArizeLogger(OpenTelemetry): dict: Health check result with status and message """ try: - config = self.get_arize_config() + config: Final = self.get_arize_config() if not config.space_id and not config.space_key: return { @@ -183,7 +183,7 @@ class ArizeLogger(OpenTelemetry): Returns: dict: A dictionary of dynamic Arize headers """ - dynamic_headers = {} + dynamic_headers: Final = {} ######################################################### # `arize-space-id` handling diff --git a/litellm/integrations/arize/arize_phoenix.py b/litellm/integrations/arize/arize_phoenix.py index ae8a6994488..41011a6ee98 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, Union +from typing import TYPE_CHECKING, Any, Final, Union from litellm._logging import verbose_logger from litellm.integrations.arize import _utils @@ -43,8 +43,8 @@ else: OpenTelemetry = None # type: ignore -ARIZE_HOSTED_PHOENIX_ENDPOINT = "https://otlp.arize.com/v1/traces" -_MAX_PROJECT_PROVIDERS = 64 +ARIZE_HOSTED_PHOENIX_ENDPOINT: Final = "https://otlp.arize.com/v1/traces" +_MAX_PROJECT_PROVIDERS: Final = 64 class ArizePhoenixLogger(OpenTelemetry): # type: ignore @@ -80,7 +80,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore self._shared_span_processor = self._get_span_processor() self.span_kind = SpanKind - default_project = self._resolve_project_name({}) + default_project: Final = self._resolve_project_name({}) self.tracer = self._get_tracer_for(default_project) verbose_logger.debug( "ArizePhoenixLogger: Initialized per-project TracerProvider cache " @@ -100,7 +100,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore if getattr(self, "_use_injected_tracer_provider", False): return - shared_processor = getattr(self, "_shared_span_processor", None) + shared_processor: Final = getattr(self, "_shared_span_processor", None) if shared_processor is not None: try: shared_processor.force_flush() @@ -111,7 +111,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore ) with getattr(self, "_project_providers_lock", threading.Lock()): - providers = list(getattr(self, "_project_providers", {}).values()) + providers: Final = list(getattr(self, "_project_providers", {}).values()) for provider in providers: try: @@ -129,24 +129,24 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore """ from opentelemetry.sdk.resources import OTELResourceDetector, Resource - project_attributes: dict[str, str] = { + project_attributes: Final[dict[str, str]] = { "openinference.project.name": project_name, "model_id": project_name, "service.name": project_name, } - deployment_environment = getattr(self.config, "deployment_environment", None) + deployment_environment: Final = getattr(self.config, "deployment_environment", None) if deployment_environment is not None: project_attributes["deployment.environment"] = deployment_environment - env_resource = OTELResourceDetector().detect() - project_resource = Resource.create(project_attributes) # type: ignore[arg-type] + env_resource: Final = OTELResourceDetector().detect() + project_resource: Final = Resource.create(project_attributes) # type: ignore[arg-type] return env_resource.merge(project_resource) def _build_tracer_provider_for_project(self, project_name: str) -> TracerProvider: """Create a TracerProvider for *project_name* (caller holds no cache lock).""" from opentelemetry.sdk.trace import TracerProvider - provider = TracerProvider(resource=self._get_litellm_resource_for_project(project_name)) + provider: Final = TracerProvider(resource=self._get_litellm_resource_for_project(project_name)) provider.add_span_processor(self._shared_span_processor) return provider @@ -162,7 +162,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore # OTELResourceDetector().detect() is synchronous; build outside the lock so # concurrent requests for other projects are not blocked on cache misses. - new_provider = self._build_tracer_provider_for_project(project_name) + new_provider: Final = self._build_tracer_provider_for_project(project_name) with self._project_providers_lock: if project_name in self._project_providers: @@ -177,7 +177,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore 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) + project_name: Final = self._resolve_project_name(kwargs) return project_name, self._get_tracer_for(project_name) def get_tracer_to_use_for_request(self, kwargs: dict) -> Tracer: @@ -204,7 +204,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore def _normalize_project_name(name: str | None) -> str | None: if name is None: return None - normalized = str(name).strip() + normalized: Final = str(name).strip() return normalized if normalized else None @staticmethod @@ -228,7 +228,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore user-supplied and would let an authenticated caller fake proxy-mode detection to route their telemetry into arbitrary Arize/Phoenix projects. """ - litellm_params = kwargs.get("litellm_params") + litellm_params: Final = kwargs.get("litellm_params") return isinstance(litellm_params, dict) and bool(litellm_params.get("proxy_server_request")) @staticmethod @@ -240,9 +240,9 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore select the project. SDK callers may still set project fields directly on ``metadata``. """ - auth_metadata = metadata.get("user_api_key_auth_metadata") + auth_metadata: Final = metadata.get("user_api_key_auth_metadata") if isinstance(auth_metadata, dict): - project = ArizePhoenixLogger._normalize_project_name(auth_metadata.get(metadata_key)) + project: Final = ArizePhoenixLogger._normalize_project_name(auth_metadata.get(metadata_key)) if project: return project @@ -252,7 +252,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore @staticmethod def _metadata_project_from_kwargs(kwargs: dict, metadata_key: str) -> str | None: - proxy_mode = ArizePhoenixLogger._is_proxy_request(kwargs) + proxy_mode: Final = 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) if project: @@ -268,15 +268,15 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore ``user_api_key_auth_metadata.phoenix_project_name``, env, then ``default``. SDK priority: request metadata fields, then env, then ``default``. """ - override = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name_override") + override: Final = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name_override") if override: return override - phoenix_name = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name") + phoenix_name: Final = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name") if phoenix_name: return phoenix_name - env_name = ArizePhoenixLogger._normalize_project_name( + env_name: Final = ArizePhoenixLogger._normalize_project_name( os.environ.get("PHOENIX_PROJECT_NAME") or os.environ.get("ARIZE_PROJECT_NAME") ) if env_name: @@ -304,23 +304,23 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore if tracer is None: tracer = self._resolve_tracer_for_kwargs(kwargs)[1] - litellm_params = kwargs.get("litellm_params", {}) or {} - proxy_server_request = litellm_params.get("proxy_server_request", {}) or {} - headers = proxy_server_request.get("headers", {}) or {} + litellm_params: Final = kwargs.get("litellm_params", {}) or {} + proxy_server_request: Final = litellm_params.get("proxy_server_request", {}) or {} + headers: Final = proxy_server_request.get("headers", {}) or {} traceparent_ctx = self.get_traceparent_from_header(headers=headers) if headers.get("traceparent") else None - is_proxy_mode = bool(proxy_server_request) + is_proxy_mode: Final = bool(proxy_server_request) if is_proxy_mode: - start_time_val = kwargs.get("start_time", kwargs.get("api_call_start_time")) - parent_span = tracer.start_span( + start_time_val: Final = kwargs.get("start_time", kwargs.get("api_call_start_time")) + parent_span: Final = tracer.start_span( name="litellm_proxy_request", start_time=(self._to_ns(start_time_val) if start_time_val is not None else None), context=traceparent_ctx, kind=self.span_kind.SERVER, ) - ctx = trace.set_span_in_context(parent_span) + ctx: Final = trace.set_span_in_context(parent_span) return ctx, parent_span return traceparent_ctx, None @@ -352,9 +352,9 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore _project_name, tracer = self._resolve_tracer_for_kwargs(kwargs) ctx, parent_span = self._get_phoenix_context(kwargs, tracer=tracer) - status = Status(StatusCode.OK if success else StatusCode.ERROR) + status: Final = Status(StatusCode.OK if success else StatusCode.ERROR) - span = tracer.start_span( + span: Final = tracer.start_span( name=self._get_span_name(kwargs), start_time=self._to_ns(start_time), context=ctx, @@ -389,13 +389,13 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore Retrieves the Arize Phoenix configuration based on environment variables. Returns: """ - api_key = os.environ.get("PHOENIX_API_KEY", None) + api_key: Final = os.environ.get("PHOENIX_API_KEY", None) collector_endpoint = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None) if not collector_endpoint: - grpc_endpoint = os.environ.get("PHOENIX_COLLECTOR_ENDPOINT", None) - http_endpoint = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None) + grpc_endpoint: Final = os.environ.get("PHOENIX_COLLECTOR_ENDPOINT", None) + http_endpoint: Final = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None) collector_endpoint = http_endpoint or grpc_endpoint endpoint = None @@ -434,7 +434,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore elif "app.phoenix.arize.com" in endpoint: raise ValueError("PHOENIX_API_KEY must be set when using Phoenix Cloud (app.phoenix.arize.com).") - project_name = os.environ.get("PHOENIX_PROJECT_NAME") or "default" + project_name: Final = os.environ.get("PHOENIX_PROJECT_NAME") or "default" return ArizePhoenixConfig( otlp_auth_headers=otlp_auth_headers, @@ -444,7 +444,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore ) async def async_health_check(self): - config = self.get_arize_phoenix_config() + config: Final = self.get_arize_phoenix_config() if not config.otlp_auth_headers: return { diff --git a/litellm/integrations/arize/arize_phoenix_client.py b/litellm/integrations/arize/arize_phoenix_client.py index 6f1787fae9e..18d35fee34b 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 +from typing import Any, Final from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -63,15 +63,15 @@ class ArizePhoenixClient: Returns: Dictionary containing prompt version data, or None if not found """ - safe_id = _sanitize_id(prompt_version_id) - url = f"{self.api_base}/v1/prompt_versions/{safe_id}" + safe_id: Final = _sanitize_id(prompt_version_id) + url: Final = f"{self.api_base}/v1/prompt_versions/{safe_id}" try: # Use the underlying httpx client directly to avoid query param extraction response = self.http_handler.get(url, headers=self.headers) response.raise_for_status() - data = response.json() + data: Final = response.json() return data.get("data") except Exception as e: @@ -100,8 +100,8 @@ class ArizePhoenixClient: """ try: # Try to access the prompt_versions endpoint to test connection - url = f"{self.api_base}/prompt_versions" - response = self.http_handler.client.get(url, headers=self.headers) + url: Final = f"{self.api_base}/prompt_versions" + response: Final = self.http_handler.client.get(url, headers=self.headers) response.raise_for_status() return True except Exception: diff --git a/litellm/integrations/arize/arize_phoenix_prompt_manager.py b/litellm/integrations/arize/arize_phoenix_prompt_manager.py index 9985bc20af0..a541817ca8e 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 +from typing import Any, Final from jinja2 import DictLoader, select_autoescape from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -97,10 +97,10 @@ class ArizePhoenixTemplateManager: """Load a specific prompt version from Arize Phoenix.""" try: # Fetch the prompt version from Arize Phoenix - prompt_data = self.arize_client.get_prompt_version(prompt_version_id) + prompt_data: Final = self.arize_client.get_prompt_version(prompt_version_id) if prompt_data: - template = self._parse_prompt_data(prompt_data, prompt_version_id) + template: Final = self._parse_prompt_data(prompt_data, prompt_version_id) self.prompts[prompt_version_id] = template else: raise ValueError(f"Prompt version '{prompt_version_id}' not found") @@ -109,11 +109,11 @@ class ArizePhoenixTemplateManager: 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", []) + template_data: Final = data.get("template", {}) + messages: Final = template_data.get("messages", []) # Extract invocation parameters - invocation_params = data.get("invocation_parameters", {}) + invocation_params: Final = data.get("invocation_parameters", {}) provider_params = {} # Extract provider-specific parameters @@ -129,7 +129,7 @@ class ArizePhoenixTemplateManager: break # Build metadata dictionary - metadata = { + metadata: Final = { "model_name": data.get("model_name"), "model_provider": data.get("model_provider"), "description": data.get("description", ""), @@ -151,8 +151,8 @@ class ArizePhoenixTemplateManager: if template_id not in self.prompts: raise ValueError(f"Template '{template_id}' not found") - template = self.prompts[template_id] - rendered_messages: list[AllMessageValues] = [] + template: Final = self.prompts[template_id] + rendered_messages: Final[list[AllMessageValues]] = [] for message in template.messages: role = message.get("role", "user") @@ -257,22 +257,22 @@ class ArizePhoenixPromptManager(CustomPromptManagement): Returns: Tuple of (rendered_messages, metadata) """ - template = self.prompt_manager.get_template(prompt_id) + template: Final = self.prompt_manager.get_template(prompt_id) if not template: raise ValueError(f"Prompt template '{prompt_id}' not found") # Render the template - rendered_messages = self.prompt_manager.render_template(prompt_id, prompt_variables or {}) + rendered_messages: Final = self.prompt_manager.render_template(prompt_id, prompt_variables or {}) # Extract metadata - metadata = { + metadata: Final = { "model": template.model, "temperature": template.temperature, "max_tokens": template.max_tokens, } # Add additional invocation parameters - invocation_params = template.invocation_parameters + invocation_params: Final = template.invocation_parameters provider_params = {} if "openai" in invocation_params: @@ -395,10 +395,10 @@ class ArizePhoenixPromptManager(CustomPromptManagement): rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) # Extract model from metadata (if specified) - template_model = prompt_metadata.get("model") + template_model: Final = prompt_metadata.get("model") # Extract optional parameters from metadata - optional_params = {} + optional_params: Final = {} for param in [ "temperature", "max_tokens", diff --git a/litellm/integrations/athina.py b/litellm/integrations/athina.py index f57c4c8b545..acc7d1003a2 100644 --- a/litellm/integrations/athina.py +++ b/litellm/integrations/athina.py @@ -1,4 +1,5 @@ import datetime +from typing import Final import litellm @@ -34,11 +35,11 @@ class AthinaLogger: import traceback try: - is_stream = kwargs.get("stream", False) + is_stream: Final = kwargs.get("stream", False) if is_stream: if "complete_streaming_response" in kwargs: # Log the completion response in streaming mode - completion_response = kwargs["complete_streaming_response"] + completion_response: Final = kwargs["complete_streaming_response"] response_json = completion_response.model_dump() if completion_response else {} else: # Skip logging if the completion response is not available @@ -46,7 +47,7 @@ class AthinaLogger: else: # Log the completion response in non streaming mode response_json = response_obj.model_dump() if response_obj else {} - data = { + data: Final = { "language_model_id": kwargs.get("model"), "request": kwargs, "response": response_json, @@ -62,16 +63,16 @@ class AthinaLogger: data["prompt"] = kwargs.get("messages", None) # Directly add tools or functions if present - optional_params = kwargs.get("optional_params", {}) + optional_params: Final = kwargs.get("optional_params", {}) data.update((k, v) for k, v in optional_params.items() if k in ["tools", "functions"]) # Add additional metadata keys - metadata = kwargs.get("litellm_params", {}).get("metadata", {}) + metadata: Final = kwargs.get("litellm_params", {}).get("metadata", {}) if metadata: for key in self.additional_keys: if key in metadata: data[key] = metadata[key] - response = litellm.module_level_client.post( + response: Final = litellm.module_level_client.post( self.athina_logging_url, headers=self.headers, data=json.dumps(data, default=str), diff --git a/litellm/integrations/azure_sentinel/azure_sentinel.py b/litellm/integrations/azure_sentinel/azure_sentinel.py index 29bbac2912a..f23317ae9df 100644 --- a/litellm/integrations/azure_sentinel/azure_sentinel.py +++ b/litellm/integrations/azure_sentinel/azure_sentinel.py @@ -16,6 +16,7 @@ import asyncio import os import time import traceback +from typing import Final from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -64,15 +65,15 @@ class AzureSentinelLogger(CustomBatchLogger): """ self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - resolved_dcr_immutable_id = dcr_immutable_id or os.getenv("AZURE_SENTINEL_DCR_IMMUTABLE_ID") - resolved_stream_name = stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM" - resolved_audit_stream_name = ( + resolved_dcr_immutable_id: Final = dcr_immutable_id or os.getenv("AZURE_SENTINEL_DCR_IMMUTABLE_ID") + resolved_stream_name: Final = stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM" + resolved_audit_stream_name: Final = ( audit_stream_name or os.getenv("AZURE_SENTINEL_AUDIT_STREAM_NAME") or resolved_stream_name ) - resolved_endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT") - resolved_tenant_id = tenant_id or os.getenv("AZURE_SENTINEL_TENANT_ID") or os.getenv("AZURE_TENANT_ID") - resolved_client_id = client_id or os.getenv("AZURE_SENTINEL_CLIENT_ID") or os.getenv("AZURE_CLIENT_ID") - resolved_client_secret = ( + resolved_endpoint: Final = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT") + resolved_tenant_id: Final = tenant_id or os.getenv("AZURE_SENTINEL_TENANT_ID") or os.getenv("AZURE_TENANT_ID") + resolved_client_id: Final = client_id or os.getenv("AZURE_SENTINEL_CLIENT_ID") or os.getenv("AZURE_CLIENT_ID") + resolved_client_secret: Final = ( client_secret or os.getenv("AZURE_SENTINEL_CLIENT_SECRET") or os.getenv("AZURE_CLIENT_SECRET") ) @@ -149,16 +150,16 @@ class AzureSentinelLogger(CustomBatchLogger): assert self.client_id is not None, "client_id is required" assert self.client_secret is not None, "client_secret is required" - token_url = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token" + token_url: Final = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token" - token_data = { + token_data: Final = { "client_id": self.client_id, "client_secret": self.client_secret, "scope": self.oauth_scope, "grant_type": "client_credentials", } - response = await self.async_httpx_client.post( + response: Final = await self.async_httpx_client.post( url=token_url, data=token_data, headers={"Content-Type": "application/x-www-form-urlencoded"}, @@ -167,9 +168,9 @@ class AzureSentinelLogger(CustomBatchLogger): if response.status_code != 200: raise Exception(f"Failed to get OAuth2 token: {response.status_code} - {response.text}") - token_response = response.json() + token_response: Final = response.json() self.oauth_token = token_response.get("access_token") - expires_in = token_response.get("expires_in", 3600) + expires_in: Final = token_response.get("expires_in", 3600) if not self.oauth_token: raise Exception("OAuth2 token response did not contain access_token") @@ -191,7 +192,7 @@ class AzureSentinelLogger(CustomBatchLogger): """ try: verbose_logger.debug("Azure Sentinel: Logging - Enters logging function for model %s", kwargs) - standard_logging_payload = kwargs.get("standard_logging_object", None) + standard_logging_payload: Final = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs") @@ -221,7 +222,7 @@ class AzureSentinelLogger(CustomBatchLogger): "Azure Sentinel: Logging - Enters failure logging function for model %s", kwargs, ) - standard_logging_payload = kwargs.get("standard_logging_object", None) + standard_logging_payload: Final = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs") @@ -294,14 +295,14 @@ class AzureSentinelLogger(CustomBatchLogger): verbose_logger.debug("Azure Sentinel - about to flush %s %s", len(log_queue), log_type) # Get OAuth2 token - bearer_token = await self._get_oauth_token() + bearer_token: Final = await self._get_oauth_token() # Convert log queue to JSON array format expected by Logs Ingestion API # Each log entry should be a JSON object in the array - body = safe_dumps(log_queue) + body: Final = safe_dumps(log_queue) # Set headers for Logs Ingestion API - headers = { + headers: Final = { "Authorization": f"Bearer {bearer_token}", "Content-Type": "application/json", } diff --git a/litellm/integrations/azure_storage/azure_storage.py b/litellm/integrations/azure_storage/azure_storage.py index 142a0a56967..d2181dbbb38 100644 --- a/litellm/integrations/azure_storage/azure_storage.py +++ b/litellm/integrations/azure_storage/azure_storage.py @@ -2,6 +2,7 @@ import asyncio import os import time from datetime import datetime, timedelta +from typing import Final from litellm._logging import verbose_logger from litellm._uuid import uuid @@ -32,11 +33,11 @@ class AzureBlobStorageLogger(CustomBatchLogger): 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") + _azure_storage_account_name: Final = os.getenv("AZURE_STORAGE_ACCOUNT_NAME") if not _azure_storage_account_name: raise ValueError("Missing required environment variable: AZURE_STORAGE_ACCOUNT_NAME") self.azure_storage_account_name: str = _azure_storage_account_name - _azure_storage_file_system = os.getenv("AZURE_STORAGE_FILE_SYSTEM") + _azure_storage_file_system: Final = os.getenv("AZURE_STORAGE_FILE_SYSTEM") if not _azure_storage_file_system: raise ValueError("Missing required environment variable: AZURE_STORAGE_FILE_SYSTEM") self.azure_storage_file_system: str = _azure_storage_file_system @@ -71,7 +72,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): "AzureBlobStorageLogger: Logging - Enters logging function for model %s", kwargs, ) - standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") + standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_payload is not set") @@ -94,7 +95,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): "AzureBlobStorageLogger: Logging - Enters logging function for model %s", kwargs, ) - standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object") + standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise ValueError("standard_logging_payload is not set") @@ -139,10 +140,10 @@ class AzureBlobStorageLogger(CustomBatchLogger): else: # Get a valid token instead of always requesting a new one await self.set_valid_azure_ad_token() - async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - json_payload = safe_dumps(payload) + "\n" # Add newline for each log entry - payload_bytes = json_payload.encode("utf-8") - filename = f"{payload.get('id') or str(uuid.uuid4())}.json" + async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) + json_payload: Final = safe_dumps(payload) + "\n" # Add newline for each log entry + payload_bytes: Final = json_payload.encode("utf-8") + filename: Final = f"{payload.get('id') or str(uuid.uuid4())}.json" base_url = f"https://{self.azure_storage_account_name}.dfs.core.windows.net/{self.azure_storage_file_system}/{filename}" # Execute the 3-step upload process @@ -160,12 +161,12 @@ class AzureBlobStorageLogger(CustomBatchLogger): """Helper method to create the file resource""" try: verbose_logger.debug("Creating file resource at: %s", base_url) - headers = { + headers: Final = { "x-ms-version": AZURE_STORAGE_MSFT_VERSION, "Content-Length": "0", "Authorization": f"Bearer {self.azure_auth_token}", } - response = await client.put(f"{base_url}?resource=file", headers=headers) + response: Final = await client.put(f"{base_url}?resource=file", headers=headers) response.raise_for_status() verbose_logger.debug("Successfully created file resource") except Exception as e: @@ -176,12 +177,12 @@ class AzureBlobStorageLogger(CustomBatchLogger): """Helper method to append data to the file""" try: verbose_logger.debug("Appending data to file: %s", base_url) - headers = { + headers: Final = { "x-ms-version": AZURE_STORAGE_MSFT_VERSION, "Content-Type": "application/json", "Authorization": f"Bearer {self.azure_auth_token}", } - response = await client.patch( + response: Final = await client.patch( f"{base_url}?action=append&position=0", headers=headers, data=json_payload, @@ -196,12 +197,12 @@ class AzureBlobStorageLogger(CustomBatchLogger): """Helper method to flush the data""" try: verbose_logger.debug("Flushing data at position %s", position) - headers = { + headers: Final = { "x-ms-version": AZURE_STORAGE_MSFT_VERSION, "Content-Length": "0", "Authorization": f"Bearer {self.azure_auth_token}", } - response = await client.patch(f"{base_url}?action=flush&position={position}", headers=headers) + response: Final = await client.patch(f"{base_url}?action=flush&position={position}", headers=headers) response.raise_for_status() verbose_logger.debug("Successfully flushed data") except Exception as e: @@ -253,13 +254,13 @@ class AzureBlobStorageLogger(CustomBatchLogger): if client_secret is None: raise ValueError("Missing required environment variable: AZURE_STORAGE_CLIENT_SECRET") - token_provider = get_azure_ad_token_from_entra_id( + token_provider: Final = get_azure_ad_token_from_entra_id( tenant_id=tenant_id, client_id=client_id, client_secret=client_secret, scope="https://storage.azure.com/.default", ) - token = token_provider() + token: Final = token_provider() verbose_logger.debug("azure auth token %s", token) @@ -310,16 +311,16 @@ class AzureBlobStorageLogger(CustomBatchLogger): # Create an async service client - service_client = await self.get_service_client() + service_client: Final = await self.get_service_client() # Get file system client - file_system_client = service_client.get_file_system_client(file_system=self.azure_storage_file_system) + file_system_client: Final = service_client.get_file_system_client(file_system=self.azure_storage_file_system) try: # Create directory with today's date from datetime import datetime - today = datetime.now().strftime("%Y-%m-%d") - directory_client = file_system_client.get_directory_client(today) + today: Final = datetime.now().strftime("%Y-%m-%d") + directory_client: Final = file_system_client.get_directory_client(today) # check if the directory exists if not await directory_client.exists(): @@ -327,14 +328,14 @@ class AzureBlobStorageLogger(CustomBatchLogger): verbose_logger.debug("Created directory: %s", today) # Create a file client - file_name = f"{payload.get('id') or str(uuid.uuid4())}.json" - file_client = directory_client.get_file_client(file_name) + file_name: Final = f"{payload.get('id') or str(uuid.uuid4())}.json" + file_client: Final = directory_client.get_file_client(file_name) # Create the file await file_client.create_file() # Content to append - content = safe_dumps(payload).encode("utf-8") + content: Final = safe_dumps(payload).encode("utf-8") # Append content to the file await file_client.append_data(data=content, offset=0, length=len(content)) diff --git a/litellm/integrations/bitbucket/__init__.py b/litellm/integrations/bitbucket/__init__.py index 28f645597e1..17ef5f65eb5 100644 --- a/litellm/integrations/bitbucket/__init__.py +++ b/litellm/integrations/bitbucket/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Final if TYPE_CHECKING: from litellm.integrations.custom_prompt_management import CustomPromptManagement @@ -11,7 +11,7 @@ from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .bitbucket_prompt_manager import BitBucketPromptManager # Global instances -global_bitbucket_config: dict | None = None +global_bitbucket_config: Final[dict | None] = None def set_global_bitbucket_config(config: dict) -> None: @@ -34,14 +34,14 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom """ Initialize a prompt from a BitBucket repository. """ - bitbucket_config = getattr(litellm_params, "bitbucket_config", None) - prompt_id = getattr(litellm_params, "prompt_id", None) + bitbucket_config: Final = getattr(litellm_params, "bitbucket_config", None) + prompt_id: Final = getattr(litellm_params, "prompt_id", None) if not bitbucket_config: raise ValueError("bitbucket_config is required for BitBucket prompt integration") try: - bitbucket_prompt_manager = BitBucketPromptManager( + bitbucket_prompt_manager: Final = BitBucketPromptManager( bitbucket_config=bitbucket_config, prompt_id=prompt_id, ) @@ -51,7 +51,7 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom raise e -prompt_initializer_registry = { +prompt_initializer_registry: Final = { SupportedPromptIntegrations.BITBUCKET.value: prompt_initializer, } diff --git a/litellm/integrations/bitbucket/bitbucket_client.py b/litellm/integrations/bitbucket/bitbucket_client.py index 756d1bed80c..e06e5ab358f 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 +from typing import Any, Final from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -13,7 +13,7 @@ def _sanitize_file_path(file_path: str) -> str: """Reject path traversal and URL-encode each path segment.""" if "#" in file_path or "?" in file_path: raise ValueError(f"Invalid file path {file_path!r}: contains URL special characters") - parts = file_path.split("/") + parts: Final = file_path.split("/") for part in parts: if part == "..": raise ValueError(f"Invalid file path {file_path!r}: path traversal detected") @@ -64,8 +64,8 @@ class BitBucketClient: if self.auth_method == "basic" and self.username: # Use basic auth with username and app password - credentials = f"{self.username}:{self.access_token}" - encoded_credentials = base64.b64encode(credentials.encode()).decode() + credentials: Final = f"{self.username}:{self.access_token}" + encoded_credentials: Final = base64.b64encode(credentials.encode()).decode() self.headers["Authorization"] = f"Basic {encoded_credentials}" else: # Use token-based authentication (default) @@ -84,11 +84,11 @@ class BitBucketClient: Returns: File content as string, or None if file not found """ - safe_path = _sanitize_file_path(file_path) - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" + safe_path: Final = _sanitize_file_path(file_path) + url: Final = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" try: - response = self.http_handler.get(url, headers=self.headers) + response: Final = self.http_handler.get(url, headers=self.headers) response.raise_for_status() # BitBucket returns file content as base64 encoded @@ -128,15 +128,15 @@ class BitBucketClient: Returns: List of file paths """ - safe_dir = _sanitize_file_path(directory_path) if directory_path else "" - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_dir}" + safe_dir: Final = _sanitize_file_path(directory_path) if directory_path else "" + url: Final = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_dir}" try: - response = self.http_handler.get(url, headers=self.headers) + response: Final = self.http_handler.get(url, headers=self.headers) response.raise_for_status() - data = response.json() - files = [] + data: Final = response.json() + files: Final = [] for item in data.get("values", []): if item.get("type") == "commit_file": @@ -169,10 +169,10 @@ class BitBucketClient: Returns: Dictionary containing repository information """ - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}" + url: Final = f"{self.base_url}/repositories/{self.workspace}/{self.repository}" try: - response = self.http_handler.get(url, headers=self.headers) + response: Final = self.http_handler.get(url, headers=self.headers) response.raise_for_status() return response.json() except Exception as e: @@ -198,13 +198,13 @@ class BitBucketClient: Returns: List of branch information dictionaries """ - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/refs/branches" + url: Final = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/refs/branches" try: - response = self.http_handler.get(url, headers=self.headers) + response: Final = self.http_handler.get(url, headers=self.headers) response.raise_for_status() - data = response.json() + data: Final = response.json() return data.get("values", []) except Exception as e: raise Exception(f"Failed to get branches: {e}") @@ -219,15 +219,15 @@ class BitBucketClient: Returns: Dictionary containing file metadata, or None if file not found """ - safe_path = _sanitize_file_path(file_path) - url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" + safe_path: Final = _sanitize_file_path(file_path) + url: Final = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}" try: # Use GET with Range header to get just the headers (HEAD equivalent) - headers = self.headers.copy() + headers: Final = self.headers.copy() headers["Range"] = "bytes=0-0" # Request only first byte to get headers - response = self.http_handler.get(url, headers=headers) + response: Final = self.http_handler.get(url, headers=headers) response.raise_for_status() return { diff --git a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py index 3a61c900600..88fd7dc55dc 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 +from typing import TYPE_CHECKING, Any, Final from jinja2 import DictLoader, select_autoescape from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -99,10 +99,10 @@ class BitBucketTemplateManager: """Load a specific .prompt file from BitBucket.""" try: # Fetch the .prompt file from BitBucket - prompt_content = self.bitbucket_client.get_file_content(f"{prompt_id}.prompt") + prompt_content: Final = self.bitbucket_client.get_file_content(f"{prompt_id}.prompt") if prompt_content: - template = self._parse_prompt_file(prompt_content, prompt_id) + template: Final = self._parse_prompt_file(prompt_content, prompt_id) self.prompts[prompt_id] = template except Exception as e: raise Exception(f"Failed to load prompt '{prompt_id}' from BitBucket: {e}") @@ -111,7 +111,7 @@ class BitBucketTemplateManager: """Parse a .prompt file content and extract metadata and template.""" # Split frontmatter and content if content.startswith("---"): - parts = content.split("---", 2) + parts: Final = content.split("---", 2) if len(parts) >= 3: frontmatter_str = parts[1].strip() template_content = parts[2].strip() @@ -143,7 +143,7 @@ class BitBucketTemplateManager: 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: Final[dict[str, Any]] = {} for line in yaml_str.split("\n"): line = line.strip() if ":" in line and not line.startswith("#"): @@ -167,8 +167,8 @@ class BitBucketTemplateManager: 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) + template: Final = self.prompts[template_id] + jinja_template: Final = self.jinja_env.from_string(template.content) return jinja_template.render(**(variables or {})) @@ -246,15 +246,15 @@ class BitBucketPromptManager(CustomPromptManagement): Returns: Tuple of (rendered_prompt, metadata) """ - template = self.prompt_manager.get_template(prompt_id) + template: Final = self.prompt_manager.get_template(prompt_id) if not template: raise ValueError(f"Prompt template '{prompt_id}' not found") # Render the template - rendered_prompt = self.prompt_manager.render_template(prompt_id, prompt_variables or {}) + rendered_prompt: Final = self.prompt_manager.render_template(prompt_id, prompt_variables or {}) # Extract metadata - metadata = { + metadata: Final = { "model": template.model, "temperature": template.temperature, "max_tokens": template.max_tokens, @@ -284,7 +284,7 @@ class BitBucketPromptManager(CustomPromptManagement): rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) # Parse the rendered prompt into messages - parsed_messages = self._parse_prompt_to_messages(rendered_prompt) + parsed_messages: Final = self._parse_prompt_to_messages(rendered_prompt) # Merge with existing messages if parsed_messages: @@ -329,7 +329,7 @@ class BitBucketPromptManager(CustomPromptManagement): Handles both simple prompts and multi-role conversations. """ messages = [] - lines = prompt_content.strip().split("\n") + lines: Final = prompt_content.strip().split("\n") current_role = None current_content = [] @@ -453,13 +453,13 @@ class BitBucketPromptManager(CustomPromptManagement): rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) # Convert rendered content to chat messages - messages = self._parse_prompt_to_messages(rendered_prompt) + messages: Final = self._parse_prompt_to_messages(rendered_prompt) # Extract model from metadata (if specified) - template_model = prompt_metadata.get("model") + template_model: Final = prompt_metadata.get("model") # Extract optional parameters from metadata - optional_params = {} + optional_params: Final = {} for param in [ "temperature", "max_tokens", diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index a4f3335809a..cc87b217dd0 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -3,6 +3,7 @@ import os from datetime import datetime +from typing import Final import httpx @@ -20,7 +21,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.utils import print_verbose -API_BASE = "https://api.braintrustdata.com/v1" +API_BASE: Final = "https://api.braintrustdata.com/v1" def get_utc_datetime(): @@ -58,7 +59,7 @@ class BraintrustLogger(CustomLogger): in the environment """ - missing_keys = [] + missing_keys: Final = [] if api_key is None and os.getenv("BRAINTRUST_API_KEY", None) is None: missing_keys.append("BRAINTRUST_API_KEY") @@ -74,13 +75,13 @@ class BraintrustLogger(CustomLogger): return self._project_id_cache[project_name] try: - response = self.global_braintrust_sync_http_handler.post( + response: Final = self.global_braintrust_sync_http_handler.post( f"{self.api_base}/project", headers=self.headers, json={"name": project_name}, ) - project_dict = response.json() - project_id = project_dict["id"] + project_dict: Final = response.json() + project_id: Final = project_dict["id"] self._project_id_cache[project_name] = project_id return project_id except httpx.HTTPStatusError as e: @@ -94,42 +95,42 @@ class BraintrustLogger(CustomLogger): return self._project_id_cache[project_name] try: - response = await self.global_braintrust_http_handler.post( + response: Final = await self.global_braintrust_http_handler.post( f"{self.api_base}/project/register", headers=self.headers, json={"name": project_name}, ) - project_dict = response.json() - project_id = project_dict["id"] + project_dict: Final = response.json() + project_id: Final = project_dict["id"] self._project_id_cache[project_name] = project_id return project_id except httpx.HTTPStatusError as e: raise Exception(f"Failed to register project: {e.response.text}") async def create_default_project_and_experiment(self): - project = await self.global_braintrust_http_handler.post( + project: Final = await self.global_braintrust_http_handler.post( f"{self.api_base}/project", headers=self.headers, json={"name": "litellm"} ) - project_dict = project.json() + project_dict: Final = project.json() self.default_project_id = project_dict["id"] def create_sync_default_project_and_experiment(self): - project = self.global_braintrust_sync_http_handler.post( + project: Final = self.global_braintrust_sync_http_handler.post( f"{self.api_base}/project", headers=self.headers, json={"name": "litellm"} ) - project_dict = project.json() + project_dict: Final = project.json() self.default_project_id = project_dict["id"] def log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: - litellm_call_id = kwargs.get("litellm_call_id") - standard_logging_object = kwargs.get("standard_logging_object", {}) - prompt = {"messages": kwargs.get("messages")} + litellm_call_id: Final = kwargs.get("litellm_call_id") + standard_logging_object: Final = kwargs.get("standard_logging_object", {}) + prompt: Final = {"messages": kwargs.get("messages")} output = None choices = [] @@ -146,13 +147,13 @@ class BraintrustLogger(CustomLogger): elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse): output = response_obj["data"] - litellm_params = kwargs.get("litellm_params", {}) or {} - dynamic_metadata = litellm_params.get("metadata", {}) or {} + litellm_params: Final = kwargs.get("litellm_params", {}) or {} + dynamic_metadata: Final = litellm_params.get("metadata", {}) or {} # Get project_id from metadata or create default if needed project_id = dynamic_metadata.get("project_id") if project_id is None: - project_name = dynamic_metadata.get("project_name") + project_name: Final = dynamic_metadata.get("project_name") project_id = self.get_project_id_sync(project_name) if project_name else None if project_id is None: @@ -160,7 +161,7 @@ class BraintrustLogger(CustomLogger): self.create_sync_default_project_and_experiment() project_id = self.default_project_id - tags = [] + tags: Final = [] if isinstance(dynamic_metadata, dict): for key, value in dynamic_metadata.items(): @@ -177,10 +178,10 @@ class BraintrustLogger(CustomLogger): ): # support logging dynamic metadata to braintrust standard_logging_object[key] = value - cost = kwargs.get("response_cost", None) + cost: Final = kwargs.get("response_cost", None) metrics: dict | None = None - usage_obj = getattr(response_obj, "usage", None) + usage_obj: Final = getattr(response_obj, "usage", None) if usage_obj and isinstance(usage_obj, litellm.Usage): litellm.utils.get_logging_id(start_time, response_obj) metrics = { @@ -194,7 +195,7 @@ class BraintrustLogger(CustomLogger): } # Allow metadata override for span name - span_name = dynamic_metadata.get("span_name", "Chat Completion") + span_name: Final = dynamic_metadata.get("span_name", "Chat Completion") # Span parents is a special case span_parents = dynamic_metadata.get("span_parents") @@ -204,13 +205,13 @@ class BraintrustLogger(CustomLogger): span_parents = [s.strip() for s in span_parents.split(",") if s.strip()] # Add optional span attributes only if present - span_attributes = { + span_attributes: Final = { "span_id": dynamic_metadata.get("span_id"), "root_span_id": dynamic_metadata.get("root_span_id"), "span_parents": span_parents, } - request_data = { + request_data: Final = { "id": litellm_call_id, "input": prompt["messages"], "metadata": standard_logging_object, @@ -253,9 +254,9 @@ class BraintrustLogger(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: - litellm_call_id = kwargs.get("litellm_call_id") - standard_logging_object = kwargs.get("standard_logging_object", {}) - prompt = {"messages": kwargs.get("messages")} + litellm_call_id: Final = kwargs.get("litellm_call_id") + standard_logging_object: Final = kwargs.get("standard_logging_object", {}) + prompt: Final = {"messages": kwargs.get("messages")} output = None choices = [] if response_obj is not None and ( @@ -271,13 +272,13 @@ class BraintrustLogger(CustomLogger): elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse): output = response_obj["data"] - litellm_params = kwargs.get("litellm_params", {}) - dynamic_metadata = litellm_params.get("metadata", {}) or {} + litellm_params: Final = kwargs.get("litellm_params", {}) + dynamic_metadata: Final = litellm_params.get("metadata", {}) or {} # Get project_id from metadata or create default if needed project_id = dynamic_metadata.get("project_id") if project_id is None: - project_name = dynamic_metadata.get("project_name") + project_name: Final = dynamic_metadata.get("project_name") project_id = await self.get_project_id_async(project_name) if project_name else None if project_id is None: @@ -285,7 +286,7 @@ class BraintrustLogger(CustomLogger): await self.create_default_project_and_experiment() project_id = self.default_project_id - tags = [] + tags: Final = [] if isinstance(dynamic_metadata, dict): for key, value in dynamic_metadata.items(): @@ -302,10 +303,10 @@ class BraintrustLogger(CustomLogger): ): # support logging dynamic metadata to braintrust standard_logging_object[key] = value - cost = kwargs.get("response_cost", None) + cost: Final = kwargs.get("response_cost", None) metrics: dict | None = None - usage_obj = getattr(response_obj, "usage", None) + usage_obj: Final = getattr(response_obj, "usage", None) if usage_obj and isinstance(usage_obj, litellm.Usage): litellm.utils.get_logging_id(start_time, response_obj) metrics = { @@ -317,14 +318,14 @@ class BraintrustLogger(CustomLogger): "end": end_time.timestamp(), } - api_call_start_time = kwargs.get("api_call_start_time") - completion_start_time = kwargs.get("completion_start_time") + api_call_start_time: Final = kwargs.get("api_call_start_time") + completion_start_time: Final = kwargs.get("completion_start_time") if api_call_start_time is not None and completion_start_time is not None: metrics["time_to_first_token"] = completion_start_time.timestamp() - api_call_start_time.timestamp() # Allow metadata override for span name - span_name = dynamic_metadata.get("span_name", "Chat Completion") + span_name: Final = dynamic_metadata.get("span_name", "Chat Completion") # Span parents is a special case span_parents = dynamic_metadata.get("span_parents") @@ -334,13 +335,13 @@ class BraintrustLogger(CustomLogger): span_parents = [s.strip() for s in span_parents.split(",") if s.strip()] # Add optional span attributes only if present - span_attributes = { + span_attributes: Final = { "span_id": dynamic_metadata.get("span_id"), "root_span_id": dynamic_metadata.get("root_span_id"), "span_parents": span_parents, } - request_data = { + request_data: Final = { "id": litellm_call_id, "input": prompt["messages"], "output": output, diff --git a/litellm/integrations/braintrust_mock_client.py b/litellm/integrations/braintrust_mock_client.py index c775a8f3ab8..07c01c58305 100644 --- a/litellm/integrations/braintrust_mock_client.py +++ b/litellm/integrations/braintrust_mock_client.py @@ -10,6 +10,7 @@ Usage: import os import time +from typing import Final from urllib.parse import urlparse from litellm._logging import verbose_logger @@ -22,7 +23,7 @@ from litellm.integrations.mock_client_factory import ( # Use factory for should_use_mock and MockResponse # Braintrust uses both HTTPHandler (sync) and AsyncHTTPHandler (async) # Braintrust needs endpoint-specific responses, so we use custom HTTPHandler.post patching -_config = MockClientConfig( +_config: Final = MockClientConfig( "BRAINTRUST", "BRAINTRUST_MOCK", default_latency_ms=100, @@ -51,7 +52,7 @@ _original_http_handler_post = None _mocks_initialized = False # Default mock latency in seconds -_MOCK_LATENCY_SECONDS = float(os.getenv("BRAINTRUST_MOCK_LATENCY_MS", "100")) / 1000.0 +_MOCK_LATENCY_SECONDS: Final = float(os.getenv("BRAINTRUST_MOCK_LATENCY_MS", "100")) / 1000.0 def _is_braintrust_url(url: str) -> bool: @@ -59,8 +60,8 @@ def _is_braintrust_url(url: str) -> bool: if not isinstance(url, str): return False - parsed = urlparse(url) - host = (parsed.hostname or "").lower() + parsed: Final = urlparse(url) + host: Final = (parsed.hostname or "").lower() if not host: return False @@ -94,7 +95,7 @@ def _mock_http_handler_post( # Return appropriate mock response based on endpoint if "/project" in url: # Project creation/retrieval/register endpoint - project_name = json.get("name", "litellm") if json else "litellm" + project_name: Final = json.get("name", "litellm") if json else "litellm" mock_data = {"id": f"mock-project-id-{project_name}", "name": project_name} elif "/project_logs" in url: # Log insertion endpoint diff --git a/litellm/integrations/cloudzero/cloudzero.py b/litellm/integrations/cloudzero/cloudzero.py index 5cab952cfea..f41bd885442 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, cast +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm._logging import verbose_logger @@ -56,7 +56,7 @@ class CloudZeroLogger(CustomLogger): ) from litellm.proxy.proxy_server import proxy_logging_obj - pod_lock_manager = proxy_logging_obj.db_spend_update_writer.pod_lock_manager + pod_lock_manager: Final = proxy_logging_obj.db_spend_update_writer.pod_lock_manager # if using redis, ensure only one pod exports the data at a time if pod_lock_manager and pod_lock_manager.redis_cache: @@ -80,9 +80,9 @@ class CloudZeroLogger(CustomLogger): from litellm.constants import CLOUDZERO_MAX_FETCHED_DATA_RECORDS - current_time_utc = datetime.now(timezone.utc) + current_time_utc: Final = datetime.now(timezone.utc) # Mitigates the possibility of missing spend if an hour is skipped due to a restart in an ephemeral environment - one_hour_ago_utc = current_time_utc - timedelta(minutes=CLOUDZERO_EXPORT_INTERVAL_MINUTES * 2) + one_hour_ago_utc: Final = current_time_utc - timedelta(minutes=CLOUDZERO_EXPORT_INTERVAL_MINUTES * 2) await self.export_usage_data( limit=CLOUDZERO_MAX_FETCHED_DATA_RECORDS, operation="replace_hourly", @@ -122,7 +122,7 @@ class CloudZeroLogger(CustomLogger): ) # Initialize database connection and load data - database = LiteLLMDatabase() + database: Final = LiteLLMDatabase() verbose_logger.debug("CloudZero Logger: Loading usage data from database") data = await database.get_usage_data(limit=limit, start_time_utc=start_time_utc, end_time_utc=end_time_utc) @@ -133,15 +133,15 @@ class CloudZeroLogger(CustomLogger): verbose_logger.debug("CloudZero Logger: Processing %s records", len(data)) # Transform data to CloudZero CBF format - transformer = CBFTransformer() - cbf_data = transformer.transform(data) + transformer: Final = CBFTransformer() + cbf_data: Final = transformer.transform(data) if cbf_data.is_empty(): verbose_logger.warning("CloudZero Logger: No valid data after transformation") return # Send data to CloudZero - streamer = CloudZeroStreamer( + streamer: Final = CloudZeroStreamer( api_key=self.api_key, connection_id=self.connection_id, user_timezone=self.timezone, @@ -173,9 +173,9 @@ class CloudZeroLogger(CustomLogger): verbose_logger.debug("CloudZero Logger: Starting dry run export") # Initialize database connection and load data - database = LiteLLMDatabase() + database: Final = LiteLLMDatabase() verbose_logger.debug("CloudZero Logger: Loading usage data for dry run") - data = await database.get_usage_data(limit=limit) + data: Final = await database.get_usage_data(limit=limit) if data.is_empty(): verbose_logger.warning("CloudZero Dry Run: No usage data found") @@ -194,11 +194,11 @@ class CloudZeroLogger(CustomLogger): verbose_logger.debug("CloudZero Dry Run: Processing %s records...", len(data)) # Convert usage data to dict format for response - usage_data_sample = data.head(50).to_dicts() # Return first 50 rows + usage_data_sample: Final = data.head(50).to_dicts() # Return first 50 rows # Transform data to CloudZero CBF format - transformer = CBFTransformer() - cbf_data = transformer.transform(data) + transformer: Final = CBFTransformer() + cbf_data: Final = transformer.transform(data) if cbf_data.is_empty(): verbose_logger.warning("CloudZero Dry Run: No valid data after transformation") @@ -217,17 +217,17 @@ class CloudZeroLogger(CustomLogger): } # Convert CBF data to dict format for response - cbf_data_dict = cbf_data.to_dicts() + cbf_data_dict: Final = cbf_data.to_dicts() # Calculate summary statistics - total_cost = sum(record.get("cost/cost", 0) for record in cbf_data_dict) - unique_accounts = len( + total_cost: Final = sum(record.get("cost/cost", 0) for record in cbf_data_dict) + unique_accounts: Final = len( set(record.get("resource/account", "") for record in cbf_data_dict if record.get("resource/account")) ) - unique_services = len( + unique_services: Final = len( set(record.get("resource/service", "") for record in cbf_data_dict if record.get("resource/service")) ) - total_tokens = sum(record.get("usage/amount", 0) for record in cbf_data_dict) + total_tokens: Final = sum(record.get("usage/amount", 0) for record in cbf_data_dict) verbose_logger.debug("CloudZero Logger: Dry run completed for %s records", len(cbf_data)) @@ -254,7 +254,7 @@ class CloudZeroLogger(CustomLogger): from rich.console import Console from rich.table import Table - console = Console() + console: Final = Console() if cbf_data.is_empty(): console.print("[yellow]No CBF data to display[/yellow]") @@ -263,10 +263,10 @@ class CloudZeroLogger(CustomLogger): console.print(f"\n[bold green]💰 CloudZero CBF Transformed Data ({len(cbf_data)} records)[/bold green]") # Convert to dicts for easier processing - records = cbf_data.to_dicts() + records: Final = cbf_data.to_dicts() # Create main CBF table - cbf_table = Table(show_header=True, header_style="bold cyan", box=SIMPLE, padding=(0, 1)) + cbf_table: Final = Table(show_header=True, header_style="bold cyan", box=SIMPLE, padding=(0, 1)) cbf_table.add_column("time/usage_start", style="blue", no_wrap=False) cbf_table.add_column("cost/cost", style="green", justify="right", no_wrap=False) cbf_table.add_column("entity_type", style="magenta", justify="right", no_wrap=False) @@ -316,16 +316,16 @@ class CloudZeroLogger(CustomLogger): console.print(cbf_table) # Show summary statistics - total_cost = sum(record.get("cost/cost", 0) for record in records) - unique_accounts = len( + total_cost: Final = sum(record.get("cost/cost", 0) for record in records) + unique_accounts: Final = len( set(record.get("resource/account", "") for record in records if record.get("resource/account")) ) - unique_services = len( + unique_services: Final = len( set(record.get("resource/service", "") for record in records if record.get("resource/service")) ) # Count total tokens from usage metrics - total_tokens = sum(record.get("usage/amount", 0) for record in records) + total_tokens: Final = sum(record.get("usage/amount", 0) for record in records) console.print("\n[bold blue]📊 CBF Summary[/bold blue]") console.print(f" Records: {len(records):,}") @@ -346,13 +346,13 @@ 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: Final[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 verbose_logger.debug("found %s cloudzero loggers", len(prometheus_loggers)) if len(prometheus_loggers) > 0: - cloudzero_logger = cast(CloudZeroLogger, prometheus_loggers[0]) + cloudzero_logger: Final = cast(CloudZeroLogger, prometheus_loggers[0]) verbose_logger.debug( "Initializing remaining budget metrics as a cron job executing every %s minutes" % CLOUDZERO_EXPORT_INTERVAL_MINUTES diff --git a/litellm/integrations/cloudzero/cz_resource_names.py b/litellm/integrations/cloudzero/cz_resource_names.py index 6147002bd7c..9b685996a41 100644 --- a/litellm/integrations/cloudzero/cz_resource_names.py +++ b/litellm/integrations/cloudzero/cz_resource_names.py @@ -18,7 +18,7 @@ import re from enum import Enum -from typing import Any, cast +from typing import Any, Final, cast import litellm @@ -48,20 +48,20 @@ class CZRNGenerator: - resource-type: 'llm-usage' (represents LLM usage/inference) - cloud-local-id: model """ - service_type = "litellm" - provider = self._normalize_provider(row.get("custom_llm_provider", "unknown")) - region = "cross-region" + service_type: Final = "litellm" + provider: Final = self._normalize_provider(row.get("custom_llm_provider", "unknown")) + region: Final = "cross-region" # Use the actual entity_id (team_id or user_id) as the owner account - team_id = row.get("team_id", "unknown") - owner_account_id = self._normalize_component(team_id) + team_id: Final = row.get("team_id", "unknown") + owner_account_id: Final = self._normalize_component(team_id) - resource_type = "llm-usage" + resource_type: Final = "llm-usage" # Create a unique identifier with just the model (entity info already in owner_account_id) - model = row.get("model", "unknown") + model: Final = row.get("model", "unknown") - cloud_local_id = model + cloud_local_id: Final = model return self.create_from_components( service_type=service_type, @@ -90,7 +90,7 @@ class CZRNGenerator: resource_type = self._normalize_component(resource_type) # cloud_local_id can contain pipes and other characters, so don't normalize it - czrn = f"czrn:{service_type}:{provider}:{region}:{owner_account_id}:{resource_type}:{cloud_local_id}" + czrn: Final = f"czrn:{service_type}:{provider}:{region}:{owner_account_id}:{resource_type}:{cloud_local_id}" if not self.is_valid(czrn): raise ValueError(f"Generated CZRN is invalid: {czrn}") @@ -106,7 +106,7 @@ class CZRNGenerator: Returns: (service_type, provider, region, owner_account_id, resource_type, cloud_local_id) """ - match = self.CZRN_REGEX.match(czrn) + match: Final = self.CZRN_REGEX.match(czrn) if not match: raise ValueError(f"Invalid CZRN format: {czrn}") @@ -115,7 +115,7 @@ class CZRNGenerator: def _normalize_provider(self, provider: str) -> str: """Normalize provider names to standard CZRN format.""" # Map common provider names to CZRN standards - provider_map = { + provider_map: Final = { litellm.LlmProviders.AZURE.value: "azure", litellm.LlmProviders.AZURE_AI.value: "azure", litellm.LlmProviders.ANTHROPIC.value: "anthropic", @@ -128,7 +128,7 @@ class CZRNGenerator: litellm.LlmProviders.TOGETHER_AI.value: "together-ai", } - normalized = provider.lower().replace("_", "-") + normalized: Final = provider.lower().replace("_", "-") # use litellm custom llm provider if not in provider_map if normalized not in provider_map: diff --git a/litellm/integrations/cloudzero/cz_stream_api.py b/litellm/integrations/cloudzero/cz_stream_api.py index 2a2507011e4..1e2fa318786 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 +from typing import Any, Final import httpx import polars as pl @@ -55,7 +55,7 @@ class CloudZeroStreamer: return # Group data by date and send each day as a batch - daily_batches = self._group_by_date(data) + daily_batches: Final = self._group_by_date(data) if not daily_batches: self.console.print("[yellow]No valid daily batches to send[/yellow]") @@ -68,7 +68,7 @@ class CloudZeroStreamer: def _group_by_date(self, data: pl.DataFrame) -> dict[str, pl.DataFrame]: """Group data by date, converting to UTC and validating dates.""" - daily_batches: dict[str, list[dict[str, Any]]] = {} + daily_batches: Final[dict[str, list[dict[str, Any]]]] = {} # Ensure we have the required columns if "time/usage_start" not in data.columns: @@ -153,22 +153,22 @@ class CloudZeroStreamer: if batch_data.is_empty(): return - headers = { + headers: Final = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } # Use the correct API endpoint format from documentation - url = f"{self.base_url}/v2/connections/billing/anycost/{self.connection_id}/billing_drops" + url: Final = f"{self.base_url}/v2/connections/billing/anycost/{self.connection_id}/billing_drops" # Prepare the batch payload according to AnyCost API format - payload = self._prepare_batch_payload(batch_date, batch_data, operation) + payload: Final = self._prepare_batch_payload(batch_date, batch_data, operation) try: with httpx.Client(timeout=30.0) as client: self.console.print(f"[blue]Sending batch for {batch_date} ({len(batch_data)} records)[/blue]") - response = client.post(url, headers=headers, json=payload) + response: Final = client.post(url, headers=headers, json=payload) response.raise_for_status() self.console.print( @@ -188,20 +188,20 @@ class CloudZeroStreamer: """Prepare batch payload according to CloudZero AnyCost API format.""" # Convert batch_date to month for the API (YYYY-MM format) try: - date_obj = datetime.strptime(batch_date, "%Y-%m-%d") + date_obj: Final = datetime.strptime(batch_date, "%Y-%m-%d") month_str = date_obj.strftime("%Y-%m") except ValueError: # Fallback to current month month_str = datetime.now().strftime("%Y-%m") # Convert DataFrame rows to API format - data_records = [] + data_records: Final = [] for row in batch_data.iter_rows(named=True): record = self._convert_cbf_to_api_format(row) if record: data_records.append(record) - payload = {"month": month_str, "operation": operation, "data": data_records} + payload: Final = {"month": month_str, "operation": operation, "data": data_records} return payload @@ -209,7 +209,7 @@ class CloudZeroStreamer: """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 - api_record = {} + api_record: Final = {} # Copy all CBF fields, converting numeric values to strings as required by CloudZero for key, value in row.items(): @@ -241,7 +241,7 @@ class CloudZeroStreamer: return datetime.now(timezone.utc).isoformat() try: - dt = self._parse_and_convert_timestamp(timestamp_str) + dt: Final = self._parse_and_convert_timestamp(timestamp_str) return dt.isoformat().replace("+00:00", "Z") except Exception: # Fallback to current time in UTC diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 2d0f81af98b..b050ee8e1ed 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 +from typing import Any, Final import polars as pl @@ -44,7 +44,7 @@ class LiteLLMDatabase: end_time_utc: datetime | None = None, ) -> pl.DataFrame: """Retrieve usage data from LiteLLM daily user spend table.""" - client = self._ensure_prisma_client() + client: Final = self._ensure_prisma_client() # Query to get user spend data with team information. Use parameter binding to # avoid SQL injection from user-supplied timestamps or limits. @@ -80,7 +80,7 @@ class LiteLLMDatabase: ORDER BY dus.date DESC, dus.created_at DESC """ - params: list[Any] = [ + params: Final[list[Any]] = [ start_time_utc, end_time_utc, ] @@ -93,7 +93,7 @@ class LiteLLMDatabase: query += " LIMIT $3" try: - db_response = await client.db.query_raw(query, *params) + db_response: Final = await client.db.query_raw(query, *params) # Convert the response to polars DataFrame with full schema inference # This prevents schema mismatch errors when data types vary across rows return pl.DataFrame(db_response, infer_schema_length=None) diff --git a/litellm/integrations/cloudzero/transform.py b/litellm/integrations/cloudzero/transform.py index 3acfd3d8451..f0d4d67fc22 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 +from typing import Any, Final import polars as pl @@ -40,7 +40,7 @@ class CBFTransformer: return pl.DataFrame() # Filter out records with zero successful_requests first - original_count = len(data) + original_count: Final = len(data) if "successful_requests" in data.columns: filtered_data = data.filter(pl.col("successful_requests") > 0) zero_requests_dropped = original_count - len(filtered_data) @@ -48,9 +48,9 @@ class CBFTransformer: filtered_data = data zero_requests_dropped = 0 - cbf_data = [] + cbf_data: Final = [] czrn_dropped_count = 0 - filtered_count = len(filtered_data) + filtered_count: Final = len(filtered_data) for row in filtered_data.iter_rows(named=True): try: @@ -65,7 +65,7 @@ class CBFTransformer: # Print summary of dropped records if any from rich.console import Console - console = Console() + console: Final = Console() if zero_requests_dropped > 0: console.print( @@ -86,35 +86,35 @@ class CBFTransformer: """Create a single CBF record from LiteLLM daily spend row.""" # Parse date (daily spend tables use date strings like '2025-04-19') - usage_date = self._parse_date(row.get("date")) + usage_date: Final = self._parse_date(row.get("date")) # Calculate total tokens - prompt_tokens = int(row.get("prompt_tokens", 0)) - completion_tokens = int(row.get("completion_tokens", 0)) - total_tokens = prompt_tokens + completion_tokens + prompt_tokens: Final = int(row.get("prompt_tokens", 0)) + completion_tokens: Final = int(row.get("completion_tokens", 0)) + total_tokens: Final = prompt_tokens + completion_tokens # Create CloudZero Resource Name (CZRN) as resource_id - resource_id = self.czrn_generator.create_from_litellm_data(row) + resource_id: Final = self.czrn_generator.create_from_litellm_data(row) # Build dimensions for CloudZero - model = str(row.get("model", "")) - api_key_hash = str(row.get("api_key", ""))[:8] # First 8 chars for identification + model: Final = str(row.get("model", "")) + api_key_hash: Final = str(row.get("api_key", ""))[:8] # First 8 chars for identification # Handle team information with fallbacks - team_id = row.get("team_id") - team_alias = row.get("team_alias") - user_email = row.get("user_email") + team_id: Final = row.get("team_id") + team_alias: Final = row.get("team_alias") + user_email: Final = row.get("user_email") # Use team_alias if available, otherwise team_id, otherwise fallback to 'unknown' - entity_id = str(team_alias) if team_alias else (str(team_id) if team_id else "unknown") + entity_id: Final = str(team_alias) if team_alias else (str(team_id) if team_id else "unknown") # Get alias fields if they exist - api_key_alias = row.get("api_key_alias") - organization_alias = row.get("organization_alias") - project_alias = row.get("project_alias") - user_alias = row.get("user_alias") + api_key_alias: Final = row.get("api_key_alias") + organization_alias: Final = row.get("organization_alias") + project_alias: Final = row.get("project_alias") + user_alias: Final = row.get("user_alias") - dimensions = { + dimensions: Final = { "entity_type": CZEntityType.TEAM.value, "entity_id": entity_id, "team_alias": str(team_alias) if team_alias else "unknown", @@ -135,7 +135,7 @@ class CBFTransformer: } # Extract CZRN components to populate corresponding CBF columns - czrn_components = self.czrn_generator.extract_components(resource_id) + czrn_components: Final = self.czrn_generator.extract_components(resource_id) ( service_type, provider, @@ -146,10 +146,10 @@ class CBFTransformer: ) = czrn_components # Build resource/account as concat of api_key_alias and api_key_prefix - resource_account = f"{api_key_alias}|{api_key_hash}" if api_key_alias else api_key_hash + resource_account: Final = f"{api_key_alias}|{api_key_hash}" if api_key_alias else api_key_hash # CloudZero CBF format with proper column names - cbf_record = { + cbf_record: Final = { # Required CBF fields "time/usage_start": ( usage_date.isoformat() if usage_date else None diff --git a/litellm/integrations/code_interpreter_interception/handler.py b/litellm/integrations/code_interpreter_interception/handler.py index db34f00b051..f142f1b88a7 100644 --- a/litellm/integrations/code_interpreter_interception/handler.py +++ b/litellm/integrations/code_interpreter_interception/handler.py @@ -9,7 +9,7 @@ captured stdout back through the typed agentic loop plan. import json import time import uuid -from typing import Any, Literal, TypedDict, cast +from typing import Any, Final, Literal, TypedDict, cast from pydantic import ValidationError @@ -37,14 +37,14 @@ from litellm.types.utils import ( ModelResponse, ) -LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution" -_INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active" -_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" -_SESSION_SCOPED_KEY = "_code_interpreter_interception_session_scoped" -_CONVERTED_STREAM_KEY = "_code_interpreter_interception_converted_stream" -_LITELLM_METADATA_KEY = "litellm_metadata" -_CACHE_TTL_SECONDS = 15 * 60 -_SESSION_SCOPED_PER_IDENTITY_CAP = 10 +LITELLM_CODE_EXECUTION_TOOL_NAME: Final = "litellm_code_execution" +_INTERCEPTION_ACTIVE_KEY: Final = "_code_interpreter_interception_active" +_SANDBOX_KEY: Final = "_code_interpreter_interception_sandbox_key" +_SESSION_SCOPED_KEY: Final = "_code_interpreter_interception_session_scoped" +_CONVERTED_STREAM_KEY: Final = "_code_interpreter_interception_converted_stream" +_LITELLM_METADATA_KEY: Final = "litellm_metadata" +_CACHE_TTL_SECONDS: Final = 15 * 60 +_SESSION_SCOPED_PER_IDENTITY_CAP: Final = 10 class CodeExecutionToolCall(TypedDict, total=False): @@ -200,16 +200,16 @@ class CodeInterpreterInterceptionLogger(CustomLogger): if self.enabled_providers is not None and self._resolve_provider(kwargs) not in self.enabled_providers: return None - tools = kwargs.get("tools") + tools: Final = kwargs.get("tools") if not isinstance(tools, list): return None if not any(isinstance(tool, dict) and tool.get("type") == "code_interpreter" for tool in tools): return None kwargs[_INTERCEPTION_ACTIVE_KEY] = True - session_id = _extract_session_id(kwargs) + session_id: Final = _extract_session_id(kwargs) if session_id: - identity = _extract_identity(kwargs) + identity: Final = _extract_identity(kwargs) kwargs[_SANDBOX_KEY] = f"{identity}:{session_id}" if identity else session_id kwargs[_SESSION_SCOPED_KEY] = True else: @@ -219,7 +219,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): kwargs[_CONVERTED_STREAM_KEY] = True self._write_interception_metadata(kwargs) - function_tool = self._get_function_tool(call_type=call_type) + function_tool: Final = self._get_function_tool(call_type=call_type) kwargs["tools"] = [ (function_tool if isinstance(tool, dict) and tool.get("type") == "code_interpreter" else tool) for tool in tools @@ -230,10 +230,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger): @staticmethod def _strip_interception_metadata(kwargs: dict[str, Any]) -> None: - metadata = kwargs.get(_LITELLM_METADATA_KEY) + metadata: Final = kwargs.get(_LITELLM_METADATA_KEY) if not isinstance(metadata, dict): return - filtered_metadata = { + filtered_metadata: Final = { key: value for key, value in metadata.items() if not is_interception_internal_key(key) @@ -264,7 +264,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): } def _get_function_tool(self, call_type: CallTypes | None) -> CodeExecutionFunctionTool: - description = "Execute python code in a sandbox and return stdout." + description: Final = "Execute python code in a sandbox and return stdout." if call_type in (CallTypes.completion, CallTypes.acompletion): return { "type": "function", @@ -299,7 +299,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool: if not isinstance(tool_choice, dict): return False - function = tool_choice.get("function") + function: Final = tool_choice.get("function") return ( tool_choice.get("type") == "code_interpreter" or tool_choice.get("name") == "code_interpreter" @@ -308,10 +308,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ) def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None: - provider = kwargs.get("custom_llm_provider") + provider: Final = kwargs.get("custom_llm_provider") if provider: return provider - model = kwargs.get("model") + model: Final = kwargs.get("model") if not isinstance(model, str): return None try: @@ -336,7 +336,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers: return False, {} - tool_calls = ( + tool_calls: Final = ( self._extract_chat_completion_code_execution_tool_calls(response=response) if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE else self._extract_code_execution_tool_calls(response=response) @@ -368,16 +368,16 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ) await self._prune_expired_cache() - tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) - sandbox_key = kwargs.get(_SANDBOX_KEY) - is_session = bool(kwargs.get(_SESSION_SCOPED_KEY)) - identity = _extract_identity(kwargs) if is_session else None + tool_calls: Final = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) + sandbox_key: Final = kwargs.get(_SANDBOX_KEY) + is_session: Final = bool(kwargs.get(_SESSION_SCOPED_KEY)) + identity: Final = _extract_identity(kwargs) if is_session else None container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity) try: - container_id = cast(str | None, getattr(container, "id", None)) - input_list = self._normalize_messages(messages) - code_interpreter_calls: list[CodeInterpreterCall] = [] + container_id: Final = cast(str | None, getattr(container, "id", None)) + input_list: Final = self._normalize_messages(messages) + code_interpreter_calls: Final[list[CodeInterpreterCall]] = [] for tool_call in tool_calls: arguments = tool_call.get("arguments", "") code = self._parse_code(arguments) @@ -411,8 +411,8 @@ class CodeInterpreterInterceptionLogger(CustomLogger): await self._delete_container_for_cache_key(sandbox_key) raise - optional_params = anthropic_messages_optional_request_params - request_patch = AgenticLoopRequestPatch( + optional_params: Final = anthropic_messages_optional_request_params + request_patch: Final = AgenticLoopRequestPatch( model=model, messages=input_list, tools=self._get_followup_tools( @@ -443,15 +443,15 @@ class CodeInterpreterInterceptionLogger(CustomLogger): kwargs: dict[str, object], ) -> AgenticLoopPlan: await self._prune_expired_cache() - tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) - sandbox_key = cast(str | None, kwargs.get(_SANDBOX_KEY)) - is_session = bool(kwargs.get(_SESSION_SCOPED_KEY)) - identity = _extract_identity(cast(dict[str, Any], kwargs)) if is_session else None + tool_calls: Final = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) + sandbox_key: Final = cast(str | None, kwargs.get(_SANDBOX_KEY)) + is_session: Final = bool(kwargs.get(_SESSION_SCOPED_KEY)) + identity: Final = _extract_identity(cast(dict[str, Any], kwargs)) if is_session else None container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity) try: - container_id = cast(str | None, getattr(container, "id", None)) - tool_results = [ + container_id: Final = cast(str | None, getattr(container, "id", None)) + tool_results: Final = [ await self._build_chat_completion_tool_result( container=container, params=params, @@ -463,10 +463,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger): except Exception: await self._delete_container_for_cache_key(sandbox_key) raise - tool_messages = [result[0] for result in tool_results] - code_interpreter_calls = [result[1] for result in tool_results] + tool_messages: Final = [result[0] for result in tool_results] + code_interpreter_calls: Final = [result[1] for result in tool_results] - request_patch = AgenticLoopRequestPatch( + request_patch: Final = AgenticLoopRequestPatch( model=model, messages=list(messages) + [self._build_chat_completion_assistant_message(tool_calls)] + tool_messages, tools=self._get_followup_tools( @@ -496,10 +496,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger): tool_call: CodeExecutionToolCall, container_id: str | None, ) -> tuple[ChatCompletionToolMessage, CodeInterpreterCall]: - arguments = tool_call.get("arguments", "") - code = self._parse_code(arguments) - stdout = await self._run_tool_call(container=container, params=params, arguments=arguments) - tool_call_id = tool_call.get("id") or tool_call.get("call_id") or uuid.uuid4().hex + arguments: Final = tool_call.get("arguments", "") + code: Final = self._parse_code(arguments) + stdout: Final = await self._run_tool_call(container=container, params=params, arguments=arguments) + tool_call_id: Final = tool_call.get("id") or tool_call.get("call_id") or uuid.uuid4().hex return ( { "role": "tool", @@ -517,7 +517,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ) async def async_agentic_loop_cleanup_hook(self, plan: AgenticLoopPlan, kwargs: dict) -> None: - metadata = plan.metadata or {} if plan else {} + metadata: Final = plan.metadata or {} if plan else {} if metadata.get("is_session_scoped"): return await self._delete_container_for_cache_key(metadata.get("sandbox_key")) @@ -544,33 +544,33 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ] def _get_followup_optional_params(self, optional_params: dict[str, object]) -> dict[str, object]: - drop_tool_choice = self._tool_choice_targets_code_interpreter(optional_params.get("tool_choice")) + drop_tool_choice: Final = self._tool_choice_targets_code_interpreter(optional_params.get("tool_choice")) return { k: v for k, v in optional_params.items() if k != "tools" and not (k == "tool_choice" and drop_tool_choice) } async def async_post_agentic_loop_response_hook(self, response: Any, plan: AgenticLoopPlan, kwargs: dict) -> Any: - metadata = plan.metadata or {} if plan else {} + metadata: Final = plan.metadata or {} if plan else {} if not metadata.get("is_session_scoped"): await self._delete_container_for_cache_key(metadata.get("sandbox_key")) - calls = metadata.get("code_interpreter_calls") + calls: Final = metadata.get("code_interpreter_calls") if not calls: return response - is_dict = isinstance(response, dict) - output = response.get("output") if is_dict else getattr(response, "output", None) + is_dict: Final = isinstance(response, dict) + output: Final = response.get("output") if is_dict else getattr(response, "output", None) if not isinstance(output, list): return response def _item_type(item: Any) -> Any: return item.get("type") if isinstance(item, dict) else getattr(item, "type", None) - insert_at = next( + insert_at: Final = next( (i for i, item in enumerate(output) if _item_type(item) == "message"), len(output), ) - new_output = output[:insert_at] + list(calls) + output[insert_at:] + new_output: Final = output[:insert_at] + list(calls) + output[insert_at:] if is_dict: response["output"] = new_output else: @@ -586,14 +586,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger): async def _run_tool_call(self, container: Any, params: dict[str, Any] | None, arguments: str) -> str: try: - code = json.loads(arguments).get("code", "") if arguments else "" + code: Final = json.loads(arguments).get("code", "") if arguments else "" except (json.JSONDecodeError, TypeError): return "[invalid tool arguments: could not parse code]" - result = await self._run_code(container=container, params=params, code=code) + result: Final = await self._run_code(container=container, params=params, code=code) if getattr(result, "error", None): - error = result.error - message = error.get("value") or error.get("name") if isinstance(error, dict) else str(error) + error: Final = result.error + message: Final = error.get("value") or error.get("name") if isinstance(error, dict) else str(error) return f"[execution error] {message}" return getattr(result, "stdout", "") or "" @@ -603,7 +603,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): identity: str | None = None, ) -> tuple[Any, dict[str, Any] | None]: if cache_key: - cached = self._container_cache.get(cache_key) + cached: Final = self._container_cache.get(cache_key) if cached is not None: self._container_cache[cache_key] = (cached[0], cached[1], time.time(), cached[3]) return cached[0], cached[1] @@ -616,7 +616,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): return container, params async def _evict_lru_session_if_over_cap(self, identity: str) -> None: - identity_entries = [(k, v) for k, v in self._container_cache.items() if v[3] == identity] + identity_entries: Final = [(k, v) for k, v in self._container_cache.items() if v[3] == identity] if len(identity_entries) < _SESSION_SCOPED_PER_IDENTITY_CAP: return lru_key, lru_entry = min(identity_entries, key=lambda item: item[1][2]) @@ -627,14 +627,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger): if self.sandbox_config is not None: return await self.sandbox_config.acreate_sandbox(), None - params = _resolve_sandbox_tool(self.sandbox_tool_name) + params: Final = _resolve_sandbox_tool(self.sandbox_tool_name) if params is None: raise ValueError( "CodeInterpreterInterception: no sandbox available. Provide a " "sandbox_config or configure a sandbox tool resolvable via " "sandbox_tool_name." ) - container = await litellm.acreate_sandbox( + container: Final = await litellm.acreate_sandbox( provider=params["sandbox_provider"], api_key=params.get("api_key"), api_base=params.get("api_base"), @@ -672,7 +672,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger): async def _delete_container_for_cache_key(self, cache_key: str | None) -> None: if not cache_key: return - cached = self._container_cache.pop(cache_key, None) + cached: Final = self._container_cache.pop(cache_key, None) if cached is None: return await self._delete_container(container=cached[0], params=cached[1]) @@ -705,14 +705,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger): def _extract_chat_completion_code_execution_tool_calls( self, response: ModelResponse | dict[str, Any] ) -> list[CodeExecutionToolCall]: - model_response = self._to_model_response(response) + model_response: Final = self._to_model_response(response) if model_response is None: return [] - choices = model_response.choices or [] + choices: Final = model_response.choices or [] if not choices: return [] - message = choices[0].message - tool_calls = message.tool_calls or [] + message: Final = choices[0].message + tool_calls: Final = message.tool_calls or [] return [ normalized @@ -783,8 +783,8 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ) async def _prune_expired_cache(self) -> None: - now = time.time() - expired = [ + now: Final = time.time() + expired: Final = [ (cache_key, container, params) for cache_key, (container, params, last_accessed, *_) in self._container_cache.items() if now - last_accessed > _CACHE_TTL_SECONDS diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index 93765001c94..ff2b1197c5f 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, cast +from typing import Any, Final, cast from litellm._logging import verbose_logger from litellm.compression import compress @@ -22,8 +22,8 @@ from litellm.types.integrations.custom_logger import ( ) from litellm.types.utils import CallTypes -LITELLM_CONTENT_RETRIEVE_TOOL_NAME = "litellm_content_retrieve" -_CACHE_TTL_SECONDS = 15 * 60 +LITELLM_CONTENT_RETRIEVE_TOOL_NAME: Final = "litellm_content_retrieve" +_CACHE_TTL_SECONDS: Final = 15 * 60 def _compression_savings_from_counts( @@ -54,7 +54,7 @@ def _record_compression_savings(kwargs: dict[str, object], savings: CompressionS to the same object; replacing it would orphan writes made through those references. """ - existing = kwargs.get("litellm_metadata") + existing: Final = kwargs.get("litellm_metadata") if isinstance(existing, dict): existing["compression_savings"] = savings return @@ -123,8 +123,8 @@ class CompressionInterceptionLogger(CustomLogger): if int(kwargs.get("_agentic_loop_depth", 0) or 0) > 0: return None - messages = kwargs.get("messages") - model = kwargs.get("model") + messages: Final = kwargs.get("messages") + model: Final = kwargs.get("model") if not isinstance(messages, list) or not isinstance(model, str): return None @@ -133,7 +133,7 @@ class CompressionInterceptionLogger(CustomLogger): self._prune_expired_cache() - compressed = compress( # type: ignore + compressed: Final = compress( # type: ignore messages=messages, model=model, call_type=CallTypes.anthropic_messages, @@ -143,9 +143,9 @@ class CompressionInterceptionLogger(CustomLogger): embedding_model_params=self.embedding_model_params, ) - 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", [])) + cache: Final = cast(dict[str, str], compressed.get("cache", {})) + skip_reason: Final = cast(str | None, compressed.get("compression_skipped_reason")) + compressed_tools: Final = 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.), @@ -164,7 +164,7 @@ class CompressionInterceptionLogger(CustomLogger): call_id = str(uuid.uuid4()) kwargs["litellm_call_id"] = call_id self._compression_cache_by_call_id[call_id] = (cache, time.time()) - savings = _compression_savings_from_counts( + savings: Final = _compression_savings_from_counts( original_tokens=compressed.get("original_tokens"), compressed_tokens=compressed.get("compressed_tokens"), ) @@ -225,14 +225,14 @@ class CompressionInterceptionLogger(CustomLogger): 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: Final = cast(list[dict[str, Any]], tools.get("tool_calls", [])) + thinking_blocks: Final = 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) - retrieval_results = [self._resolve_retrieval_content(tc, cache) for tc in tool_calls] + call_id: Final = self._resolve_call_id(logging_obj=logging_obj, kwargs=kwargs) + cache: Final = self._get_cache(call_id=call_id) + retrieval_results: Final = [self._resolve_retrieval_content(tc, cache) for tc in tool_calls] - assistant_message = { + assistant_message: Final = { "role": "assistant", "content": thinking_blocks + [ @@ -245,7 +245,7 @@ class CompressionInterceptionLogger(CustomLogger): for tc in tool_calls ], } - user_message = { + user_message: Final = { "role": "user", "content": [ { @@ -256,22 +256,22 @@ class CompressionInterceptionLogger(CustomLogger): for i in range(len(tool_calls)) ], } - follow_up_messages = messages + [assistant_message, user_message] + follow_up_messages: Final = messages + [assistant_message, user_message] - max_tokens = cast( + max_tokens: Final = cast( int | None, anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get("max_tokens"), ) - optional_params_without_max_tokens = { + optional_params_without_max_tokens: Final = { k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens" } full_model_name = model if logging_obj is not None: - agentic_params = logging_obj.model_call_details.get("agentic_loop_params", {}) + agentic_params: Final = logging_obj.model_call_details.get("agentic_loop_params", {}) full_model_name = cast(str, agentic_params.get("model", model)) - request_patch = AgenticLoopRequestPatch( + request_patch: Final = AgenticLoopRequestPatch( model=full_model_name, messages=follow_up_messages, max_tokens=max_tokens, @@ -286,7 +286,7 @@ class CompressionInterceptionLogger(CustomLogger): ) def _prune_expired_cache(self) -> None: - now = time.time() + now: Final = time.time() self._compression_cache_by_call_id = { call_id: (cache, created_at) for call_id, ( @@ -299,21 +299,21 @@ class CompressionInterceptionLogger(CustomLogger): 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) + cache_entry: Final = self._compression_cache_by_call_id.get(call_id) if cache_entry is None: return {} return cache_entry[0] 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) + logging_call_id: Final = 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") + kwargs_call_id: Final = kwargs.get("litellm_call_id") 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: - raw_input = tool_call.get("input", {}) + raw_input: Final = tool_call.get("input", {}) key = "" if isinstance(raw_input, dict): key = str(raw_input.get("key", "") or "") @@ -332,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: Final[list[dict[str, Any]]] = [] + thinking_blocks: Final[list[dict[str, Any]]] = [] for block in content: if isinstance(block, dict): @@ -381,7 +381,7 @@ class CompressionInterceptionLogger(CustomLogger): return tool_calls, thinking_blocks def _prepare_followup_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]: - internal_keys = {"litellm_logging_obj"} + internal_keys: Final = {"litellm_logging_obj"} return { k: v for k, v in kwargs.items() if not k.startswith("_compression_interception") and k not in internal_keys } @@ -405,7 +405,7 @@ class CompressionInterceptionLogger(CustomLogger): existing_tools: list[dict[str, Any]] | None, compressed_tools: list[dict[str, Any]], ) -> list[dict[str, Any]]: - merged = list(existing_tools or []) + merged: Final = list(existing_tools or []) if self._has_retrieval_tool(merged): return merged merged.extend(compressed_tools) diff --git a/litellm/integrations/custom_batch_logger.py b/litellm/integrations/custom_batch_logger.py index 7559bc83cfa..c9e24913900 100644 --- a/litellm/integrations/custom_batch_logger.py +++ b/litellm/integrations/custom_batch_logger.py @@ -6,6 +6,7 @@ Use this if you want your logs to be stored in memory and flushed periodically. import asyncio import time +from typing import Final import litellm from litellm._logging import verbose_logger @@ -56,7 +57,7 @@ class CustomBatchLogger(CustomLogger): async with self.flush_lock: if self.log_queue: - log_queue_length = len(self.log_queue) + log_queue_length: Final = len(self.log_queue) verbose_logger.debug("CustomLogger: Flushing batch of %s events", len(self.log_queue)) try: await self.async_send_batch() @@ -73,7 +74,7 @@ class CustomBatchLogger(CustomLogger): # Guard against unbounded queue growth if the destination # is persistently unreachable. Drop the oldest events # beyond ``max_queue_size``. - overflow = len(self.log_queue) - self.max_queue_size + overflow: Final = len(self.log_queue) - self.max_queue_size if overflow > 0: del self.log_queue[:overflow] verbose_logger.warning( diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index d9d65375ea8..a80b3ff5364 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -3,14 +3,7 @@ import hashlib import os import secrets from datetime import datetime -from typing import ( - TYPE_CHECKING, - Any, - ClassVar, - Literal, - Optional, - get_args, -) +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args from litellm._logging import verbose_logger from litellm.caching import DualCache @@ -45,7 +38,7 @@ except ImportError: if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -dc = DualCache() +dc: Final = DualCache() from litellm.constants import ( @@ -63,11 +56,11 @@ from litellm.exceptions import ( # honors markers carrying this token, so a caller cannot forge the metadata # field to suppress a guardrail on the direct-SDK path that never reaches the # proxy's metadata sanitizer. -_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16) +_PRE_CALL_EXECUTED_TOKEN: Final = secrets.token_hex(16) -_GUARDRAIL_BLOCK_STATUS_CODES = frozenset({400, 403, 422}) +_GUARDRAIL_BLOCK_STATUS_CODES: Final = frozenset({400, 403, 422}) -_guardrail_self_recorded: contextvars.ContextVar[bool] = contextvars.ContextVar( +_guardrail_self_recorded: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar( "litellm_guardrail_self_recorded", default=False ) @@ -79,10 +72,10 @@ def _strict_guardrail_modes_enabled() -> bool: for guardrails whose supported_event_hooks list newly includes their configured mode: log the mismatch and continue instead of raising at boot. """ - raw = os.environ.get("LITELLM_STRICT_GUARDRAIL_MODES") + raw: Final = os.environ.get("LITELLM_STRICT_GUARDRAIL_MODES") if raw is None: return True - parsed = str_to_bool(raw) + parsed: Final = str_to_bool(raw) return True if parsed is None else parsed @@ -92,12 +85,12 @@ def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None if session_id: return str(session_id) - metadata = request_data.get("metadata") or {} + metadata: Final = request_data.get("metadata") or {} session_id = metadata.get("session_id") if session_id: return str(session_id) - litellm_metadata = request_data.get("litellm_metadata") or {} + litellm_metadata: Final = request_data.get("litellm_metadata") or {} session_id = litellm_metadata.get("session_id") if session_id: return str(session_id) @@ -187,7 +180,7 @@ class CustomGuardrail(CustomLogger): if not self.violation_message_template: return default - format_context: dict[str, Any] = {"default_message": default} + format_context: Final[dict[str, Any]] = {"default_message": default} if context: format_context.update(context) try: @@ -234,7 +227,7 @@ class CustomGuardrail(CustomLogger): detection_info=detection_info ) """ - model = request_data.get("model", "unknown") + model: Final = request_data.get("model", "unknown") raise ModifyResponseException( message=violation_message, @@ -270,7 +263,7 @@ class CustomGuardrail(CustomLogger): Raises: SensitiveDataRouteException: Always raises to trigger rerouting """ - session_id = self._get_session_id_from_request_data(request_data) + session_id: Final = self._get_session_id_from_request_data(request_data) if not session_id: raise ValueError( "Cannot route sensitive data without a session_id. " @@ -325,7 +318,7 @@ class CustomGuardrail(CustomLogger): ) return None - session_id = get_session_id_from_request_data(request_data) + session_id: Final = get_session_id_from_request_data(request_data) if not session_id: verbose_logger.debug( "Guardrail %s: only_scan_new_messages enabled but request has no session id; scanning full context.", @@ -334,7 +327,7 @@ class CustomGuardrail(CustomLogger): return None try: - cached: object = await cache.async_get_cache(key=self._scanned_texts_cache_key(session_id)) + cached: Final[object] = await cache.async_get_cache(key=self._scanned_texts_cache_key(session_id)) except Exception as e: # noqa: BLE001 # cache is best-effort; any failure must fall back to a full scan verbose_logger.warning( "Guardrail %s: failed to read scanned-message cache (%s); scanning full context.", @@ -343,7 +336,7 @@ class CustomGuardrail(CustomLogger): ) return None - seen: set[str] = {str(h) for h in cached} if isinstance(cached, list) else set() + seen: Final[set[str]] = {str(h) for h in cached} if isinstance(cached, list) else set() return [text for text in texts if self._scanned_text_hash(text) not in seen] async def mark_texts_scanned( @@ -361,16 +354,16 @@ class CustomGuardrail(CustomLogger): return if self.mask_request_content or self.mask_response_content: return - session_id = get_session_id_from_request_data(request_data) + session_id: Final = get_session_id_from_request_data(request_data) if not session_id: return - cache_key = self._scanned_texts_cache_key(session_id) - current_hashes = [self._scanned_text_hash(text) for text in texts] + cache_key: Final = self._scanned_texts_cache_key(session_id) + current_hashes: Final = [self._scanned_text_hash(text) for text in texts] try: - existing: object = await cache.async_get_cache(key=cache_key) - existing_hashes: list[str] = [str(h) for h in existing] if isinstance(existing, list) else [] - merged: list[str] = list(dict.fromkeys(existing_hashes + current_hashes)) + existing: Final[object] = await cache.async_get_cache(key=cache_key) + existing_hashes: Final[list[str]] = [str(h) for h in existing] if isinstance(existing, list) else [] + merged: Final[list[str]] = list(dict.fromkeys(existing_hashes + current_hashes)) await cache.async_set_cache( key=cache_key, value=merged, @@ -477,7 +470,7 @@ class CustomGuardrail(CustomLogger): if isinstance(event_hook, list): _validate_event_hook_list_is_in_supported_event_hooks(event_hook, supported_event_hooks) elif isinstance(event_hook, Mode): - tag_values_flat: list = [] + tag_values_flat: Final[list] = [] for v in event_hook.tags.values(): if isinstance(v, list): tag_values_flat.extend(v) @@ -528,7 +521,7 @@ class CustomGuardrail(CustomLogger): Reads from admin-configured key/team metadata only. """ - value = self._get_admin_metadata(data).get("opted_out_global_guardrails") + value: Final = self._get_admin_metadata(data).get("opted_out_global_guardrails") return value if isinstance(value, list) else [] def _is_valid_response_type(self, result: Any) -> bool: @@ -544,7 +537,7 @@ class CustomGuardrail(CustomLogger): try: # Try isinstance check on valid types that support it - response_types = get_args(LLMResponseTypes) + response_types: Final = get_args(LLMResponseTypes) return isinstance(result, response_types) except TypeError as e: # TypedDict types don't support isinstance checks @@ -587,7 +580,7 @@ class CustomGuardrail(CustomLogger): return False def _pre_call_marker(self) -> str | None: - name = self.guardrail_name + name: Final = self.guardrail_name if not name: return None return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}" @@ -602,7 +595,7 @@ class CustomGuardrail(CustomLogger): top-level request kwargs, which would otherwise re-trigger the same hook from ``async_pre_call_deployment_hook``. """ - marker = self._pre_call_marker() + marker: Final = self._pre_call_marker() if marker is None: return for meta_key in ("metadata", "litellm_metadata"): @@ -618,7 +611,7 @@ class CustomGuardrail(CustomLogger): data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]} def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool: - marker = self._pre_call_marker() + marker: Final = self._pre_call_marker() if marker is None: return False for meta_key in ("metadata", "litellm_metadata"): @@ -649,7 +642,7 @@ class CustomGuardrail(CustomLogger): from litellm.proxy._types import UserAPIKeyAuth # should run guardrail - litellm_guardrails = kwargs.get("guardrails") + litellm_guardrails: Final = kwargs.get("guardrails") if litellm_guardrails is None or not isinstance(litellm_guardrails, list): return kwargs @@ -661,10 +654,10 @@ class CustomGuardrail(CustomLogger): # CHECK IF GUARDRAIL REJECTS THE REQUEST if call_type == CallTypes.completion or call_type == CallTypes.acompletion: - target = self._deployment_pre_call_target() + target: Final = self._deployment_pre_call_target() if target is not self: kwargs["guardrail_to_apply"] = self - result = await target.async_pre_call_hook( + result: Final = await target.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( user_id=kwargs.get("user_api_key_user_id"), team_id=kwargs.get("user_api_key_team_id"), @@ -678,7 +671,7 @@ class CustomGuardrail(CustomLogger): ) if result is not None and isinstance(result, dict): - result_messages = result.get("messages") + result_messages: Final = result.get("messages") if result_messages is not None: # update for any pii / masking logic kwargs["messages"] = result_messages @@ -696,7 +689,7 @@ class CustomGuardrail(CustomLogger): from litellm.proxy._types import UserAPIKeyAuth # should run guardrail - litellm_guardrails = request_data.get("guardrails") + litellm_guardrails: Final = request_data.get("guardrails") if litellm_guardrails is None or not isinstance(litellm_guardrails, list): return response @@ -704,7 +697,7 @@ class CustomGuardrail(CustomLogger): return response # CHECK IF GUARDRAIL REJECTS THE REQUEST - result = await self.async_post_call_success_hook( + result: Final = await self.async_post_call_success_hook( user_api_key_dict=UserAPIKeyAuth( user_id=request_data.get("user_api_key_user_id"), team_id=request_data.get("user_api_key_team_id"), @@ -729,9 +722,9 @@ class CustomGuardrail(CustomLogger): """ Returns True if the guardrail should be run on the event_type """ - requested_guardrails = self.get_guardrail_from_metadata(data) - disable_global_guardrail = self.get_disable_global_guardrail(data) - opted_out_global_guardrails = self.get_opted_out_global_guardrails_from_metadata(data) + requested_guardrails: Final = self.get_guardrail_from_metadata(data) + disable_global_guardrail: Final = self.get_disable_global_guardrail(data) + opted_out_global_guardrails: Final = self.get_opted_out_global_guardrails_from_metadata(data) verbose_logger.debug( "inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s self.default_on= %s", self.guardrail_name, @@ -809,7 +802,7 @@ class CustomGuardrail(CustomLogger): elif event_type.value == tag_value: return True if self.event_hook.default: - default_list = ( + default_list: Final = ( self.event_hook.default if isinstance(self.event_hook.default, list) else [self.event_hook.default] ) return event_type.value in default_list @@ -834,7 +827,7 @@ class CustomGuardrail(CustomLogger): Args: request_data: The original `request_data` passed to LiteLLM Proxy """ - requested_guardrails = self.get_guardrail_from_metadata(request_data) + requested_guardrails: Final = self.get_guardrail_from_metadata(request_data) # Look for the guardrail configuration matching self.guardrail_name for guardrail in requested_guardrails: @@ -932,7 +925,7 @@ class CustomGuardrail(CustomLogger): clean_guardrail_response = mask_credentials_in_payload(clean_guardrail_response) - slg = StandardLoggingGuardrailInformation( + slg: Final = StandardLoggingGuardrailInformation( guardrail_name=self.guardrail_name, guardrail_provider=guardrail_provider, guardrail_mode=guardrail_mode, @@ -946,8 +939,8 @@ class CustomGuardrail(CustomLogger): ) def _append_guardrail_info(container: dict) -> None: - key = "standard_logging_guardrail_information" - existing = container.get(key) + key: Final = "standard_logging_guardrail_information" + existing: Final = container.get(key) if existing is None: container[key] = [slg] elif isinstance(existing, list): @@ -1094,7 +1087,7 @@ class CustomGuardrail(CustomLogger): This gets logged on downsteam Langfuse, DataDog, etc. """ - guardrail_status: GuardrailStatus = ( + guardrail_status: Final[GuardrailStatus] = ( "guardrail_intervened" if self._is_guardrail_intervention(e) else "guardrail_failed_to_respond" ) # For custom_code_guardrail scenario, log as "deny" instead of full exception @@ -1121,7 +1114,7 @@ class CustomGuardrail(CustomLogger): Returns True if the inputs were modified (mask scenario), False otherwise (allow scenario). """ # Get all keys from both dictionaries - all_keys = set(original_inputs.keys()) | set(response.keys()) + all_keys: Final = set(original_inputs.keys()) | set(response.keys()) # Compare each key's value for key in all_keys: @@ -1191,11 +1184,11 @@ class CustomGuardrail(CustomLogger): LiteLLMCompletionResponsesConfig, ) - input_data = data.get("input") + input_data: Final = data.get("input") if input_data is None: return None - messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + messages: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( input=input_data, responses_api_request=data, ) @@ -1209,7 +1202,7 @@ def _append_slg_to_litellm_params(lp: object, entries: list) -> None: return if lp.get("metadata") is None: lp["metadata"] = {} - existing = lp["metadata"].setdefault("standard_logging_guardrail_information", []) + existing: Final = lp["metadata"].setdefault("standard_logging_guardrail_information", []) for entry in entries: if entry not in existing: existing.append(entry) @@ -1228,12 +1221,12 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object) """ if logging_obj is None: return - meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} - slg_info = meta_src.get("standard_logging_guardrail_information") + meta_src: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} + slg_info: Final = meta_src.get("standard_logging_guardrail_information") if not slg_info: return - entries: list = slg_info if isinstance(slg_info, list) else [slg_info] - mcd = getattr(logging_obj, "model_call_details", None) or {} + entries: Final[list] = slg_info if isinstance(slg_info, list) else [slg_info] + mcd: Final = getattr(logging_obj, "model_call_details", None) or {} _append_slg_to_litellm_params(getattr(logging_obj, "litellm_params", None), entries) _append_slg_to_litellm_params(mcd.get("litellm_params"), entries) @@ -1290,20 +1283,20 @@ def log_guardrail_information(func): @functools.wraps(func) async def async_wrapper(*args, **kwargs): - start_time = datetime.now() # Move start_time inside the wrapper - self: CustomGuardrail = args[0] - request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {} - event_type = _infer_event_type_from_function_name(func.__name__) + start_time: Final = datetime.now() # Move start_time inside the wrapper + self: Final[CustomGuardrail] = args[0] + request_data: Final[dict] = kwargs.get("data") or kwargs.get("request_data") or {} + event_type: Final = _infer_event_type_from_function_name(func.__name__) # Store original inputs for comparison (for apply_guardrail functions) original_inputs = None if func.__name__ == "apply_guardrail" and "inputs" in kwargs: original_inputs = kwargs.get("inputs") - logging_obj = kwargs.get("logging_obj") - self_recorded_token = _guardrail_self_recorded.set(False) + logging_obj: Final = kwargs.get("logging_obj") + self_recorded_token: Final = _guardrail_self_recorded.set(False) try: - response = await func(*args, **kwargs) + response: Final = await func(*args, **kwargs) if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( @@ -1332,20 +1325,20 @@ def log_guardrail_information(func): @functools.wraps(func) def sync_wrapper(*args, **kwargs): - start_time = datetime.now() # Move start_time inside the wrapper - self: CustomGuardrail = args[0] - request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {} - event_type = _infer_event_type_from_function_name(func.__name__) + start_time: Final = datetime.now() # Move start_time inside the wrapper + self: Final[CustomGuardrail] = args[0] + request_data: Final[dict] = kwargs.get("data") or kwargs.get("request_data") or {} + event_type: Final = _infer_event_type_from_function_name(func.__name__) # Store original inputs for comparison (for apply_guardrail functions) original_inputs = None if func.__name__ == "apply_guardrail" and "inputs" in kwargs: original_inputs = kwargs.get("inputs") - logging_obj = kwargs.get("logging_obj") - self_recorded_token = _guardrail_self_recorded.set(False) + logging_obj: Final = kwargs.get("logging_obj") + self_recorded_token: Final = _guardrail_self_recorded.set(False) try: - response = func(*args, **kwargs) + response: Final = func(*args, **kwargs) if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 9df0cf6e84d..29ef04af123 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -3,12 +3,7 @@ import re import traceback from collections.abc import AsyncGenerator -from typing import ( - TYPE_CHECKING, - Any, - Optional, - Union, -) +from typing import TYPE_CHECKING, Any, Final, Optional, Union from pydantic import BaseModel @@ -52,12 +47,12 @@ else: MCPPostCallResponseObject = Any MCPPreCallRequestObject = Any MCPPreCallResponseObject = Any - MCPDuringCallRequestObject = Any - MCPDuringCallResponseObject = Any + MCPDuringCallRequestObject: Final = Any + MCPDuringCallResponseObject: Final = Any PreRoutingHookResponse = Any -_BASE64_INLINE_PATTERN = re.compile( +_BASE64_INLINE_PATTERN: Final = re.compile( r"data:(?:application|image|audio|video)/[a-zA-Z0-9.+-]+;base64,[A-Za-z0-9+/=\s]+", re.MULTILINE, ) @@ -95,24 +90,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac if callback_name is None: return [] - normalized_name = callback_name.lower() + normalized_name: Final = callback_name.lower() - alias_map = { + alias_map: Final = { "langfuse_otel": "langfuse", } - lookup_name = alias_map.get(normalized_name, normalized_name) + lookup_name: Final = alias_map.get(normalized_name, normalized_name) try: from litellm.proxy._types import AllCallbacks except Exception: return [] - callbacks = AllCallbacks() - callback_info = getattr(callbacks, lookup_name, None) + callbacks: Final = AllCallbacks() + callback_info: Final = getattr(callbacks, lookup_name, None) if callback_info is None: return [] - params = getattr(callback_info, "litellm_callback_params", None) + params: Final = getattr(callback_info, "litellm_callback_params", None) if not params: return [] @@ -761,10 +756,10 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac This function truncates the error string and the message content if they exceed a certain length. """ - MAX_STR_LENGTH = 10_000 + MAX_STR_LENGTH: Final = 10_000 # Truncate fields that might exceed max length - fields_to_truncate = ["error_str", "messages", "response"] + fields_to_truncate: Final = ["error_str", "messages", "response"] for field in fields_to_truncate: self._truncate_field( standard_logging_object=standard_logging_object, @@ -788,9 +783,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac - Converting to string and then truncating the logged content catches this 2. We want to avoid modifying the original `messages`, `response`, and `error_str` in the logging payload since these are in kwargs and could be returned to the user """ - field_value = standard_logging_object.get(field_name) # type: ignore + field_value: Final = standard_logging_object.get(field_name) # type: ignore if field_value: - str_value = str(field_value) + str_value: Final = str(field_value) if len(str_value) > max_length: standard_logging_object[field_name] = self._truncate_text( # type: ignore text=str_value, max_length=max_length @@ -836,8 +831,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac import litellm from litellm import Choices, Message, ModelResponse - turn_off_message_logging: bool = getattr(self, "turn_off_message_logging", False) - excluded_fields: list[str] | None = getattr(litellm, "standard_logging_payload_excluded_fields", None) + turn_off_message_logging: Final[bool] = getattr(self, "turn_off_message_logging", False) + excluded_fields: Final[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: @@ -845,13 +840,13 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac # Only make a shallow copy of the top-level dict to avoid deepcopy issues # with complex objects like AuthenticationError that may be present - model_call_details_copy = copy(model_call_details) - standard_logging_object = model_call_details.get("standard_logging_object") + model_call_details_copy: Final = copy(model_call_details) + standard_logging_object: Final = model_call_details.get("standard_logging_object") if standard_logging_object is None: return model_call_details_copy # Make a copy of just the standard_logging_object to avoid modifying the original - standard_logging_object_copy = copy(standard_logging_object) + standard_logging_object_copy: Final = copy(standard_logging_object) # Handle excluded fields - remove them entirely from the payload if excluded_fields: @@ -861,19 +856,19 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac # Handle turn_off_message_logging - redact messages and responses (if not already excluded) if turn_off_message_logging: - redacted_str = "redacted-by-litellm" + redacted_str: Final = "redacted-by-litellm" if "messages" not in (excluded_fields or []) and standard_logging_object_copy.get("messages") is not None: standard_logging_object_copy["messages"] = [Message(content=redacted_str).model_dump()] if "response" not in (excluded_fields or []) and standard_logging_object_copy.get("response") is not None: - response = standard_logging_object_copy["response"] + response: Final = standard_logging_object_copy["response"] # Check if this is a ResponsesAPIResponse (has "output" field) if isinstance(response, dict) and "output" in response: # Make a copy to avoid modifying the original from copy import deepcopy - response_copy = deepcopy(response) + response_copy: Final = deepcopy(response) # Redact content in output array if isinstance(response_copy.get("output"), list): for output_item in response_copy["output"]: @@ -886,8 +881,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac standard_logging_object_copy["response"] = response_copy else: # Standard ModelResponse format - model_response = ModelResponse(choices=[Choices(message=Message(content=redacted_str))]) - model_response_dict = model_response.model_dump() + model_response: Final = ModelResponse(choices=[Choices(message=Message(content=redacted_str))]) + model_response_dict: Final = model_response.model_dump() standard_logging_object_copy["response"] = model_response_dict model_call_details_copy["standard_logging_object"] = standard_logging_object_copy @@ -911,7 +906,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac import litellm from litellm._logging import verbose_logger - all_callbacks = litellm.logging_callback_manager._get_all_callbacks() + all_callbacks: Final = litellm.logging_callback_manager._get_all_callbacks() for callback_obj in all_callbacks: if hasattr(callback_obj, "increment_callback_logging_failure"): @@ -944,8 +939,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac • Keep untyped or text content. • 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 [] + raw_messages: Final[Any] = payload.get("messages", []) + messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else [] verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages)) if messages: @@ -976,8 +971,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac • Keep untyped or text content. • 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 [] + raw_messages: Final[Any] = payload.get("messages", []) + messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else [] verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages)) if messages: @@ -1024,7 +1019,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac return True if "file" in content: return False - ctype = content.get("type") + ctype: Final = content.get("type") return not (isinstance(ctype, str) and ctype != "text") def _process_messages( @@ -1032,7 +1027,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER, ) -> list[dict[str, Any]]: - filtered_messages: list[dict[str, Any]] = [] + filtered_messages: Final[list[dict[str, Any]]] = [] for msg in messages: if not isinstance(msg, dict): continue diff --git a/litellm/integrations/custom_sso_handler.py b/litellm/integrations/custom_sso_handler.py index 202e488e0e4..345a9051c06 100644 --- a/litellm/integrations/custom_sso_handler.py +++ b/litellm/integrations/custom_sso_handler.py @@ -1,3 +1,5 @@ +from typing import Final + from fastapi import Request from fastapi_sso.sso.base import OpenID @@ -29,7 +31,7 @@ class CustomSSOLoginHandler(CustomLogger): feature_name="Custom UI SSO", ) - request_headers_dict = dict(request.headers) + request_headers_dict: Final = dict(request.headers) return OpenID( id=request_headers_dict.get("x-litellm-user-id"), email=request_headers_dict.get("x-litellm-user-email"), diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index ce6dda96820..fd4faeed41a 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 +from typing import Any, Final import httpx from httpx import Response @@ -66,17 +66,17 @@ from ..additional_logging_utils import AdditionalLoggingUtils # specify what ServiceTypes are logged as success events to DD. (We don't want to spam DD traces with large number of service types) -DD_LOGGED_SUCCESS_SERVICE_TYPES = [ +DD_LOGGED_SUCCESS_SERVICE_TYPES: Final = [ ServiceTypes.RESET_BUDGET_JOB, ] def _resolve_dd_batch_size() -> int: - raw = os.getenv("DD_BATCH_SIZE") + raw: Final = os.getenv("DD_BATCH_SIZE") if raw is None: return DD_MAX_BATCH_SIZE try: - value = int(raw) + value: Final = int(raw) except ValueError: verbose_logger.warning( "Datadog: ignoring invalid DD_BATCH_SIZE=%r, using %s", @@ -136,14 +136,14 @@ class DataDogLogger( ######################################################### # Handle datadog_params set as litellm.datadog_params ######################################################### - dict_datadog_params = self._get_datadog_params() + dict_datadog_params: Final = self._get_datadog_params() kwargs.update(dict_datadog_params) self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) # Configure DataDog endpoint (Agent or Direct API) # Prefer explicit kwargs, then fall back to env vars - resolved_agent_host = dd_agent_host or os.getenv("LITELLM_DD_AGENT_HOST") + resolved_agent_host: Final = dd_agent_host or os.getenv("LITELLM_DD_AGENT_HOST") if resolved_agent_host: self._configure_dd_agent( dd_agent_host=resolved_agent_host, @@ -159,7 +159,7 @@ class DataDogLogger( ) # Optional override for testing - dd_base_url = get_datadog_base_url_from_env() + dd_base_url: Final = get_datadog_base_url_from_env() if dd_base_url: self.intake_url = f"{dd_base_url}/api/v2/logs" self.sync_client = _get_httpx_client() @@ -205,7 +205,7 @@ class DataDogLogger( dd_api_key: Datadog API key. Falls back to DD_API_KEY env var when allow_env_credentials is True. Optional when using agent. allow_env_credentials: When False, never read the API key from DD_API_KEY env var. """ - resolved_port = dd_agent_port or os.getenv("LITELLM_DD_AGENT_PORT", "10518") # default port for logs + resolved_port: Final = dd_agent_port or os.getenv("LITELLM_DD_AGENT_PORT", "10518") # default port for logs self.intake_url = f"http://{dd_agent_host}:{resolved_port}/api/v2/logs" self.DD_API_KEY = dd_api_key or ( os.getenv("DD_API_KEY") if allow_env_credentials else None @@ -229,8 +229,8 @@ class DataDogLogger( Raises: Exception: If required credentials are not provided via args or env vars """ - resolved_api_key = dd_api_key or (os.getenv("DD_API_KEY") if allow_env_credentials else None) - resolved_site = dd_site or os.getenv("DD_SITE") + resolved_api_key: Final = dd_api_key or (os.getenv("DD_API_KEY") if allow_env_credentials else None) + resolved_site: Final = dd_site or os.getenv("DD_SITE") if resolved_api_key is None: raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>") @@ -287,11 +287,11 @@ class DataDogLogger( ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - error_information = StandardLoggingPayloadSetup.get_error_information( + error_information: Final = StandardLoggingPayloadSetup.get_error_information( original_exception=original_exception, traceback_str=traceback_str, ) - _code = error_information.get("error_code") or "" + _code: Final = error_information.get("error_code") or "" status_code: int | None = None if _code and str(_code).strip().isdigit(): status_code = int(_code) @@ -303,7 +303,7 @@ class DataDogLogger( LiteLLMProxyRequestSetup, ) - _meta = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( + _meta: Final = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( user_api_key_dict=user_api_key_dict ) user_context = dict(_meta) if isinstance(_meta, dict) else _meta @@ -318,7 +318,7 @@ class DataDogLogger( if hasattr(user_api_key_dict, "end_user_id"): user_context["end_user_id"] = getattr(user_api_key_dict, "end_user_id", None) - message_payload: DatadogProxyFailureHookJsonMessage = { + message_payload: Final[DatadogProxyFailureHookJsonMessage] = { "exception": error_information.get("error_message") or str(original_exception), "error_class": error_information.get("error_class") or original_exception.__class__.__name__, "status_code": status_code, @@ -326,7 +326,7 @@ class DataDogLogger( "user_api_key_dict": user_context, } - dd_payload = DatadogPayload( + dd_payload: Final = DatadogPayload( ddsource=get_datadog_source(), ddtags=",".join(get_datadog_tags()), hostname=get_datadog_hostname(), @@ -358,7 +358,7 @@ class DataDogLogger( verbose_logger.exception("Datadog: log_queue does not exist") return - batch_to_send = self.log_queue[:] + batch_to_send: Final = self.log_queue[:] self.log_queue = [] try: @@ -371,7 +371,7 @@ class DataDogLogger( if self.is_mock_mode: verbose_logger.debug("[DATADOG MOCK] Mock mode enabled - API calls will be intercepted") - undelivered = await self._send_with_413_split(batch_to_send) + undelivered: Final = await self._send_with_413_split(batch_to_send) if undelivered: self.log_queue = undelivered + self.log_queue @@ -395,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: Final[list[list]] = [batch] while pending: chunk = pending.pop() if not chunk: @@ -455,7 +455,7 @@ class DataDogLogger( if len(chunk) > DD_MAX_BATCH_SIZE: return True - payload_size_bytes = len(safe_dumps(chunk).encode("utf-8")) + payload_size_bytes: Final = len(safe_dumps(chunk).encode("utf-8")) return payload_size_bytes > DD_MAX_PAYLOAD_SIZE_BYTES async def flush_queue(self): @@ -493,12 +493,12 @@ class DataDogLogger( ) # Build headers - headers = {} + headers: Final = {} # Add API key if available (required for direct API, optional for agent) if self.DD_API_KEY: headers["DD-API-KEY"] = self.DD_API_KEY - response = self.sync_client.post( + response: Final = self.sync_client.post( url=self.intake_url, json=dd_payload, # type: ignore headers=headers, @@ -518,7 +518,7 @@ class DataDogLogger( verbose_logger.exception("Datadog Layer Error - %s\n%s", e, traceback.format_exc()) async def _log_async_event(self, kwargs, response_obj, start_time, end_time): - dd_payload = self.create_datadog_logging_payload( + dd_payload: Final = self.create_datadog_logging_payload( kwargs=kwargs, response_obj=response_obj, start_time=start_time, @@ -538,9 +538,9 @@ class DataDogLogger( ) -> DatadogPayload: from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - json_payload = safe_dumps(standard_logging_object) + json_payload: Final = safe_dumps(standard_logging_object) verbose_logger.debug("Datadog: Logger - Logging payload = %s", json_payload) - dd_payload = DatadogPayload( + dd_payload: Final = DatadogPayload( ddsource=get_datadog_source(), ddtags=",".join(get_datadog_tags(standard_logging_object=standard_logging_object)), hostname=get_datadog_hostname(), @@ -571,7 +571,7 @@ class DataDogLogger( DatadogPayload: defined in types.py """ - standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: raise ValueError("standard_logging_object not found in kwargs") @@ -582,7 +582,7 @@ class DataDogLogger( # Build the initial payload self.truncate_standard_logging_payload_content(standard_logging_object) - dd_payload = self._create_datadog_logging_payload_helper( + dd_payload: Final = self._create_datadog_logging_payload_helper( standard_logging_object=standard_logging_object, status=status, ) @@ -602,10 +602,10 @@ class DataDogLogger( from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - compressed_data = gzip.compress(safe_dumps(data).encode("utf-8")) + compressed_data: Final = gzip.compress(safe_dumps(data).encode("utf-8")) # Build headers - headers = { + headers: Final = { "Content-Encoding": "gzip", "Content-Type": "application/json", } @@ -614,7 +614,7 @@ class DataDogLogger( if self.DD_API_KEY: headers["DD-API-KEY"] = self.DD_API_KEY - response = await self.async_client.post( + response: Final = await self.async_client.post( url=self.intake_url, data=compressed_data, # type: ignore headers=headers, @@ -636,12 +636,12 @@ class DataDogLogger( - example - Redis is failing / erroring, will be logged on DataDog """ try: - _payload_dict = payload.model_dump() + _payload_dict: Final = payload.model_dump() _payload_dict.update(event_metadata or {}) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - _dd_message_str = safe_dumps(_payload_dict) - _dd_payload = DatadogPayload( + _dd_message_str: Final = safe_dumps(_payload_dict) + _dd_payload: Final = DatadogPayload( ddsource=get_datadog_source(), ddtags=",".join(get_datadog_tags()), hostname=get_datadog_hostname(), @@ -674,13 +674,13 @@ class DataDogLogger( if payload.service not in DD_LOGGED_SUCCESS_SERVICE_TYPES: return - _payload_dict = payload.model_dump() + _payload_dict: Final = payload.model_dump() _payload_dict.update(event_metadata or {}) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - _dd_message_str = safe_dumps(_payload_dict) - _dd_payload = DatadogPayload( + _dd_message_str: Final = safe_dumps(_payload_dict) + _dd_payload: Final = DatadogPayload( ddsource=get_datadog_source(), ddtags=",".join(get_datadog_tags()), hostname=get_datadog_hostname(), @@ -708,14 +708,14 @@ class DataDogLogger( (Not Recommended) If you want this to get logged set `litellm.datadog_use_v1 = True` """ - litellm_params = kwargs.get("litellm_params", {}) - metadata = litellm_params.get("metadata", {}) or {} # if litellm_params['metadata'] == None - messages = kwargs.get("messages") - optional_params = kwargs.get("optional_params", {}) - call_type = kwargs.get("call_type", "litellm.completion") - cache_hit = kwargs.get("cache_hit", False) + litellm_params: Final = kwargs.get("litellm_params", {}) + metadata: Final = litellm_params.get("metadata", {}) or {} # if litellm_params['metadata'] == None + messages: Final = kwargs.get("messages") + optional_params: Final = kwargs.get("optional_params", {}) + call_type: Final = kwargs.get("call_type", "litellm.completion") + cache_hit: Final = kwargs.get("cache_hit", False) usage = response_obj["usage"] - id = response_obj.get("id", str(uuid.uuid4())) + id: Final = response_obj.get("id", str(uuid.uuid4())) usage = dict(usage) try: response_time = (end_time - start_time).total_seconds() * 1000 @@ -730,7 +730,7 @@ class DataDogLogger( # 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 = {} + clean_metadata: Final = {} if isinstance(metadata, dict): for key, value in metadata.items(): # clean litellm metadata before logging @@ -744,7 +744,7 @@ class DataDogLogger( clean_metadata[key] = value # Build the initial payload - payload = { + payload: Final = { "id": id, "call_type": call_type, "cache_hit": cache_hit, @@ -763,11 +763,11 @@ class DataDogLogger( from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - json_payload = safe_dumps(payload) + json_payload: Final = safe_dumps(payload) verbose_logger.debug("Datadog: Logger - Logging payload = %s", json_payload) - dd_payload = DatadogPayload( + dd_payload: Final = DatadogPayload( ddsource=get_datadog_source(), ddtags=",".join(get_datadog_tags()), hostname=get_datadog_hostname(), @@ -784,12 +784,12 @@ class DataDogLogger( """Attach Datadog APM trace context if one is active.""" try: - trace_context = self._get_active_trace_context() + trace_context: Final = self._get_active_trace_context() if trace_context is None: return dd_payload["dd.trace_id"] = trace_context["trace_id"] - span_id = trace_context.get("span_id") + span_id: Final = trace_context.get("span_id") if span_id is not None: dd_payload["dd.span_id"] = span_id except Exception: @@ -798,24 +798,24 @@ class DataDogLogger( def _get_active_trace_context(self) -> dict[str, str] | None: try: current_span = None - current_span_fn = getattr(tracer, "current_span", None) + current_span_fn: Final = getattr(tracer, "current_span", None) if callable(current_span_fn): current_span = current_span_fn() if current_span is None: - current_root_span_fn = getattr(tracer, "current_root_span", None) + current_root_span_fn: Final = getattr(tracer, "current_root_span", None) if callable(current_root_span_fn): current_span = current_root_span_fn() if current_span is None: return None - trace_id = getattr(current_span, "trace_id", None) + trace_id: Final = getattr(current_span, "trace_id", None) if trace_id is None: return None - span_id = getattr(current_span, "span_id", None) - trace_context: dict[str, str] = {"trace_id": str(trace_id)} + span_id: Final = getattr(current_span, "span_id", None) + trace_context: Final[dict[str, str]] = {"trace_id": str(trace_id)} if span_id is not None: trace_context["span_id"] = str(span_id) return trace_context @@ -831,13 +831,13 @@ class DataDogLogger( create_dummy_standard_logging_payload, ) - standard_logging_object = create_dummy_standard_logging_payload() - dd_payload = self._create_datadog_logging_payload_helper( + standard_logging_object: Final = create_dummy_standard_logging_payload() + dd_payload: Final = self._create_datadog_logging_payload_helper( standard_logging_object=standard_logging_object, status=DataDogStatus.INFO, ) - log_queue = [dd_payload] - response = await self.async_send_compressed_data(log_queue) + log_queue: Final = [dd_payload] + response: Final = await self.async_send_compressed_data(log_queue) try: response.raise_for_status() return IntegrationHealthCheckStatus( diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index 21a289877c2..b30700e98f2 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, cast +from typing import Any, Final, cast from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -27,7 +27,7 @@ from litellm.types.utils import StandardLoggingPayload # request_tags / metadata cannot overwrite these, even when the key is # allowlisted via cost_tag_keys, because that would let an authenticated caller # spoof cost attribution (e.g. request_tags=["team:victim-team"]). -_RESERVED_TAG_KEYS: frozenset = frozenset( +_RESERVED_TAG_KEYS: Final[frozenset] = frozenset( { "env", "service", @@ -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: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -90,11 +90,11 @@ class DatadogCostManagementLogger(CustomBatchLogger): if not self.log_queue: return - batch_to_send = self.log_queue[:] + batch_to_send: Final = self.log_queue[:] self.log_queue = [] try: - aggregated_entries = self._aggregate_costs(batch_to_send) + aggregated_entries: Final = self._aggregate_costs(batch_to_send) if not aggregated_entries: verbose_logger.debug( "Datadog Cost Management: batch produced no aggregable entries; dropping %d log(s) from queue.", @@ -111,7 +111,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): Aggregates costs by Provider, Model, and Date. Returns a list of DatadogFOCUSCostEntry. """ - aggregator: dict[tuple[str, str, str, tuple[tuple[str, str], ...]], DatadogFOCUSCostEntry] = {} + aggregator: Final[dict[tuple[str, str, str, tuple[tuple[str, str], ...]], DatadogFOCUSCostEntry]] = {} for log in logs: try: @@ -165,7 +165,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): return list(aggregator.values()) def _extract_tags(self, log: StandardLoggingPayload) -> dict[str, str]: - tags: dict[str, str] = { + tags: Final[dict[str, str]] = { "env": get_datadog_env(), "service": get_datadog_service(), "host": get_datadog_hostname(), @@ -180,12 +180,12 @@ 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: Final[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"): tags["user"] = str(metadata["user_api_key_alias"]) - team_tag = ( + team_tag: Final = ( metadata.get("user_api_key_team_alias") or metadata.get("team_alias") or metadata.get("user_api_key_team_id") @@ -200,7 +200,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): # Reserved keys are hard-blocked here regardless of allowlist membership — # see _RESERVED_TAG_KEYS for the rationale. if self.cost_tag_keys: - allow = set(self.cost_tag_keys) + allow: Final = set(self.cost_tag_keys) for rt in log.get("request_tags") or []: if not isinstance(rt, str) or ":" not in rt: continue @@ -240,16 +240,16 @@ class DatadogCostManagementLogger(CustomBatchLogger): if not self.dd_api_key or not self.dd_app_key: return - headers = { + headers: Final = { "Content-Type": "application/json", "DD-API-KEY": self.dd_api_key, "DD-APPLICATION-KEY": self.dd_app_key, } # The API endpoint expects a list of objects directly in the body (file content behavior) - data_json = safe_dumps(payload) + data_json: Final = safe_dumps(payload) - response = await self.async_client.put(self.upload_url, content=data_json, headers=headers) + response: Final = await self.async_client.put(self.upload_url, content=data_json, headers=headers) response.raise_for_status() diff --git a/litellm/integrations/datadog/datadog_handler.py b/litellm/integrations/datadog/datadog_handler.py index 6a86803ed46..2450382a192 100644 --- a/litellm/integrations/datadog/datadog_handler.py +++ b/litellm/integrations/datadog/datadog_handler.py @@ -3,6 +3,7 @@ from __future__ import annotations import os +from typing import Final from litellm.types.utils import StandardLoggingPayload @@ -45,7 +46,7 @@ def get_datadog_tags( comma: ",".join(get_datadog_tags(...)). """ - base_tags = { + base_tags: Final = { "env": get_datadog_env(), "service": get_datadog_service(), "version": os.getenv("DD_VERSION", "unknown"), @@ -53,15 +54,15 @@ def get_datadog_tags( "POD_NAME": get_datadog_pod_name(), } - tags: list[str] = [f"{k}:{v}" for k, v in base_tags.items()] + tags: Final[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 [] + request_tags: Final = standard_logging_object.get("request_tags", []) or [] tags.extend(f"request_tag:{tag}" for tag in request_tags) # Add Team Tag - metadata = standard_logging_object.get("metadata", {}) or {} - team_tag = ( + metadata: Final = standard_logging_object.get("metadata", {}) or {} + team_tag: Final = ( metadata.get("user_api_key_team_alias") or metadata.get("team_alias") or metadata.get("user_api_key_team_id") diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 8d7ed415315..704f0323e95 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -10,7 +10,7 @@ import asyncio import json import os from datetime import datetime -from typing import Any, Literal +from typing import Any, Final, Literal import httpx @@ -58,7 +58,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): # Configure DataDog endpoint (Agent or Direct API) # Use LITELLM_DD_AGENT_HOST to avoid conflicts with ddtrace's DD_AGENT_HOST # Check for agent mode FIRST - agent mode doesn't require DD_API_KEY or DD_SITE - dd_agent_host = os.getenv("LITELLM_DD_AGENT_HOST") + dd_agent_host: Final = os.getenv("LITELLM_DD_AGENT_HOST") self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) self.DD_API_KEY = os.getenv("DD_API_KEY") @@ -74,7 +74,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): self._configure_dd_direct_api() # Optional override for testing - dd_base_url = get_datadog_base_url_from_env() + dd_base_url: Final = get_datadog_base_url_from_env() if dd_base_url: self.intake_url = f"{dd_base_url}/api/intake/llm-obs/v1/trace/spans" @@ -85,7 +85,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): ######################################################### # Handle datadog_llm_observability_params set as litellm.datadog_llm_observability_params ######################################################### - dict_datadog_llm_obs_params = self._get_datadog_llm_obs_params() + dict_datadog_llm_obs_params: Final = self._get_datadog_llm_obs_params() kwargs.update(dict_datadog_llm_obs_params) CustomBatchLogger.__init__(self, **kwargs, flush_lock=self.flush_lock) except Exception as e: @@ -100,7 +100,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): # Reference: https://docs.datadoghq.com/llm_observability/setup/sdk/#agent-setup # Use specific port for LLM Obs (Trace Agent) to avoid conflict with Logs Agent (10518) - agent_port = os.getenv("LITELLM_DD_LLM_OBS_PORT", "8126") + agent_port: Final = os.getenv("LITELLM_DD_LLM_OBS_PORT", "8126") self.DD_SITE = "localhost" # Not used for URL construction in agent mode self.intake_url = f"http://{dd_agent_host}:{agent_port}/api/intake/llm-obs/v1/trace/spans" verbose_logger.debug("DataDogLLMObs: Using DD Agent at %s", self.intake_url) @@ -138,7 +138,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: verbose_logger.debug("DataDogLLMObs: Logging success event for model %s", kwargs.get("model", "unknown")) - payload = self.create_llm_obs_payload(kwargs, start_time, end_time) + payload: Final = self.create_llm_obs_payload(kwargs, start_time, end_time) verbose_logger.debug("DataDogLLMObs: Payload: %s", payload) self.log_queue.append(payload) @@ -150,7 +150,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: verbose_logger.debug("DataDogLLMObs: Logging failure event for model %s", kwargs.get("model", "unknown")) - payload = self.create_llm_obs_payload(kwargs, start_time, end_time) + payload: Final = self.create_llm_obs_payload(kwargs, start_time, end_time) verbose_logger.debug("DataDogLLMObs: Payload: %s", payload) self.log_queue.append(payload) @@ -170,7 +170,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): verbose_logger.debug("[DATADOG MOCK] Mock mode enabled - API calls will be intercepted") # Prepare the payload - payload = { + payload: Final = { "data": DDIntakePayload( type="span", attributes=DDSpanAttributes( @@ -189,13 +189,13 @@ class DataDogLLMObsLogger(CustomBatchLogger): except Exception as debug_error: verbose_logger.debug("payload serialization failed: %s", str(debug_error)) - json_payload = safe_dumps(payload) + json_payload: Final = safe_dumps(payload) - headers = {"Content-Type": "application/json"} + headers: Final = {"Content-Type": "application/json"} if self.DD_API_KEY: headers["DD-API-KEY"] = self.DD_API_KEY - response = await self.async_client.post( + response: Final = await self.async_client.post( url=self.intake_url, content=json_payload, headers=headers, @@ -217,30 +217,30 @@ class DataDogLLMObsLogger(CustomBatchLogger): verbose_logger.exception("DataDogLLMObs: Error sending batch - %s", e) 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") + standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") if standard_logging_payload is None: raise Exception("DataDogLLMObs: standard_logging_object is not set") messages = standard_logging_payload["messages"] messages = self._ensure_string_content(messages=messages) - metadata = kwargs.get("litellm_params", {}).get("metadata", {}) + metadata: Final = kwargs.get("litellm_params", {}).get("metadata", {}) - input_meta = InputMeta(messages=handle_any_messages_to_chat_completion_str_messages_conversion(messages)) - output_meta = OutputMeta( + input_meta: Final = InputMeta(messages=handle_any_messages_to_chat_completion_str_messages_conversion(messages)) + output_meta: Final = OutputMeta( messages=self._get_response_messages( standard_logging_payload=standard_logging_payload, call_type=standard_logging_payload.get("call_type"), ) ) - error_info = self._assemble_error_info(standard_logging_payload) + error_info: Final = self._assemble_error_info(standard_logging_payload) metadata_parent_id: str | None = None if isinstance(metadata, dict): metadata_parent_id = metadata.get("parent_id") - meta = Meta( + meta: Final = Meta( kind=self._get_datadog_span_kind(standard_logging_payload.get("call_type"), metadata_parent_id), input=input_meta, output=output_meta, @@ -249,7 +249,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): ) # Calculate metrics (you may need to adjust these based on available data) - metrics = LLMMetrics( + metrics: Final = LLMMetrics( input_tokens=float(standard_logging_payload.get("prompt_tokens", 0)), output_tokens=float(standard_logging_payload.get("completion_tokens", 0)), total_tokens=float(standard_logging_payload.get("total_tokens", 0)), @@ -257,7 +257,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): time_to_first_token=self._get_time_to_first_token_seconds(standard_logging_payload), ) - payload: LLMObsPayload = LLMObsPayload( + payload: Final[LLMObsPayload] = LLMObsPayload( parent_id=metadata_parent_id if metadata_parent_id else "undefined", trace_id=standard_logging_payload.get("trace_id", str(uuid.uuid4())), span_id=metadata.get("span_id", str(uuid.uuid4())), @@ -270,7 +270,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): tags=get_datadog_tags(standard_logging_object=standard_logging_payload), ) - apm_trace_id = self._get_apm_trace_id() + apm_trace_id: Final = self._get_apm_trace_id() if apm_trace_id is not None: payload["apm_id"] = apm_trace_id @@ -279,11 +279,11 @@ class DataDogLLMObsLogger(CustomBatchLogger): 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) + current_span_fn: Final = getattr(tracer, "current_span", None) if callable(current_span_fn): - current_span = current_span_fn() + current_span: Final = current_span_fn() if current_span is not None: - trace_id = getattr(current_span, "trace_id", None) + trace_id: Final = getattr(current_span, "trace_id", None) if trace_id is not None: return str(trace_id) except Exception: @@ -299,7 +299,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): if standard_logging_payload.get("status") == "failure": # Try to get structured error information first - error_information: StandardLoggingPayloadErrorInformation | None = standard_logging_payload.get( + error_information: Final[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: 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") + start_time: Final[float | None] = standard_logging_payload.get("startTime") + completion_start_time: Final[float | None] = standard_logging_payload.get("completionStartTime") + end_time: Final[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 @@ -372,7 +372,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): try: # Safely extract message from response_obj, handle failure cases if isinstance(response_obj, dict) and "choices" in response_obj: - choices = response_obj["choices"] + choices: Final = response_obj["choices"] if choices and len(choices) > 0 and "message" in choices[0]: return [choices[0]["message"]] return [] @@ -499,7 +499,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): """ Fields to track in DD LLM Observability metadata from litellm standard logging payload """ - _metadata: dict[str, Any] = { + _metadata: Final[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"), @@ -514,20 +514,20 @@ class DataDogLLMObsLogger(CustomBatchLogger): ######################################################### # Add latency metrics to metadata ######################################################### - latency_metrics = self._get_latency_metrics(standard_logging_payload) + latency_metrics: Final = self._get_latency_metrics(standard_logging_payload) _metadata.update({"latency_metrics": dict(latency_metrics)}) ######################################################### # Add spend metrics to metadata ######################################################### - spend_metrics = self._get_spend_metrics(standard_logging_payload) + spend_metrics: Final = self._get_spend_metrics(standard_logging_payload) _metadata.update({"spend_metrics": dict(spend_metrics)}) ## extract tool calls and add to metadata - tool_call_metadata = self._extract_tool_call_metadata(standard_logging_payload) + tool_call_metadata: Final = self._extract_tool_call_metadata(standard_logging_payload) _metadata.update(tool_call_metadata) - _standard_logging_metadata: dict = dict(standard_logging_payload.get("metadata", {})) or {} + _standard_logging_metadata: Final[dict] = dict(standard_logging_payload.get("metadata", {})) or {} _metadata.update(_standard_logging_metadata) return _metadata @@ -535,21 +535,21 @@ class DataDogLLMObsLogger(CustomBatchLogger): """ Get the latency metrics from the standard logging payload """ - latency_metrics: DDLLMObsLatencyMetrics = DDLLMObsLatencyMetrics() + latency_metrics: Final[DDLLMObsLatencyMetrics] = DDLLMObsLatencyMetrics() # Add latency metrics to metadata # Time to first token (convert from seconds to milliseconds for consistency) - time_to_first_token_seconds = self._get_time_to_first_token_seconds(standard_logging_payload) + time_to_first_token_seconds: Final = self._get_time_to_first_token_seconds(standard_logging_payload) if time_to_first_token_seconds > 0: latency_metrics["time_to_first_token_ms"] = time_to_first_token_seconds * 1000 # LiteLLM overhead time - hidden_params = standard_logging_payload.get("hidden_params", {}) - litellm_overhead_ms = hidden_params.get("litellm_overhead_time_ms") + hidden_params: Final = standard_logging_payload.get("hidden_params", {}) + litellm_overhead_ms: Final = hidden_params.get("litellm_overhead_time_ms") if litellm_overhead_ms is not None: latency_metrics["litellm_overhead_time_ms"] = litellm_overhead_ms # Guardrail overhead latency - guardrail_info: list[StandardLoggingGuardrailInformation] | None = standard_logging_payload.get( + guardrail_info: Final[list[StandardLoggingGuardrailInformation] | None] = standard_logging_payload.get( "guardrail_information" ) if guardrail_info is not None: @@ -581,7 +581,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): return True # Fallback to model_parameters.stream for original request parameters - model_params = standard_logging_payload.get("model_parameters", {}) + model_params: Final = standard_logging_payload.get("model_parameters", {}) if isinstance(model_params, dict): stream_value = model_params.get("stream") if stream_value is True: @@ -594,21 +594,21 @@ class DataDogLLMObsLogger(CustomBatchLogger): """ Get the spend metrics from the standard logging payload """ - spend_metrics: DDLLMObsSpendMetrics = DDLLMObsSpendMetrics() + spend_metrics: Final[DDLLMObsSpendMetrics] = DDLLMObsSpendMetrics() # send response cost spend_metrics["response_cost"] = standard_logging_payload.get("response_cost", 0.0) # Get budget information from metadata - metadata = standard_logging_payload.get("metadata", {}) + metadata: Final = standard_logging_payload.get("metadata", {}) # API key max budget - user_api_key_max_budget = metadata.get("user_api_key_max_budget") + user_api_key_max_budget: Final = metadata.get("user_api_key_max_budget") if user_api_key_max_budget is not None: spend_metrics["user_api_key_max_budget"] = float(user_api_key_max_budget) # API key spend - user_api_key_spend = metadata.get("user_api_key_spend") + user_api_key_spend: Final = metadata.get("user_api_key_spend") if user_api_key_spend is not None: try: spend_metrics["user_api_key_spend"] = float(user_api_key_spend) @@ -616,7 +616,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): verbose_logger.debug("Invalid user_api_key_spend value: %s", user_api_key_spend) # API key budget reset datetime - user_api_key_budget_reset_at = metadata.get("user_api_key_budget_reset_at") + user_api_key_budget_reset_at: Final = metadata.get("user_api_key_budget_reset_at") if user_api_key_budget_reset_at is not None: try: from datetime import datetime, timezone @@ -654,7 +654,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): This bypasses the lossy string conversion when tool calls are present, allowing complex nested tool_calls objects to be preserved for Datadog. """ - processed = [] + processed: Final = [] for msg in messages: if isinstance(msg, dict): # Preserve messages with tool_calls or tool role as-is @@ -677,7 +677,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): Similar to OpenTelemetry's implementation but adapted for Datadog's format. """ - kv_pairs: dict[str, Any] = {} + kv_pairs: Final[dict[str, Any]] = {} for idx, tool_call in enumerate(tool_calls): try: # Extract tool call ID @@ -716,11 +716,11 @@ class DataDogLLMObsLogger(CustomBatchLogger): """ Extract tool call information from both input messages and response for Datadog metadata. """ - tool_call_metadata: dict[str, Any] = {} + tool_call_metadata: Final[dict[str, Any]] = {} try: # Extract tool calls from input messages - messages = standard_logging_payload.get("messages", []) + messages: Final = standard_logging_payload.get("messages", []) if messages and isinstance(messages, list): for message in messages: if isinstance(message, dict) and "tool_calls" in message: @@ -732,9 +732,9 @@ class DataDogLLMObsLogger(CustomBatchLogger): tool_call_metadata[f"input_{key}"] = value # Extract tool calls from response - response_obj = standard_logging_payload.get("response") + response_obj: Final = standard_logging_payload.get("response") if response_obj and isinstance(response_obj, dict): - choices = response_obj.get("choices", []) + choices: Final = response_obj.get("choices", []) for choice in choices: if isinstance(choice, dict): message = choice.get("message") diff --git a/litellm/integrations/datadog/datadog_metrics.py b/litellm/integrations/datadog/datadog_metrics.py index c33c44e4249..37421126985 100644 --- a/litellm/integrations/datadog/datadog_metrics.py +++ b/litellm/integrations/datadog/datadog_metrics.py @@ -3,6 +3,7 @@ import gzip import os import time from datetime import datetime +from typing import Final from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -65,7 +66,7 @@ class DatadogMetricsLogger(CustomBatchLogger): Builds the list of tags for a Datadog metric point """ # Base tags - tags = [ + tags: Final = [ f"env:{get_datadog_env()}", f"service:{get_datadog_service()}", f"version:{os.getenv('DD_VERSION', 'unknown')}", @@ -87,8 +88,8 @@ class DatadogMetricsLogger(CustomBatchLogger): tags.append(f"status_code:{status_code}") # Extract team tag - metadata = log.get("metadata", {}) or {} - team_tag = ( + metadata: Final = log.get("metadata", {}) or {} + team_tag: Final = ( metadata.get("user_api_key_team_alias") or metadata.get("team_alias") # type: ignore or metadata.get("user_api_key_team_id") @@ -109,17 +110,17 @@ class DatadogMetricsLogger(CustomBatchLogger): """ Extracts latencies and appends Datadog metric series to the queue """ - tags = self._extract_tags(log, status_code=status_code) + tags: Final = self._extract_tags(log, status_code=status_code) # We record metrics with the end_time as the timestamp for the point - end_time_dt = kwargs.get("end_time") or datetime.now() - timestamp = int(end_time_dt.timestamp()) + end_time_dt: Final = kwargs.get("end_time") or datetime.now() + timestamp: Final = int(end_time_dt.timestamp()) # 1. Total Request Latency Metric (End to End) - start_time_dt = kwargs.get("start_time") + start_time_dt: Final = kwargs.get("start_time") if start_time_dt and end_time_dt: - total_duration = (end_time_dt - start_time_dt).total_seconds() - series_total_latency: DatadogMetricSeries = { + total_duration: Final = (end_time_dt - start_time_dt).total_seconds() + series_total_latency: Final[DatadogMetricSeries] = { "metric": "litellm.request.total_latency", "type": 3, # gauge "points": [{"timestamp": timestamp, "value": total_duration}], @@ -128,10 +129,10 @@ class DatadogMetricsLogger(CustomBatchLogger): self.log_queue.append(series_total_latency) # 2. LLM API Latency Metric (Provider alone) - api_call_start_time = kwargs.get("api_call_start_time") + api_call_start_time: Final = kwargs.get("api_call_start_time") if api_call_start_time and end_time_dt: - llm_api_duration = (end_time_dt - api_call_start_time).total_seconds() - series_llm_latency: DatadogMetricSeries = { + llm_api_duration: Final = (end_time_dt - api_call_start_time).total_seconds() + series_llm_latency: Final[DatadogMetricSeries] = { "metric": "litellm.llm_api.latency", "type": 3, # gauge "points": [{"timestamp": timestamp, "value": llm_api_duration}], @@ -140,11 +141,11 @@ class DatadogMetricsLogger(CustomBatchLogger): self.log_queue.append(series_llm_latency) # 3. LiteLLM Overhead Latency Metric (total - llm_api time) - hidden_params = log.get("hidden_params", {}) or {} - litellm_overhead_time_ms = hidden_params.get("litellm_overhead_time_ms") + hidden_params: Final = log.get("hidden_params", {}) or {} + litellm_overhead_time_ms: Final = hidden_params.get("litellm_overhead_time_ms") if litellm_overhead_time_ms is not None: - overhead_tags = self._extract_tags(log) # no status_code on latency metric - series_overhead: DatadogMetricSeries = { + overhead_tags: Final = self._extract_tags(log) # no status_code on latency metric + series_overhead: Final[DatadogMetricSeries] = { "metric": "litellm.overhead.latency", "type": 3, # gauge "points": [ @@ -158,7 +159,7 @@ class DatadogMetricsLogger(CustomBatchLogger): self.log_queue.append(series_overhead) # 4. Request Count / Status Code - series_count: DatadogMetricSeries = { + series_count: Final[DatadogMetricSeries] = { "metric": "litellm.llm_api.request_count", "type": 1, # count "points": [{"timestamp": timestamp, "value": 1.0}], @@ -169,7 +170,7 @@ class DatadogMetricsLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: - standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return @@ -184,15 +185,15 @@ class DatadogMetricsLogger(CustomBatchLogger): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: - standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None) + standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: return # Extract status code from error information status_code = "500" # default - error_information = standard_logging_object.get("error_information", {}) or {} - error_code = error_information.get("error_code") # type: ignore + error_information: Final = standard_logging_object.get("error_information", {}) or {} + error_code: Final = error_information.get("error_code") # type: ignore if error_code is not None: status_code = str(error_code) @@ -208,8 +209,8 @@ class DatadogMetricsLogger(CustomBatchLogger): if not self.log_queue: return - batch = self.log_queue.copy() - payload_data: DatadogMetricsPayload = {"series": batch} + batch: Final = self.log_queue.copy() + payload_data: Final[DatadogMetricsPayload] = {"series": batch} try: await self._upload_to_datadog(payload_data) @@ -221,7 +222,7 @@ class DatadogMetricsLogger(CustomBatchLogger): if not self.dd_api_key: return - headers = { + headers: Final = { "Content-Type": "application/json", "DD-API-KEY": self.dd_api_key, } @@ -229,11 +230,11 @@ class DatadogMetricsLogger(CustomBatchLogger): if self.dd_app_key: headers["DD-APPLICATION-KEY"] = self.dd_app_key - json_data = safe_dumps(payload) - compressed_data = gzip.compress(json_data.encode("utf-8")) + json_data: Final = safe_dumps(payload) + compressed_data: Final = gzip.compress(json_data.encode("utf-8")) headers["Content-Encoding"] = "gzip" - response = await self.async_client.post( + response: Final = await self.async_client.post( self.upload_url, content=compressed_data, headers=headers, # type: ignore @@ -251,18 +252,18 @@ class DatadogMetricsLogger(CustomBatchLogger): """ try: # Send a test metric point to Datadog - test_metric_point: DatadogMetricPoint = { + test_metric_point: Final[DatadogMetricPoint] = { "timestamp": int(time.time()), "value": 1.0, } - test_metric_series: DatadogMetricSeries = { + test_metric_series: Final[DatadogMetricSeries] = { "metric": "litellm.health_check", "type": 3, # Gauge "points": [test_metric_point], "tags": ["env:health_check"], } - payload_data: DatadogMetricsPayload = {"series": [test_metric_series]} + payload_data: Final[DatadogMetricsPayload] = {"series": [test_metric_series]} await self._upload_to_datadog(payload_data) diff --git a/litellm/integrations/datadog/datadog_mock_client.py b/litellm/integrations/datadog/datadog_mock_client.py index c50cdc6a019..a90ffbc6512 100644 --- a/litellm/integrations/datadog/datadog_mock_client.py +++ b/litellm/integrations/datadog/datadog_mock_client.py @@ -8,13 +8,15 @@ Usage: Set DATADOG_MOCK=true in environment variables or config to enable mock mode. """ +from typing import Final + from litellm.integrations.mock_client_factory import ( MockClientConfig, create_mock_client_factory, ) # Create mock client using factory -_config = MockClientConfig( +_config: Final = MockClientConfig( name="DATADOG", env_var="DATADOG_MOCK", default_latency_ms=100, diff --git a/litellm/integrations/datadog/datadog_team_handler.py b/litellm/integrations/datadog/datadog_team_handler.py index 53eebe7c505..dd848108968 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, TypedDict +from typing import TYPE_CHECKING, Any, Final, TypedDict from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams @@ -42,10 +42,10 @@ class DataDogHandler: The global (env-var based) DataDogLogger is managed separately by _init_custom_logger_compatible_class via _in_memory_loggers. """ - _credentials = DataDogHandler.get_dynamic_datadog_logging_config( + _credentials: Final = DataDogHandler.get_dynamic_datadog_logging_config( standard_callback_dynamic_params=standard_callback_dynamic_params, ) - credentials_dict = dict(_credentials) + credentials_dict: Final = dict(_credentials) # check if datadog logger is already cached temp_datadog_logger = in_memory_dynamic_logger_cache.get_cache( @@ -71,8 +71,8 @@ class DataDogHandler: """ # When the destination is caller-supplied (dd_agent_host/dd_site), never fall back to the # proxy's DD_API_KEY env var, otherwise it would be sent to a team-controlled host. - allow_env_credentials = credentials.get("dd_agent_host") is None and credentials.get("dd_site") is None - datadog_logger = DataDogLogger( + allow_env_credentials: Final = credentials.get("dd_agent_host") is None and credentials.get("dd_site") is None + datadog_logger: Final = DataDogLogger( dd_api_key=credentials.get("dd_api_key"), dd_site=credentials.get("dd_site"), dd_agent_host=credentials.get("dd_agent_host"), diff --git a/litellm/integrations/deepeval/api.py b/litellm/integrations/deepeval/api.py index 512c74e035c..60639fe2941 100644 --- a/litellm/integrations/deepeval/api.py +++ b/litellm/integrations/deepeval/api.py @@ -1,16 +1,17 @@ # duplicate -> https://github.com/confident-ai/deepeval/blob/main/deepeval/confident/api.py import logging from enum import Enum +from typing import Final import httpx from litellm._logging import verbose_logger -DEEPEVAL_BASE_URL = "https://deepeval.confident-ai.com" -DEEPEVAL_BASE_URL_EU = "https://eu.deepeval.confident-ai.com" -API_BASE_URL = "https://api.confident-ai.com" -API_BASE_URL_EU = "https://eu.api.confident-ai.com" -retryable_exceptions = httpx.HTTPError +DEEPEVAL_BASE_URL: Final = "https://deepeval.confident-ai.com" +DEEPEVAL_BASE_URL_EU: Final = "https://eu.deepeval.confident-ai.com" +API_BASE_URL: Final = "https://api.confident-ai.com" +API_BASE_URL_EU: Final = "https://eu.api.confident-ai.com" +retryable_exceptions: Final = httpx.HTTPError from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, @@ -20,8 +21,8 @@ from litellm.llms.custom_httpx.http_handler import ( def log_retry_error(details): - exception = details.get("exception") - tries = details.get("tries") + exception: Final = details.get("exception") + tries: Final = details.get("tries") if exception: logging.error("Confident AI Error: %s. Retrying: %s time(s)...", exception, tries) else: @@ -78,8 +79,8 @@ class Api: raise e def send_request(self, method: HttpMethods, endpoint: Endpoints, body=None, params=None): - url = f"{self.base_api_url}{endpoint.value}" - res = self._http_request( + url: Final = f"{self.base_api_url}{endpoint.value}" + res: Final = self._http_request( method=method.value, url=url, headers=self._headers, @@ -100,7 +101,7 @@ class Api: if method != HttpMethods.POST: raise Exception("Only POST requests are supported") - url = f"{self.base_api_url}{endpoint.value}" + url: Final = f"{self.base_api_url}{endpoint.value}" try: await self.async_http_handler.post( url=url, diff --git a/litellm/integrations/deepeval/deepeval.py b/litellm/integrations/deepeval/deepeval.py index e194d351b8c..71d149e4014 100644 --- a/litellm/integrations/deepeval/deepeval.py +++ b/litellm/integrations/deepeval/deepeval.py @@ -1,4 +1,5 @@ import os +from typing import Final from litellm._logging import verbose_logger from litellm._uuid import uuid @@ -22,7 +23,7 @@ class DeepEvalLogger(CustomLogger): """Logs litellm traces to DeepEval's platform.""" def __init__(self, *args, **kwargs): - api_key = os.getenv("CONFIDENT_API_KEY") + api_key: Final = os.getenv("CONFIDENT_API_KEY") self.litellm_environment = os.getenv("LITELM_ENVIRONMENT", "development") validate_environment(self.litellm_environment) if not api_key: @@ -47,17 +48,17 @@ class DeepEvalLogger(CustomLogger): await self._async_event_handler(kwargs, response_obj, start_time, end_time, is_success=True) def _prepare_trace_api(self, kwargs, response_obj, start_time, end_time, is_success): - _start_time = to_zod_compatible_iso(start_time) - _end_time = to_zod_compatible_iso(end_time) - _standard_logging_object = kwargs.get("standard_logging_object", {}) - base_api_span = self._create_base_api_span( + _start_time: Final = to_zod_compatible_iso(start_time) + _end_time: Final = to_zod_compatible_iso(end_time) + _standard_logging_object: Final = kwargs.get("standard_logging_object", {}) + base_api_span: Final = self._create_base_api_span( kwargs, standard_logging_object=_standard_logging_object, start_time=_start_time, end_time=_end_time, is_success=is_success, ) - trace_api = self._create_trace_api( + trace_api: Final = self._create_trace_api( base_api_span, standard_logging_object=_standard_logging_object, start_time=_start_time, @@ -75,9 +76,9 @@ class DeepEvalLogger(CustomLogger): return body def _sync_event_handler(self, kwargs, response_obj, start_time, end_time, is_success): - body = self._prepare_trace_api(kwargs, response_obj, start_time, end_time, is_success) + body: Final = self._prepare_trace_api(kwargs, response_obj, start_time, end_time, is_success) try: - response = self.api.send_request( + response: Final = self.api.send_request( method=HttpMethods.POST, endpoint=Endpoints.TRACING_ENDPOINT, body=body, @@ -87,8 +88,8 @@ class DeepEvalLogger(CustomLogger): verbose_logger.debug("DeepEvalLogger: sync_log_failure_event: Api response %s", response) async def _async_event_handler(self, kwargs, response_obj, start_time, end_time, is_success): - body = self._prepare_trace_api(kwargs, response_obj, start_time, end_time, is_success) - response = await self.api.a_send_request( + body: Final = self._prepare_trace_api(kwargs, response_obj, start_time, end_time, is_success) + response: Final = await self.api.a_send_request( method=HttpMethods.POST, endpoint=Endpoints.TRACING_ENDPOINT, body=body, @@ -98,7 +99,7 @@ class DeepEvalLogger(CustomLogger): def _create_base_api_span(self, kwargs, standard_logging_object, start_time, end_time, is_success): # extract usage - usage = standard_logging_object.get("response", {}).get("usage", {}) + usage: Final = standard_logging_object.get("response", {}).get("usage", {}) if is_success: output = ( standard_logging_object.get("response", {}) diff --git a/litellm/integrations/deepeval/utils.py b/litellm/integrations/deepeval/utils.py index 9d65b509fe2..92e2600a3f1 100644 --- a/litellm/integrations/deepeval/utils.py +++ b/litellm/integrations/deepeval/utils.py @@ -1,4 +1,5 @@ from datetime import datetime, timezone +from typing import Final from litellm.integrations.deepeval.types import Environment @@ -9,5 +10,5 @@ def to_zod_compatible_iso(dt: datetime) -> str: def validate_environment(environment: str): if environment not in [env.value for env in Environment]: - valid_values = ", ".join(f'"{env.value}"' for env in Environment) + valid_values: Final = ", ".join(f'"{env.value}"' for env in Environment) raise ValueError(f"Invalid environment: {environment}. Please use one of the following instead: {valid_values}") diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py index b254c0315af..578b7c63871 100644 --- a/litellm/integrations/dotprompt/__init__.py +++ b/litellm/integrations/dotprompt/__init__.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, Final, Optional if TYPE_CHECKING: from litellm.integrations.custom_prompt_management import CustomPromptManagement @@ -11,8 +11,8 @@ from litellm.types.prompts.init_prompts import SupportedPromptIntegrations from .dotprompt_manager import DotpromptManager # Global instances -global_prompt_directory: str | None = None -global_prompt_manager: Optional["PromptManager"] = None +global_prompt_directory: Final[str | None] = None +global_prompt_manager: Final[Optional["PromptManager"]] = None def set_global_prompt_directory(directory: str) -> None: @@ -36,7 +36,7 @@ def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict: from .prompt_manager import PromptManager # Parse the dotprompt content to extract frontmatter and content - temp_manager = PromptManager() + temp_manager: Final = PromptManager() metadata, content = temp_manager._parse_frontmatter(dotprompt_content) # Convert to prompt_data format @@ -47,23 +47,23 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom """ Initialize a prompt from a .prompt file. """ - prompt_directory = getattr(litellm_params, "prompt_directory", None) + prompt_directory: Final = getattr(litellm_params, "prompt_directory", None) prompt_data = getattr(litellm_params, "prompt_data", None) - prompt_id = getattr(litellm_params, "prompt_id", None) + prompt_id: Final = getattr(litellm_params, "prompt_id", None) if prompt_directory: raise ValueError( "Cannot set prompt_directory when working with prompt_initializer. Needs to be a specific dotprompt file" ) - prompt_file = getattr(litellm_params, "prompt_file", None) + prompt_file: Final = getattr(litellm_params, "prompt_file", None) # Handle dotprompt_content from database - dotprompt_content = getattr(litellm_params, "dotprompt_content", None) + dotprompt_content: Final = getattr(litellm_params, "dotprompt_content", None) if dotprompt_content and not prompt_data and not prompt_file: prompt_data = _get_prompt_data_from_dotprompt_content(dotprompt_content) try: - dot_prompt_manager = DotpromptManager( + dot_prompt_manager: Final = DotpromptManager( prompt_directory=prompt_directory, prompt_data=prompt_data, prompt_file=prompt_file, @@ -75,7 +75,7 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom raise e -prompt_initializer_registry = { +prompt_initializer_registry: Final = { SupportedPromptIntegrations.DOT_PROMPT.value: prompt_initializer, } diff --git a/litellm/integrations/dotprompt/dotprompt_manager.py b/litellm/integrations/dotprompt/dotprompt_manager.py index 588ef442378..bedeb803c27 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 +from typing import TYPE_CHECKING, Any, Final from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.integrations.prompt_management_base import PromptManagementClient @@ -125,26 +125,26 @@ class DotpromptManager(CustomPromptManagement): try: # Get the prompt template (versioned or base) - template = self.prompt_manager.get_prompt(prompt_id=prompt_id, version=prompt_version) + template: Final = self.prompt_manager.get_prompt(prompt_id=prompt_id, version=prompt_version) if template is None: - version_str = f" (version {prompt_version})" if prompt_version else "" + version_str: Final = f" (version {prompt_version})" if prompt_version else "" raise ValueError(f"Prompt '{prompt_id}'{version_str} not found in prompt directory") # Render the template with variables (pass version for proper lookup) - rendered_content = self.prompt_manager.render( + rendered_content: Final = self.prompt_manager.render( prompt_id=prompt_id, prompt_variables=prompt_variables, version=prompt_version, ) # Convert rendered content to chat messages - messages = self._convert_to_messages(rendered_content) + messages: Final = self._convert_to_messages(rendered_content) # Extract model from metadata (if specified) - template_model = template.model + template_model: Final = template.model # Extract optional parameters from metadata - optional_params = self._extract_optional_params(template) + optional_params: Final = self._extract_optional_params(template) return PromptManagementClient( prompt_id=prompt_id, @@ -259,14 +259,14 @@ class DotpromptManager(CustomPromptManagement): 3. Already formatted as a single message """ # Clean up the content - content = rendered_content.strip() + content: Final = rendered_content.strip() # Try to parse role-based format (System: ..., User: ..., etc.) - messages = [] + messages: Final = [] current_role = None current_content = [] - lines = content.split("\n") + lines: Final = content.split("\n") for line in lines: line = line.strip() @@ -298,7 +298,7 @@ class DotpromptManager(CustomPromptManagement): # Add the last message if current_role and current_content: - content_text = "\n".join(current_content).strip() + content_text: Final = "\n".join(current_content).strip() if content_text: # Only add if there's actual content messages.append(self._create_message(current_role, content_text)) @@ -321,7 +321,7 @@ class DotpromptManager(CustomPromptManagement): Includes parameters like temperature, max_tokens, etc. """ - optional_params = {} + optional_params: Final = {} # Extract common parameters from metadata if template.optional_params is not None: @@ -341,8 +341,8 @@ class DotpromptManager(CustomPromptManagement): 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", {}) + content: Final = json_data.get("content", "") + metadata: Final = 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: diff --git a/litellm/integrations/dotprompt/prompt_manager.py b/litellm/integrations/dotprompt/prompt_manager.py index 5bfe63e0f41..70bad2f7290 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 +from typing import Any, Final import yaml from jinja2 import DictLoader, select_autoescape @@ -25,7 +25,7 @@ class PromptTemplate: self.template_id = template_id # Extract common metadata fields - restricted_keys = ["model", "input", "output"] + restricted_keys: Final = ["model", "input", "output"] self.model = self.metadata.get("model") self.input_schema = self.metadata.get("input", {}).get("schema", {}) self.output_format = self.metadata.get("output", {}).get("format") @@ -83,7 +83,7 @@ class PromptManager: if not prompt_id: raise ValueError("prompt_id is required when prompt_file is provided") - template = self._load_prompt_file(self.prompt_file, prompt_id) + template: Final = self._load_prompt_file(self.prompt_file, prompt_id) self.prompts[prompt_id] = template # Load prompts from JSON data if provided @@ -95,7 +95,7 @@ class PromptManager: if not self.prompt_directory or not self.prompt_directory.exists(): raise ValueError(f"Prompt directory does not exist: {self.prompt_directory}") - prompt_files = list(self.prompt_directory.glob("*.prompt")) + prompt_files: Final = list(self.prompt_directory.glob("*.prompt")) for prompt_file in prompt_files: try: @@ -148,7 +148,7 @@ class PromptManager: if isinstance(file_path, str): file_path = Path(file_path) - content = file_path.read_text(encoding="utf-8") + content: Final = file_path.read_text(encoding="utf-8") # Split frontmatter and content frontmatter, template_content = self._parse_frontmatter(content) @@ -162,11 +162,11 @@ class PromptManager: 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(.*)$" - match = re.match(frontmatter_pattern, content, re.DOTALL) + frontmatter_pattern: Final = r"^---\s*\n(.*?)\n---\s*\n(.*)$" + match: Final = re.match(frontmatter_pattern, content, re.DOTALL) if match: - frontmatter_yaml = match.group(1) + frontmatter_yaml: Final = match.group(1) template_content = match.group(2) try: @@ -202,14 +202,14 @@ class PromptManager: ValueError: If template rendering fails """ # Get the template (versioned or base) - template = self.get_prompt(prompt_id=prompt_id, version=version) + template: Final = self.get_prompt(prompt_id=prompt_id, version=version) if template is None: - available_prompts = list(self.prompts.keys()) - version_str = f" (version {version})" if version else "" + available_prompts: Final = list(self.prompts.keys()) + version_str: Final = f" (version {version})" if version else "" raise KeyError(f"Prompt '{prompt_id}'{version_str} not found. Available prompts: {available_prompts}") - variables = prompt_variables or {} + variables: Final = prompt_variables or {} # Validate input variables against schema if defined if template.input_schema: @@ -217,8 +217,8 @@ class PromptManager: try: # Create Jinja2 template and render - jinja_template = self.jinja_env.from_string(template.content) - rendered = jinja_template.render(**variables) + jinja_template: Final = self.jinja_env.from_string(template.content) + rendered: Final = jinja_template.render(**variables) return rendered except Exception as e: raise ValueError(f"Error rendering template '{prompt_id}': {e}") @@ -238,7 +238,7 @@ class PromptManager: def _get_python_type(self, schema_type: str) -> type | tuple: """Convert schema type string to Python type.""" - type_mapping: dict[str, type | tuple] = { + type_mapping: Final[dict[str, type | tuple]] = { "string": str, "str": str, "number": (int, float), @@ -268,7 +268,7 @@ class PromptManager: """ if version is not None: # Try versioned prompt first: prompt_id.v{version} - versioned_id = f"{prompt_id}.v{version}" + versioned_id: Final = f"{prompt_id}.v{version}" if versioned_id in self.prompts: return self.prompts[versioned_id] @@ -281,7 +281,7 @@ class PromptManager: def get_prompt_metadata(self, prompt_id: str) -> dict[str, Any] | None: """Get metadata for a specific prompt.""" - template = self.prompts.get(prompt_id) + template: Final = self.prompts.get(prompt_id) return template.metadata if template else None def reload_prompts(self) -> None: @@ -292,7 +292,7 @@ class PromptManager: 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) + template: Final = PromptTemplate(content=content, metadata=metadata or {}, template_id=prompt_id) self.prompts[prompt_id] = template def prompt_file_to_json(self, file_path: str | Path) -> dict[str, Any]: @@ -305,7 +305,7 @@ class PromptManager: Dictionary with 'content' and 'metadata' keys """ file_path = Path(file_path) - content = file_path.read_text(encoding="utf-8") + content: Final = file_path.read_text(encoding="utf-8") # Parse frontmatter and content frontmatter, template_content = self._parse_frontmatter(content) @@ -321,8 +321,8 @@ class PromptManager: Returns: String content in .prompt file format """ - content = prompt_data.get("content", "") - metadata = prompt_data.get("metadata", {}) + content: Final = prompt_data.get("content", "") + metadata: Final = prompt_data.get("metadata", {}) if not metadata: # No metadata, return just the content @@ -331,7 +331,7 @@ class PromptManager: # Convert metadata to YAML frontmatter import yaml - frontmatter_yaml = yaml.dump(metadata, default_flow_style=False) + frontmatter_yaml: Final = yaml.dump(metadata, default_flow_style=False) return f"---\n{frontmatter_yaml}---\n{content}" @@ -341,7 +341,7 @@ class PromptManager: Returns: Dictionary mapping prompt_id to prompt data """ - result = {} + result: Final = {} for prompt_id, template in self.prompts.items(): result[prompt_id] = { "content": template.content, diff --git a/litellm/integrations/dynamodb.py b/litellm/integrations/dynamodb.py index a41130cbab1..38f5924a233 100644 --- a/litellm/integrations/dynamodb.py +++ b/litellm/integrations/dynamodb.py @@ -3,7 +3,7 @@ import os import traceback -from typing import Any +from typing import Any, Final import litellm from litellm._uuid import uuid @@ -32,16 +32,16 @@ class DyanmoDBLogger: # construct payload to send to DynamoDB # follows the same params as langfuse.py - litellm_params = kwargs.get("litellm_params", {}) - metadata = litellm_params.get("metadata", {}) or {} # if litellm_params['metadata'] == None - messages = kwargs.get("messages") - optional_params = kwargs.get("optional_params", {}) - call_type = kwargs.get("call_type", "litellm.completion") - usage = response_obj["usage"] - id = response_obj.get("id", str(uuid.uuid4())) + litellm_params: Final = kwargs.get("litellm_params", {}) + metadata: Final = litellm_params.get("metadata", {}) or {} # if litellm_params['metadata'] == None + messages: Final = kwargs.get("messages") + optional_params: Final = kwargs.get("optional_params", {}) + call_type: Final = kwargs.get("call_type", "litellm.completion") + usage: Final = response_obj["usage"] + id: Final = response_obj.get("id", str(uuid.uuid4())) # Build the initial payload - payload = { + payload: Final = { "id": id, "call_type": call_type, "startTime": start_time, @@ -66,9 +66,9 @@ class DyanmoDBLogger: print_verbose(f"\nDynamoDB Logger - Logging payload = {payload}") # put data in dyanmo DB - table = self.dynamodb.Table(self.table_name) + table: Final = self.dynamodb.Table(self.table_name) # Assuming log_data is a dictionary with log information - response = table.put_item(Item=payload) + response: Final = table.put_item(Item=payload) print_verbose(f"Response from DynamoDB:{response}") diff --git a/litellm/integrations/email_alerting.py b/litellm/integrations/email_alerting.py index 92a56eaaf75..351896425bb 100644 --- a/litellm/integrations/email_alerting.py +++ b/litellm/integrations/email_alerting.py @@ -3,14 +3,15 @@ Functions for sending Email Alerts """ import os +from typing import Final from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.proxy._types import WebhookEvent from litellm.repositories.team_repository import TeamRepository # we use this for the email header, please send a test email if you change this. verify it looks good on email -LITELLM_LOGO_URL = "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" -LITELLM_SUPPORT_CONTACT = "support@berri.ai" +LITELLM_LOGO_URL: Final = "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" +LITELLM_SUPPORT_CONTACT: Final = "support@berri.ai" async def get_all_team_member_emails(team_id: str | None = None) -> list: @@ -22,7 +23,7 @@ async def get_all_team_member_emails(team_id: str | None = None) -> list: if prisma_client is None: raise Exception("Not connected to DB!") - team_row = await TeamRepository(prisma_client).table.find_unique( + team_row: Final = await TeamRepository(prisma_client).table.find_unique( where={ "team_id": team_id, } @@ -31,33 +32,33 @@ async def get_all_team_member_emails(team_id: str | None = None) -> list: if team_row is None: return [] - _team_members = team_row.members_with_roles + _team_members: Final = team_row.members_with_roles verbose_logger.debug( "Email Alerting: Got team members for team_id=%s Team Members: %s", team_id, _team_members, ) - _team_member_user_ids: list[str] = [] + _team_member_user_ids: Final[list[str]] = [] for member in _team_members: if member and isinstance(member, dict): _user_id = member.get("user_id") if _user_id and isinstance(_user_id, str): _team_member_user_ids.append(_user_id) - sql_query = """ + sql_query: Final = """ SELECT user_email FROM "LiteLLM_UserTable" WHERE user_id = ANY($1::TEXT[]); """ - _result = await prisma_client.db.query_raw(sql_query, _team_member_user_ids) + _result: Final = await prisma_client.db.query_raw(sql_query, _team_member_user_ids) verbose_logger.debug("Email Alerting: Got all Emails for team, emails=%s", _result) if _result is None: return [] - emails = [] + emails: Final = [] for user in _result: if user and isinstance(user, dict) and user.get("user_email", None) is not None: emails.append(user.get("user_email")) @@ -71,8 +72,8 @@ async def send_team_budget_alert(webhook_event: WebhookEvent) -> bool: """ from litellm.proxy.utils import send_email - _team_id = webhook_event.team_id - team_alias = webhook_event.team_alias + _team_id: Final = webhook_event.team_id + team_alias: Final = webhook_event.team_alias verbose_logger.debug("Email Alerting: Sending Team Budget Alert for team=%s", team_alias) email_logo_url = os.getenv("SMTP_SENDER_LOGO", os.getenv("EMAIL_LOGO_URL", None)) @@ -86,12 +87,12 @@ async def send_team_budget_alert(webhook_event: WebhookEvent) -> bool: email_logo_url = LITELLM_LOGO_URL if email_support_contact is None: email_support_contact = LITELLM_SUPPORT_CONTACT - recipient_emails = await get_all_team_member_emails(_team_id) - recipient_emails_str: str = ",".join(recipient_emails) + recipient_emails: Final = await get_all_team_member_emails(_team_id) + recipient_emails_str: Final[str] = ",".join(recipient_emails) verbose_logger.debug("Email Alerting: Sending team budget alert to %s", recipient_emails_str) - event_name = webhook_event.event_message - max_budget = webhook_event.max_budget + event_name: Final = webhook_event.event_message + max_budget: Final = webhook_event.max_budget email_html_content = "Alert from LiteLLM Server" if recipient_emails_str is None: @@ -115,7 +116,7 @@ async def send_team_budget_alert(webhook_event: WebhookEvent) -> bool: The LiteLLM team
""" - email_event = { + email_event: Final = { "to": recipient_emails_str, "subject": f"LiteLLM {event_name} for Team {team_alias}", "html": email_html_content, diff --git a/litellm/integrations/email_templates/email_footer.py b/litellm/integrations/email_templates/email_footer.py index feb692354a0..950496a5c07 100644 --- a/litellm/integrations/email_templates/email_footer.py +++ b/litellm/integrations/email_templates/email_footer.py @@ -1,4 +1,6 @@ -EMAIL_FOOTER = """ +from typing import Final + +EMAIL_FOOTER: Final = """